Files
2026-08-16 21:18:45 -05:00

655 lines
20 KiB
Go

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