initial commit
This commit is contained in:
@@ -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
|
||||
}
|
||||
@@ -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:])
|
||||
}
|
||||
@@ -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';
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
Reference in New Issue
Block a user