Files
vibedns/internal/database/database_test.go
T
2026-08-16 21:18:45 -05:00

444 lines
12 KiB
Go

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:])
}