// 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) }