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

443 lines
13 KiB
Go

package app_test
import (
"context"
"io"
"log/slog"
"net/netip"
"path/filepath"
"strings"
"testing"
"github.com/miekg/dns"
"github.com/owen/vibedns/internal/app"
"github.com/owen/vibedns/internal/auditlog"
"github.com/owen/vibedns/internal/config"
"github.com/owen/vibedns/internal/database"
"github.com/owen/vibedns/internal/models"
)
// newTestApp builds a fully wired application against a temporary database,
// without binding any listener.
func newTestApp(t *testing.T) *app.App {
t.Helper()
ctx := context.Background()
path := filepath.Join(t.TempDir(), "test.db")
db, err := database.Open(path)
if err != nil {
t.Fatalf("open database: %v", err)
}
t.Cleanup(func() { db.Close() })
if _, err := db.Migrate(ctx); err != nil {
t.Fatalf("migrate: %v", err)
}
log := slog.New(slog.NewTextHandler(io.Discard, nil))
boot := config.DefaultBootstrap()
boot.DBPath = path
a, err := app.New(ctx, boot, db, log)
if err != nil {
t.Fatalf("build app: %v", err)
}
return a
}
func testActor() auditlog.Actor {
return auditlog.Actor{Name: "test", Source: auditlog.SourceCLI, ClientIP: "127.0.0.1"}
}
// TestZoneLifecycleAndResolution walks the path an operator actually takes:
// create a zone, add records, and confirm the DNS engine serves them.
func TestZoneLifecycleAndResolution(t *testing.T) {
a := newTestApp(t)
ctx := context.Background()
zone, err := a.CreateZone(ctx, testActor(), app.ZoneInput{
Name: "example.com",
Description: "integration test zone",
AdminEmail: "hostmaster@example.com",
})
if err != nil {
t.Fatalf("create zone: %v", err)
}
if zone.Name != "example.com." {
t.Errorf("zone name = %q, want the normalised form", zone.Name)
}
records := []app.RecordInput{
{Name: "@", Type: "A", Data: "192.0.2.10"},
{Name: "www", Type: "CNAME", Data: "example.com."},
{Name: "mail", Type: "A", Data: "192.0.2.20"},
{Name: "@", Type: "MX", Data: "10 mail.example.com."},
{Name: "txt", Type: "TXT", Fields: map[string]string{"text": "hello world"}},
}
for _, in := range records {
if _, err := a.CreateRecord(ctx, testActor(), zone.ID, in); err != nil {
t.Fatalf("create record %s %s: %v", in.Name, in.Type, err)
}
}
// The snapshot is rebuilt on demand rather than waiting for the debounce.
if err := a.Reload(ctx); err != nil {
t.Fatalf("reload: %v", err)
}
snap := a.Snapshot()
if snap.ZoneCount != 1 {
t.Fatalf("indexed zones = %d, want 1", snap.ZoneCount)
}
if len(snap.Problems) != 0 {
t.Errorf("build problems: %v", snap.Problems)
}
client := netip.MustParseAddr("127.0.0.1")
tests := []struct {
name string
qname string
qtype uint16
rcode int
wantIn string
answers int
}{
{"apex A", "example.com.", dns.TypeA, dns.RcodeSuccess, "192.0.2.10", 1},
{"host A", "mail.example.com.", dns.TypeA, dns.RcodeSuccess, "192.0.2.20", 1},
{"CNAME is followed", "www.example.com.", dns.TypeA, dns.RcodeSuccess, "192.0.2.10", 2},
{"MX", "example.com.", dns.TypeMX, dns.RcodeSuccess, "mail.example.com.", 1},
{"TXT", "txt.example.com.", dns.TypeTXT, dns.RcodeSuccess, "hello world", 1},
{"NODATA", "mail.example.com.", dns.TypeTXT, dns.RcodeSuccess, "", 0},
{"NXDOMAIN", "missing.example.com.", dns.TypeA, dns.RcodeNameError, "", 0},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
msg, source, err := a.DNS.Resolve(ctx, tc.qname, tc.qtype, client, false)
if err != nil {
t.Fatalf("resolve: %v", err)
}
if msg.Rcode != tc.rcode {
t.Errorf("rcode = %s, want %s", dns.RcodeToString[msg.Rcode], dns.RcodeToString[tc.rcode])
}
if len(msg.Answer) != tc.answers {
t.Errorf("answers = %d, want %d: %v", len(msg.Answer), tc.answers, msg.Answer)
}
if source != models.SourceAuthoritative {
t.Errorf("source = %q, want authoritative", source)
}
if tc.wantIn != "" {
var found bool
for _, rr := range msg.Answer {
if strings.Contains(rr.String(), tc.wantIn) {
found = true
}
}
if !found {
t.Errorf("no answer contains %q: %v", tc.wantIn, msg.Answer)
}
}
})
}
}
func TestCNAMEConflictIsRejected(t *testing.T) {
a := newTestApp(t)
ctx := context.Background()
zone, err := a.CreateZone(ctx, testActor(), app.ZoneInput{Name: "example.com"})
if err != nil {
t.Fatalf("create zone: %v", err)
}
if _, err := a.CreateRecord(ctx, testActor(), zone.ID,
app.RecordInput{Name: "www", Type: "A", Data: "192.0.2.1"}); err != nil {
t.Fatalf("create A: %v", err)
}
// A CNAME cannot coexist with the A record already at that name.
_, err = a.CreateRecord(ctx, testActor(), zone.ID,
app.RecordInput{Name: "www", Type: "CNAME", Data: "other.example.com."})
if err == nil {
t.Fatal("expected the conflicting CNAME to be rejected")
}
if app.StatusOf(err) != 400 {
t.Errorf("status = %d, want 400", app.StatusOf(err))
}
// And a CNAME at the apex is always wrong.
_, err = a.CreateRecord(ctx, testActor(), zone.ID,
app.RecordInput{Name: "@", Type: "CNAME", Data: "other.example.com."})
if err == nil {
t.Error("expected an apex CNAME to be rejected")
}
}
func TestReverseZoneCreationFromCIDR(t *testing.T) {
a := newTestApp(t)
ctx := context.Background()
zone, err := a.CreateZone(ctx, testActor(), app.ZoneInput{
CIDR: "192.168.1.0/24",
Kind: "reverse4",
})
if err != nil {
t.Fatalf("create reverse zone: %v", err)
}
if zone.Name != "1.168.192.in-addr.arpa." {
t.Errorf("zone name = %q, want 1.168.192.in-addr.arpa.", zone.Name)
}
if _, err := a.CreateRecord(ctx, testActor(), zone.ID,
app.RecordInput{Name: "10", Type: "PTR", Data: "host.example.com."}); err != nil {
t.Fatalf("create PTR: %v", err)
}
if err := a.Reload(ctx); err != nil {
t.Fatalf("reload: %v", err)
}
msg, _, err := a.DNS.Resolve(ctx, "10.1.168.192.in-addr.arpa.", dns.TypePTR,
netip.MustParseAddr("127.0.0.1"), false)
if err != nil {
t.Fatalf("resolve PTR: %v", err)
}
if len(msg.Answer) != 1 {
t.Fatalf("PTR answers = %d, want 1", len(msg.Answer))
}
if ptr, ok := msg.Answer[0].(*dns.PTR); !ok || ptr.Ptr != "host.example.com." {
t.Errorf("PTR answer = %v", msg.Answer[0])
}
}
// TestPolicyBlocking exercises the filtering path end to end: import a
// blocklist, attach it to a policy and a network, and confirm the DNS engine
// blocks a matching query from a client in that network.
func TestPolicyBlocking(t *testing.T) {
a := newTestApp(t)
ctx := context.Background()
list, err := a.CreateDomainList(ctx, testActor(), app.ListInput{
Kind: models.KindBlacklist, Name: "Test Blocks",
})
if err != nil {
t.Fatalf("create list: %v", err)
}
summary, err := a.ImportDomains(ctx, testActor(), list.ID,
strings.NewReader("0.0.0.0 ads.example\ntracker.example.net\n# a comment\n"), true)
if err != nil {
t.Fatalf("import: %v", err)
}
if summary.Imported != 2 {
t.Fatalf("imported = %d, want 2", summary.Imported)
}
policy, err := a.CreatePolicy(ctx, testActor(), app.PolicyInput{
Name: "Test Policy", BlockAction: "nxdomain", ListIDs: []int64{list.ID},
})
if err != nil {
t.Fatalf("create policy: %v", err)
}
if _, err := a.CreateNetwork(ctx, testActor(), app.NetworkInput{
Name: "Test Net", CIDR: "100.64.30.0/24", PolicyIDs: []int64{policy.ID},
}); err != nil {
t.Fatalf("create network: %v", err)
}
if err := a.Reload(ctx); err != nil {
t.Fatalf("reload: %v", err)
}
inNetwork := netip.MustParseAddr("100.64.30.5")
outside := netip.MustParseAddr("192.0.2.1")
// A blocked name from inside the network.
msg, source, err := a.DNS.Resolve(ctx, "ads.example.", dns.TypeA, inNetwork, false)
if err != nil {
t.Fatalf("resolve: %v", err)
}
if msg.Rcode != dns.RcodeNameError {
t.Errorf("rcode = %s, want NXDOMAIN", dns.RcodeToString[msg.Rcode])
}
if source != models.SourceBlocked {
t.Errorf("source = %q, want blocked", source)
}
// A subdomain of a blocked name is covered without being stored.
msg, _, _ = a.DNS.Resolve(ctx, "cdn.ads.example.", dns.TypeA, inNetwork, false)
if msg.Rcode != dns.RcodeNameError {
t.Errorf("subdomain rcode = %s, want NXDOMAIN", dns.RcodeToString[msg.Rcode])
}
// The same name from a client outside the network is not blocked. With no
// upstream reachable in a test it will fail to resolve, but it must not be
// blocked, and it must not be refused for a private client.
_, source, _ = a.DNS.Resolve(ctx, "ads.example.", dns.TypeA, outside, false)
if source == models.SourceBlocked {
t.Error("a client outside the configured network was filtered")
}
}
// TestRecursionIsRefusedByDefaultForPublicClients is the open-resolver guard.
func TestRecursionIsRefusedForClientsOutsideTheACL(t *testing.T) {
a := newTestApp(t)
ctx := context.Background()
// A public address is not in the default private-network ACL.
public := netip.MustParseAddr("203.0.113.50")
msg, source, err := a.DNS.Resolve(ctx, "example.org.", dns.TypeA, public, false)
if err != nil {
t.Fatalf("resolve: %v", err)
}
if msg.Rcode != dns.RcodeRefused {
t.Errorf("rcode = %s, want REFUSED for a client outside the recursion ACL",
dns.RcodeToString[msg.Rcode])
}
if source != models.SourceRefused {
t.Errorf("source = %q, want refused", source)
}
}
// TestAuthoritativeAnswersSurviveRecursionDenial: a client that may not
// recurse must still get answers for zones we are authoritative for.
func TestAuthoritativeAnswersWorkWithoutRecursionRights(t *testing.T) {
a := newTestApp(t)
ctx := context.Background()
zone, err := a.CreateZone(ctx, testActor(), app.ZoneInput{Name: "internal.example"})
if err != nil {
t.Fatalf("create zone: %v", err)
}
if _, err := a.CreateRecord(ctx, testActor(), zone.ID,
app.RecordInput{Name: "@", Type: "A", Data: "192.0.2.1"}); err != nil {
t.Fatalf("create record: %v", err)
}
if err := a.Reload(ctx); err != nil {
t.Fatalf("reload: %v", err)
}
public := netip.MustParseAddr("203.0.113.50")
msg, source, err := a.DNS.Resolve(ctx, "internal.example.", dns.TypeA, public, false)
if err != nil {
t.Fatalf("resolve: %v", err)
}
if msg.Rcode != dns.RcodeSuccess {
t.Errorf("rcode = %s, want NOERROR: an authoritative zone must answer "+
"even when the client may not recurse", dns.RcodeToString[msg.Rcode])
}
if source != models.SourceAuthoritative {
t.Errorf("source = %q, want authoritative", source)
}
}
func TestSettingsValidationRejectsOpenResolver(t *testing.T) {
a := newTestApp(t)
ctx := context.Background()
next := a.Settings()
next.Resolver.AllowNetworks = nil // would deny everyone
if err := a.SaveSettings(ctx, testActor(), app.GroupResolver, next); err == nil {
t.Error("an empty recursion ACL with recursion enabled should be rejected")
}
next = a.Settings()
next.Resolver.Upstreams = nil
if err := a.SaveSettings(ctx, testActor(), app.GroupResolver, next); err == nil {
t.Error("enabling recursion with no upstreams should be rejected")
}
}
func TestConfigExportOmitsSecrets(t *testing.T) {
a := newTestApp(t)
ctx := context.Background()
if _, err := a.CreateAPIToken(ctx, testActor(), "test-token", ""); err != nil {
t.Fatalf("create token: %v", err)
}
var buf strings.Builder
if err := a.WriteConfigExport(ctx, &buf, false); err != nil {
t.Fatalf("export: %v", err)
}
out := buf.String()
for _, forbidden := range []string{"argon2id", "password_hash", "token_hash", "csrf_key", "vibedns_"} {
if strings.Contains(out, forbidden) {
t.Errorf("the configuration export contains %q, which must never leave the server", forbidden)
}
}
}
func TestAuditLogRecordsChanges(t *testing.T) {
a := newTestApp(t)
ctx := context.Background()
if _, err := a.CreateZone(ctx, testActor(), app.ZoneInput{Name: "audited.example"}); err != nil {
t.Fatalf("create zone: %v", err)
}
entries, total, err := a.AuditLogs(ctx, database.AuditFilter{Limit: 10})
if err != nil {
t.Fatalf("read audit log: %v", err)
}
if total == 0 {
t.Fatal("no audit entry was recorded for a zone creation")
}
var found bool
for _, e := range entries {
if e.Action == "zone.create" && e.ObjectName == "audited.example." {
found = true
}
}
if !found {
t.Errorf("no zone.create entry found in %v", entries)
}
}
func TestBulkRecordOperations(t *testing.T) {
a := newTestApp(t)
ctx := context.Background()
zone, err := a.CreateZone(ctx, testActor(), app.ZoneInput{Name: "bulk.example"})
if err != nil {
t.Fatalf("create zone: %v", err)
}
var ids []int64
for _, name := range []string{"a", "b", "c"} {
rec, err := a.CreateRecord(ctx, testActor(), zone.ID,
app.RecordInput{Name: name, Type: "A", Data: "192.0.2.1"})
if err != nil {
t.Fatalf("create %s: %v", name, err)
}
ids = append(ids, rec.ID)
}
n, err := a.BulkRecords(ctx, testActor(), zone.ID, ids, app.BulkDisable)
if err != nil {
t.Fatalf("bulk disable: %v", err)
}
if n != 3 {
t.Errorf("disabled %d, want 3", n)
}
// Disabled records must not be served.
if err := a.Reload(ctx); err != nil {
t.Fatalf("reload: %v", err)
}
msg, _, _ := a.DNS.Resolve(ctx, "a.bulk.example.", dns.TypeA, netip.MustParseAddr("127.0.0.1"), false)
if len(msg.Answer) != 0 {
t.Errorf("a disabled record was still served: %v", msg.Answer)
}
n, err = a.BulkRecords(ctx, testActor(), zone.ID, ids, app.BulkDelete)
if err != nil {
t.Fatalf("bulk delete: %v", err)
}
if n != 3 {
t.Errorf("deleted %d, want 3", n)
}
}