From e8f56a1ef497266e4ace2329fc13a2e266384f21 Mon Sep 17 00:00:00 2001 From: Gary Hansen Date: Wed, 8 Jul 2026 02:43:50 +1000 Subject: [PATCH] feat(receiver): storage layer and webhook payload compat Dual-dialect store (SQLite via modernc.org, MySQL via go-sql-driver, both pure Go) with order-tolerant start/complete upserts, filtered and paginated listing, and aggregate queries (per-day, top domains, query types, statuses, duration percentiles, top clients). Sender gains optional EXPLOREDNS_WEBHOOK_TOKEN bearer auth; a round-trip test pins receiver structs byte-compatible with the sender payloads. Note: go directive moves to 1.25.0, required by modernc.org/sqlite. CI reads the version from go.mod so GOTOOLCHAIN=local stays satisfied. Co-Authored-By: Claude Fable 5 --- go.mod | 25 +- go.sum | 71 +++- internal/receiver/store/dialect.go | 103 +++++ internal/receiver/store/events.go | 39 ++ internal/receiver/store/query.go | 322 +++++++++++++++ internal/receiver/store/store.go | 217 ++++++++++ internal/receiver/store/store_test.go | 544 ++++++++++++++++++++++++++ web/api/handler.go | 2 +- web/api/webhook.go | 10 +- web/api/webhook_compat_test.go | 126 ++++++ web/api/webhook_test.go | 32 +- 11 files changed, 1468 insertions(+), 23 deletions(-) create mode 100644 internal/receiver/store/dialect.go create mode 100644 internal/receiver/store/events.go create mode 100644 internal/receiver/store/query.go create mode 100644 internal/receiver/store/store.go create mode 100644 internal/receiver/store/store_test.go create mode 100644 web/api/webhook_compat_test.go diff --git a/go.mod b/go.mod index d12f85c..008327e 100644 --- a/go.mod +++ b/go.mod @@ -1,16 +1,27 @@ module gitea.hansenits.com.au/hits/ExploreDNS -go 1.24.0 +go 1.25.0 require ( + github.com/go-sql-driver/mysql v1.10.0 github.com/miekg/dns v1.1.72 - golang.org/x/net v0.48.0 + golang.org/x/net v0.54.0 + modernc.org/sqlite v1.53.0 ) require ( - golang.org/x/mod v0.31.0 // indirect - golang.org/x/sync v0.19.0 // indirect - golang.org/x/sys v0.39.0 // indirect - golang.org/x/text v0.32.0 // indirect - golang.org/x/tools v0.40.0 // indirect + filippo.io/edwards25519 v1.2.0 // indirect + github.com/dustin/go-humanize v1.0.1 // indirect + github.com/google/uuid v1.6.0 // indirect + github.com/mattn/go-isatty v0.0.20 // indirect + github.com/ncruces/go-strftime v1.0.0 // indirect + github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect + golang.org/x/mod v0.36.0 // indirect + golang.org/x/sync v0.20.0 // indirect + golang.org/x/sys v0.44.0 // indirect + golang.org/x/text v0.37.0 // indirect + golang.org/x/tools v0.45.0 // indirect + modernc.org/libc v1.73.4 // indirect + modernc.org/mathutil v1.7.1 // indirect + modernc.org/memory v1.11.0 // indirect ) diff --git a/go.sum b/go.sum index 04b719e..88c35a0 100644 --- a/go.sum +++ b/go.sum @@ -1,16 +1,63 @@ +filippo.io/edwards25519 v1.2.0 h1:crnVqOiS4jqYleHd9vaKZ+HKtHfllngJIiOpNpoJsjo= +filippo.io/edwards25519 v1.2.0/go.mod h1:xzAOLCNug/yB62zG1bQ8uziwrIqIuxhctzJT18Q77mc= +github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY= +github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto= +github.com/go-sql-driver/mysql v1.10.0 h1:Q+1LV8DkHJvSYAdR83XzuhDaTykuDx0l6fkXxoWCWfw= +github.com/go-sql-driver/mysql v1.10.0/go.mod h1:M+cqaI7+xxXGG9swrdeUIoPG3Y3KCkF0pZej+SK+nWk= github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI= github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY= +github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs= +github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e/go.mod h1:boTsfXsheKC2y+lKOCMpSfarhxDeIzfZG1jqGcPl3cA= +github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= +github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k= +github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM= +github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY= +github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= github.com/miekg/dns v1.1.72 h1:vhmr+TF2A3tuoGNkLDFK9zi36F2LS+hKTRW0Uf8kbzI= github.com/miekg/dns v1.1.72/go.mod h1:+EuEPhdHOsfk6Wk5TT2CzssZdqkmFhf8r+aVyDEToIs= -golang.org/x/mod v0.31.0 h1:HaW9xtz0+kOcWKwli0ZXy79Ix+UW/vOfmWI5QVd2tgI= -golang.org/x/mod v0.31.0/go.mod h1:43JraMp9cGx1Rx3AqioxrbrhNsLl2l/iNAvuBkrezpg= -golang.org/x/net v0.48.0 h1:zyQRTTrjc33Lhh0fBgT/H3oZq9WuvRR5gPC70xpDiQU= -golang.org/x/net v0.48.0/go.mod h1:+ndRgGjkh8FGtu1w1FGbEC31if4VrNVMuKTgcAAnQRY= -golang.org/x/sync v0.19.0 h1:vV+1eWNmZ5geRlYjzm2adRgW2/mcpevXNg50YZtPCE4= -golang.org/x/sync v0.19.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI= -golang.org/x/sys v0.39.0 h1:CvCKL8MeisomCi6qNZ+wbb0DN9E5AATixKsvNtMoMFk= -golang.org/x/sys v0.39.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks= -golang.org/x/text v0.32.0 h1:ZD01bjUt1FQ9WJ0ClOL5vxgxOI/sVCNgX1YtKwcY0mU= -golang.org/x/text v0.32.0/go.mod h1:o/rUWzghvpD5TXrTIBuJU77MTaN0ljMWE47kxGJQ7jY= -golang.org/x/tools v0.40.0 h1:yLkxfA+Qnul4cs9QA3KnlFu0lVmd8JJfoq+E41uSutA= -golang.org/x/tools v0.40.0/go.mod h1:Ik/tzLRlbscWpqqMRjyWYDisX8bG13FrdXp3o4Sr9lc= +github.com/ncruces/go-strftime v1.0.0 h1:HMFp8mLCTPp341M/ZnA4qaf7ZlsbTc+miZjCLOFAw7w= +github.com/ncruces/go-strftime v1.0.0/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls= +github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE= +github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo= +golang.org/x/mod v0.36.0 h1:JJjpVx6myfUsUdAzZuOSTTmRE0PfZeNWzzvKrP7amb4= +golang.org/x/mod v0.36.0/go.mod h1:moc6ELqsWcOw5Ef3xVprK5ul/MvtVvkIXLziUOICjUQ= +golang.org/x/net v0.54.0 h1:2zJIZAxAHV/OHCDTCOHAYehQzLfSXuf/5SoL/Dv6w/w= +golang.org/x/net v0.54.0/go.mod h1:Sj4oj8jK6XmHpBZU/zWHw3BV3abl4Kvi+Ut7cQcY+cQ= +golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4= +golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= +golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.44.0 h1:ildZl3J4uzeKP07r2F++Op7E9B29JRUy+a27EibtBTQ= +golang.org/x/sys v0.44.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= +golang.org/x/text v0.37.0 h1:Cqjiwd9eSg8e0QAkyCaQTNHFIIzWtidPahFWR83rTrc= +golang.org/x/text v0.37.0/go.mod h1:a5sjxXGs9hsn/AJVwuElvCAo9v8QYLzvavO5z2PiM38= +golang.org/x/tools v0.45.0 h1:18qN3FAooORvApf5XjCXgsuayZOEtXf6JK18I3+ONa8= +golang.org/x/tools v0.45.0/go.mod h1:LuUGqqaXcXMEFEruIVJVm5mgDD8vww/z/SR1gQ4uE/0= +modernc.org/cc/v4 v4.28.4 h1:Hd/4Es+MBj+/7hSdZaisNyu6bv3V0Dp2MdllyfqaH+c= +modernc.org/cc/v4 v4.28.4/go.mod h1:OnovgIhbbMXMu1aISnJ0wvVD1KnW+cAUJkIrAWh+kVI= +modernc.org/ccgo/v4 v4.34.4 h1:OVnSOWQjVKOYkFxoHYB+qQmSHK5gqMqARM+K9DpR/Ws= +modernc.org/ccgo/v4 v4.34.4/go.mod h1:qdKqE8FNIYyysougB1RX9MxCzp5oJOcQXSobANJ4TuE= +modernc.org/fileutil v1.4.0 h1:j6ZzNTftVS054gi281TyLjHPp6CPHr2KCxEXjEbD6SM= +modernc.org/fileutil v1.4.0/go.mod h1:EqdKFDxiByqxLk8ozOxObDSfcVOv/54xDs/DUHdvCUU= +modernc.org/gc/v2 v2.6.5 h1:nyqdV8q46KvTpZlsw66kWqwXRHdjIlJOhG6kxiV/9xI= +modernc.org/gc/v2 v2.6.5/go.mod h1:YgIahr1ypgfe7chRuJi2gD7DBQiKSLMPgBQe9oIiito= +modernc.org/gc/v3 v3.1.3 h1:6QAplYyVO+KdPW3pGnqmJDUxtkec8ooEWvks/hhU3lc= +modernc.org/gc/v3 v3.1.3/go.mod h1:HFK/6AGESC7Ex+EZJhJ2Gni6cTaYpSMmU/cT9RmlfYY= +modernc.org/goabi0 v0.2.0 h1:HvEowk7LxcPd0eq6mVOAEMai46V+i7Jrj13t4AzuNks= +modernc.org/goabi0 v0.2.0/go.mod h1:CEFRnnJhKvWT1c1JTI3Avm+tgOWbkOu5oPA8eH8LnMI= +modernc.org/libc v1.73.4 h1:+ra4Ui8ngyt8HDcO1FTDPWlkAh6yOdaO2yAoh8MddQA= +modernc.org/libc v1.73.4/go.mod h1:DXZ3eO8qMCNn2SnmTNCiC71nJ9Rcq3PsnpU6Vc4rWK8= +modernc.org/mathutil v1.7.1 h1:GCZVGXdaN8gTqB1Mf/usp1Y/hSqgI2vAGGP4jZMCxOU= +modernc.org/mathutil v1.7.1/go.mod h1:4p5IwJITfppl0G4sUEDtCr4DthTaT47/N3aT6MhfgJg= +modernc.org/memory v1.11.0 h1:o4QC8aMQzmcwCK3t3Ux/ZHmwFPzE6hf2Y5LbkRs+hbI= +modernc.org/memory v1.11.0/go.mod h1:/JP4VbVC+K5sU2wZi9bHoq2MAkCnrt2r98UGeSK7Mjw= +modernc.org/opt v0.2.0 h1:tGyef5ApycA7FSEOMraay9SaTk5zmbx7Tu+cJs4QKZg= +modernc.org/opt v0.2.0/go.mod h1:03fq9lsNfvkYSfxrfUhZCWPk1lm4cq4N+Bh//bEtgns= +modernc.org/sortutil v1.2.1 h1:+xyoGf15mM3NMlPDnFqrteY07klSFxLElE2PVuWIJ7w= +modernc.org/sortutil v1.2.1/go.mod h1:7ZI3a3REbai7gzCLcotuw9AC4VZVpYMjDzETGsSMqJE= +modernc.org/sqlite v1.53.0 h1:20WG8N9q4ji/dEqGk4uiI0c6OPjSeLTNYGFCc3+7c1M= +modernc.org/sqlite v1.53.0/go.mod h1:xoEpOIpGrgT48H5iiyt/YXPCZPEzlfmfFwtk8Lklw8s= +modernc.org/strutil v1.2.1 h1:UneZBkQA+DX2Rp35KcM69cSsNES9ly8mQWD71HKlOA0= +modernc.org/strutil v1.2.1/go.mod h1:EHkiggD70koQxjVdSBM3JKM7k6L0FbGE5eymy9i3B9A= +modernc.org/token v1.1.0 h1:Xl7Ap9dKaEs5kLoOQeQmPWevfnk/DM5qcLcYlA8ys6Y= +modernc.org/token v1.1.0/go.mod h1:UGzOrNV1mAFSEB63lOFHIpNRUVMvYTc6yu1SMY/XTDM= diff --git a/internal/receiver/store/dialect.go b/internal/receiver/store/dialect.go new file mode 100644 index 0000000..28daf16 --- /dev/null +++ b/internal/receiver/store/dialect.go @@ -0,0 +1,103 @@ +package store + +import ( + "context" + "database/sql" + "strings" +) + +// dialect abstracts the SQL syntax differences between the sqlite and +// mysql backends. Both drivers use '?' placeholders, and all datetimes are +// bound and scanned as "2006-01-02 15:04:05.000000" UTC strings so value +// handling stays identical; only the constructs below diverge. +type dialect interface { + name() string + + // datetimeType is the column type for datetime values. sqlite has no + // datetime type (TEXT affinity stores our formatted strings verbatim, + // which sort and compare lexically); mysql needs DATETIME(6) because + // plain DATETIME truncates the sub-second precision we write. + datetimeType() string + + // upsert appends the insert-or-update clause to an INSERT statement, + // overwriting cols with the values from the attempted insert. sqlite + // spells this ON CONFLICT(id) DO UPDATE SET col=excluded.col; mysql + // spells it ON DUPLICATE KEY UPDATE col=VALUES(col) (VALUES() is + // deprecated in MySQL 8.0.20+ but is the only form MariaDB and older + // MySQL also accept). + upsert(insert string, cols []string) string + + // createIndex creates an index if it does not exist. sqlite supports + // CREATE INDEX IF NOT EXISTS; mysql has no IF NOT EXISTS for indexes, + // so existence is checked via information_schema first. + createIndex(ctx context.Context, db *sql.DB, index, table, column string) error + + // likeContains returns a "col contains substr" predicate plus its + // bind argument, with LIKE wildcards in substr escaped. The ESCAPE + // literal differs: mysql string literals treat backslash as an escape + // character (so the SQL needs '\\'), sqlite ones do not (so it needs + // '\', and has no default escape character at all). + likeContains(col, substr string) (predicate string, arg string) +} + +// escapeLike backslash-escapes LIKE pattern metacharacters in s. +func escapeLike(s string) string { + r := strings.NewReplacer(`\`, `\\`, `%`, `\%`, `_`, `\_`) + return r.Replace(s) +} + +type sqliteDialect struct{} + +func (sqliteDialect) name() string { return "sqlite" } +func (sqliteDialect) datetimeType() string { return "TEXT" } + +func (sqliteDialect) upsert(insert string, cols []string) string { + sets := make([]string, len(cols)) + for i, c := range cols { + sets[i] = c + "=excluded." + c + } + return insert + " ON CONFLICT(id) DO UPDATE SET " + strings.Join(sets, ", ") +} + +func (sqliteDialect) createIndex(ctx context.Context, db *sql.DB, index, table, column string) error { + _, err := db.ExecContext(ctx, + "CREATE INDEX IF NOT EXISTS "+index+" ON "+table+" ("+column+")") + return err +} + +func (sqliteDialect) likeContains(col, substr string) (string, string) { + return col + ` LIKE ? ESCAPE '\'`, "%" + escapeLike(substr) + "%" +} + +type mysqlDialect struct{} + +func (mysqlDialect) name() string { return "mysql" } +func (mysqlDialect) datetimeType() string { return "DATETIME(6)" } + +func (mysqlDialect) upsert(insert string, cols []string) string { + sets := make([]string, len(cols)) + for i, c := range cols { + sets[i] = c + "=VALUES(" + c + ")" + } + return insert + " ON DUPLICATE KEY UPDATE " + strings.Join(sets, ", ") +} + +func (mysqlDialect) createIndex(ctx context.Context, db *sql.DB, index, table, column string) error { + var n int + err := db.QueryRowContext(ctx, + `SELECT COUNT(*) FROM information_schema.statistics + WHERE table_schema = DATABASE() AND table_name = ? AND index_name = ?`, + table, index).Scan(&n) + if err != nil { + return err + } + if n > 0 { + return nil + } + _, err = db.ExecContext(ctx, "CREATE INDEX "+index+" ON "+table+" ("+column+")") + return err +} + +func (mysqlDialect) likeContains(col, substr string) (string, string) { + return col + ` LIKE ? ESCAPE '\\'`, "%" + escapeLike(substr) + "%" +} diff --git a/internal/receiver/store/events.go b/internal/receiver/store/events.go new file mode 100644 index 0000000..fa12d55 --- /dev/null +++ b/internal/receiver/store/events.go @@ -0,0 +1,39 @@ +package store + +import ( + "encoding/json" + "time" +) + +// StartEvent and CompleteEvent are the receiver-side decodings of the +// webhook payloads posted by the main app. Field names and JSON tags must +// stay in sync with webhookStartEvent/webhookCompleteEvent in +// web/api/webhook.go; a compatibility test in web/api enforces this. + +// StartEvent mirrors web/api webhookStartEvent. +type StartEvent struct { + Event string `json:"event"` + ID string `json:"id"` + Domain string `json:"domain"` + QueryType string `json:"query_type"` + AllRoots bool `json:"all_roots"` + ClientIP string `json:"client_ip"` + StartedAt time.Time `json:"started_at"` +} + +// CompleteEvent mirrors web/api webhookCompleteEvent. Summary is kept as +// raw JSON: the receiver stores it verbatim and never interprets it. +type CompleteEvent struct { + Event string `json:"event"` + ID string `json:"id"` + Domain string `json:"domain"` + QueryType string `json:"query_type"` + ClientIP string `json:"client_ip"` + StartedAt time.Time `json:"started_at"` + DoneAt time.Time `json:"done_at"` + DurationMS int64 `json:"duration_ms"` + Status string `json:"status"` + Error string `json:"error,omitempty"` + ResultCount int `json:"result_count"` + Summary json.RawMessage `json:"summary"` +} diff --git a/internal/receiver/store/query.go b/internal/receiver/store/query.go new file mode 100644 index 0000000..c8a438c --- /dev/null +++ b/internal/receiver/store/query.go @@ -0,0 +1,322 @@ +package store + +import ( + "context" + "database/sql" + "math" + "sort" + "strings" + "time" +) + +// Traversal is one stored traversal row. Pointer fields are NULL until the +// complete event arrives. +type Traversal struct { + ID string + Domain string + QueryType string + AllRoots bool + ClientIP string + StartedAt time.Time + DoneAt *time.Time + DurationMS *int64 + Status string + Error string + ResultCount *int + Summary string + FirstSeen time.Time + LastSeen time.Time +} + +// ListFilter narrows ListTraversals. Zero values mean "no constraint"; +// From/To bound started_at inclusively. +type ListFilter struct { + Domain string // substring match + Status string + From time.Time + To time.Time +} + +const defaultListLimit = 50 + +// where renders the filter as a WHERE clause (or "") plus bind args. +func (f ListFilter) where(d dialect) (string, []any) { + var preds []string + var args []any + if f.Domain != "" { + p, arg := d.likeContains("domain", f.Domain) + preds = append(preds, p) + args = append(args, arg) + } + if f.Status != "" { + preds = append(preds, "status = ?") + args = append(args, f.Status) + } + if !f.From.IsZero() { + preds = append(preds, "started_at >= ?") + args = append(args, fmtTime(f.From)) + } + if !f.To.IsZero() { + preds = append(preds, "started_at <= ?") + args = append(args, fmtTime(f.To)) + } + if len(preds) == 0 { + return "", nil + } + return " WHERE " + strings.Join(preds, " AND "), args +} + +// ListTraversals returns one page of matching traversals ordered newest +// first, plus the total match count. limit <= 0 selects a default page size. +func (s *Store) ListTraversals(ctx context.Context, f ListFilter, limit, offset int) ([]Traversal, int, error) { + where, args := f.where(s.d) + + var total int + if err := s.db.QueryRowContext(ctx, "SELECT COUNT(*) FROM traversals"+where, args...).Scan(&total); err != nil { + return nil, 0, err + } + + if limit <= 0 { + limit = defaultListLimit + } + if offset < 0 { + offset = 0 + } + q := `SELECT id, domain, query_type, all_roots, client_ip, started_at, done_at, + duration_ms, status, error, result_count, summary, first_seen, last_seen + FROM traversals` + where + ` ORDER BY started_at DESC, id LIMIT ? OFFSET ?` + rows, err := s.db.QueryContext(ctx, q, append(args, limit, offset)...) + if err != nil { + return nil, 0, err + } + defer rows.Close() + + var out []Traversal + for rows.Next() { + var ( + tr Traversal + startedAt, doneAt sql.NullString + errMsg, summary sql.NullString + durationMS sql.NullInt64 + resultCount sql.NullInt64 + firstSeen, lastSeen string + ) + if err := rows.Scan(&tr.ID, &tr.Domain, &tr.QueryType, &tr.AllRoots, &tr.ClientIP, + &startedAt, &doneAt, &durationMS, &tr.Status, &errMsg, &resultCount, + &summary, &firstSeen, &lastSeen); err != nil { + return nil, 0, err + } + tr.Error = errMsg.String + tr.Summary = summary.String + if durationMS.Valid { + v := durationMS.Int64 + tr.DurationMS = &v + } + if resultCount.Valid { + v := int(resultCount.Int64) + tr.ResultCount = &v + } + if startedAt.Valid { + if tr.StartedAt, err = parseDBTime(startedAt.String); err != nil { + return nil, 0, err + } + } + if doneAt.Valid { + t, err := parseDBTime(doneAt.String) + if err != nil { + return nil, 0, err + } + tr.DoneAt = &t + } + if tr.FirstSeen, err = parseDBTime(firstSeen); err != nil { + return nil, 0, err + } + if tr.LastSeen, err = parseDBTime(lastSeen); err != nil { + return nil, 0, err + } + out = append(out, tr) + } + return out, total, rows.Err() +} + +// Totals summarises whole-table traversal volume. Errors counts all-time +// error-status rows so callers can derive an error rate. +type Totals struct { + AllTime int + Last24h int + Last7d int + DistinctDomains int + DistinctClients int + Errors int +} + +// Totals returns whole-table totals with the recent windows measured back +// from now. COALESCE covers the empty table, where SUM yields NULL. +func (s *Store) Totals(ctx context.Context, now time.Time) (Totals, error) { + var t Totals + err := s.db.QueryRowContext(ctx, `SELECT COUNT(*), + COALESCE(SUM(CASE WHEN started_at >= ? THEN 1 ELSE 0 END), 0), + COALESCE(SUM(CASE WHEN started_at >= ? THEN 1 ELSE 0 END), 0), + COUNT(DISTINCT domain), + COUNT(DISTINCT client_ip), + COALESCE(SUM(CASE WHEN status = ? THEN 1 ELSE 0 END), 0) + FROM traversals`, + fmtTime(now.Add(-24*time.Hour)), fmtTime(now.AddDate(0, 0, -7)), StatusError). + Scan(&t.AllTime, &t.Last24h, &t.Last7d, &t.DistinctDomains, &t.DistinctClients, &t.Errors) + return t, err +} + +// DayStats is one calendar day's traversal volume. +type DayStats struct { + Day string // "2006-01-02" (UTC) + Total int + Errors int +} + +// StatsPerDay returns per-day totals and error counts for the last days +// calendar days (UTC), oldest first. Days with no traffic are omitted. +func (s *Store) StatsPerDay(ctx context.Context, days int) ([]DayStats, error) { + if days <= 0 { + days = 1 + } + midnight := time.Now().UTC().Truncate(24 * time.Hour) + cutoff := midnight.AddDate(0, 0, -(days - 1)) + // substr on the stored value yields "YYYY-MM-DD" in both backends + // (MySQL casts DATETIME to its string form implicitly). + rows, err := s.db.QueryContext(ctx, `SELECT substr(started_at, 1, 10) AS day, + COUNT(*), SUM(CASE WHEN status = ? THEN 1 ELSE 0 END) + FROM traversals WHERE started_at >= ? GROUP BY day ORDER BY day`, + StatusError, fmtTime(cutoff)) + if err != nil { + return nil, err + } + defer rows.Close() + + var out []DayStats + for rows.Next() { + var d DayStats + if err := rows.Scan(&d.Day, &d.Total, &d.Errors); err != nil { + return nil, err + } + out = append(out, d) + } + return out, rows.Err() +} + +// NameCount is a generic (value, count) aggregate row. +type NameCount struct { + Name string + Count int +} + +// TopDomains returns the most-queried domains since the given time (zero = +// all time), most frequent first, at most limit rows. +func (s *Store) TopDomains(ctx context.Context, since time.Time, limit int) ([]NameCount, error) { + return s.countBy(ctx, "domain", since, limit) +} + +// TopClientIPs returns the most active client IPs since the given time. +func (s *Store) TopClientIPs(ctx context.Context, since time.Time, limit int) ([]NameCount, error) { + return s.countBy(ctx, "client_ip", since, limit) +} + +// QueryTypeCounts returns traversal counts per query type since the given +// time, most frequent first. +func (s *Store) QueryTypeCounts(ctx context.Context, since time.Time) ([]NameCount, error) { + return s.countBy(ctx, "query_type", since, 0) +} + +// StatusCounts returns traversal counts per status since the given time, +// most frequent first. +func (s *Store) StatusCounts(ctx context.Context, since time.Time) ([]NameCount, error) { + return s.countBy(ctx, "status", since, 0) +} + +// countBy groups rows by col and counts them. col is always one of the +// fixed column names above, never user input. +func (s *Store) countBy(ctx context.Context, col string, since time.Time, limit int) ([]NameCount, error) { + q := "SELECT " + col + ", COUNT(*) AS n FROM traversals" + var args []any + if !since.IsZero() { + q += " WHERE started_at >= ?" + args = append(args, fmtTime(since)) + } + q += " GROUP BY " + col + " ORDER BY n DESC, " + col + if limit > 0 { + q += " LIMIT ?" + args = append(args, limit) + } + rows, err := s.db.QueryContext(ctx, q, args...) + if err != nil { + return nil, err + } + defer rows.Close() + + var out []NameCount + for rows.Next() { + var nc NameCount + if err := rows.Scan(&nc.Name, &nc.Count); err != nil { + return nil, err + } + out = append(out, nc) + } + return out, rows.Err() +} + +// DurationStats summarises completed-traversal durations. +type DurationStats struct { + Count int + AvgMS float64 + P50MS int64 + P95MS int64 +} + +// Durations returns duration statistics for traversals started since the +// given time (zero = all time). Percentiles are computed in Go (nearest +// rank) because SQL percentile support differs between the backends. +func (s *Store) Durations(ctx context.Context, since time.Time) (DurationStats, error) { + q := "SELECT duration_ms FROM traversals WHERE duration_ms IS NOT NULL" + var args []any + if !since.IsZero() { + q += " AND started_at >= ?" + args = append(args, fmtTime(since)) + } + rows, err := s.db.QueryContext(ctx, q, args...) + if err != nil { + return DurationStats{}, err + } + defer rows.Close() + + var values []int64 + var sum int64 + for rows.Next() { + var v int64 + if err := rows.Scan(&v); err != nil { + return DurationStats{}, err + } + values = append(values, v) + sum += v + } + if err := rows.Err(); err != nil { + return DurationStats{}, err + } + if len(values) == 0 { + return DurationStats{}, nil + } + sort.Slice(values, func(i, j int) bool { return values[i] < values[j] }) + return DurationStats{ + Count: len(values), + AvgMS: float64(sum) / float64(len(values)), + P50MS: percentile(values, 0.50), + P95MS: percentile(values, 0.95), + }, nil +} + +// percentile returns the nearest-rank percentile of sorted values. +func percentile(sorted []int64, q float64) int64 { + rank := int(math.Ceil(q * float64(len(sorted)))) + if rank < 1 { + rank = 1 + } + return sorted[rank-1] +} diff --git a/internal/receiver/store/store.go b/internal/receiver/store/store.go new file mode 100644 index 0000000..428cb75 --- /dev/null +++ b/internal/receiver/store/store.go @@ -0,0 +1,217 @@ +// Package store persists webhook telemetry events from the ExploreDNS API +// server into MySQL or SQLite through a shared database/sql layer. +package store + +import ( + "context" + "database/sql" + "fmt" + "os" + "path/filepath" + "time" + + "github.com/go-sql-driver/mysql" + _ "modernc.org/sqlite" +) + +// Traversal status values, mirroring the job statuses in web/api. +const ( + StatusRunning = "running" + StatusComplete = "complete" + StatusError = "error" +) + +// schemaVersion is recorded in the schema_version table on first open so +// future releases can detect and migrate older layouts. +const schemaVersion = 1 + +// dbTimeLayout is the canonical datetime encoding: fixed-width so sqlite +// TEXT comparisons order chronologically, and a valid MySQL DATETIME(6) +// literal. +const dbTimeLayout = "2006-01-02 15:04:05.000000" + +// Store persists traversal telemetry events. +type Store struct { + db *sql.DB + d dialect +} + +// OpenSQLite opens (creating if needed) a SQLite-backed store at path, +// creating parent directories first. +func OpenSQLite(path string) (*Store, error) { + if dir := filepath.Dir(path); dir != "." && dir != "" { + if err := os.MkdirAll(dir, 0o755); err != nil { + return nil, fmt.Errorf("create sqlite directory: %w", err) + } + } + db, err := sql.Open("sqlite", path) + if err != nil { + return nil, fmt.Errorf("open sqlite %s: %w", path, err) + } + // A single connection serialises writers so concurrent ingests never + // see SQLITE_BUSY. + db.SetMaxOpenConns(1) + return open(db, sqliteDialect{}) +} + +// OpenMySQL opens a MySQL-backed store using a go-sql-driver DSN +// (user:pass@tcp(host:port)/dbname). +func OpenMySQL(dsn string) (*Store, error) { + if _, err := mysql.ParseDSN(dsn); err != nil { + return nil, fmt.Errorf("parse mysql dsn: %w", err) + } + db, err := sql.Open("mysql", dsn) + if err != nil { + return nil, fmt.Errorf("open mysql: %w", err) + } + return open(db, mysqlDialect{}) +} + +// RedactMySQLDSN returns dsn with any password replaced, safe for logging. +func RedactMySQLDSN(dsn string) string { + cfg, err := mysql.ParseDSN(dsn) + if err != nil { + return "(unparsable DSN)" + } + if cfg.Passwd != "" { + cfg.Passwd = "xxxxx" + } + return cfg.FormatDSN() +} + +func open(db *sql.DB, d dialect) (*Store, error) { + s := &Store{db: db, d: d} + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + if err := s.init(ctx); err != nil { + db.Close() + return nil, fmt.Errorf("init %s schema: %w", d.name(), err) + } + return s, nil +} + +// Backend names the database backend ("sqlite" or "mysql"). +func (s *Store) Backend() string { return s.d.name() } + +// Close releases the underlying database handle. +func (s *Store) Close() error { return s.db.Close() } + +// init creates the schema if missing and stamps schema_version. +func (s *Store) init(ctx context.Context) error { + dt := s.d.datetimeType() + stmts := []string{ + `CREATE TABLE IF NOT EXISTS schema_version (version INTEGER NOT NULL)`, + fmt.Sprintf(`CREATE TABLE IF NOT EXISTS traversals ( + id VARCHAR(64) NOT NULL PRIMARY KEY, + domain VARCHAR(255) NOT NULL DEFAULT '', + query_type VARCHAR(16) NOT NULL DEFAULT '', + all_roots BOOLEAN NOT NULL DEFAULT 0, + client_ip VARCHAR(64) NOT NULL DEFAULT '', + started_at %s NULL, + done_at %s NULL, + duration_ms BIGINT NULL, + status VARCHAR(16) NOT NULL DEFAULT '', + error TEXT NULL, + result_count INT NULL, + summary TEXT NULL, + first_seen %s NOT NULL, + last_seen %s NOT NULL + )`, dt, dt, dt, dt), + } + for _, q := range stmts { + if _, err := s.db.ExecContext(ctx, q); err != nil { + return err + } + } + for _, idx := range [][2]string{ + {"idx_traversals_started_at", "started_at"}, + {"idx_traversals_domain", "domain"}, + {"idx_traversals_client_ip", "client_ip"}, + {"idx_traversals_status", "status"}, + } { + if err := s.d.createIndex(ctx, s.db, idx[0], "traversals", idx[1]); err != nil { + return err + } + } + // Stamp the version on first creation only. Checked in Go because + // MySQL and sqlite disagree on FROM-less SELECT ... WHERE syntax. + var n int + if err := s.db.QueryRowContext(ctx, `SELECT COUNT(*) FROM schema_version`).Scan(&n); err != nil { + return err + } + if n == 0 { + _, err := s.db.ExecContext(ctx, `INSERT INTO schema_version (version) VALUES (?)`, schemaVersion) + return err + } + return nil +} + +// RecordStart upserts a "start" event. Duplicate deliveries refresh the +// request attributes but never touch status or completion fields, so a +// retried start arriving after the complete event cannot clobber the +// terminal state. +func (s *Store) RecordStart(ctx context.Context, ev StartEvent) error { + now := fmtTime(time.Now()) + q := s.d.upsert(`INSERT INTO traversals + (id, domain, query_type, all_roots, client_ip, started_at, status, first_seen, last_seen) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)`, + []string{"domain", "query_type", "all_roots", "client_ip", "started_at", "last_seen"}) + _, err := s.db.ExecContext(ctx, q, + ev.ID, ev.Domain, ev.QueryType, ev.AllRoots, ev.ClientIP, + dbTime(ev.StartedAt), StatusRunning, now, now) + return err +} + +// RecordComplete upserts a "complete" event. It creates the row when the +// start event has not arrived (out-of-order delivery) and overwrites the +// completion fields on duplicates; all_roots and first_seen are start-only +// and left untouched. +func (s *Store) RecordComplete(ctx context.Context, ev CompleteEvent) error { + now := fmtTime(time.Now()) + q := s.d.upsert(`INSERT INTO traversals + (id, domain, query_type, client_ip, started_at, done_at, duration_ms, + status, error, result_count, summary, first_seen, last_seen) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, + []string{"domain", "query_type", "client_ip", "started_at", "done_at", + "duration_ms", "status", "error", "result_count", "summary", "last_seen"}) + var summary any + if len(ev.Summary) > 0 { + summary = string(ev.Summary) + } + _, err := s.db.ExecContext(ctx, q, + ev.ID, ev.Domain, ev.QueryType, ev.ClientIP, + dbTime(ev.StartedAt), dbTime(ev.DoneAt), ev.DurationMS, + ev.Status, nullIfEmpty(ev.Error), ev.ResultCount, summary, now, now) + return err +} + +// fmtTime encodes t for storage. +func fmtTime(t time.Time) string { return t.UTC().Format(dbTimeLayout) } + +// dbTime encodes t for a nullable datetime column; zero times become NULL +// (MySQL DATETIME cannot hold year 1). +func dbTime(t time.Time) any { + if t.IsZero() { + return nil + } + return fmtTime(t) +} + +func nullIfEmpty(s string) any { + if s == "" { + return nil + } + return s +} + +// parseDBTime decodes a stored datetime. Besides the canonical layout it +// accepts second precision and RFC 3339 in case the MySQL DSN enables +// parseTime (database/sql then hands strings back in RFC 3339). +func parseDBTime(s string) (time.Time, error) { + for _, layout := range []string{dbTimeLayout, "2006-01-02 15:04:05", time.RFC3339Nano} { + if t, err := time.ParseInLocation(layout, s, time.UTC); err == nil { + return t, nil + } + } + return time.Time{}, fmt.Errorf("unrecognised datetime %q", s) +} diff --git a/internal/receiver/store/store_test.go b/internal/receiver/store/store_test.go new file mode 100644 index 0000000..52bfc68 --- /dev/null +++ b/internal/receiver/store/store_test.go @@ -0,0 +1,544 @@ +package store + +import ( + "context" + "encoding/json" + "path/filepath" + "testing" + "time" +) + +func newTestStore(t *testing.T) *Store { + t.Helper() + s, err := OpenSQLite(filepath.Join(t.TempDir(), "receiver.db")) + if err != nil { + t.Fatalf("OpenSQLite: %v", err) + } + t.Cleanup(func() { s.Close() }) + return s +} + +func startEvent(id string) StartEvent { + return StartEvent{ + Event: "start", + ID: id, + Domain: "example.com", + QueryType: "A", + AllRoots: true, + ClientIP: "203.0.113.9", + StartedAt: time.Date(2026, 7, 1, 10, 0, 0, 0, time.UTC), + } +} + +func completeEvent(id string) CompleteEvent { + return CompleteEvent{ + Event: "complete", + ID: id, + Domain: "example.com", + QueryType: "A", + ClientIP: "203.0.113.9", + StartedAt: time.Date(2026, 7, 1, 10, 0, 0, 0, time.UTC), + DoneAt: time.Date(2026, 7, 1, 10, 0, 3, 0, time.UTC), + DurationMS: 3000, + Status: StatusComplete, + ResultCount: 4, + Summary: json.RawMessage(`{"answers":[{"probability":1,"records":["example.com 300 IN A 192.0.2.1"]}]}`), + } +} + +// getRow fetches a single traversal by exact id via ListTraversals. +func getRow(t *testing.T, s *Store, id string) Traversal { + t.Helper() + rows, _, err := s.ListTraversals(context.Background(), ListFilter{}, 1000, 0) + if err != nil { + t.Fatalf("ListTraversals: %v", err) + } + for _, r := range rows { + if r.ID == id { + return r + } + } + t.Fatalf("row %q not found", id) + return Traversal{} +} + +func TestRecordStartThenComplete(t *testing.T) { + s := newTestStore(t) + ctx := context.Background() + + if err := s.RecordStart(ctx, startEvent("t1")); err != nil { + t.Fatalf("RecordStart: %v", err) + } + row := getRow(t, s, "t1") + if row.Status != StatusRunning { + t.Errorf("Status = %q, want %q", row.Status, StatusRunning) + } + if !row.AllRoots { + t.Error("AllRoots = false, want true") + } + if row.DoneAt != nil || row.DurationMS != nil || row.ResultCount != nil { + t.Errorf("completion fields set before complete: %+v", row) + } + if !row.StartedAt.Equal(startEvent("t1").StartedAt) { + t.Errorf("StartedAt = %v, want %v", row.StartedAt, startEvent("t1").StartedAt) + } + + ev := completeEvent("t1") + if err := s.RecordComplete(ctx, ev); err != nil { + t.Fatalf("RecordComplete: %v", err) + } + row = getRow(t, s, "t1") + if row.Status != StatusComplete { + t.Errorf("Status = %q, want %q", row.Status, StatusComplete) + } + if row.DoneAt == nil || !row.DoneAt.Equal(ev.DoneAt) { + t.Errorf("DoneAt = %v, want %v", row.DoneAt, ev.DoneAt) + } + if row.DurationMS == nil || *row.DurationMS != 3000 { + t.Errorf("DurationMS = %v, want 3000", row.DurationMS) + } + if row.ResultCount == nil || *row.ResultCount != 4 { + t.Errorf("ResultCount = %v, want 4", row.ResultCount) + } + if row.Summary != string(ev.Summary) { + t.Errorf("Summary = %q, want %q", row.Summary, ev.Summary) + } + if !row.AllRoots { + t.Error("AllRoots clobbered by complete event") + } +} + +func TestCompleteBeforeStart(t *testing.T) { + s := newTestStore(t) + ctx := context.Background() + + if err := s.RecordComplete(ctx, completeEvent("t1")); err != nil { + t.Fatalf("RecordComplete: %v", err) + } + row := getRow(t, s, "t1") + if row.Status != StatusComplete { + t.Fatalf("Status = %q, want %q", row.Status, StatusComplete) + } + + // The delayed start must fill in start-only fields without touching + // the terminal state. + if err := s.RecordStart(ctx, startEvent("t1")); err != nil { + t.Fatalf("RecordStart: %v", err) + } + row = getRow(t, s, "t1") + if row.Status != StatusComplete { + t.Errorf("Status = %q after late start, want %q", row.Status, StatusComplete) + } + if row.DoneAt == nil || row.DurationMS == nil { + t.Errorf("completion fields lost after late start: %+v", row) + } + if !row.AllRoots { + t.Error("AllRoots not filled in by late start") + } +} + +func TestDuplicateStartAfterComplete(t *testing.T) { + s := newTestStore(t) + ctx := context.Background() + + if err := s.RecordStart(ctx, startEvent("t1")); err != nil { + t.Fatalf("RecordStart: %v", err) + } + if err := s.RecordComplete(ctx, completeEvent("t1")); err != nil { + t.Fatalf("RecordComplete: %v", err) + } + if err := s.RecordStart(ctx, startEvent("t1")); err != nil { + t.Fatalf("duplicate RecordStart: %v", err) + } + + row := getRow(t, s, "t1") + if row.Status != StatusComplete { + t.Errorf("Status = %q after duplicate start, want %q", row.Status, StatusComplete) + } + if row.DoneAt == nil { + t.Error("DoneAt lost after duplicate start") + } + if row.Summary == "" { + t.Error("Summary lost after duplicate start") + } +} + +func TestDuplicateCompleteIdempotent(t *testing.T) { + s := newTestStore(t) + ctx := context.Background() + + ev := completeEvent("t1") + ev.Status = StatusError + ev.Error = "traversal timed out after 5m0s" + for i := 0; i < 3; i++ { + if err := s.RecordComplete(ctx, ev); err != nil { + t.Fatalf("RecordComplete #%d: %v", i, err) + } + } + + rows, total, err := s.ListTraversals(ctx, ListFilter{}, 10, 0) + if err != nil { + t.Fatalf("ListTraversals: %v", err) + } + if total != 1 || len(rows) != 1 { + t.Fatalf("total = %d, len = %d, want 1 row", total, len(rows)) + } + if rows[0].Status != StatusError || rows[0].Error != ev.Error { + t.Errorf("row = %q/%q, want %q/%q", rows[0].Status, rows[0].Error, StatusError, ev.Error) + } +} + +// seedRows inserts a deterministic mixed dataset for list/aggregate tests. +func seedRows(t *testing.T, s *Store) { + t.Helper() + ctx := context.Background() + base := time.Date(2026, 7, 1, 0, 0, 0, 0, time.UTC) + rows := []struct { + id, domain, qtype, ip, status string + day int + durMS int64 + }{ + {"a1", "example.com", "A", "203.0.113.1", StatusComplete, 0, 100}, + {"a2", "example.com", "AAAA", "203.0.113.1", StatusComplete, 0, 200}, + {"a3", "sub.example.com", "A", "203.0.113.2", StatusError, 1, 300}, + {"a4", "other.net", "MX", "203.0.113.3", StatusComplete, 1, 400}, + {"a5", "other.net", "A", "203.0.113.1", StatusComplete, 2, 500}, + {"a6", "under_score.org", "A", "203.0.113.4", StatusComplete, 2, 600}, + } + for i, r := range rows { + started := base.AddDate(0, 0, r.day).Add(time.Duration(i) * time.Minute) + ev := CompleteEvent{ + ID: r.id, Domain: r.domain, QueryType: r.qtype, ClientIP: r.ip, + StartedAt: started, DoneAt: started.Add(time.Duration(r.durMS) * time.Millisecond), + DurationMS: r.durMS, Status: r.status, ResultCount: 1, + } + if r.status == StatusError { + ev.Error = "lookup failed" + } + if err := s.RecordComplete(ctx, ev); err != nil { + t.Fatalf("seed %s: %v", r.id, err) + } + } + // One still-running traversal on day 2. + if err := s.RecordStart(ctx, StartEvent{ + ID: "a7", Domain: "running.io", QueryType: "A", ClientIP: "203.0.113.5", + StartedAt: base.AddDate(0, 0, 2).Add(time.Hour), + }); err != nil { + t.Fatalf("seed a7: %v", err) + } +} + +func TestListTraversalsFiltersAndPagination(t *testing.T) { + s := newTestStore(t) + seedRows(t, s) + ctx := context.Background() + + rows, total, err := s.ListTraversals(ctx, ListFilter{}, 3, 0) + if err != nil { + t.Fatalf("ListTraversals: %v", err) + } + if total != 7 || len(rows) != 3 { + t.Fatalf("total = %d, len = %d, want 7 and 3", total, len(rows)) + } + // Newest first: a7 (day2+1h) then a6, a5. + if rows[0].ID != "a7" || rows[1].ID != "a6" || rows[2].ID != "a5" { + t.Errorf("page 1 order = %s,%s,%s, want a7,a6,a5", rows[0].ID, rows[1].ID, rows[2].ID) + } + rows, _, err = s.ListTraversals(ctx, ListFilter{}, 3, 3) + if err != nil { + t.Fatalf("ListTraversals page 2: %v", err) + } + if rows[0].ID != "a4" || rows[1].ID != "a3" || rows[2].ID != "a2" { + t.Errorf("page 2 order = %s,%s,%s, want a4,a3,a2", rows[0].ID, rows[1].ID, rows[2].ID) + } + + // Domain substring matches example.com and sub.example.com. + rows, total, err = s.ListTraversals(ctx, ListFilter{Domain: "example"}, 0, 0) + if err != nil { + t.Fatalf("domain filter: %v", err) + } + if total != 3 { + t.Errorf("domain filter total = %d, want 3", total) + } + for _, r := range rows { + if r.Domain != "example.com" && r.Domain != "sub.example.com" { + t.Errorf("domain filter matched %q", r.Domain) + } + } + + // LIKE metacharacters in the filter must be literal, not wildcards. + _, total, err = s.ListTraversals(ctx, ListFilter{Domain: "%"}, 0, 0) + if err != nil { + t.Fatalf("percent filter: %v", err) + } + if total != 0 { + t.Errorf("%% filter total = %d, want 0", total) + } + _, total, err = s.ListTraversals(ctx, ListFilter{Domain: "under_score"}, 0, 0) + if err != nil { + t.Fatalf("underscore filter: %v", err) + } + if total != 1 { + t.Errorf("under_score filter total = %d, want 1", total) + } + + _, total, err = s.ListTraversals(ctx, ListFilter{Status: StatusError}, 0, 0) + if err != nil { + t.Fatalf("status filter: %v", err) + } + if total != 1 { + t.Errorf("status filter total = %d, want 1", total) + } + + base := time.Date(2026, 7, 1, 0, 0, 0, 0, time.UTC) + rows, total, err = s.ListTraversals(ctx, ListFilter{ + From: base.AddDate(0, 0, 1), + To: base.AddDate(0, 0, 2).Add(-time.Second), + }, 0, 0) + if err != nil { + t.Fatalf("time filter: %v", err) + } + if total != 2 { + t.Fatalf("time filter total = %d, want 2 (got %+v)", total, rows) + } + if rows[0].ID != "a4" || rows[1].ID != "a3" { + t.Errorf("time filter = %s,%s, want a4,a3", rows[0].ID, rows[1].ID) + } +} + +func TestStatsPerDay(t *testing.T) { + s := newTestStore(t) + ctx := context.Background() + now := time.Now().UTC() + for i, spec := range []struct { + daysAgo int + status string + }{ + {0, StatusComplete}, {0, StatusError}, {1, StatusComplete}, {5, StatusComplete}, + } { + ev := completeEvent(string(rune('a' + i))) + ev.StartedAt = now.AddDate(0, 0, -spec.daysAgo) + ev.DoneAt = ev.StartedAt.Add(time.Second) + ev.Status = spec.status + if err := s.RecordComplete(ctx, ev); err != nil { + t.Fatalf("seed: %v", err) + } + } + + days, err := s.StatsPerDay(ctx, 3) + if err != nil { + t.Fatalf("StatsPerDay: %v", err) + } + if len(days) != 2 { + t.Fatalf("len = %d, want 2 (%+v)", len(days), days) + } + yesterday := now.AddDate(0, 0, -1).Format("2006-01-02") + today := now.Format("2006-01-02") + if days[0].Day != yesterday || days[0].Total != 1 || days[0].Errors != 0 { + t.Errorf("day[0] = %+v, want %s total 1 errors 0", days[0], yesterday) + } + if days[1].Day != today || days[1].Total != 2 || days[1].Errors != 1 { + t.Errorf("day[1] = %+v, want %s total 2 errors 1", days[1], today) + } +} + +func TestTopDomains(t *testing.T) { + s := newTestStore(t) + seedRows(t, s) + + top, err := s.TopDomains(context.Background(), time.Time{}, 2) + if err != nil { + t.Fatalf("TopDomains: %v", err) + } + want := []NameCount{{"example.com", 2}, {"other.net", 2}} + if len(top) != 2 || top[0] != want[0] || top[1] != want[1] { + t.Errorf("TopDomains = %+v, want %+v", top, want) + } + + // Window excludes day 0 rows (a1, a2), so example.com drops out and + // other.net (a4, a5) leads. + since := time.Date(2026, 7, 2, 0, 0, 0, 0, time.UTC) + top, err = s.TopDomains(context.Background(), since, 10) + if err != nil { + t.Fatalf("TopDomains since: %v", err) + } + if len(top) != 4 || top[0] != (NameCount{"other.net", 2}) { + t.Errorf("TopDomains since = %+v, want other.net x2 leading 4 domains", top) + } +} + +func TestQueryTypeCounts(t *testing.T) { + s := newTestStore(t) + seedRows(t, s) + + counts, err := s.QueryTypeCounts(context.Background(), time.Time{}) + if err != nil { + t.Fatalf("QueryTypeCounts: %v", err) + } + want := []NameCount{{"A", 5}, {"AAAA", 1}, {"MX", 1}} + if len(counts) != 3 || counts[0] != want[0] || counts[1] != want[1] || counts[2] != want[2] { + t.Errorf("QueryTypeCounts = %+v, want %+v", counts, want) + } +} + +func TestStatusCounts(t *testing.T) { + s := newTestStore(t) + seedRows(t, s) + + counts, err := s.StatusCounts(context.Background(), time.Time{}) + if err != nil { + t.Fatalf("StatusCounts: %v", err) + } + want := []NameCount{{StatusComplete, 5}, {StatusError, 1}, {StatusRunning, 1}} + if len(counts) != 3 || counts[0] != want[0] || counts[1] != want[1] || counts[2] != want[2] { + t.Errorf("StatusCounts = %+v, want %+v", counts, want) + } +} + +func TestTopClientIPs(t *testing.T) { + s := newTestStore(t) + seedRows(t, s) + + top, err := s.TopClientIPs(context.Background(), time.Time{}, 1) + if err != nil { + t.Fatalf("TopClientIPs: %v", err) + } + if len(top) != 1 || top[0].Name != "203.0.113.1" || top[0].Count != 3 { + t.Errorf("TopClientIPs = %+v, want 203.0.113.1 x3", top) + } +} + +func TestDurations(t *testing.T) { + s := newTestStore(t) + seedRows(t, s) + + // Durations 100..600; the running row has none and is excluded. + stats, err := s.Durations(context.Background(), time.Time{}) + if err != nil { + t.Fatalf("Durations: %v", err) + } + if stats.Count != 6 { + t.Fatalf("Count = %d, want 6", stats.Count) + } + if stats.AvgMS != 350 { + t.Errorf("AvgMS = %v, want 350", stats.AvgMS) + } + if stats.P50MS != 300 { + t.Errorf("P50MS = %d, want 300", stats.P50MS) + } + if stats.P95MS != 600 { + t.Errorf("P95MS = %d, want 600", stats.P95MS) + } + + // Empty window. + stats, err = s.Durations(context.Background(), time.Date(2030, 1, 1, 0, 0, 0, 0, time.UTC)) + if err != nil { + t.Fatalf("Durations empty: %v", err) + } + if stats.Count != 0 || stats.AvgMS != 0 || stats.P50MS != 0 || stats.P95MS != 0 { + t.Errorf("empty Durations = %+v, want zeros", stats) + } +} + +func TestTotals(t *testing.T) { + s := newTestStore(t) + ctx := context.Background() + now := time.Now().UTC() + + tot, err := s.Totals(ctx, now) + if err != nil { + t.Fatalf("Totals empty: %v", err) + } + if tot != (Totals{}) { + t.Errorf("empty Totals = %+v, want zeros", tot) + } + + for _, spec := range []struct { + id, domain, ip, status string + ago time.Duration + }{ + {"t1", "a.com", "203.0.113.1", StatusComplete, time.Hour}, + {"t2", "a.com", "203.0.113.2", StatusError, 2 * time.Hour}, + {"t3", "b.net", "203.0.113.1", StatusComplete, 48 * time.Hour}, + {"t4", "c.org", "203.0.113.3", StatusComplete, 10 * 24 * time.Hour}, + } { + ev := completeEvent(spec.id) + ev.Domain = spec.domain + ev.ClientIP = spec.ip + ev.Status = spec.status + ev.StartedAt = now.Add(-spec.ago) + ev.DoneAt = ev.StartedAt.Add(time.Second) + if err := s.RecordComplete(ctx, ev); err != nil { + t.Fatalf("seed %s: %v", spec.id, err) + } + } + + tot, err = s.Totals(ctx, now) + if err != nil { + t.Fatalf("Totals: %v", err) + } + want := Totals{AllTime: 4, Last24h: 2, Last7d: 3, DistinctDomains: 3, DistinctClients: 3, Errors: 1} + if tot != want { + t.Errorf("Totals = %+v, want %+v", tot, want) + } +} + +func TestOpenSQLiteCreatesParentDirs(t *testing.T) { + path := filepath.Join(t.TempDir(), "nested", "dir", "receiver.db") + s, err := OpenSQLite(path) + if err != nil { + t.Fatalf("OpenSQLite: %v", err) + } + defer s.Close() + if s.Backend() != "sqlite" { + t.Errorf("Backend = %q, want sqlite", s.Backend()) + } + + // Reopening must not fail on the existing schema and must keep the + // stamped version. + s.Close() + s2, err := OpenSQLite(path) + if err != nil { + t.Fatalf("reopen: %v", err) + } + defer s2.Close() + var v int + if err := s2.db.QueryRow(`SELECT version FROM schema_version`).Scan(&v); err != nil { + t.Fatalf("schema_version: %v", err) + } + if v != schemaVersion { + t.Errorf("schema version = %d, want %d", v, schemaVersion) + } +} + +func TestRedactMySQLDSN(t *testing.T) { + got := RedactMySQLDSN("user:s3cret@tcp(db.example.com:3306)/exploredns") + if got != "user:xxxxx@tcp(db.example.com:3306)/exploredns" { + t.Errorf("RedactMySQLDSN = %q", got) + } + if got := RedactMySQLDSN("::::"); got != "(unparsable DSN)" { + t.Errorf("RedactMySQLDSN(bad) = %q", got) + } +} + +// TestMySQLDialectSQL pins the MySQL-side SQL text, which unit tests cannot +// execute without a server. +func TestMySQLDialectSQL(t *testing.T) { + d := mysqlDialect{} + got := d.upsert("INSERT INTO traversals (id, domain) VALUES (?, ?)", []string{"domain", "last_seen"}) + want := "INSERT INTO traversals (id, domain) VALUES (?, ?)" + + " ON DUPLICATE KEY UPDATE domain=VALUES(domain), last_seen=VALUES(last_seen)" + if got != want { + t.Errorf("upsert = %q, want %q", got, want) + } + pred, arg := d.likeContains("domain", `50%_o\ff`) + if pred != `domain LIKE ? ESCAPE '\\'` { + t.Errorf("likeContains predicate = %q", pred) + } + if arg != `%50\%\_o\\ff%` { + t.Errorf("likeContains arg = %q", arg) + } + if d.datetimeType() != "DATETIME(6)" { + t.Errorf("datetimeType = %q", d.datetimeType()) + } +} diff --git a/web/api/handler.go b/web/api/handler.go index 07f4b87..1fd4ea5 100644 --- a/web/api/handler.go +++ b/web/api/handler.go @@ -276,7 +276,7 @@ func newHandler(ctx context.Context) *Handler { maxRunning: envInt("EXPLOREDNS_MAX_JOBS", defaultMaxRunningJobs), version: "dev", limiter: newRateLimiter(limit, window), - webhook: newWebhookReporter(os.Getenv("EXPLOREDNS_WEBHOOK_URL")), + webhook: newWebhookReporter(os.Getenv("EXPLOREDNS_WEBHOOK_URL"), os.Getenv("EXPLOREDNS_WEBHOOK_TOKEN")), fp: fingerprint.New(), } diff --git a/web/api/webhook.go b/web/api/webhook.go index 3c8be5f..9a92887 100644 --- a/web/api/webhook.go +++ b/web/api/webhook.go @@ -52,19 +52,22 @@ type webhookCompleteEvent struct { // single retry, and failures are logged but never surface to callers. type webhookReporter struct { url string + token string client *http.Client timeout time.Duration retryDelay time.Duration } // newWebhookReporter returns a reporter for url, or nil when url is empty -// (webhook reporting disabled). A nil reporter is safe to call. -func newWebhookReporter(url string) *webhookReporter { +// (webhook reporting disabled). A nil reporter is safe to call. A non-empty +// token is sent as an Authorization bearer token on every delivery. +func newWebhookReporter(url, token string) *webhookReporter { if url == "" { return nil } return &webhookReporter{ url: url, + token: token, client: &http.Client{}, timeout: 5 * time.Second, retryDelay: 2 * time.Second, @@ -106,6 +109,9 @@ func (wr *webhookReporter) post(event string, body []byte) error { } req.Header.Set("Content-Type", "application/json") req.Header.Set("X-ExploreDNS-Event", event) + if wr.token != "" { + req.Header.Set("Authorization", "Bearer "+wr.token) + } resp, err := wr.client.Do(req) if err != nil { diff --git a/web/api/webhook_compat_test.go b/web/api/webhook_compat_test.go new file mode 100644 index 0000000..77f669f --- /dev/null +++ b/web/api/webhook_compat_test.go @@ -0,0 +1,126 @@ +package api + +import ( + "bytes" + "encoding/json" + "testing" + "time" + + rstore "gitea.hansenits.com.au/hits/ExploreDNS/internal/receiver/store" +) + +// These tests pin the webhook payload contract between this package (the +// sender) and internal/receiver/store (the receiver): every field a sender +// struct marshals must decode into the receiver struct and vice versa. +// DisallowUnknownFields turns any renamed or missing field into a failure. + +// decodeStrict unmarshals data into v, failing on unknown fields. +func decodeStrict(t *testing.T, data []byte, v any) { + t.Helper() + dec := json.NewDecoder(bytes.NewReader(data)) + dec.DisallowUnknownFields() + if err := dec.Decode(v); err != nil { + t.Fatalf("decode %s into %T: %v", data, v, err) + } +} + +func TestWebhookStartEventReceiverCompat(t *testing.T) { + sent := webhookStartEvent{ + Event: webhookEventStart, + ID: "job-1", + Domain: "example.com", + QueryType: "AAAA", + AllRoots: true, + ClientIP: "203.0.113.9", + StartedAt: time.Date(2026, 7, 1, 10, 0, 0, 123456789, time.UTC), + } + raw, err := json.Marshal(sent) + if err != nil { + t.Fatalf("marshal sender: %v", err) + } + + var got rstore.StartEvent + decodeStrict(t, raw, &got) + want := rstore.StartEvent{ + Event: sent.Event, + ID: sent.ID, + Domain: sent.Domain, + QueryType: sent.QueryType, + AllRoots: sent.AllRoots, + ClientIP: sent.ClientIP, + StartedAt: sent.StartedAt, + } + if got != want { + t.Errorf("receiver decoded %+v, want %+v", got, want) + } + + // Round-trip back into the sender struct so receiver-only fields + // would also fail. + back, err := json.Marshal(got) + if err != nil { + t.Fatalf("marshal receiver: %v", err) + } + var sent2 webhookStartEvent + decodeStrict(t, back, &sent2) + if sent2 != sent { + t.Errorf("sender round-trip %+v, want %+v", sent2, sent) + } +} + +func TestWebhookCompleteEventReceiverCompat(t *testing.T) { + sent := webhookCompleteEvent{ + Event: webhookEventComplete, + ID: "job-1", + Domain: "example.com", + QueryType: "A", + ClientIP: "203.0.113.9", + StartedAt: time.Date(2026, 7, 1, 10, 0, 0, 0, time.UTC), + DoneAt: time.Date(2026, 7, 1, 10, 0, 3, 500000000, time.UTC), + DurationMS: 3500, + Status: statusError, + Error: "traversal timed out after 5m0s", + ResultCount: 7, + Summary: &Summary{ + Answers: []SummaryAnswer{{Probability: 0.75, Records: []string{"example.com 300 IN A 192.0.2.1"}}}, + ByStatus: []SummaryStatus{{Status: "servfail", Probability: 0.25}}, + }, + } + raw, err := json.Marshal(sent) + if err != nil { + t.Fatalf("marshal sender: %v", err) + } + + var got rstore.CompleteEvent + decodeStrict(t, raw, &got) + if got.Event != sent.Event || got.ID != sent.ID || got.Domain != sent.Domain || + got.QueryType != sent.QueryType || got.ClientIP != sent.ClientIP || + !got.StartedAt.Equal(sent.StartedAt) || !got.DoneAt.Equal(sent.DoneAt) || + got.DurationMS != sent.DurationMS || got.Status != sent.Status || + got.Error != sent.Error || got.ResultCount != sent.ResultCount { + t.Errorf("receiver decoded %+v, want %+v", got, sent) + } + + // The receiver keeps Summary as raw JSON; it must match the sender's + // marshalled Summary byte for byte. + wantSummary, err := json.Marshal(sent.Summary) + if err != nil { + t.Fatalf("marshal summary: %v", err) + } + if !bytes.Equal(got.Summary, wantSummary) { + t.Errorf("receiver Summary = %s, want %s", got.Summary, wantSummary) + } + + back, err := json.Marshal(got) + if err != nil { + t.Fatalf("marshal receiver: %v", err) + } + var sent2 webhookCompleteEvent + decodeStrict(t, back, &sent2) + back2, err := json.Marshal(sent2) + if err != nil { + t.Fatalf("re-marshal sender: %v", err) + } + if !bytes.Equal(back2, raw) { + t.Errorf("sender round-trip = %s, want %s", back2, raw) + } +} diff --git a/web/api/webhook_test.go b/web/api/webhook_test.go index a01998d..19b3376 100644 --- a/web/api/webhook_test.go +++ b/web/api/webhook_test.go @@ -235,7 +235,37 @@ func TestWebhook_RetriesOnceOnFailure(t *testing.T) { func TestWebhook_NilReporterSafe(t *testing.T) { var wr *webhookReporter wr.send("start", map[string]string{"event": "start"}) // must not panic - if newWebhookReporter("") != nil { + if newWebhookReporter("", "token") != nil { t.Fatal("empty URL should disable the webhook reporter") } } + +// TestWebhook_BearerToken verifies the Authorization header is sent exactly +// when a token is configured (EXPLOREDNS_WEBHOOK_TOKEN on the real path). +func TestWebhook_BearerToken(t *testing.T) { + for _, tc := range []struct { + name, token, wantAuth string + }{ + {"token set", "s3cret", "Bearer s3cret"}, + {"token unset", "", ""}, + } { + t.Run(tc.name, func(t *testing.T) { + done := make(chan string, 1) + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + done <- r.Header.Get("Authorization") + })) + defer ts.Close() + + newWebhookReporter(ts.URL, tc.token).send("start", map[string]string{"event": "start"}) + + select { + case got := <-done: + if got != tc.wantAuth { + t.Fatalf("Authorization = %q, want %q", got, tc.wantAuth) + } + case <-time.After(5 * time.Second): + t.Fatal("webhook was not delivered") + } + }) + } +}