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 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5
parent
d96f25b1d6
commit
e8f56a1ef4
@@ -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) + "%"
|
||||
}
|
||||
@@ -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"`
|
||||
}
|
||||
@@ -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]
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user