initial commit

This commit is contained in:
2026-08-16 21:18:45 -05:00
commit 1e05a01bcf
122 changed files with 29178 additions and 0 deletions
+205
View File
@@ -0,0 +1,205 @@
// Package database owns the SQLite connection, the migration runner and every
// SQL statement in the application. Higher layers talk to *DB and never build
// SQL themselves, which keeps parameterisation and transaction handling in one
// auditable place.
package database
import (
"context"
"database/sql"
"errors"
"fmt"
"net/url"
"os"
"path/filepath"
"time"
_ "modernc.org/sqlite" // pure-Go driver: no cgo, single static binary
)
// DB wraps the SQLite handle with the helpers the rest of the app needs.
type DB struct {
*sql.DB
path string
}
// Common storage errors surfaced to the HTTP layer as 404/409 responses.
var (
ErrNotFound = errors.New("not found")
ErrConflict = errors.New("already exists")
)
// Open opens (creating if necessary) the SQLite database at path and applies
// the connection pragmas the application relies on.
//
// The file is created with 0600 and its parent directory with 0750: the
// database holds the administrator password hash and API token hashes, so it
// must not be world readable.
func Open(path string) (*DB, error) {
if path == "" {
return nil, errors.New("database path is empty")
}
dir := filepath.Dir(path)
if err := os.MkdirAll(dir, 0o750); err != nil {
return nil, fmt.Errorf("create database directory %s: %w", dir, err)
}
// _txlock=immediate makes database/sql issue BEGIN IMMEDIATE, so SQLite's
// busy handler can actually resolve writer contention instead of failing
// with SQLITE_BUSY when a deferred transaction tries to upgrade.
dsn := "file:" + url.PathEscape(path) + "?" +
"_pragma=journal_mode(WAL)" +
"&_pragma=foreign_keys(1)" +
"&_pragma=busy_timeout(15000)" +
"&_pragma=synchronous(NORMAL)" +
"&_txlock=immediate"
sqlDB, err := sql.Open("sqlite", dsn)
if err != nil {
return nil, fmt.Errorf("open database: %w", err)
}
// SQLite serialises writes; a small pool avoids piling up blocked writers
// while still allowing concurrent WAL readers.
sqlDB.SetMaxOpenConns(8)
sqlDB.SetMaxIdleConns(8)
sqlDB.SetConnMaxLifetime(time.Hour)
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
defer cancel()
if err := sqlDB.PingContext(ctx); err != nil {
sqlDB.Close()
return nil, fmt.Errorf("connect to database: %w", err)
}
db := &DB{DB: sqlDB, path: path}
if err := db.hardenPermissions(); err != nil {
sqlDB.Close()
return nil, err
}
return db, nil
}
// Path returns the on-disk location of the database.
func (db *DB) Path() string { return db.path }
// hardenPermissions restricts the database and its WAL sidecars to the owner.
func (db *DB) hardenPermissions() error {
for _, suffix := range []string{"", "-wal", "-shm"} {
p := db.path + suffix
if _, err := os.Stat(p); err != nil {
continue // sidecars may not exist yet
}
if err := os.Chmod(p, 0o600); err != nil {
return fmt.Errorf("secure %s: %w", p, err)
}
}
return nil
}
// Checkpoint flushes the write-ahead log into the main database file. It runs
// before backups so the copied file is complete.
func (db *DB) Checkpoint(ctx context.Context) error {
_, err := db.ExecContext(ctx, `PRAGMA wal_checkpoint(TRUNCATE)`)
return err
}
// Vacuum rebuilds the database, reclaiming space after large deletions.
func (db *DB) Vacuum(ctx context.Context) error {
_, err := db.ExecContext(ctx, `VACUUM`)
return err
}
// Stats describes database size and row counts for the settings UI.
type Stats struct {
Path string `json:"path"`
SizeBytes int64 `json:"size_bytes"`
WALBytes int64 `json:"wal_bytes"`
PageSize int64 `json:"page_size"`
PageCount int64 `json:"page_count"`
FreePages int64 `json:"free_pages"`
Zones int64 `json:"zones"`
Records int64 `json:"records"`
Domains int64 `json:"domains"`
QueryLogs int64 `json:"query_logs"`
AuditLogs int64 `json:"audit_logs"`
APITokens int64 `json:"api_tokens"`
SchemaVer int `json:"schema_version"`
}
// Stats collects database size and row-count information.
func (db *DB) Stats(ctx context.Context) (Stats, error) {
s := Stats{Path: db.path}
if fi, err := os.Stat(db.path); err == nil {
s.SizeBytes = fi.Size()
}
if fi, err := os.Stat(db.path + "-wal"); err == nil {
s.WALBytes = fi.Size()
}
_ = db.QueryRowContext(ctx, `PRAGMA page_size`).Scan(&s.PageSize)
_ = db.QueryRowContext(ctx, `PRAGMA page_count`).Scan(&s.PageCount)
_ = db.QueryRowContext(ctx, `PRAGMA freelist_count`).Scan(&s.FreePages)
counts := []struct {
table string
dst *int64
}{
{"zones", &s.Zones},
{"records", &s.Records},
{"domain_entries", &s.Domains},
{"query_logs", &s.QueryLogs},
{"audit_logs", &s.AuditLogs},
{"api_tokens", &s.APITokens},
}
for _, c := range counts {
if err := db.QueryRowContext(ctx, `SELECT COUNT(*) FROM `+c.table).Scan(c.dst); err != nil {
return s, fmt.Errorf("count %s: %w", c.table, err)
}
}
v, err := db.SchemaVersion(ctx)
if err != nil {
return s, err
}
s.SchemaVer = v
return s, nil
}
// InTx runs fn inside a transaction, committing on success and rolling back on
// error or panic.
func (db *DB) InTx(ctx context.Context, fn func(*sql.Tx) error) error {
tx, err := db.BeginTx(ctx, nil)
if err != nil {
return fmt.Errorf("begin transaction: %w", err)
}
defer func() {
if p := recover(); p != nil {
_ = tx.Rollback()
panic(p)
}
}()
if err := fn(tx); err != nil {
_ = tx.Rollback()
return err
}
if err := tx.Commit(); err != nil {
return fmt.Errorf("commit transaction: %w", err)
}
return nil
}
// unixPtr converts a nullable epoch-seconds column into a *time.Time.
func unixPtr(n sql.NullInt64) *time.Time {
if !n.Valid {
return nil
}
t := time.Unix(n.Int64, 0)
return &t
}
// nullInt64 converts a *int64 into a driver-friendly nullable value.
func nullInt64(v *int64) any {
if v == nil {
return nil
}
return *v
}
+443
View File
@@ -0,0 +1,443 @@
package database
import (
"context"
"errors"
"path/filepath"
"testing"
"time"
"github.com/owen/vibedns/internal/models"
)
// newTestDB opens a migrated database in a temporary directory.
func newTestDB(t *testing.T) *DB {
t.Helper()
path := filepath.Join(t.TempDir(), "test.db")
db, err := Open(path)
if err != nil {
t.Fatalf("open: %v", err)
}
t.Cleanup(func() { db.Close() })
if _, err := db.Migrate(context.Background()); err != nil {
t.Fatalf("migrate: %v", err)
}
return db
}
func TestMigrationsApplyAndAreIdempotent(t *testing.T) {
ctx := context.Background()
path := filepath.Join(t.TempDir(), "migrate.db")
db, err := Open(path)
if err != nil {
t.Fatalf("open: %v", err)
}
defer db.Close()
n, err := db.Migrate(ctx)
if err != nil {
t.Fatalf("first migrate: %v", err)
}
if n == 0 {
t.Fatal("expected migrations to be applied on a fresh database")
}
// A second run must be a no-op, which is what makes it safe to run on
// every start.
again, err := db.Migrate(ctx)
if err != nil {
t.Fatalf("second migrate: %v", err)
}
if again != 0 {
t.Errorf("second migrate applied %d migrations, want 0", again)
}
v, err := db.SchemaVersion(ctx)
if err != nil {
t.Fatalf("schema version: %v", err)
}
if v < 1 {
t.Errorf("schema version = %d, want at least 1", v)
}
statuses, err := db.MigrationStatuses(ctx)
if err != nil {
t.Fatalf("statuses: %v", err)
}
for _, s := range statuses {
if !s.Applied {
t.Errorf("migration %d_%s was not applied", s.Version, s.Name)
}
if s.Drifted {
t.Errorf("migration %d_%s reports drift on a fresh database", s.Version, s.Name)
}
}
}
func TestSeedDataExists(t *testing.T) {
db := newTestDB(t)
ctx := context.Background()
lists, err := db.DomainLists(ctx, models.KindBlacklist, "")
if err != nil {
t.Fatalf("list blacklists: %v", err)
}
if len(lists) == 0 {
t.Error("expected the seed migration to create starter blacklists")
}
nets, err := db.Networks(ctx, "", true)
if err != nil {
t.Fatalf("list networks: %v", err)
}
if len(nets) == 0 {
t.Error("expected the seed migration to create private-range networks")
}
// The seeded networks must be private ranges, never the whole Internet.
for _, n := range nets {
if n.CIDR == "0.0.0.0/0" || n.CIDR == "::/0" {
t.Errorf("seed data contains an internet-wide network %q", n.CIDR)
}
}
}
func TestZoneCRUD(t *testing.T) {
db := newTestDB(t)
ctx := context.Background()
zone := models.Zone{
Name: "example.com.", Kind: models.ZoneForward, Enabled: true,
DefaultTTL: 3600, PrimaryNS: "ns1.example.com.", AdminEmail: "hostmaster@example.com",
Serial: 1, Refresh: 7200, Retry: 3600, Expire: 1209600, Minimum: 3600, AutoSerial: true,
}
created, err := db.CreateZone(ctx, zone)
if err != nil {
t.Fatalf("create: %v", err)
}
if created.ID == 0 {
t.Fatal("created zone has no ID")
}
// Duplicate names must be rejected.
if _, err := db.CreateZone(ctx, zone); !errors.Is(err, ErrConflict) {
t.Errorf("duplicate create error = %v, want ErrConflict", err)
}
loaded, err := db.Zone(ctx, created.ID)
if err != nil {
t.Fatalf("load: %v", err)
}
if loaded.Name != "example.com." {
t.Errorf("name = %q, want example.com.", loaded.Name)
}
if _, err := db.Zone(ctx, 99999); !errors.Is(err, ErrNotFound) {
t.Errorf("missing zone error = %v, want ErrNotFound", err)
}
loaded.Description = "updated"
if _, err := db.UpdateZone(ctx, loaded); err != nil {
t.Fatalf("update: %v", err)
}
if err := db.DeleteZone(ctx, created.ID); err != nil {
t.Fatalf("delete: %v", err)
}
if err := db.DeleteZone(ctx, created.ID); !errors.Is(err, ErrNotFound) {
t.Errorf("second delete error = %v, want ErrNotFound", err)
}
}
func TestRecordsCascadeAndSerialBump(t *testing.T) {
db := newTestDB(t)
ctx := context.Background()
zone, err := db.CreateZone(ctx, models.Zone{
Name: "example.com.", Kind: models.ZoneForward, Enabled: true,
DefaultTTL: 3600, PrimaryNS: "ns1.example.com.", AdminEmail: "a@example.com",
Serial: 1, AutoSerial: true,
})
if err != nil {
t.Fatalf("create zone: %v", err)
}
if _, err := db.CreateRecord(ctx, models.Record{
ZoneID: zone.ID, Name: "www", Type: "A", Data: "192.0.2.1", Enabled: true,
}); err != nil {
t.Fatalf("create record: %v", err)
}
// Adding a record must advance the serial, which is how secondaries learn
// the zone changed.
after, err := db.Zone(ctx, zone.ID)
if err != nil {
t.Fatalf("reload zone: %v", err)
}
if after.Serial <= zone.Serial {
t.Errorf("serial = %d, want greater than %d after a record change", after.Serial, zone.Serial)
}
// Deleting the zone must take its records with it.
if err := db.DeleteZone(ctx, zone.ID); err != nil {
t.Fatalf("delete zone: %v", err)
}
recs, total, err := db.Records(ctx, RecordFilter{ZoneID: zone.ID})
if err != nil {
t.Fatalf("list records: %v", err)
}
if total != 0 || len(recs) != 0 {
t.Errorf("records remained after the zone was deleted: %d", total)
}
}
func TestManualSerialIsPreserved(t *testing.T) {
db := newTestDB(t)
ctx := context.Background()
zone, err := db.CreateZone(ctx, models.Zone{
Name: "manual.example.", Kind: models.ZoneForward, Enabled: true,
DefaultTTL: 300, PrimaryNS: "ns1.manual.example.", AdminEmail: "a@manual.example",
Serial: 2024010101, AutoSerial: false,
})
if err != nil {
t.Fatalf("create: %v", err)
}
if _, err := db.CreateRecord(ctx, models.Record{
ZoneID: zone.ID, Name: "@", Type: "A", Data: "192.0.2.1", Enabled: true,
}); err != nil {
t.Fatalf("create record: %v", err)
}
after, _ := db.Zone(ctx, zone.ID)
if after.Serial != 2024010101 {
t.Errorf("serial = %d, want the manual value to be left alone", after.Serial)
}
}
func TestImportDomainsCountsDuplicates(t *testing.T) {
db := newTestDB(t)
ctx := context.Background()
list, err := db.CreateDomainList(ctx, models.DomainList{
Kind: models.KindBlacklist, Name: "Import Test", Enabled: true,
})
if err != nil {
t.Fatalf("create list: %v", err)
}
rows := []ImportDomain{
{Domain: "a.example", MatchSubdomains: true},
{Domain: "b.example", MatchSubdomains: true},
{Domain: "c.example", MatchSubdomains: true},
}
imported, dupes, err := db.ImportDomains(ctx, list.ID, rows)
if err != nil {
t.Fatalf("import: %v", err)
}
if imported != 3 || dupes != 0 {
t.Errorf("first import = %d imported, %d duplicates; want 3, 0", imported, dupes)
}
// Re-importing the same rows plus one new one.
rows = append(rows, ImportDomain{Domain: "d.example", MatchSubdomains: true})
imported, dupes, err = db.ImportDomains(ctx, list.ID, rows)
if err != nil {
t.Fatalf("second import: %v", err)
}
if imported != 1 || dupes != 3 {
t.Errorf("second import = %d imported, %d duplicates; want 1, 3", imported, dupes)
}
loaded, err := db.DomainList(ctx, list.ID)
if err != nil {
t.Fatalf("reload list: %v", err)
}
if loaded.DomainCount != 4 {
t.Errorf("domain count = %d, want 4", loaded.DomainCount)
}
}
// TestImportLargeBatchIsOneTransaction exercises the path a real blocklist
// takes. If this were one transaction per domain it would take minutes.
func TestImportLargeBatch(t *testing.T) {
db := newTestDB(t)
ctx := context.Background()
list, err := db.CreateDomainList(ctx, models.DomainList{
Kind: models.KindBlacklist, Name: "Large", Enabled: true,
})
if err != nil {
t.Fatalf("create list: %v", err)
}
const n = 20000
rows := make([]ImportDomain, 0, n)
for i := 0; i < n; i++ {
rows = append(rows, ImportDomain{
Domain: "host" + itoa(i) + ".example.com",
MatchSubdomains: true,
})
}
imported, _, err := db.ImportDomains(ctx, list.ID, rows)
if err != nil {
t.Fatalf("bulk import: %v", err)
}
if imported != n {
t.Errorf("imported = %d, want %d", imported, n)
}
// The snapshot query must return them all for the in-memory matcher.
count := 0
if err := db.SnapshotDomains(ctx, func(SnapshotDomainEntry) { count++ }); err != nil {
t.Fatalf("snapshot: %v", err)
}
if count != n {
t.Errorf("snapshot returned %d domains, want %d", count, n)
}
}
func TestSettingsRoundTrip(t *testing.T) {
db := newTestDB(t)
ctx := context.Background()
if err := db.SetSetting(ctx, "test.key", "value"); err != nil {
t.Fatalf("set: %v", err)
}
v, ok, err := db.Setting(ctx, "test.key")
if err != nil || !ok || v != "value" {
t.Errorf("get = %q, %v, %v; want \"value\", true, nil", v, ok, err)
}
// Writing again must update rather than fail on the primary key.
if err := db.SetSetting(ctx, "test.key", "changed"); err != nil {
t.Fatalf("overwrite: %v", err)
}
v, _, _ = db.Setting(ctx, "test.key")
if v != "changed" {
t.Errorf("after overwrite = %q, want \"changed\"", v)
}
if _, ok, _ := db.Setting(ctx, "missing.key"); ok {
t.Error("a missing key should report ok=false")
}
}
func TestQueryLogPruning(t *testing.T) {
db := newTestDB(t)
ctx := context.Background()
var entries []models.QueryLogEntry
for i := 0; i < 100; i++ {
entries = append(entries, models.QueryLogEntry{
Timestamp: time.Now(), ClientIP: "192.0.2.1",
QName: "example.com.", QType: "A", Rcode: "NOERROR",
Source: models.SourceCache, Protocol: "udp",
})
}
if err := db.InsertQueryLogs(ctx, entries); err != nil {
t.Fatalf("insert: %v", err)
}
n, err := db.QueryLogCount(ctx)
if err != nil || n != 100 {
t.Fatalf("count = %d, %v; want 100", n, err)
}
// Trim to the newest 40 rows.
removed, err := db.PruneQueryLogs(ctx, 0, 40)
if err != nil {
t.Fatalf("prune: %v", err)
}
if removed != 60 {
t.Errorf("pruned %d rows, want 60", removed)
}
n, _ = db.QueryLogCount(ctx)
if n != 40 {
t.Errorf("rows remaining = %d, want 40", n)
}
}
func TestAuditLog(t *testing.T) {
db := newTestDB(t)
ctx := context.Background()
err := db.InsertAudit(ctx, models.AuditEntry{
Actor: "admin", Source: "web", ClientIP: "192.0.2.1",
Action: "zone.create", ObjectType: "zone", ObjectName: "example.com.",
})
if err != nil {
t.Fatalf("insert: %v", err)
}
entries, total, err := db.AuditLogs(ctx, AuditFilter{Limit: 10})
if err != nil {
t.Fatalf("read: %v", err)
}
if total != 1 || len(entries) != 1 {
t.Fatalf("got %d of %d, want 1 of 1", len(entries), total)
}
if entries[0].Action != "zone.create" {
t.Errorf("action = %q", entries[0].Action)
}
}
func TestAPITokenLookup(t *testing.T) {
db := newTestDB(t)
ctx := context.Background()
tok, err := db.CreateAPIToken(ctx, "test", "a description", "abcd1234", "hash-value")
if err != nil {
t.Fatalf("create: %v", err)
}
if tok.ID == 0 {
t.Fatal("no ID assigned")
}
candidates, err := db.APITokensByPrefix(ctx, "abcd1234")
if err != nil {
t.Fatalf("lookup: %v", err)
}
if len(candidates) != 1 || candidates[0].Hash != "hash-value" {
t.Errorf("lookup returned %v", candidates)
}
// A disabled token must not be returned as a candidate at all.
if err := db.SetAPITokenEnabled(ctx, tok.ID, false); err != nil {
t.Fatalf("disable: %v", err)
}
candidates, _ = db.APITokensByPrefix(ctx, "abcd1234")
if len(candidates) != 0 {
t.Error("a disabled token was returned by the prefix lookup")
}
}
func TestForeignKeysAreEnforced(t *testing.T) {
db := newTestDB(t)
ctx := context.Background()
// A record pointing at a zone that does not exist must be rejected;
// without PRAGMA foreign_keys this would silently succeed.
_, err := db.CreateRecord(ctx, models.Record{
ZoneID: 99999, Name: "www", Type: "A", Data: "192.0.2.1", Enabled: true,
})
if err == nil {
t.Error("expected a foreign key violation for an orphaned record")
}
}
func itoa(i int) string {
if i == 0 {
return "0"
}
var buf [12]byte
pos := len(buf)
for i > 0 {
pos--
buf[pos] = byte('0' + i%10)
i /= 10
}
return string(buf[pos:])
}
+202
View File
@@ -0,0 +1,202 @@
package database
import (
"context"
"crypto/sha256"
"database/sql"
"embed"
"encoding/hex"
"fmt"
"io/fs"
"sort"
"strconv"
"strings"
)
//go:embed migrations/*.sql
var migrationFS embed.FS
// Migration is one versioned schema change, embedded in the binary.
type Migration struct {
Version int
Name string
SQL string
}
// MigrationStatus reports whether a migration has been applied.
type MigrationStatus struct {
Version int `json:"version"`
Name string `json:"name"`
Applied bool `json:"applied"`
AppliedAt int64 `json:"applied_at,omitempty"`
Checksum string `json:"checksum"`
Drifted bool `json:"drifted"`
}
// loadMigrations reads and orders the embedded migration files. File names must
// look like "0001_description.sql".
func loadMigrations() ([]Migration, error) {
entries, err := fs.ReadDir(migrationFS, "migrations")
if err != nil {
return nil, fmt.Errorf("read embedded migrations: %w", err)
}
var out []Migration
for _, e := range entries {
if e.IsDir() || !strings.HasSuffix(e.Name(), ".sql") {
continue
}
base := strings.TrimSuffix(e.Name(), ".sql")
parts := strings.SplitN(base, "_", 2)
if len(parts) != 2 {
return nil, fmt.Errorf("migration %q: expected NNNN_name.sql", e.Name())
}
v, err := strconv.Atoi(parts[0])
if err != nil {
return nil, fmt.Errorf("migration %q: bad version prefix: %w", e.Name(), err)
}
body, err := migrationFS.ReadFile("migrations/" + e.Name())
if err != nil {
return nil, fmt.Errorf("read migration %q: %w", e.Name(), err)
}
out = append(out, Migration{Version: v, Name: parts[1], SQL: string(body)})
}
sort.Slice(out, func(i, j int) bool { return out[i].Version < out[j].Version })
for i := 1; i < len(out); i++ {
if out[i].Version == out[i-1].Version {
return nil, fmt.Errorf("duplicate migration version %d", out[i].Version)
}
}
return out, nil
}
func checksum(s string) string {
sum := sha256.Sum256([]byte(s))
return hex.EncodeToString(sum[:])
}
// ensureMigrationTable creates the migration bookkeeping table.
func (db *DB) ensureMigrationTable(ctx context.Context) error {
_, err := db.ExecContext(ctx, `
CREATE TABLE IF NOT EXISTS schema_migrations (
version INTEGER PRIMARY KEY,
name TEXT NOT NULL,
checksum TEXT NOT NULL,
applied_at INTEGER NOT NULL DEFAULT (unixepoch())
)`)
if err != nil {
return fmt.Errorf("create schema_migrations: %w", err)
}
return nil
}
type appliedMigration struct {
name string
checksum string
appliedAt int64
}
func (db *DB) appliedMigrations(ctx context.Context) (map[int]appliedMigration, error) {
rows, err := db.QueryContext(ctx, `SELECT version, name, checksum, applied_at FROM schema_migrations`)
if err != nil {
return nil, fmt.Errorf("read schema_migrations: %w", err)
}
defer rows.Close()
out := map[int]appliedMigration{}
for rows.Next() {
var v int
var a appliedMigration
if err := rows.Scan(&v, &a.name, &a.checksum, &a.appliedAt); err != nil {
return nil, err
}
out[v] = a
}
return out, rows.Err()
}
// Migrate applies every pending migration in version order. It returns the
// number of migrations that were applied.
func (db *DB) Migrate(ctx context.Context) (int, error) {
if err := db.ensureMigrationTable(ctx); err != nil {
return 0, err
}
migrations, err := loadMigrations()
if err != nil {
return 0, err
}
applied, err := db.appliedMigrations(ctx)
if err != nil {
return 0, err
}
count := 0
for _, m := range migrations {
sum := checksum(m.SQL)
if prev, ok := applied[m.Version]; ok {
if prev.checksum != sum {
return count, fmt.Errorf(
"migration %04d_%s was modified after being applied (expected checksum %s, found %s); "+
"roll the change into a new migration instead of editing history",
m.Version, m.Name, prev.checksum, sum)
}
continue
}
// Each migration is one transaction: a failure leaves no partial schema.
err := db.InTx(ctx, func(tx *sql.Tx) error {
if _, err := tx.ExecContext(ctx, m.SQL); err != nil {
return fmt.Errorf("apply migration %04d_%s: %w", m.Version, m.Name, err)
}
_, err := tx.ExecContext(ctx,
`INSERT INTO schema_migrations (version, name, checksum) VALUES (?, ?, ?)`,
m.Version, m.Name, sum)
return err
})
if err != nil {
return count, err
}
count++
}
return count, nil
}
// SchemaVersion returns the highest applied migration version, or 0.
func (db *DB) SchemaVersion(ctx context.Context) (int, error) {
if err := db.ensureMigrationTable(ctx); err != nil {
return 0, err
}
var v int
err := db.QueryRowContext(ctx, `SELECT COALESCE(MAX(version), 0) FROM schema_migrations`).Scan(&v)
if err != nil {
return 0, fmt.Errorf("read schema version: %w", err)
}
return v, nil
}
// MigrationStatuses lists every known migration and whether it is applied.
func (db *DB) MigrationStatuses(ctx context.Context) ([]MigrationStatus, error) {
if err := db.ensureMigrationTable(ctx); err != nil {
return nil, err
}
migrations, err := loadMigrations()
if err != nil {
return nil, err
}
applied, err := db.appliedMigrations(ctx)
if err != nil {
return nil, err
}
out := make([]MigrationStatus, 0, len(migrations))
for _, m := range migrations {
sum := checksum(m.SQL)
st := MigrationStatus{Version: m.Version, Name: m.Name, Checksum: sum[:12]}
if a, ok := applied[m.Version]; ok {
st.Applied = true
st.AppliedAt = a.appliedAt
st.Drifted = a.checksum != sum
}
out = append(out, st)
}
return out, nil
}
@@ -0,0 +1,182 @@
-- Initial schema.
--
-- Design notes:
-- * Every searchable entity gets a real relational table. Only genuinely
-- list-shaped configuration (upstream resolvers, ACL networks) lives as a
-- JSON value inside the settings key/value table.
-- * Timestamps are stored as unix epoch seconds (INTEGER) so that they sort
-- and range-scan cheaply and carry no timezone ambiguity.
-- * Booleans are INTEGER 0/1.
CREATE TABLE settings (
key TEXT PRIMARY KEY,
value TEXT NOT NULL,
updated_at INTEGER NOT NULL DEFAULT (unixepoch())
);
CREATE TABLE admin_user (
id INTEGER PRIMARY KEY CHECK (id = 1),
username TEXT NOT NULL,
password_hash TEXT NOT NULL,
must_change_password INTEGER NOT NULL DEFAULT 0,
created_at INTEGER NOT NULL DEFAULT (unixepoch()),
updated_at INTEGER NOT NULL DEFAULT (unixepoch()),
last_login_at INTEGER
);
CREATE TABLE api_tokens (
id INTEGER PRIMARY KEY AUTOINCREMENT,
name TEXT NOT NULL UNIQUE,
description TEXT NOT NULL DEFAULT '',
token_prefix TEXT NOT NULL,
token_hash TEXT NOT NULL,
enabled INTEGER NOT NULL DEFAULT 1,
created_at INTEGER NOT NULL DEFAULT (unixepoch()),
last_used_at INTEGER
);
CREATE INDEX idx_api_tokens_prefix ON api_tokens (token_prefix);
CREATE TABLE zones (
id INTEGER PRIMARY KEY AUTOINCREMENT,
name TEXT NOT NULL UNIQUE, -- normalised FQDN with trailing dot
kind TEXT NOT NULL DEFAULT 'forward'
CHECK (kind IN ('forward', 'reverse4', 'reverse6')),
description TEXT NOT NULL DEFAULT '',
enabled INTEGER NOT NULL DEFAULT 1,
default_ttl INTEGER NOT NULL DEFAULT 3600,
primary_ns TEXT NOT NULL,
admin_email TEXT NOT NULL,
serial INTEGER NOT NULL DEFAULT 1,
refresh INTEGER NOT NULL DEFAULT 7200,
retry INTEGER NOT NULL DEFAULT 3600,
expire INTEGER NOT NULL DEFAULT 1209600,
minimum INTEGER NOT NULL DEFAULT 3600,
auto_serial INTEGER NOT NULL DEFAULT 1,
created_at INTEGER NOT NULL DEFAULT (unixepoch()),
updated_at INTEGER NOT NULL DEFAULT (unixepoch())
);
CREATE INDEX idx_zones_enabled ON zones (enabled);
CREATE INDEX idx_zones_kind ON zones (kind);
CREATE TABLE records (
id INTEGER PRIMARY KEY AUTOINCREMENT,
zone_id INTEGER NOT NULL REFERENCES zones (id) ON DELETE CASCADE,
name TEXT NOT NULL, -- relative to the apex; '@' is the apex itself
type TEXT NOT NULL,
data TEXT NOT NULL, -- rdata in zone-file presentation format
ttl INTEGER, -- NULL inherits zones.default_ttl
enabled INTEGER NOT NULL DEFAULT 1,
comment TEXT NOT NULL DEFAULT '',
created_at INTEGER NOT NULL DEFAULT (unixepoch()),
updated_at INTEGER NOT NULL DEFAULT (unixepoch())
);
CREATE INDEX idx_records_zone ON records (zone_id);
CREATE INDEX idx_records_zone_name ON records (zone_id, name);
CREATE INDEX idx_records_zone_name_type ON records (zone_id, name, type);
CREATE INDEX idx_records_type ON records (type);
CREATE TABLE networks (
id INTEGER PRIMARY KEY AUTOINCREMENT,
name TEXT NOT NULL UNIQUE,
cidr TEXT NOT NULL,
description TEXT NOT NULL DEFAULT '',
enabled INTEGER NOT NULL DEFAULT 1,
created_at INTEGER NOT NULL DEFAULT (unixepoch()),
updated_at INTEGER NOT NULL DEFAULT (unixepoch())
);
CREATE INDEX idx_networks_enabled ON networks (enabled);
CREATE TABLE policies (
id INTEGER PRIMARY KEY AUTOINCREMENT,
name TEXT NOT NULL UNIQUE,
description TEXT NOT NULL DEFAULT '',
enabled INTEGER NOT NULL DEFAULT 1,
block_action TEXT NOT NULL DEFAULT 'nxdomain'
CHECK (block_action IN ('nxdomain', 'refused', 'sinkhole')),
sinkhole_ipv4 TEXT NOT NULL DEFAULT '0.0.0.0',
sinkhole_ipv6 TEXT NOT NULL DEFAULT '::',
block_ttl INTEGER NOT NULL DEFAULT 60,
created_at INTEGER NOT NULL DEFAULT (unixepoch()),
updated_at INTEGER NOT NULL DEFAULT (unixepoch())
);
CREATE TABLE domain_lists (
id INTEGER PRIMARY KEY AUTOINCREMENT,
kind TEXT NOT NULL CHECK (kind IN ('blacklist', 'allowlist')),
name TEXT NOT NULL,
description TEXT NOT NULL DEFAULT '',
enabled INTEGER NOT NULL DEFAULT 1,
source_url TEXT NOT NULL DEFAULT '',
created_at INTEGER NOT NULL DEFAULT (unixepoch()),
updated_at INTEGER NOT NULL DEFAULT (unixepoch()),
UNIQUE (kind, name)
);
CREATE TABLE domain_entries (
id INTEGER PRIMARY KEY AUTOINCREMENT,
list_id INTEGER NOT NULL REFERENCES domain_lists (id) ON DELETE CASCADE,
domain TEXT NOT NULL, -- normalised: lowercase, no trailing dot
match_subdomains INTEGER NOT NULL DEFAULT 1,
enabled INTEGER NOT NULL DEFAULT 1,
comment TEXT NOT NULL DEFAULT '',
created_at INTEGER NOT NULL DEFAULT (unixepoch()),
UNIQUE (list_id, domain)
);
CREATE INDEX idx_domain_entries_list ON domain_entries (list_id);
CREATE INDEX idx_domain_entries_domain ON domain_entries (domain);
-- Policy assignments.
CREATE TABLE network_policies (
network_id INTEGER NOT NULL REFERENCES networks (id) ON DELETE CASCADE,
policy_id INTEGER NOT NULL REFERENCES policies (id) ON DELETE CASCADE,
PRIMARY KEY (network_id, policy_id)
);
CREATE INDEX idx_network_policies_policy ON network_policies (policy_id);
CREATE TABLE policy_lists (
policy_id INTEGER NOT NULL REFERENCES policies (id) ON DELETE CASCADE,
list_id INTEGER NOT NULL REFERENCES domain_lists (id) ON DELETE CASCADE,
PRIMARY KEY (policy_id, list_id)
);
CREATE INDEX idx_policy_lists_list ON policy_lists (list_id);
CREATE TABLE query_logs (
id INTEGER PRIMARY KEY AUTOINCREMENT,
ts INTEGER NOT NULL, -- unix milliseconds
client_ip TEXT NOT NULL,
network_id INTEGER,
network_name TEXT NOT NULL DEFAULT '',
qname TEXT NOT NULL,
qtype TEXT NOT NULL,
rcode TEXT NOT NULL,
source TEXT NOT NULL,
cache_hit INTEGER NOT NULL DEFAULT 0,
blocked INTEGER NOT NULL DEFAULT 0,
policy_id INTEGER,
policy_name TEXT NOT NULL DEFAULT '',
blacklist_id INTEGER,
blacklist_name TEXT NOT NULL DEFAULT '',
matched_rule TEXT NOT NULL DEFAULT '',
protocol TEXT NOT NULL DEFAULT 'udp',
duration_us INTEGER NOT NULL DEFAULT 0,
answer_count INTEGER NOT NULL DEFAULT 0
);
CREATE INDEX idx_query_logs_ts ON query_logs (ts);
CREATE INDEX idx_query_logs_client ON query_logs (client_ip, ts);
CREATE INDEX idx_query_logs_qname ON query_logs (qname, ts);
CREATE INDEX idx_query_logs_blocked ON query_logs (blocked, ts);
CREATE TABLE audit_logs (
id INTEGER PRIMARY KEY AUTOINCREMENT,
ts INTEGER NOT NULL, -- unix milliseconds
actor TEXT NOT NULL DEFAULT '',
source TEXT NOT NULL DEFAULT 'web',
client_ip TEXT NOT NULL DEFAULT '',
action TEXT NOT NULL,
object_type TEXT NOT NULL DEFAULT '',
object_id TEXT NOT NULL DEFAULT '',
object_name TEXT NOT NULL DEFAULT '',
details TEXT NOT NULL DEFAULT ''
);
CREATE INDEX idx_audit_logs_ts ON audit_logs (ts);
CREATE INDEX idx_audit_logs_object ON audit_logs (object_type, ts);
@@ -0,0 +1,31 @@
-- Baseline data that makes a fresh install immediately useful without
-- creating an open resolver: a private-network ACL, an empty malware
-- blacklist, and a default policy bound to RFC1918 / ULA space.
INSERT INTO domain_lists (kind, name, description) VALUES
('blacklist', 'Malware', 'Known malware and command-and-control domains.'),
('blacklist', 'Advertising', 'Advertising and tracking domains.'),
('blacklist', 'Adult Content', 'Adult content domains.'),
('blacklist', 'Gambling', 'Gambling and betting domains.'),
('blacklist', 'Guest Network Custom Blocks', 'Locally maintained blocks for guest networks.'),
('allowlist', 'Global Allowlist', 'Domains that must never be blocked.');
INSERT INTO policies (name, description, block_action) VALUES
('Default Protection', 'Malware blocking applied to local networks.', 'nxdomain');
INSERT INTO policy_lists (policy_id, list_id)
SELECT p.id, l.id
FROM policies p, domain_lists l
WHERE p.name = 'Default Protection'
AND l.name IN ('Malware', 'Global Allowlist');
INSERT INTO networks (name, cidr, description) VALUES
('Private IPv4 10.0.0.0/8', '10.0.0.0/8', 'RFC1918 private address space.'),
('Private IPv4 172.16.0.0/12', '172.16.0.0/12', 'RFC1918 private address space.'),
('Private IPv4 192.168.0.0/16', '192.168.0.0/16', 'RFC1918 private address space.'),
('Loopback IPv4', '127.0.0.0/8', 'Local host.'),
('Loopback IPv6', '::1/128', 'Local host.'),
('Unique Local IPv6', 'fc00::/7', 'RFC4193 unique local addresses.');
INSERT INTO network_policies (network_id, policy_id)
SELECT n.id, p.id FROM networks n, policies p WHERE p.name = 'Default Protection';
+207
View File
@@ -0,0 +1,207 @@
package database
import (
"context"
"database/sql"
"errors"
"fmt"
"strings"
"time"
"github.com/owen/vibedns/internal/models"
)
// Admin returns the administrator account, or ErrNotFound before first setup.
func (db *DB) Admin(ctx context.Context) (models.Admin, error) {
var a models.Admin
var created, updated int64
var lastLogin sql.NullInt64
var mustChange int
err := db.QueryRowContext(ctx, `
SELECT username, password_hash, must_change_password, created_at, updated_at, last_login_at
FROM admin_user WHERE id = 1`).
Scan(&a.Username, &a.PasswordHash, &mustChange, &created, &updated, &lastLogin)
switch {
case errors.Is(err, sql.ErrNoRows):
return a, ErrNotFound
case err != nil:
return a, fmt.Errorf("load administrator: %w", err)
}
a.MustChangePassword = mustChange != 0
a.CreatedAt = time.Unix(created, 0)
a.UpdatedAt = time.Unix(updated, 0)
a.LastLoginAt = unixPtr(lastLogin)
return a, nil
}
// CreateAdmin inserts the single administrator row. It fails if one exists.
func (db *DB) CreateAdmin(ctx context.Context, username, passwordHash string, mustChange bool) error {
_, err := db.ExecContext(ctx, `
INSERT INTO admin_user (id, username, password_hash, must_change_password)
VALUES (1, ?, ?, ?)`, username, passwordHash, boolInt(mustChange))
if err != nil {
if isUniqueViolation(err) {
return ErrConflict
}
return fmt.Errorf("create administrator: %w", err)
}
return nil
}
// UpdateAdminCredentials replaces the username and/or password hash.
func (db *DB) UpdateAdminCredentials(ctx context.Context, username, passwordHash string, mustChange bool) error {
res, err := db.ExecContext(ctx, `
UPDATE admin_user
SET username = ?, password_hash = ?, must_change_password = ?, updated_at = unixepoch()
WHERE id = 1`, username, passwordHash, boolInt(mustChange))
if err != nil {
return fmt.Errorf("update administrator: %w", err)
}
if n, _ := res.RowsAffected(); n == 0 {
return ErrNotFound
}
return nil
}
// TouchAdminLogin records a successful authentication.
func (db *DB) TouchAdminLogin(ctx context.Context) error {
_, err := db.ExecContext(ctx, `UPDATE admin_user SET last_login_at = unixepoch() WHERE id = 1`)
return err
}
// --- API tokens ---------------------------------------------------------
// CreateAPIToken stores a new token. Only the prefix and hash are persisted.
func (db *DB) CreateAPIToken(ctx context.Context, name, description, prefix, hash string) (models.APIToken, error) {
res, err := db.ExecContext(ctx, `
INSERT INTO api_tokens (name, description, token_prefix, token_hash)
VALUES (?, ?, ?, ?)`, name, description, prefix, hash)
if err != nil {
if isUniqueViolation(err) {
return models.APIToken{}, ErrConflict
}
return models.APIToken{}, fmt.Errorf("create API token: %w", err)
}
id, _ := res.LastInsertId()
return db.APIToken(ctx, id)
}
// APIToken loads one token by ID.
func (db *DB) APIToken(ctx context.Context, id int64) (models.APIToken, error) {
rows, err := db.queryTokens(ctx, `WHERE id = ?`, id)
if err != nil {
return models.APIToken{}, err
}
if len(rows) == 0 {
return models.APIToken{}, ErrNotFound
}
return rows[0], nil
}
// APITokens lists all tokens, newest first.
func (db *DB) APITokens(ctx context.Context) ([]models.APIToken, error) {
return db.queryTokens(ctx, `ORDER BY created_at DESC, id DESC`)
}
func (db *DB) queryTokens(ctx context.Context, where string, args ...any) ([]models.APIToken, error) {
q := `SELECT id, name, description, token_prefix, enabled, created_at, last_used_at FROM api_tokens ` + where
rows, err := db.QueryContext(ctx, q, args...)
if err != nil {
return nil, fmt.Errorf("list API tokens: %w", err)
}
defer rows.Close()
var out []models.APIToken
for rows.Next() {
var t models.APIToken
var enabled int
var created int64
var lastUsed sql.NullInt64
if err := rows.Scan(&t.ID, &t.Name, &t.Description, &t.Prefix, &enabled, &created, &lastUsed); err != nil {
return nil, err
}
t.Enabled = enabled != 0
t.CreatedAt = time.Unix(created, 0)
t.LastUsedAt = unixPtr(lastUsed)
out = append(out, t)
}
return out, rows.Err()
}
// APITokenCandidate is a stored token hash keyed by its lookup prefix.
type APITokenCandidate struct {
ID int64
Name string
Hash string
}
// APITokensByPrefix returns enabled tokens whose prefix matches. The prefix
// narrows the search; the caller still verifies the hash in constant time.
func (db *DB) APITokensByPrefix(ctx context.Context, prefix string) ([]APITokenCandidate, error) {
rows, err := db.QueryContext(ctx,
`SELECT id, name, token_hash FROM api_tokens WHERE token_prefix = ? AND enabled = 1`, prefix)
if err != nil {
return nil, fmt.Errorf("lookup API token: %w", err)
}
defer rows.Close()
var out []APITokenCandidate
for rows.Next() {
var c APITokenCandidate
if err := rows.Scan(&c.ID, &c.Name, &c.Hash); err != nil {
return nil, err
}
out = append(out, c)
}
return out, rows.Err()
}
// TouchAPIToken records that a token was just used. Errors are non-fatal to the
// request path, so callers may ignore them.
func (db *DB) TouchAPIToken(ctx context.Context, id int64) error {
_, err := db.ExecContext(ctx, `UPDATE api_tokens SET last_used_at = unixepoch() WHERE id = ?`, id)
return err
}
// SetAPITokenEnabled enables or disables a token without deleting it.
func (db *DB) SetAPITokenEnabled(ctx context.Context, id int64, enabled bool) error {
res, err := db.ExecContext(ctx, `UPDATE api_tokens SET enabled = ? WHERE id = ?`, boolInt(enabled), id)
if err != nil {
return fmt.Errorf("update API token: %w", err)
}
if n, _ := res.RowsAffected(); n == 0 {
return ErrNotFound
}
return nil
}
// DeleteAPIToken permanently revokes a token.
func (db *DB) DeleteAPIToken(ctx context.Context, id int64) error {
res, err := db.ExecContext(ctx, `DELETE FROM api_tokens WHERE id = ?`, id)
if err != nil {
return fmt.Errorf("delete API token: %w", err)
}
if n, _ := res.RowsAffected(); n == 0 {
return ErrNotFound
}
return nil
}
func boolInt(b bool) int {
if b {
return 1
}
return 0
}
// isUniqueViolation detects SQLite UNIQUE/PRIMARY KEY constraint failures
// without depending on driver-specific error types.
func isUniqueViolation(err error) bool {
if err == nil {
return false
}
msg := strings.ToLower(err.Error())
return strings.Contains(msg, "unique constraint failed") ||
strings.Contains(msg, "constraint failed: unique")
}
+409
View File
@@ -0,0 +1,409 @@
package database
import (
"context"
"database/sql"
"fmt"
"strings"
"time"
"github.com/owen/vibedns/internal/models"
)
// --- Query log ----------------------------------------------------------
// InsertQueryLogs writes a batch of query log rows in one transaction. The
// query logger buffers rows in memory and calls this periodically so DNS
// resolution never waits on disk.
func (db *DB) InsertQueryLogs(ctx context.Context, entries []models.QueryLogEntry) error {
if len(entries) == 0 {
return nil
}
return db.InTx(ctx, func(tx *sql.Tx) error {
stmt, err := tx.PrepareContext(ctx, `
INSERT INTO query_logs (ts, client_ip, network_id, network_name, qname, qtype, rcode,
source, cache_hit, blocked, policy_id, policy_name, blacklist_id, blacklist_name,
matched_rule, protocol, duration_us, answer_count)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`)
if err != nil {
return err
}
defer stmt.Close()
for _, e := range entries {
_, err := stmt.ExecContext(ctx,
e.Timestamp.UnixMilli(), e.ClientIP, nullInt64(e.NetworkID), e.NetworkName,
e.QName, e.QType, e.Rcode, e.Source, boolInt(e.CacheHit), boolInt(e.Blocked),
nullInt64(e.PolicyID), e.PolicyName, nullInt64(e.BlacklistID), e.BlacklistName,
e.MatchedRule, e.Protocol, e.DurationUS, e.AnswerCount)
if err != nil {
return fmt.Errorf("write query log: %w", err)
}
}
return nil
})
}
// QueryLogFilter narrows a query log search.
type QueryLogFilter struct {
Domain string
ClientIP string
QType string
Rcode string
Source string
Blocked string // "", "blocked", "allowed"
NetworkID int64
From time.Time
To time.Time
Limit int
Offset int
}
func (f QueryLogFilter) where() (string, []any) {
var conds []string
var args []any
if s := strings.TrimSpace(f.Domain); s != "" {
conds = append(conds, "qname LIKE ?")
args = append(args, "%"+strings.ToLower(s)+"%")
}
if s := strings.TrimSpace(f.ClientIP); s != "" {
conds = append(conds, "client_ip LIKE ?")
args = append(args, "%"+s+"%")
}
if s := strings.ToUpper(strings.TrimSpace(f.QType)); s != "" {
conds = append(conds, "qtype = ?")
args = append(args, s)
}
if s := strings.ToUpper(strings.TrimSpace(f.Rcode)); s != "" {
conds = append(conds, "rcode = ?")
args = append(args, s)
}
if s := strings.TrimSpace(f.Source); s != "" {
conds = append(conds, "source = ?")
args = append(args, s)
}
switch f.Blocked {
case "blocked":
conds = append(conds, "blocked = 1")
case "allowed":
conds = append(conds, "blocked = 0")
}
if f.NetworkID > 0 {
conds = append(conds, "network_id = ?")
args = append(args, f.NetworkID)
}
if !f.From.IsZero() {
conds = append(conds, "ts >= ?")
args = append(args, f.From.UnixMilli())
}
if !f.To.IsZero() {
conds = append(conds, "ts <= ?")
args = append(args, f.To.UnixMilli())
}
if len(conds) == 0 {
return "", nil
}
return " WHERE " + strings.Join(conds, " AND "), args
}
// QueryLogs returns matching rows newest first, plus the total match count.
func (db *DB) QueryLogs(ctx context.Context, f QueryLogFilter) ([]models.QueryLogEntry, int, error) {
whereSQL, args := f.where()
var total int
if err := db.QueryRowContext(ctx, `SELECT COUNT(*) FROM query_logs`+whereSQL, args...).Scan(&total); err != nil {
return nil, 0, fmt.Errorf("count query logs: %w", err)
}
q := `SELECT id, ts, client_ip, network_id, network_name, qname, qtype, rcode, source,
cache_hit, blocked, policy_id, policy_name, blacklist_id, blacklist_name, matched_rule,
protocol, duration_us, answer_count
FROM query_logs` + whereSQL + ` ORDER BY ts DESC, id DESC`
qargs := args
if f.Limit > 0 {
q += ` LIMIT ? OFFSET ?`
qargs = append(append([]any{}, args...), f.Limit, f.Offset)
}
rows, err := db.QueryContext(ctx, q, qargs...)
if err != nil {
return nil, 0, fmt.Errorf("read query logs: %w", err)
}
defer rows.Close()
var out []models.QueryLogEntry
for rows.Next() {
var e models.QueryLogEntry
var ts int64
var netID, polID, blID sql.NullInt64
var cacheHit, blocked int
err := rows.Scan(&e.ID, &ts, &e.ClientIP, &netID, &e.NetworkName, &e.QName, &e.QType,
&e.Rcode, &e.Source, &cacheHit, &blocked, &polID, &e.PolicyName, &blID,
&e.BlacklistName, &e.MatchedRule, &e.Protocol, &e.DurationUS, &e.AnswerCount)
if err != nil {
return nil, 0, err
}
e.Timestamp = time.UnixMilli(ts)
e.CacheHit = cacheHit != 0
e.Blocked = blocked != 0
if netID.Valid {
v := netID.Int64
e.NetworkID = &v
}
if polID.Valid {
v := polID.Int64
e.PolicyID = &v
}
if blID.Valid {
v := blID.Int64
e.BlacklistID = &v
}
out = append(out, e)
}
return out, total, rows.Err()
}
// PruneQueryLogs enforces the retention policy. Rows older than retentionDays
// are removed first, then the table is trimmed to maxRows newest entries.
// Either limit may be zero to disable it.
func (db *DB) PruneQueryLogs(ctx context.Context, retentionDays, maxRows int) (int64, error) {
var deleted int64
if retentionDays > 0 {
cutoff := time.Now().AddDate(0, 0, -retentionDays).UnixMilli()
res, err := db.ExecContext(ctx, `DELETE FROM query_logs WHERE ts < ?`, cutoff)
if err != nil {
return deleted, fmt.Errorf("prune query logs by age: %w", err)
}
n, _ := res.RowsAffected()
deleted += n
}
if maxRows > 0 {
res, err := db.ExecContext(ctx, `
DELETE FROM query_logs WHERE id NOT IN (
SELECT id FROM query_logs ORDER BY ts DESC, id DESC LIMIT ?
)`, maxRows)
if err != nil {
return deleted, fmt.Errorf("prune query logs by count: %w", err)
}
n, _ := res.RowsAffected()
deleted += n
}
return deleted, nil
}
// TruncateQueryLogs empties the query log table.
func (db *DB) TruncateQueryLogs(ctx context.Context) (int64, error) {
res, err := db.ExecContext(ctx, `DELETE FROM query_logs`)
if err != nil {
return 0, fmt.Errorf("clear query logs: %w", err)
}
n, _ := res.RowsAffected()
return n, nil
}
// NameCount is a domain/client aggregate used by the dashboard top-N lists.
type NameCount struct {
Name string `json:"name"`
Count int64 `json:"count"`
Extra string `json:"extra,omitempty"`
}
// TopQueried returns the most frequently queried names since `since`.
func (db *DB) TopQueried(ctx context.Context, since time.Time, limit int) ([]NameCount, error) {
return db.topBy(ctx, `SELECT qname, COUNT(*) c FROM query_logs WHERE ts >= ? GROUP BY qname ORDER BY c DESC LIMIT ?`,
since.UnixMilli(), limit)
}
// TopBlocked returns the most frequently blocked names since `since`.
func (db *DB) TopBlocked(ctx context.Context, since time.Time, limit int) ([]NameCount, error) {
return db.topBy(ctx, `SELECT qname, COUNT(*) c FROM query_logs WHERE blocked = 1 AND ts >= ? GROUP BY qname ORDER BY c DESC LIMIT ?`,
since.UnixMilli(), limit)
}
// TopClients returns the busiest clients since `since`.
func (db *DB) TopClients(ctx context.Context, since time.Time, limit int) ([]NameCount, error) {
rows, err := db.QueryContext(ctx, `
SELECT client_ip, COUNT(*) c, COALESCE(MAX(network_name), '')
FROM query_logs WHERE ts >= ? GROUP BY client_ip ORDER BY c DESC LIMIT ?`,
since.UnixMilli(), limit)
if err != nil {
return nil, fmt.Errorf("top clients: %w", err)
}
defer rows.Close()
var out []NameCount
for rows.Next() {
var n NameCount
if err := rows.Scan(&n.Name, &n.Count, &n.Extra); err != nil {
return nil, err
}
out = append(out, n)
}
return out, rows.Err()
}
func (db *DB) topBy(ctx context.Context, q string, args ...any) ([]NameCount, error) {
rows, err := db.QueryContext(ctx, q, args...)
if err != nil {
return nil, fmt.Errorf("aggregate query logs: %w", err)
}
defer rows.Close()
var out []NameCount
for rows.Next() {
var n NameCount
if err := rows.Scan(&n.Name, &n.Count); err != nil {
return nil, err
}
out = append(out, n)
}
return out, rows.Err()
}
// TimeBucket is one point on the dashboard activity chart.
type TimeBucket struct {
Start time.Time `json:"start"`
Total int64 `json:"total"`
Blocked int64 `json:"blocked"`
Cached int64 `json:"cached"`
}
// ActivityBuckets groups query log rows into fixed-width time buckets covering
// the window [since, now].
func (db *DB) ActivityBuckets(ctx context.Context, since time.Time, bucket time.Duration, count int) ([]TimeBucket, error) {
if bucket <= 0 || count <= 0 {
return nil, nil
}
width := bucket.Milliseconds()
start := since.UnixMilli()
rows, err := db.QueryContext(ctx, `
SELECT (ts - ?) / ? AS b, COUNT(*), SUM(blocked), SUM(cache_hit)
FROM query_logs WHERE ts >= ?
GROUP BY b ORDER BY b`, start, width, start)
if err != nil {
return nil, fmt.Errorf("bucket query logs: %w", err)
}
defer rows.Close()
buckets := make([]TimeBucket, count)
for i := range buckets {
buckets[i].Start = time.UnixMilli(start + int64(i)*width)
}
for rows.Next() {
var idx, total int64
var blocked, cached sql.NullInt64
if err := rows.Scan(&idx, &total, &blocked, &cached); err != nil {
return nil, err
}
if idx < 0 || idx >= int64(count) {
continue
}
buckets[idx].Total = total
buckets[idx].Blocked = blocked.Int64
buckets[idx].Cached = cached.Int64
}
return buckets, rows.Err()
}
// QueryLogCount returns the number of stored rows.
func (db *DB) QueryLogCount(ctx context.Context) (int64, error) {
var n int64
err := db.QueryRowContext(ctx, `SELECT COUNT(*) FROM query_logs`).Scan(&n)
return n, err
}
// --- Audit log ----------------------------------------------------------
// InsertAudit appends an administrative audit entry.
func (db *DB) InsertAudit(ctx context.Context, e models.AuditEntry) error {
if e.Timestamp.IsZero() {
e.Timestamp = time.Now()
}
_, err := db.ExecContext(ctx, `
INSERT INTO audit_logs (ts, actor, source, client_ip, action, object_type, object_id, object_name, details)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)`,
e.Timestamp.UnixMilli(), e.Actor, e.Source, e.ClientIP, e.Action,
e.ObjectType, e.ObjectID, e.ObjectName, e.Details)
if err != nil {
return fmt.Errorf("write audit log: %w", err)
}
return nil
}
// AuditFilter narrows an audit log search.
type AuditFilter struct {
Search string
ObjectType string
Source string
Limit int
Offset int
}
// AuditLogs returns audit entries newest first, plus the total match count.
func (db *DB) AuditLogs(ctx context.Context, f AuditFilter) ([]models.AuditEntry, int, error) {
var conds []string
var args []any
if s := strings.TrimSpace(f.Search); s != "" {
conds = append(conds, "(action LIKE ? OR object_name LIKE ? OR details LIKE ? OR actor LIKE ?)")
pat := "%" + s + "%"
args = append(args, pat, pat, pat, pat)
}
if f.ObjectType != "" {
conds = append(conds, "object_type = ?")
args = append(args, f.ObjectType)
}
if f.Source != "" {
conds = append(conds, "source = ?")
args = append(args, f.Source)
}
whereSQL := ""
if len(conds) > 0 {
whereSQL = " WHERE " + strings.Join(conds, " AND ")
}
var total int
if err := db.QueryRowContext(ctx, `SELECT COUNT(*) FROM audit_logs`+whereSQL, args...).Scan(&total); err != nil {
return nil, 0, fmt.Errorf("count audit logs: %w", err)
}
q := `SELECT id, ts, actor, source, client_ip, action, object_type, object_id, object_name, details
FROM audit_logs` + whereSQL + ` ORDER BY ts DESC, id DESC`
qargs := args
if f.Limit > 0 {
q += ` LIMIT ? OFFSET ?`
qargs = append(append([]any{}, args...), f.Limit, f.Offset)
}
rows, err := db.QueryContext(ctx, q, qargs...)
if err != nil {
return nil, 0, fmt.Errorf("read audit logs: %w", err)
}
defer rows.Close()
var out []models.AuditEntry
for rows.Next() {
var e models.AuditEntry
var ts int64
if err := rows.Scan(&e.ID, &ts, &e.Actor, &e.Source, &e.ClientIP, &e.Action,
&e.ObjectType, &e.ObjectID, &e.ObjectName, &e.Details); err != nil {
return nil, 0, err
}
e.Timestamp = time.UnixMilli(ts)
out = append(out, e)
}
return out, total, rows.Err()
}
// PruneAuditLogs trims the audit log to the newest maxRows entries.
func (db *DB) PruneAuditLogs(ctx context.Context, maxRows int) (int64, error) {
if maxRows <= 0 {
return 0, nil
}
res, err := db.ExecContext(ctx, `
DELETE FROM audit_logs WHERE id NOT IN (
SELECT id FROM audit_logs ORDER BY ts DESC, id DESC LIMIT ?
)`, maxRows)
if err != nil {
return 0, fmt.Errorf("prune audit logs: %w", err)
}
n, _ := res.RowsAffected()
return n, nil
}
+891
View File
@@ -0,0 +1,891 @@
package database
import (
"context"
"database/sql"
"errors"
"fmt"
"strings"
"time"
"github.com/owen/vibedns/internal/models"
)
// --- Networks -----------------------------------------------------------
// Networks lists client networks. When withPolicies is true each network is
// populated with the policies assigned to it.
func (db *DB) Networks(ctx context.Context, search string, withPolicies bool) ([]models.Network, error) {
q := `SELECT id, name, cidr, description, enabled, created_at, updated_at FROM networks`
var args []any
if s := strings.TrimSpace(search); s != "" {
q += ` WHERE name LIKE ? OR cidr LIKE ? OR description LIKE ?`
pat := "%" + s + "%"
args = append(args, pat, pat, pat)
}
q += ` ORDER BY name`
rows, err := db.QueryContext(ctx, q, args...)
if err != nil {
return nil, fmt.Errorf("list networks: %w", err)
}
defer rows.Close()
var out []models.Network
index := map[int64]int{}
for rows.Next() {
var n models.Network
var enabled int
var created, updated int64
if err := rows.Scan(&n.ID, &n.Name, &n.CIDR, &n.Description, &enabled, &created, &updated); err != nil {
return nil, err
}
n.Enabled = enabled != 0
n.CreatedAt = time.Unix(created, 0)
n.UpdatedAt = time.Unix(updated, 0)
index[n.ID] = len(out)
out = append(out, n)
}
if err := rows.Err(); err != nil {
return nil, err
}
if !withPolicies || len(out) == 0 {
return out, nil
}
// One extra query joins in every assignment rather than N+1 lookups.
prows, err := db.QueryContext(ctx, `
SELECT np.network_id, p.id, p.name, p.description, p.enabled, p.block_action,
p.sinkhole_ipv4, p.sinkhole_ipv6, p.block_ttl
FROM network_policies np JOIN policies p ON p.id = np.policy_id
ORDER BY p.name`)
if err != nil {
return nil, fmt.Errorf("list network policies: %w", err)
}
defer prows.Close()
for prows.Next() {
var nid int64
var p models.Policy
var enabled int
if err := prows.Scan(&nid, &p.ID, &p.Name, &p.Description, &enabled, &p.BlockAction,
&p.SinkholeIPv4, &p.SinkholeIPv6, &p.BlockTTL); err != nil {
return nil, err
}
p.Enabled = enabled != 0
if i, ok := index[nid]; ok {
out[i].Policies = append(out[i].Policies, p)
}
}
return out, prows.Err()
}
// Network loads one network with its policy assignments.
func (db *DB) Network(ctx context.Context, id int64) (models.Network, error) {
var n models.Network
var enabled int
var created, updated int64
err := db.QueryRowContext(ctx,
`SELECT id, name, cidr, description, enabled, created_at, updated_at FROM networks WHERE id = ?`, id).
Scan(&n.ID, &n.Name, &n.CIDR, &n.Description, &enabled, &created, &updated)
if errors.Is(err, sql.ErrNoRows) {
return n, ErrNotFound
}
if err != nil {
return n, fmt.Errorf("load network: %w", err)
}
n.Enabled = enabled != 0
n.CreatedAt = time.Unix(created, 0)
n.UpdatedAt = time.Unix(updated, 0)
ids, err := db.networkPolicyIDs(ctx, id)
if err != nil {
return n, err
}
for _, pid := range ids {
p, err := db.Policy(ctx, pid)
if err != nil {
return n, err
}
n.Policies = append(n.Policies, p)
}
return n, nil
}
func (db *DB) networkPolicyIDs(ctx context.Context, networkID int64) ([]int64, error) {
rows, err := db.QueryContext(ctx,
`SELECT policy_id FROM network_policies WHERE network_id = ?`, networkID)
if err != nil {
return nil, fmt.Errorf("load policy assignments: %w", err)
}
defer rows.Close()
var out []int64
for rows.Next() {
var id int64
if err := rows.Scan(&id); err != nil {
return nil, err
}
out = append(out, id)
}
return out, rows.Err()
}
// CreateNetwork inserts a network and its policy assignments.
func (db *DB) CreateNetwork(ctx context.Context, n models.Network, policyIDs []int64) (models.Network, error) {
var id int64
err := db.InTx(ctx, func(tx *sql.Tx) error {
res, err := tx.ExecContext(ctx,
`INSERT INTO networks (name, cidr, description, enabled) VALUES (?, ?, ?, ?)`,
n.Name, n.CIDR, n.Description, boolInt(n.Enabled))
if err != nil {
if isUniqueViolation(err) {
return ErrConflict
}
return fmt.Errorf("create network: %w", err)
}
id, _ = res.LastInsertId()
return setNetworkPoliciesTx(ctx, tx, id, policyIDs)
})
if err != nil {
return models.Network{}, err
}
return db.Network(ctx, id)
}
// UpdateNetwork saves a network and replaces its policy assignments.
func (db *DB) UpdateNetwork(ctx context.Context, n models.Network, policyIDs []int64) (models.Network, error) {
err := db.InTx(ctx, func(tx *sql.Tx) error {
res, err := tx.ExecContext(ctx, `
UPDATE networks SET name = ?, cidr = ?, description = ?, enabled = ?, updated_at = unixepoch()
WHERE id = ?`, n.Name, n.CIDR, n.Description, boolInt(n.Enabled), n.ID)
if err != nil {
if isUniqueViolation(err) {
return ErrConflict
}
return fmt.Errorf("update network: %w", err)
}
if k, _ := res.RowsAffected(); k == 0 {
return ErrNotFound
}
return setNetworkPoliciesTx(ctx, tx, n.ID, policyIDs)
})
if err != nil {
return models.Network{}, err
}
return db.Network(ctx, n.ID)
}
func setNetworkPoliciesTx(ctx context.Context, tx *sql.Tx, networkID int64, policyIDs []int64) error {
if _, err := tx.ExecContext(ctx, `DELETE FROM network_policies WHERE network_id = ?`, networkID); err != nil {
return fmt.Errorf("clear policy assignments: %w", err)
}
if len(policyIDs) == 0 {
return nil
}
stmt, err := tx.PrepareContext(ctx,
`INSERT OR IGNORE INTO network_policies (network_id, policy_id) VALUES (?, ?)`)
if err != nil {
return err
}
defer stmt.Close()
for _, pid := range policyIDs {
if _, err := stmt.ExecContext(ctx, networkID, pid); err != nil {
return fmt.Errorf("assign policy %d: %w", pid, err)
}
}
return nil
}
// DeleteNetwork removes a network; assignments cascade.
func (db *DB) DeleteNetwork(ctx context.Context, id int64) error {
res, err := db.ExecContext(ctx, `DELETE FROM networks WHERE id = ?`, id)
if err != nil {
return fmt.Errorf("delete network: %w", err)
}
if n, _ := res.RowsAffected(); n == 0 {
return ErrNotFound
}
return nil
}
// SetNetworkEnabled toggles a network.
func (db *DB) SetNetworkEnabled(ctx context.Context, id int64, enabled bool) error {
res, err := db.ExecContext(ctx,
`UPDATE networks SET enabled = ?, updated_at = unixepoch() WHERE id = ?`, boolInt(enabled), id)
if err != nil {
return fmt.Errorf("update network: %w", err)
}
if n, _ := res.RowsAffected(); n == 0 {
return ErrNotFound
}
return nil
}
// --- Policies -----------------------------------------------------------
const policyColumns = `id, name, description, enabled, block_action, sinkhole_ipv4, sinkhole_ipv6,
block_ttl, created_at, updated_at`
func scanPolicy(sc interface{ Scan(...any) error }) (models.Policy, error) {
var p models.Policy
var enabled int
var created, updated int64
err := sc.Scan(&p.ID, &p.Name, &p.Description, &enabled, &p.BlockAction,
&p.SinkholeIPv4, &p.SinkholeIPv6, &p.BlockTTL, &created, &updated)
if err != nil {
return p, err
}
p.Enabled = enabled != 0
p.CreatedAt = time.Unix(created, 0)
p.UpdatedAt = time.Unix(updated, 0)
return p, nil
}
// Policies lists every policy with its list assignments and network usage.
func (db *DB) Policies(ctx context.Context) ([]models.Policy, error) {
rows, err := db.QueryContext(ctx, `SELECT `+policyColumns+` FROM policies ORDER BY name`)
if err != nil {
return nil, fmt.Errorf("list policies: %w", err)
}
defer rows.Close()
var out []models.Policy
index := map[int64]int{}
for rows.Next() {
p, err := scanPolicy(rows)
if err != nil {
return nil, err
}
index[p.ID] = len(out)
out = append(out, p)
}
if err := rows.Err(); err != nil {
return nil, err
}
if len(out) == 0 {
return out, nil
}
lrows, err := db.QueryContext(ctx, `
SELECT pl.policy_id, l.id, l.kind, l.name
FROM policy_lists pl JOIN domain_lists l ON l.id = pl.list_id
ORDER BY l.name`)
if err != nil {
return nil, fmt.Errorf("list policy lists: %w", err)
}
defer lrows.Close()
for lrows.Next() {
var pid, lid int64
var kind, name string
if err := lrows.Scan(&pid, &lid, &kind, &name); err != nil {
return nil, err
}
i, ok := index[pid]
if !ok {
continue
}
if kind == models.KindAllowlist {
out[i].AllowlistIDs = append(out[i].AllowlistIDs, lid)
out[i].AllowlistName = append(out[i].AllowlistName, name)
} else {
out[i].BlacklistIDs = append(out[i].BlacklistIDs, lid)
out[i].BlacklistName = append(out[i].BlacklistName, name)
}
}
if err := lrows.Err(); err != nil {
return nil, err
}
nrows, err := db.QueryContext(ctx,
`SELECT policy_id, COUNT(*) FROM network_policies GROUP BY policy_id`)
if err != nil {
return nil, err
}
defer nrows.Close()
for nrows.Next() {
var pid int64
var c int
if err := nrows.Scan(&pid, &c); err != nil {
return nil, err
}
if i, ok := index[pid]; ok {
out[i].NetworkCount = c
}
}
return out, nrows.Err()
}
// Policy loads one policy with its list assignments.
func (db *DB) Policy(ctx context.Context, id int64) (models.Policy, error) {
row := db.QueryRowContext(ctx, `SELECT `+policyColumns+` FROM policies WHERE id = ?`, id)
p, err := scanPolicy(row)
if errors.Is(err, sql.ErrNoRows) {
return p, ErrNotFound
}
if err != nil {
return p, fmt.Errorf("load policy: %w", err)
}
rows, err := db.QueryContext(ctx, `
SELECT l.id, l.kind, l.name FROM policy_lists pl
JOIN domain_lists l ON l.id = pl.list_id
WHERE pl.policy_id = ? ORDER BY l.name`, id)
if err != nil {
return p, fmt.Errorf("load policy lists: %w", err)
}
defer rows.Close()
for rows.Next() {
var lid int64
var kind, name string
if err := rows.Scan(&lid, &kind, &name); err != nil {
return p, err
}
if kind == models.KindAllowlist {
p.AllowlistIDs = append(p.AllowlistIDs, lid)
p.AllowlistName = append(p.AllowlistName, name)
} else {
p.BlacklistIDs = append(p.BlacklistIDs, lid)
p.BlacklistName = append(p.BlacklistName, name)
}
}
return p, rows.Err()
}
// CreatePolicy inserts a policy with its list assignments.
func (db *DB) CreatePolicy(ctx context.Context, p models.Policy, listIDs []int64) (models.Policy, error) {
var id int64
err := db.InTx(ctx, func(tx *sql.Tx) error {
res, err := tx.ExecContext(ctx, `
INSERT INTO policies (name, description, enabled, block_action, sinkhole_ipv4, sinkhole_ipv6, block_ttl)
VALUES (?, ?, ?, ?, ?, ?, ?)`,
p.Name, p.Description, boolInt(p.Enabled), string(p.BlockAction),
p.SinkholeIPv4, p.SinkholeIPv6, p.BlockTTL)
if err != nil {
if isUniqueViolation(err) {
return ErrConflict
}
return fmt.Errorf("create policy: %w", err)
}
id, _ = res.LastInsertId()
return setPolicyListsTx(ctx, tx, id, listIDs)
})
if err != nil {
return models.Policy{}, err
}
return db.Policy(ctx, id)
}
// UpdatePolicy saves a policy and replaces its list assignments.
func (db *DB) UpdatePolicy(ctx context.Context, p models.Policy, listIDs []int64) (models.Policy, error) {
err := db.InTx(ctx, func(tx *sql.Tx) error {
res, err := tx.ExecContext(ctx, `
UPDATE policies SET name = ?, description = ?, enabled = ?, block_action = ?,
sinkhole_ipv4 = ?, sinkhole_ipv6 = ?, block_ttl = ?, updated_at = unixepoch()
WHERE id = ?`,
p.Name, p.Description, boolInt(p.Enabled), string(p.BlockAction),
p.SinkholeIPv4, p.SinkholeIPv6, p.BlockTTL, p.ID)
if err != nil {
if isUniqueViolation(err) {
return ErrConflict
}
return fmt.Errorf("update policy: %w", err)
}
if n, _ := res.RowsAffected(); n == 0 {
return ErrNotFound
}
return setPolicyListsTx(ctx, tx, p.ID, listIDs)
})
if err != nil {
return models.Policy{}, err
}
return db.Policy(ctx, p.ID)
}
func setPolicyListsTx(ctx context.Context, tx *sql.Tx, policyID int64, listIDs []int64) error {
if _, err := tx.ExecContext(ctx, `DELETE FROM policy_lists WHERE policy_id = ?`, policyID); err != nil {
return fmt.Errorf("clear policy lists: %w", err)
}
if len(listIDs) == 0 {
return nil
}
stmt, err := tx.PrepareContext(ctx,
`INSERT OR IGNORE INTO policy_lists (policy_id, list_id) VALUES (?, ?)`)
if err != nil {
return err
}
defer stmt.Close()
for _, lid := range listIDs {
if _, err := stmt.ExecContext(ctx, policyID, lid); err != nil {
return fmt.Errorf("assign list %d: %w", lid, err)
}
}
return nil
}
// DeletePolicy removes a policy; assignments cascade.
func (db *DB) DeletePolicy(ctx context.Context, id int64) error {
res, err := db.ExecContext(ctx, `DELETE FROM policies WHERE id = ?`, id)
if err != nil {
return fmt.Errorf("delete policy: %w", err)
}
if n, _ := res.RowsAffected(); n == 0 {
return ErrNotFound
}
return nil
}
// SetPolicyEnabled toggles a policy.
func (db *DB) SetPolicyEnabled(ctx context.Context, id int64, enabled bool) error {
res, err := db.ExecContext(ctx,
`UPDATE policies SET enabled = ?, updated_at = unixepoch() WHERE id = ?`, boolInt(enabled), id)
if err != nil {
return fmt.Errorf("update policy: %w", err)
}
if n, _ := res.RowsAffected(); n == 0 {
return ErrNotFound
}
return nil
}
// --- Domain lists -------------------------------------------------------
// DomainLists returns blacklists or allowlists (kind may be "" for both) with
// domain counts and the policies that reference them.
func (db *DB) DomainLists(ctx context.Context, kind, search string) ([]models.DomainList, error) {
q := `SELECT l.id, l.kind, l.name, l.description, l.enabled, l.source_url, l.created_at, l.updated_at,
(SELECT COUNT(*) FROM domain_entries e WHERE e.list_id = l.id) AS domain_count
FROM domain_lists l`
var where []string
var args []any
if kind != "" {
where = append(where, "l.kind = ?")
args = append(args, kind)
}
if s := strings.TrimSpace(search); s != "" {
where = append(where, "(l.name LIKE ? OR l.description LIKE ?)")
pat := "%" + s + "%"
args = append(args, pat, pat)
}
if len(where) > 0 {
q += " WHERE " + strings.Join(where, " AND ")
}
q += " ORDER BY l.name"
rows, err := db.QueryContext(ctx, q, args...)
if err != nil {
return nil, fmt.Errorf("list domain lists: %w", err)
}
defer rows.Close()
var out []models.DomainList
index := map[int64]int{}
for rows.Next() {
var l models.DomainList
var enabled int
var created, updated int64
if err := rows.Scan(&l.ID, &l.Kind, &l.Name, &l.Description, &enabled, &l.SourceURL,
&created, &updated, &l.DomainCount); err != nil {
return nil, err
}
l.Enabled = enabled != 0
l.CreatedAt = time.Unix(created, 0)
l.UpdatedAt = time.Unix(updated, 0)
index[l.ID] = len(out)
out = append(out, l)
}
if err := rows.Err(); err != nil {
return nil, err
}
if len(out) == 0 {
return out, nil
}
urows, err := db.QueryContext(ctx, `
SELECT pl.list_id, p.name FROM policy_lists pl
JOIN policies p ON p.id = pl.policy_id ORDER BY p.name`)
if err != nil {
return nil, err
}
defer urows.Close()
for urows.Next() {
var lid int64
var name string
if err := urows.Scan(&lid, &name); err != nil {
return nil, err
}
if i, ok := index[lid]; ok {
out[i].UsedBy = append(out[i].UsedBy, name)
}
}
return out, urows.Err()
}
// DomainList loads one list with its domain count and referencing policies.
func (db *DB) DomainList(ctx context.Context, id int64) (models.DomainList, error) {
var l models.DomainList
var enabled int
var created, updated int64
err := db.QueryRowContext(ctx, `
SELECT l.id, l.kind, l.name, l.description, l.enabled, l.source_url, l.created_at, l.updated_at,
(SELECT COUNT(*) FROM domain_entries e WHERE e.list_id = l.id)
FROM domain_lists l WHERE l.id = ?`, id).
Scan(&l.ID, &l.Kind, &l.Name, &l.Description, &enabled, &l.SourceURL, &created, &updated, &l.DomainCount)
if errors.Is(err, sql.ErrNoRows) {
return l, ErrNotFound
}
if err != nil {
return l, fmt.Errorf("load domain list: %w", err)
}
l.Enabled = enabled != 0
l.CreatedAt = time.Unix(created, 0)
l.UpdatedAt = time.Unix(updated, 0)
rows, err := db.QueryContext(ctx, `
SELECT p.name FROM policy_lists pl JOIN policies p ON p.id = pl.policy_id
WHERE pl.list_id = ? ORDER BY p.name`, id)
if err != nil {
return l, err
}
defer rows.Close()
for rows.Next() {
var n string
if err := rows.Scan(&n); err != nil {
return l, err
}
l.UsedBy = append(l.UsedBy, n)
}
return l, rows.Err()
}
// CreateDomainList inserts a blacklist or allowlist.
func (db *DB) CreateDomainList(ctx context.Context, l models.DomainList) (models.DomainList, error) {
res, err := db.ExecContext(ctx, `
INSERT INTO domain_lists (kind, name, description, enabled, source_url)
VALUES (?, ?, ?, ?, ?)`,
l.Kind, l.Name, l.Description, boolInt(l.Enabled), l.SourceURL)
if err != nil {
if isUniqueViolation(err) {
return models.DomainList{}, ErrConflict
}
return models.DomainList{}, fmt.Errorf("create domain list: %w", err)
}
id, _ := res.LastInsertId()
return db.DomainList(ctx, id)
}
// UpdateDomainList saves list metadata.
func (db *DB) UpdateDomainList(ctx context.Context, l models.DomainList) (models.DomainList, error) {
res, err := db.ExecContext(ctx, `
UPDATE domain_lists SET name = ?, description = ?, enabled = ?, source_url = ?,
updated_at = unixepoch()
WHERE id = ?`, l.Name, l.Description, boolInt(l.Enabled), l.SourceURL, l.ID)
if err != nil {
if isUniqueViolation(err) {
return models.DomainList{}, ErrConflict
}
return models.DomainList{}, fmt.Errorf("update domain list: %w", err)
}
if n, _ := res.RowsAffected(); n == 0 {
return models.DomainList{}, ErrNotFound
}
return db.DomainList(ctx, l.ID)
}
// DeleteDomainList removes a list and every domain in it.
func (db *DB) DeleteDomainList(ctx context.Context, id int64) error {
res, err := db.ExecContext(ctx, `DELETE FROM domain_lists WHERE id = ?`, id)
if err != nil {
return fmt.Errorf("delete domain list: %w", err)
}
if n, _ := res.RowsAffected(); n == 0 {
return ErrNotFound
}
return nil
}
// SetDomainListEnabled toggles a list.
func (db *DB) SetDomainListEnabled(ctx context.Context, id int64, enabled bool) error {
res, err := db.ExecContext(ctx,
`UPDATE domain_lists SET enabled = ?, updated_at = unixepoch() WHERE id = ?`, boolInt(enabled), id)
if err != nil {
return fmt.Errorf("update domain list: %w", err)
}
if n, _ := res.RowsAffected(); n == 0 {
return ErrNotFound
}
return nil
}
// --- Domain entries -----------------------------------------------------
// DomainEntries pages through the domains of one list.
func (db *DB) DomainEntries(ctx context.Context, listID int64, search string, limit, offset int) ([]models.DomainEntry, int, error) {
where := " WHERE list_id = ?"
args := []any{listID}
if s := strings.TrimSpace(search); s != "" {
where += " AND domain LIKE ?"
args = append(args, "%"+strings.ToLower(s)+"%")
}
var total int
if err := db.QueryRowContext(ctx, `SELECT COUNT(*) FROM domain_entries`+where, args...).Scan(&total); err != nil {
return nil, 0, fmt.Errorf("count domains: %w", err)
}
q := `SELECT id, list_id, domain, match_subdomains, enabled, comment, created_at
FROM domain_entries` + where + ` ORDER BY domain`
qargs := args
if limit > 0 {
q += " LIMIT ? OFFSET ?"
qargs = append(append([]any{}, args...), limit, offset)
}
rows, err := db.QueryContext(ctx, q, qargs...)
if err != nil {
return nil, 0, fmt.Errorf("list domains: %w", err)
}
defer rows.Close()
var out []models.DomainEntry
for rows.Next() {
var e models.DomainEntry
var sub, enabled int
var created int64
if err := rows.Scan(&e.ID, &e.ListID, &e.Domain, &sub, &enabled, &e.Comment, &created); err != nil {
return nil, 0, err
}
e.MatchSubdomains = sub != 0
e.Enabled = enabled != 0
e.CreatedAt = time.Unix(created, 0)
out = append(out, e)
}
return out, total, rows.Err()
}
// AddDomain inserts a single domain. It returns ErrConflict when the domain is
// already present in that list.
func (db *DB) AddDomain(ctx context.Context, e models.DomainEntry) (models.DomainEntry, error) {
res, err := db.ExecContext(ctx, `
INSERT INTO domain_entries (list_id, domain, match_subdomains, enabled, comment)
VALUES (?, ?, ?, ?, ?)`,
e.ListID, e.Domain, boolInt(e.MatchSubdomains), boolInt(e.Enabled), e.Comment)
if err != nil {
if isUniqueViolation(err) {
return models.DomainEntry{}, ErrConflict
}
return models.DomainEntry{}, fmt.Errorf("add domain: %w", err)
}
id, _ := res.LastInsertId()
e.ID = id
e.CreatedAt = time.Now()
db.touchList(ctx, e.ListID)
return e, nil
}
// UpdateDomain saves an existing domain entry.
func (db *DB) UpdateDomain(ctx context.Context, e models.DomainEntry) error {
res, err := db.ExecContext(ctx, `
UPDATE domain_entries SET domain = ?, match_subdomains = ?, enabled = ?, comment = ?
WHERE id = ?`, e.Domain, boolInt(e.MatchSubdomains), boolInt(e.Enabled), e.Comment, e.ID)
if err != nil {
if isUniqueViolation(err) {
return ErrConflict
}
return fmt.Errorf("update domain: %w", err)
}
if n, _ := res.RowsAffected(); n == 0 {
return ErrNotFound
}
db.touchList(ctx, e.ListID)
return nil
}
// DeleteDomain removes one domain entry.
func (db *DB) DeleteDomain(ctx context.Context, id int64) error {
var listID int64
_ = db.QueryRowContext(ctx, `SELECT list_id FROM domain_entries WHERE id = ?`, id).Scan(&listID)
res, err := db.ExecContext(ctx, `DELETE FROM domain_entries WHERE id = ?`, id)
if err != nil {
return fmt.Errorf("delete domain: %w", err)
}
if n, _ := res.RowsAffected(); n == 0 {
return ErrNotFound
}
db.touchList(ctx, listID)
return nil
}
// ClearDomains removes every domain from a list and returns how many went.
func (db *DB) ClearDomains(ctx context.Context, listID int64) (int64, error) {
res, err := db.ExecContext(ctx, `DELETE FROM domain_entries WHERE list_id = ?`, listID)
if err != nil {
return 0, fmt.Errorf("clear domains: %w", err)
}
n, _ := res.RowsAffected()
db.touchList(ctx, listID)
return n, nil
}
func (db *DB) touchList(ctx context.Context, listID int64) {
if listID == 0 {
return
}
_, _ = db.ExecContext(ctx, `UPDATE domain_lists SET updated_at = unixepoch() WHERE id = ?`, listID)
}
// ImportDomains bulk-inserts normalised domains into a list.
//
// Everything happens inside one transaction with a single prepared statement,
// so importing a few hundred thousand domains is one commit rather than one
// commit per domain. INSERT OR IGNORE gives duplicate detection for free.
func (db *DB) ImportDomains(ctx context.Context, listID int64, domains []ImportDomain) (imported, duplicates int, err error) {
if len(domains) == 0 {
return 0, 0, nil
}
err = db.InTx(ctx, func(tx *sql.Tx) error {
stmt, err := tx.PrepareContext(ctx, `
INSERT OR IGNORE INTO domain_entries (list_id, domain, match_subdomains, enabled, comment)
VALUES (?, ?, ?, 1, ?)`)
if err != nil {
return err
}
defer stmt.Close()
for _, d := range domains {
res, err := stmt.ExecContext(ctx, listID, d.Domain, boolInt(d.MatchSubdomains), d.Comment)
if err != nil {
return fmt.Errorf("import %q: %w", d.Domain, err)
}
if n, _ := res.RowsAffected(); n > 0 {
imported++
} else {
duplicates++
}
}
_, err = tx.ExecContext(ctx, `UPDATE domain_lists SET updated_at = unixepoch() WHERE id = ?`, listID)
return err
})
if err != nil {
return 0, 0, err
}
return imported, duplicates, nil
}
// ImportDomain is one normalised domain destined for a list.
type ImportDomain struct {
Domain string
MatchSubdomains bool
Comment string
}
// SnapshotDomainEntry is the minimal shape the in-memory matcher needs.
type SnapshotDomainEntry struct {
ListID int64
Domain string
MatchSubdomains bool
}
// SnapshotDomains streams every enabled domain of every enabled list. It is
// called on configuration change, never on the DNS query path.
func (db *DB) SnapshotDomains(ctx context.Context, fn func(SnapshotDomainEntry)) error {
rows, err := db.QueryContext(ctx, `
SELECT e.list_id, e.domain, e.match_subdomains
FROM domain_entries e
JOIN domain_lists l ON l.id = e.list_id
WHERE e.enabled = 1 AND l.enabled = 1`)
if err != nil {
return fmt.Errorf("snapshot domains: %w", err)
}
defer rows.Close()
for rows.Next() {
var e SnapshotDomainEntry
var sub int
if err := rows.Scan(&e.ListID, &e.Domain, &sub); err != nil {
return err
}
e.MatchSubdomains = sub != 0
fn(e)
}
return rows.Err()
}
// CountDomainLists returns blacklist and total-domain counts for the dashboard.
func (db *DB) CountDomainLists(ctx context.Context) (blacklists, blacklistDomains, allowlists, allowlistDomains int, err error) {
err = db.QueryRowContext(ctx, `
SELECT
(SELECT COUNT(*) FROM domain_lists WHERE kind = 'blacklist'),
(SELECT COUNT(*) FROM domain_entries e JOIN domain_lists l ON l.id = e.list_id WHERE l.kind = 'blacklist'),
(SELECT COUNT(*) FROM domain_lists WHERE kind = 'allowlist'),
(SELECT COUNT(*) FROM domain_entries e JOIN domain_lists l ON l.id = e.list_id WHERE l.kind = 'allowlist')`).
Scan(&blacklists, &blacklistDomains, &allowlists, &allowlistDomains)
if err != nil {
return 0, 0, 0, 0, fmt.Errorf("count domain lists: %w", err)
}
return
}
// ExportDomains streams every domain of a list in sorted order.
func (db *DB) ExportDomains(ctx context.Context, listID int64, fn func(domain string, matchSubdomains bool)) error {
rows, err := db.QueryContext(ctx,
`SELECT domain, match_subdomains FROM domain_entries WHERE list_id = ? ORDER BY domain`, listID)
if err != nil {
return fmt.Errorf("export domains: %w", err)
}
defer rows.Close()
for rows.Next() {
var d string
var sub int
if err := rows.Scan(&d, &sub); err != nil {
return err
}
fn(d, sub != 0)
}
return rows.Err()
}
// SnapshotNetworks loads enabled networks with the IDs of their enabled
// policies, for building the CIDR index.
func (db *DB) SnapshotNetworks(ctx context.Context) ([]models.Network, map[int64][]int64, error) {
rows, err := db.QueryContext(ctx, `
SELECT id, name, cidr, description, enabled, created_at, updated_at
FROM networks WHERE enabled = 1`)
if err != nil {
return nil, nil, fmt.Errorf("snapshot networks: %w", err)
}
defer rows.Close()
var nets []models.Network
for rows.Next() {
var n models.Network
var enabled int
var created, updated int64
if err := rows.Scan(&n.ID, &n.Name, &n.CIDR, &n.Description, &enabled, &created, &updated); err != nil {
return nil, nil, err
}
n.Enabled = enabled != 0
n.CreatedAt = time.Unix(created, 0)
n.UpdatedAt = time.Unix(updated, 0)
nets = append(nets, n)
}
if err := rows.Err(); err != nil {
return nil, nil, err
}
arows, err := db.QueryContext(ctx, `
SELECT np.network_id, np.policy_id FROM network_policies np
JOIN policies p ON p.id = np.policy_id
WHERE p.enabled = 1`)
if err != nil {
return nil, nil, fmt.Errorf("snapshot policy assignments: %w", err)
}
defer arows.Close()
assign := map[int64][]int64{}
for arows.Next() {
var nid, pid int64
if err := arows.Scan(&nid, &pid); err != nil {
return nil, nil, err
}
assign[nid] = append(assign[nid], pid)
}
return nets, assign, arows.Err()
}
+79
View File
@@ -0,0 +1,79 @@
package database
import (
"context"
"database/sql"
"fmt"
)
// Settings returns every stored setting as a key/value map.
func (db *DB) Settings(ctx context.Context) (map[string]string, error) {
rows, err := db.QueryContext(ctx, `SELECT key, value FROM settings`)
if err != nil {
return nil, fmt.Errorf("load settings: %w", err)
}
defer rows.Close()
out := map[string]string{}
for rows.Next() {
var k, v string
if err := rows.Scan(&k, &v); err != nil {
return nil, err
}
out[k] = v
}
return out, rows.Err()
}
// Setting reads a single setting. It returns ("", false, nil) when unset.
func (db *DB) Setting(ctx context.Context, key string) (string, bool, error) {
var v string
err := db.QueryRowContext(ctx, `SELECT value FROM settings WHERE key = ?`, key).Scan(&v)
switch {
case err == sql.ErrNoRows:
return "", false, nil
case err != nil:
return "", false, fmt.Errorf("read setting %s: %w", key, err)
}
return v, true, nil
}
// SetSetting writes one setting.
func (db *DB) SetSetting(ctx context.Context, key, value string) error {
_, err := db.ExecContext(ctx, `
INSERT INTO settings (key, value, updated_at) VALUES (?, ?, unixepoch())
ON CONFLICT (key) DO UPDATE SET value = excluded.value, updated_at = unixepoch()`,
key, value)
if err != nil {
return fmt.Errorf("write setting %s: %w", key, err)
}
return nil
}
// SetSettings writes several settings atomically.
func (db *DB) SetSettings(ctx context.Context, values map[string]string) error {
if len(values) == 0 {
return nil
}
return db.InTx(ctx, func(tx *sql.Tx) error {
stmt, err := tx.PrepareContext(ctx, `
INSERT INTO settings (key, value, updated_at) VALUES (?, ?, unixepoch())
ON CONFLICT (key) DO UPDATE SET value = excluded.value, updated_at = unixepoch()`)
if err != nil {
return fmt.Errorf("prepare setting write: %w", err)
}
defer stmt.Close()
for k, v := range values {
if _, err := stmt.ExecContext(ctx, k, v); err != nil {
return fmt.Errorf("write setting %s: %w", k, err)
}
}
return nil
})
}
// DeleteSetting removes a setting, reverting it to its built-in default.
func (db *DB) DeleteSetting(ctx context.Context, key string) error {
_, err := db.ExecContext(ctx, `DELETE FROM settings WHERE key = ?`, key)
return err
}
+654
View File
@@ -0,0 +1,654 @@
package database
import (
"context"
"database/sql"
"errors"
"fmt"
"strings"
"time"
"github.com/owen/vibedns/internal/models"
)
const zoneColumns = `id, name, kind, description, enabled, default_ttl, primary_ns, admin_email,
serial, refresh, retry, expire, minimum, auto_serial, created_at, updated_at`
func scanZone(sc interface{ Scan(...any) error }) (models.Zone, error) {
var z models.Zone
var enabled, autoSerial int
var created, updated int64
err := sc.Scan(&z.ID, &z.Name, &z.Kind, &z.Description, &enabled, &z.DefaultTTL,
&z.PrimaryNS, &z.AdminEmail, &z.Serial, &z.Refresh, &z.Retry, &z.Expire,
&z.Minimum, &autoSerial, &created, &updated)
if err != nil {
return z, err
}
z.Enabled = enabled != 0
z.AutoSerial = autoSerial != 0
z.CreatedAt = time.Unix(created, 0)
z.UpdatedAt = time.Unix(updated, 0)
return z, nil
}
// ZoneFilter narrows a zone listing.
type ZoneFilter struct {
Kind string // "", "forward", "reverse4", "reverse6", or "reverse" for both
Search string
}
// Zones lists zones with their record counts, ordered by name.
func (db *DB) Zones(ctx context.Context, f ZoneFilter) ([]models.Zone, error) {
var where []string
var args []any
switch f.Kind {
case "":
// no filter
case "reverse":
where = append(where, "z.kind IN ('reverse4','reverse6')")
default:
where = append(where, "z.kind = ?")
args = append(args, f.Kind)
}
if s := strings.TrimSpace(f.Search); s != "" {
where = append(where, "(z.name LIKE ? OR z.description LIKE ?)")
pat := "%" + s + "%"
args = append(args, pat, pat)
}
q := `SELECT ` + prefixCols(zoneColumns, "z") + `,
(SELECT COUNT(*) FROM records r WHERE r.zone_id = z.id) AS record_count
FROM zones z`
if len(where) > 0 {
q += " WHERE " + strings.Join(where, " AND ")
}
q += " ORDER BY z.name"
rows, err := db.QueryContext(ctx, q, args...)
if err != nil {
return nil, fmt.Errorf("list zones: %w", err)
}
defer rows.Close()
var out []models.Zone
for rows.Next() {
var z models.Zone
var enabled, autoSerial int
var created, updated int64
err := rows.Scan(&z.ID, &z.Name, &z.Kind, &z.Description, &enabled, &z.DefaultTTL,
&z.PrimaryNS, &z.AdminEmail, &z.Serial, &z.Refresh, &z.Retry, &z.Expire,
&z.Minimum, &autoSerial, &created, &updated, &z.RecordCount)
if err != nil {
return nil, err
}
z.Enabled = enabled != 0
z.AutoSerial = autoSerial != 0
z.CreatedAt = time.Unix(created, 0)
z.UpdatedAt = time.Unix(updated, 0)
out = append(out, z)
}
return out, rows.Err()
}
// prefixCols qualifies a comma separated column list with a table alias.
func prefixCols(cols, alias string) string {
parts := strings.Split(cols, ",")
for i, p := range parts {
parts[i] = alias + "." + strings.TrimSpace(p)
}
return strings.Join(parts, ", ")
}
// Zone loads a single zone by ID.
func (db *DB) Zone(ctx context.Context, id int64) (models.Zone, error) {
row := db.QueryRowContext(ctx, `SELECT `+zoneColumns+` FROM zones WHERE id = ?`, id)
z, err := scanZone(row)
if errors.Is(err, sql.ErrNoRows) {
return z, ErrNotFound
}
if err != nil {
return z, fmt.Errorf("load zone: %w", err)
}
return z, nil
}
// ZoneByName loads a zone by its normalised FQDN.
func (db *DB) ZoneByName(ctx context.Context, name string) (models.Zone, error) {
row := db.QueryRowContext(ctx, `SELECT `+zoneColumns+` FROM zones WHERE name = ?`, name)
z, err := scanZone(row)
if errors.Is(err, sql.ErrNoRows) {
return z, ErrNotFound
}
if err != nil {
return z, fmt.Errorf("load zone: %w", err)
}
return z, nil
}
// CreateZone inserts a zone. The caller is responsible for having validated and
// normalised the zone name.
func (db *DB) CreateZone(ctx context.Context, z models.Zone) (models.Zone, error) {
res, err := db.ExecContext(ctx, `
INSERT INTO zones (name, kind, description, enabled, default_ttl, primary_ns, admin_email,
serial, refresh, retry, expire, minimum, auto_serial)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`,
z.Name, string(z.Kind), z.Description, boolInt(z.Enabled), z.DefaultTTL, z.PrimaryNS,
z.AdminEmail, z.Serial, z.Refresh, z.Retry, z.Expire, z.Minimum, boolInt(z.AutoSerial))
if err != nil {
if isUniqueViolation(err) {
return models.Zone{}, ErrConflict
}
return models.Zone{}, fmt.Errorf("create zone: %w", err)
}
id, _ := res.LastInsertId()
return db.Zone(ctx, id)
}
// UpdateZone saves zone metadata. Records are managed separately.
func (db *DB) UpdateZone(ctx context.Context, z models.Zone) (models.Zone, error) {
res, err := db.ExecContext(ctx, `
UPDATE zones SET name = ?, kind = ?, description = ?, enabled = ?, default_ttl = ?,
primary_ns = ?, admin_email = ?, serial = ?, refresh = ?, retry = ?, expire = ?,
minimum = ?, auto_serial = ?, updated_at = unixepoch()
WHERE id = ?`,
z.Name, string(z.Kind), z.Description, boolInt(z.Enabled), z.DefaultTTL, z.PrimaryNS,
z.AdminEmail, z.Serial, z.Refresh, z.Retry, z.Expire, z.Minimum, boolInt(z.AutoSerial), z.ID)
if err != nil {
if isUniqueViolation(err) {
return models.Zone{}, ErrConflict
}
return models.Zone{}, fmt.Errorf("update zone: %w", err)
}
if n, _ := res.RowsAffected(); n == 0 {
return models.Zone{}, ErrNotFound
}
return db.Zone(ctx, z.ID)
}
// SetZoneEnabled toggles a zone without touching its records.
func (db *DB) SetZoneEnabled(ctx context.Context, id int64, enabled bool) error {
res, err := db.ExecContext(ctx,
`UPDATE zones SET enabled = ?, updated_at = unixepoch() WHERE id = ?`, boolInt(enabled), id)
if err != nil {
return fmt.Errorf("update zone: %w", err)
}
if n, _ := res.RowsAffected(); n == 0 {
return ErrNotFound
}
return nil
}
// DeleteZone removes a zone and, by foreign key cascade, all of its records.
func (db *DB) DeleteZone(ctx context.Context, id int64) error {
res, err := db.ExecContext(ctx, `DELETE FROM zones WHERE id = ?`, id)
if err != nil {
return fmt.Errorf("delete zone: %w", err)
}
if n, _ := res.RowsAffected(); n == 0 {
return ErrNotFound
}
return nil
}
// CloneZone copies a zone and every record into a new zone name. Record names
// are copied verbatim; rdata that referenced the old apex is rewritten so the
// clone is self-consistent.
func (db *DB) CloneZone(ctx context.Context, srcID int64, newName, description string) (models.Zone, error) {
var newID int64
err := db.InTx(ctx, func(tx *sql.Tx) error {
row := tx.QueryRowContext(ctx, `SELECT `+zoneColumns+` FROM zones WHERE id = ?`, srcID)
src, err := scanZone(row)
if errors.Is(err, sql.ErrNoRows) {
return ErrNotFound
}
if err != nil {
return fmt.Errorf("load source zone: %w", err)
}
res, err := tx.ExecContext(ctx, `
INSERT INTO zones (name, kind, description, enabled, default_ttl, primary_ns, admin_email,
serial, refresh, retry, expire, minimum, auto_serial)
VALUES (?, ?, ?, ?, ?, ?, ?, 1, ?, ?, ?, ?, ?)`,
newName, string(src.Kind), description, boolInt(src.Enabled), src.DefaultTTL,
src.PrimaryNS, src.AdminEmail, src.Refresh, src.Retry, src.Expire, src.Minimum,
boolInt(src.AutoSerial))
if err != nil {
if isUniqueViolation(err) {
return ErrConflict
}
return fmt.Errorf("create cloned zone: %w", err)
}
newID, _ = res.LastInsertId()
// REPLACE rewrites references to the source apex inside rdata so that
// e.g. "www CNAME example.com." becomes "www CNAME clone.example."
_, err = tx.ExecContext(ctx, `
INSERT INTO records (zone_id, name, type, data, ttl, enabled, comment)
SELECT ?, name, type, REPLACE(data, ?, ?), ttl, enabled, comment
FROM records WHERE zone_id = ? AND type <> 'SOA'`,
newID, src.Name, newName, srcID)
if err != nil {
return fmt.Errorf("copy records: %w", err)
}
return nil
})
if err != nil {
return models.Zone{}, err
}
return db.Zone(ctx, newID)
}
// BumpSerial increments a zone's SOA serial if auto-serial is enabled.
// Serials wrap according to RFC 1982 arithmetic, which SQLite's modulo gives us
// for free by wrapping past 2^32-1 back to 1.
func (db *DB) BumpSerial(ctx context.Context, zoneID int64) error {
_, err := db.ExecContext(ctx, `
UPDATE zones
SET serial = CASE WHEN serial >= 4294967295 THEN 1 ELSE serial + 1 END,
updated_at = unixepoch()
WHERE id = ? AND auto_serial = 1`, zoneID)
return err
}
func bumpSerialTx(ctx context.Context, tx *sql.Tx, zoneID int64) error {
_, err := tx.ExecContext(ctx, `
UPDATE zones
SET serial = CASE WHEN serial >= 4294967295 THEN 1 ELSE serial + 1 END,
updated_at = unixepoch()
WHERE id = ? AND auto_serial = 1`, zoneID)
return err
}
// --- Records ------------------------------------------------------------
const recordColumns = `id, zone_id, name, type, data, ttl, enabled, comment, created_at, updated_at`
func scanRecord(sc interface{ Scan(...any) error }) (models.Record, error) {
var r models.Record
var ttl sql.NullInt64
var enabled int
var created, updated int64
err := sc.Scan(&r.ID, &r.ZoneID, &r.Name, &r.Type, &r.Data, &ttl, &enabled, &r.Comment, &created, &updated)
if err != nil {
return r, err
}
if ttl.Valid {
v := uint32(ttl.Int64)
r.TTL = &v
}
r.Enabled = enabled != 0
r.CreatedAt = time.Unix(created, 0)
r.UpdatedAt = time.Unix(updated, 0)
return r, nil
}
// RecordFilter narrows a record listing.
type RecordFilter struct {
ZoneID int64 // 0 means all zones
Search string // matches name or data
Type string
Enabled string // "", "enabled", "disabled"
Limit int
Offset int
}
func (f RecordFilter) whereClause() (string, []any) {
var where []string
var args []any
if f.ZoneID > 0 {
where = append(where, "r.zone_id = ?")
args = append(args, f.ZoneID)
}
if s := strings.TrimSpace(f.Search); s != "" {
where = append(where, "(r.name LIKE ? OR r.data LIKE ? OR r.comment LIKE ?)")
pat := "%" + s + "%"
args = append(args, pat, pat, pat)
}
if t := strings.ToUpper(strings.TrimSpace(f.Type)); t != "" {
where = append(where, "r.type = ?")
args = append(args, t)
}
switch f.Enabled {
case "enabled":
where = append(where, "r.enabled = 1")
case "disabled":
where = append(where, "r.enabled = 0")
}
if len(where) == 0 {
return "", nil
}
return " WHERE " + strings.Join(where, " AND "), args
}
// Records lists records matching a filter, together with the total match count
// (ignoring limit/offset) so the UI can paginate.
func (db *DB) Records(ctx context.Context, f RecordFilter) ([]models.Record, int, error) {
whereSQL, args := f.whereClause()
var total int
if err := db.QueryRowContext(ctx, `SELECT COUNT(*) FROM records r`+whereSQL, args...).Scan(&total); err != nil {
return nil, 0, fmt.Errorf("count records: %w", err)
}
q := `SELECT ` + prefixCols(recordColumns, "r") + `, z.name
FROM records r JOIN zones z ON z.id = r.zone_id` + whereSQL +
` ORDER BY z.name, CASE r.name WHEN '@' THEN 0 ELSE 1 END, r.name, r.type`
qargs := args
if f.Limit > 0 {
q += " LIMIT ? OFFSET ?"
qargs = append(append([]any{}, args...), f.Limit, f.Offset)
}
rows, err := db.QueryContext(ctx, q, qargs...)
if err != nil {
return nil, 0, fmt.Errorf("list records: %w", err)
}
defer rows.Close()
var out []models.Record
for rows.Next() {
var r models.Record
var ttl sql.NullInt64
var enabled int
var created, updated int64
err := rows.Scan(&r.ID, &r.ZoneID, &r.Name, &r.Type, &r.Data, &ttl, &enabled,
&r.Comment, &created, &updated, &r.ZoneName)
if err != nil {
return nil, 0, err
}
if ttl.Valid {
v := uint32(ttl.Int64)
r.TTL = &v
}
r.Enabled = enabled != 0
r.CreatedAt = time.Unix(created, 0)
r.UpdatedAt = time.Unix(updated, 0)
out = append(out, r)
}
return out, total, rows.Err()
}
// Record loads a single record.
func (db *DB) Record(ctx context.Context, id int64) (models.Record, error) {
row := db.QueryRowContext(ctx, `SELECT `+recordColumns+` FROM records WHERE id = ?`, id)
r, err := scanRecord(row)
if errors.Is(err, sql.ErrNoRows) {
return r, ErrNotFound
}
if err != nil {
return r, fmt.Errorf("load record: %w", err)
}
return r, nil
}
// ZoneRecordsRaw returns every record of a zone in insertion order, used by the
// zone-file exporter.
func (db *DB) ZoneRecordsRaw(ctx context.Context, zoneID int64) ([]models.Record, error) {
rows, err := db.QueryContext(ctx,
`SELECT `+recordColumns+` FROM records WHERE zone_id = ?
ORDER BY CASE type WHEN 'SOA' THEN 0 WHEN 'NS' THEN 1 ELSE 2 END,
CASE name WHEN '@' THEN 0 ELSE 1 END, name, type`, zoneID)
if err != nil {
return nil, fmt.Errorf("list zone records: %w", err)
}
defer rows.Close()
var out []models.Record
for rows.Next() {
r, err := scanRecord(rows)
if err != nil {
return nil, err
}
out = append(out, r)
}
return out, rows.Err()
}
// CreateRecord inserts a record and bumps the zone serial.
func (db *DB) CreateRecord(ctx context.Context, r models.Record) (models.Record, error) {
var id int64
err := db.InTx(ctx, func(tx *sql.Tx) error {
res, err := tx.ExecContext(ctx, `
INSERT INTO records (zone_id, name, type, data, ttl, enabled, comment)
VALUES (?, ?, ?, ?, ?, ?, ?)`,
r.ZoneID, r.Name, r.Type, r.Data, ttlArg(r.TTL), boolInt(r.Enabled), r.Comment)
if err != nil {
return fmt.Errorf("create record: %w", err)
}
id, _ = res.LastInsertId()
return bumpSerialTx(ctx, tx, r.ZoneID)
})
if err != nil {
return models.Record{}, err
}
return db.Record(ctx, id)
}
// UpdateRecord saves a record and bumps the zone serial.
func (db *DB) UpdateRecord(ctx context.Context, r models.Record) (models.Record, error) {
err := db.InTx(ctx, func(tx *sql.Tx) error {
res, err := tx.ExecContext(ctx, `
UPDATE records SET name = ?, type = ?, data = ?, ttl = ?, enabled = ?, comment = ?,
updated_at = unixepoch()
WHERE id = ?`,
r.Name, r.Type, r.Data, ttlArg(r.TTL), boolInt(r.Enabled), r.Comment, r.ID)
if err != nil {
return fmt.Errorf("update record: %w", err)
}
if n, _ := res.RowsAffected(); n == 0 {
return ErrNotFound
}
return bumpSerialTx(ctx, tx, r.ZoneID)
})
if err != nil {
return models.Record{}, err
}
return db.Record(ctx, r.ID)
}
// SetRecordEnabled toggles a single record.
func (db *DB) SetRecordEnabled(ctx context.Context, id int64, enabled bool) error {
return db.InTx(ctx, func(tx *sql.Tx) error {
var zoneID int64
err := tx.QueryRowContext(ctx, `SELECT zone_id FROM records WHERE id = ?`, id).Scan(&zoneID)
if errors.Is(err, sql.ErrNoRows) {
return ErrNotFound
}
if err != nil {
return err
}
if _, err := tx.ExecContext(ctx,
`UPDATE records SET enabled = ?, updated_at = unixepoch() WHERE id = ?`,
boolInt(enabled), id); err != nil {
return fmt.Errorf("update record: %w", err)
}
return bumpSerialTx(ctx, tx, zoneID)
})
}
// DeleteRecord removes a record and bumps the zone serial.
func (db *DB) DeleteRecord(ctx context.Context, id int64) error {
return db.InTx(ctx, func(tx *sql.Tx) error {
var zoneID int64
err := tx.QueryRowContext(ctx, `SELECT zone_id FROM records WHERE id = ?`, id).Scan(&zoneID)
if errors.Is(err, sql.ErrNoRows) {
return ErrNotFound
}
if err != nil {
return err
}
if _, err := tx.ExecContext(ctx, `DELETE FROM records WHERE id = ?`, id); err != nil {
return fmt.Errorf("delete record: %w", err)
}
return bumpSerialTx(ctx, tx, zoneID)
})
}
// DeleteRecords removes several records belonging to one zone in a single
// transaction. It returns the number of rows deleted.
func (db *DB) DeleteRecords(ctx context.Context, zoneID int64, ids []int64) (int, error) {
if len(ids) == 0 {
return 0, nil
}
var deleted int
err := db.InTx(ctx, func(tx *sql.Tx) error {
stmt, err := tx.PrepareContext(ctx, `DELETE FROM records WHERE id = ? AND zone_id = ?`)
if err != nil {
return err
}
defer stmt.Close()
for _, id := range ids {
res, err := stmt.ExecContext(ctx, id, zoneID)
if err != nil {
return fmt.Errorf("delete record %d: %w", id, err)
}
n, _ := res.RowsAffected()
deleted += int(n)
}
return bumpSerialTx(ctx, tx, zoneID)
})
return deleted, err
}
// SetRecordsEnabled toggles several records of one zone at once.
func (db *DB) SetRecordsEnabled(ctx context.Context, zoneID int64, ids []int64, enabled bool) (int, error) {
if len(ids) == 0 {
return 0, nil
}
var updated int
err := db.InTx(ctx, func(tx *sql.Tx) error {
stmt, err := tx.PrepareContext(ctx,
`UPDATE records SET enabled = ?, updated_at = unixepoch() WHERE id = ? AND zone_id = ?`)
if err != nil {
return err
}
defer stmt.Close()
for _, id := range ids {
res, err := stmt.ExecContext(ctx, boolInt(enabled), id, zoneID)
if err != nil {
return fmt.Errorf("update record %d: %w", id, err)
}
n, _ := res.RowsAffected()
updated += int(n)
}
return bumpSerialTx(ctx, tx, zoneID)
})
return updated, err
}
// ReplaceZoneRecords swaps a zone's entire record set in one transaction. It
// backs the zone-file import "replace" mode.
func (db *DB) ReplaceZoneRecords(ctx context.Context, zoneID int64, recs []models.Record) error {
return db.InTx(ctx, func(tx *sql.Tx) error {
if _, err := tx.ExecContext(ctx, `DELETE FROM records WHERE zone_id = ?`, zoneID); err != nil {
return fmt.Errorf("clear zone records: %w", err)
}
return insertRecordsTx(ctx, tx, zoneID, recs)
})
}
// AppendZoneRecords adds records to a zone in one transaction.
func (db *DB) AppendZoneRecords(ctx context.Context, zoneID int64, recs []models.Record) error {
return db.InTx(ctx, func(tx *sql.Tx) error {
return insertRecordsTx(ctx, tx, zoneID, recs)
})
}
func insertRecordsTx(ctx context.Context, tx *sql.Tx, zoneID int64, recs []models.Record) error {
stmt, err := tx.PrepareContext(ctx, `
INSERT INTO records (zone_id, name, type, data, ttl, enabled, comment)
VALUES (?, ?, ?, ?, ?, ?, ?)`)
if err != nil {
return err
}
defer stmt.Close()
for _, r := range recs {
if _, err := stmt.ExecContext(ctx, zoneID, r.Name, r.Type, r.Data,
ttlArg(r.TTL), boolInt(r.Enabled), r.Comment); err != nil {
return fmt.Errorf("insert record %s %s: %w", r.Name, r.Type, err)
}
}
return bumpSerialTx(ctx, tx, zoneID)
}
func ttlArg(ttl *uint32) any {
if ttl == nil {
return nil
}
return int64(*ttl)
}
// ZoneSnapshotRow is one row of the bulk snapshot query that feeds the
// in-memory authoritative index.
type ZoneSnapshotRow struct {
Zone models.Zone
Record *models.Record // nil for a zone with no records
}
// SnapshotZones loads every enabled zone with its enabled records in a single
// query. This is the only place the DNS data path touches SQLite, and it runs
// on configuration change rather than per query.
func (db *DB) SnapshotZones(ctx context.Context) ([]models.Zone, map[int64][]models.Record, error) {
zones, err := db.Zones(ctx, ZoneFilter{})
if err != nil {
return nil, nil, err
}
rows, err := db.QueryContext(ctx, `
SELECT `+recordColumns+`
FROM records
WHERE enabled = 1 AND zone_id IN (SELECT id FROM zones WHERE enabled = 1)`)
if err != nil {
return nil, nil, fmt.Errorf("snapshot records: %w", err)
}
defer rows.Close()
byZone := map[int64][]models.Record{}
for rows.Next() {
r, err := scanRecord(rows)
if err != nil {
return nil, nil, err
}
byZone[r.ZoneID] = append(byZone[r.ZoneID], r)
}
if err := rows.Err(); err != nil {
return nil, nil, err
}
return zones, byZone, nil
}
// CountZonesAndRecords returns totals for the dashboard.
func (db *DB) CountZonesAndRecords(ctx context.Context) (zones, records int, err error) {
err = db.QueryRowContext(ctx,
`SELECT (SELECT COUNT(*) FROM zones), (SELECT COUNT(*) FROM records)`).Scan(&zones, &records)
if err != nil {
return 0, 0, fmt.Errorf("count zones and records: %w", err)
}
return zones, records, nil
}
// RecordTypesInUse lists the distinct record types present, for filter menus.
func (db *DB) RecordTypesInUse(ctx context.Context, zoneID int64) ([]string, error) {
q := `SELECT DISTINCT type FROM records`
var args []any
if zoneID > 0 {
q += ` WHERE zone_id = ?`
args = append(args, zoneID)
}
q += ` ORDER BY type`
rows, err := db.QueryContext(ctx, q, args...)
if err != nil {
return nil, err
}
defer rows.Close()
var out []string
for rows.Next() {
var t string
if err := rows.Scan(&t); err != nil {
return nil, err
}
out = append(out, t)
}
return out, rows.Err()
}