package main
import (
"context"
"encoding/json"
"net/http"
"os"
"sync"
"time"
"github.com/jackc/pgx/v5"
)
type traceKey struct{}
type queryKey struct{}
type SQLStatement struct {
SQL string `json:"sql"`
Duration int64 `json:"duration_ns"`
Error *string `json:"error"`
Calls int `json:"calls"`
}
type SQLRequest struct {
ID string `json:"request_id"`
Method string `json:"method"`
Path string `json:"path"`
Status int `json:"status"`
Count int `json:"count"`
Statements []SQLStatement `json:"statements"`
mu sync.Mutex
}
type queryStart struct {
at time.Time
sql string
}
type SQLTracer struct{}
func (SQLTracer) TraceQueryStart(ctx context.Context, _ *pgx.Conn, data pgx.TraceQueryStartData) context.Context {
return context.WithValue(ctx, queryKey{}, queryStart{time.Now(), data.SQL})
}
func (SQLTracer) TraceQueryEnd(ctx context.Context, _ *pgx.Conn, data pgx.TraceQueryEndData) {
record, ok := ctx.Value(traceKey{}).(*SQLRequest)
if !ok {
return
}
start := ctx.Value(queryKey{}).(queryStart)
var failure *string
if data.Err != nil {
msg := data.Err.Error()
failure = &msg
}
record.mu.Lock()
record.Statements = append(record.Statements, SQLStatement{start.sql, time.Since(start.at).Nanoseconds(), failure, 1})
record.mu.Unlock()
}
type statusWriter struct {
http.ResponseWriter
status int
}
func (w *statusWriter) WriteHeader(status int) {
w.status = status
w.ResponseWriter.WriteHeader(status)
}
func traceHTTP(next http.Handler, path string) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
record := &SQLRequest{ID: r.Header.Get("X-Harness-Request"), Method: r.Method, Path: r.URL.Path, Statements: []SQLStatement{}}
writer := &statusWriter{w, 200}
next.ServeHTTP(writer, r.WithContext(context.WithValue(r.Context(), traceKey{}, record)))
record.Status = writer.status
record.Count = len(record.Statements)
raw, err := json.Marshal(record)
if err != nil {
return
}
f, err := os.OpenFile(path, os.O_WRONLY|os.O_CREATE|os.O_APPEND, 0600)
if err != nil {
return
}
defer f.Close()
f.Write(append(raw, '\n'))
})
}