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

648 lines
20 KiB
Go

package app
import (
"context"
"errors"
"fmt"
"io"
"strings"
"github.com/owen/vibedns/internal/auditlog"
"github.com/owen/vibedns/internal/blacklist"
"github.com/owen/vibedns/internal/config"
"github.com/owen/vibedns/internal/database"
"github.com/owen/vibedns/internal/models"
"github.com/owen/vibedns/internal/validate"
)
// --- Client networks ----------------------------------------------------
// NetworkInput is the editable surface of a client network.
type NetworkInput struct {
Name string `json:"name"`
CIDR string `json:"cidr"`
Description string `json:"description"`
Enabled *bool `json:"enabled"`
PolicyIDs []int64 `json:"policy_ids"`
}
// Networks lists client networks with their policy assignments.
func (a *App) Networks(ctx context.Context, search string) ([]models.Network, error) {
nets, err := a.DB.Networks(ctx, search, true)
if err != nil {
return nil, Internal(err, "The network list could not be loaded.")
}
return nets, nil
}
// Network loads one client network.
func (a *App) Network(ctx context.Context, id int64) (models.Network, error) {
n, err := a.DB.Network(ctx, id)
if errors.Is(err, database.ErrNotFound) {
return n, NotFound("Network %d was not found.", id)
}
if err != nil {
return n, Internal(err, "The network could not be loaded.")
}
return n, nil
}
func (a *App) normaliseNetwork(in NetworkInput, base models.Network) (models.Network, error) {
n := base
if name := strings.TrimSpace(in.Name); name != "" {
n.Name = name
}
if n.Name == "" {
return n, Invalid("A network name is required.")
}
if cidr := strings.TrimSpace(in.CIDR); cidr != "" {
p, err := config.ParseCIDROrIP(cidr)
if err != nil {
return n, Invalid("Subnet %q: %s", cidr, err.Error())
}
n.CIDR = p.String()
}
if n.CIDR == "" {
return n, Invalid("A subnet in CIDR notation is required, for example 192.168.1.0/24.")
}
n.Description = strings.TrimSpace(in.Description)
if in.Enabled != nil {
n.Enabled = *in.Enabled
} else if base.ID == 0 {
n.Enabled = true
}
return n, nil
}
// CreateNetwork stores a client network.
func (a *App) CreateNetwork(ctx context.Context, actor auditlog.Actor, in NetworkInput) (models.Network, error) {
n, err := a.normaliseNetwork(in, models.Network{})
if err != nil {
return models.Network{}, err
}
created, err := a.DB.CreateNetwork(ctx, n, in.PolicyIDs)
if err != nil {
return models.Network{}, translate(err, "Network not found.",
fmt.Sprintf("A network named %q already exists.", n.Name))
}
a.Audit.RecordID(ctx, actor, "network.create", auditlog.ObjectNetwork, created.ID, created.Name,
auditlog.Changes("cidr", created.CIDR, "policies", fmt.Sprint(len(in.PolicyIDs))))
a.Runtime.RequestReload()
return created, nil
}
// UpdateNetwork saves a client network and its policy assignments.
func (a *App) UpdateNetwork(ctx context.Context, actor auditlog.Actor, id int64, in NetworkInput) (models.Network, error) {
existing, err := a.Network(ctx, id)
if err != nil {
return models.Network{}, err
}
n, err := a.normaliseNetwork(in, existing)
if err != nil {
return models.Network{}, err
}
n.ID = id
updated, err := a.DB.UpdateNetwork(ctx, n, in.PolicyIDs)
if err != nil {
return models.Network{}, translate(err, fmt.Sprintf("Network %d was not found.", id),
fmt.Sprintf("A network named %q already exists.", n.Name))
}
a.Audit.RecordID(ctx, actor, "network.update", auditlog.ObjectNetwork, id, updated.Name,
auditlog.Changes("cidr", updated.CIDR, "policies", fmt.Sprint(len(in.PolicyIDs))))
a.Runtime.RequestReload()
return updated, nil
}
// SetNetworkEnabled toggles a client network.
func (a *App) SetNetworkEnabled(ctx context.Context, actor auditlog.Actor, id int64, enabled bool) error {
n, err := a.Network(ctx, id)
if err != nil {
return err
}
if err := a.DB.SetNetworkEnabled(ctx, id, enabled); err != nil {
return translate(err, fmt.Sprintf("Network %d was not found.", id), "")
}
action := "network.disable"
if enabled {
action = "network.enable"
}
a.Audit.RecordID(ctx, actor, action, auditlog.ObjectNetwork, id, n.Name, "")
a.Runtime.RequestReload()
return nil
}
// DeleteNetwork removes a client network.
func (a *App) DeleteNetwork(ctx context.Context, actor auditlog.Actor, id int64) error {
n, err := a.Network(ctx, id)
if err != nil {
return err
}
if err := a.DB.DeleteNetwork(ctx, id); err != nil {
return translate(err, fmt.Sprintf("Network %d was not found.", id), "")
}
a.Audit.RecordID(ctx, actor, "network.delete", auditlog.ObjectNetwork, id, n.Name,
auditlog.Changes("cidr", n.CIDR))
a.Runtime.RequestReload()
return nil
}
// --- Policies -----------------------------------------------------------
// PolicyInput is the editable surface of a policy.
type PolicyInput struct {
Name string `json:"name"`
Description string `json:"description"`
Enabled *bool `json:"enabled"`
BlockAction string `json:"block_action"`
SinkholeIPv4 string `json:"sinkhole_ipv4"`
SinkholeIPv6 string `json:"sinkhole_ipv6"`
BlockTTL uint32 `json:"block_ttl"`
ListIDs []int64 `json:"list_ids"`
}
// Policies lists every policy.
func (a *App) Policies(ctx context.Context) ([]models.Policy, error) {
p, err := a.DB.Policies(ctx)
if err != nil {
return nil, Internal(err, "The policy list could not be loaded.")
}
return p, nil
}
// Policy loads one policy.
func (a *App) Policy(ctx context.Context, id int64) (models.Policy, error) {
p, err := a.DB.Policy(ctx, id)
if errors.Is(err, database.ErrNotFound) {
return p, NotFound("Policy %d was not found.", id)
}
if err != nil {
return p, Internal(err, "The policy could not be loaded.")
}
return p, nil
}
func (a *App) normalisePolicy(in PolicyInput, base models.Policy) (models.Policy, error) {
p := base
if name := strings.TrimSpace(in.Name); name != "" {
p.Name = name
}
if p.Name == "" {
return p, Invalid("A policy name is required.")
}
p.Description = strings.TrimSpace(in.Description)
action := models.BlockAction(strings.ToLower(strings.TrimSpace(in.BlockAction)))
if action == "" {
action = base.BlockAction
}
if action == "" {
action = models.BlockNXDOMAIN
}
if !action.Valid() {
return p, Invalid("Block action %q must be nxdomain, refused or sinkhole.", in.BlockAction)
}
p.BlockAction = action
p.SinkholeIPv4 = strings.TrimSpace(in.SinkholeIPv4)
if p.SinkholeIPv4 == "" {
p.SinkholeIPv4 = "0.0.0.0"
}
p.SinkholeIPv6 = strings.TrimSpace(in.SinkholeIPv6)
if p.SinkholeIPv6 == "" {
p.SinkholeIPv6 = "::"
}
if action == models.BlockSinkhole {
if err := requireIP(p.SinkholeIPv4, true); err != nil {
return p, Invalid("Sinkhole IPv4 address: %s", err.Error())
}
if err := requireIP(p.SinkholeIPv6, false); err != nil {
return p, Invalid("Sinkhole IPv6 address: %s", err.Error())
}
}
p.BlockTTL = in.BlockTTL
if p.BlockTTL == 0 {
p.BlockTTL = base.BlockTTL
}
if p.BlockTTL == 0 {
p.BlockTTL = 60
}
if p.BlockTTL > 86400 {
return p, Invalid("The block TTL must be 86400 seconds or less.")
}
if in.Enabled != nil {
p.Enabled = *in.Enabled
} else if base.ID == 0 {
p.Enabled = true
}
return p, nil
}
func requireIP(s string, wantV4 bool) error {
p, err := config.ParseCIDROrIP(s)
if err != nil {
return errors.New("must be a valid IP address")
}
if p.Addr().Is4() != wantV4 {
if wantV4 {
return errors.New("must be an IPv4 address")
}
return errors.New("must be an IPv6 address")
}
return nil
}
// CreatePolicy stores a policy.
func (a *App) CreatePolicy(ctx context.Context, actor auditlog.Actor, in PolicyInput) (models.Policy, error) {
p, err := a.normalisePolicy(in, models.Policy{})
if err != nil {
return models.Policy{}, err
}
created, err := a.DB.CreatePolicy(ctx, p, in.ListIDs)
if err != nil {
return models.Policy{}, translate(err, "Policy not found.",
fmt.Sprintf("A policy named %q already exists.", p.Name))
}
a.Audit.RecordID(ctx, actor, "policy.create", auditlog.ObjectPolicy, created.ID, created.Name,
auditlog.Changes("action", string(created.BlockAction), "lists", fmt.Sprint(len(in.ListIDs))))
a.Runtime.RequestReload()
return created, nil
}
// UpdatePolicy saves a policy.
func (a *App) UpdatePolicy(ctx context.Context, actor auditlog.Actor, id int64, in PolicyInput) (models.Policy, error) {
existing, err := a.Policy(ctx, id)
if err != nil {
return models.Policy{}, err
}
p, err := a.normalisePolicy(in, existing)
if err != nil {
return models.Policy{}, err
}
p.ID = id
updated, err := a.DB.UpdatePolicy(ctx, p, in.ListIDs)
if err != nil {
return models.Policy{}, translate(err, fmt.Sprintf("Policy %d was not found.", id),
fmt.Sprintf("A policy named %q already exists.", p.Name))
}
a.Audit.RecordID(ctx, actor, "policy.update", auditlog.ObjectPolicy, id, updated.Name,
auditlog.Changes("action", string(updated.BlockAction), "lists", fmt.Sprint(len(in.ListIDs))))
a.Runtime.RequestReload()
return updated, nil
}
// SetPolicyEnabled toggles a policy.
func (a *App) SetPolicyEnabled(ctx context.Context, actor auditlog.Actor, id int64, enabled bool) error {
p, err := a.Policy(ctx, id)
if err != nil {
return err
}
if err := a.DB.SetPolicyEnabled(ctx, id, enabled); err != nil {
return translate(err, fmt.Sprintf("Policy %d was not found.", id), "")
}
action := "policy.disable"
if enabled {
action = "policy.enable"
}
a.Audit.RecordID(ctx, actor, action, auditlog.ObjectPolicy, id, p.Name, "")
a.Runtime.RequestReload()
return nil
}
// DeletePolicy removes a policy.
func (a *App) DeletePolicy(ctx context.Context, actor auditlog.Actor, id int64) error {
p, err := a.Policy(ctx, id)
if err != nil {
return err
}
if err := a.DB.DeletePolicy(ctx, id); err != nil {
return translate(err, fmt.Sprintf("Policy %d was not found.", id), "")
}
a.Audit.RecordID(ctx, actor, "policy.delete", auditlog.ObjectPolicy, id, p.Name, "")
a.Runtime.RequestReload()
return nil
}
// --- Domain lists -------------------------------------------------------
// ListInput is the editable surface of a blacklist or allowlist.
type ListInput struct {
Kind string `json:"kind"`
Name string `json:"name"`
Description string `json:"description"`
Enabled *bool `json:"enabled"`
SourceURL string `json:"source_url"`
}
// DomainLists returns blacklists, allowlists, or both when kind is empty.
func (a *App) DomainLists(ctx context.Context, kind, search string) ([]models.DomainList, error) {
lists, err := a.DB.DomainLists(ctx, kind, search)
if err != nil {
return nil, Internal(err, "The list could not be loaded.")
}
return lists, nil
}
// DomainList loads one list.
func (a *App) DomainList(ctx context.Context, id int64) (models.DomainList, error) {
l, err := a.DB.DomainList(ctx, id)
if errors.Is(err, database.ErrNotFound) {
return l, NotFound("List %d was not found.", id)
}
if err != nil {
return l, Internal(err, "The list could not be loaded.")
}
return l, nil
}
// CreateDomainList stores a blacklist or allowlist.
func (a *App) CreateDomainList(ctx context.Context, actor auditlog.Actor, in ListInput) (models.DomainList, error) {
kind := strings.ToLower(strings.TrimSpace(in.Kind))
if kind != models.KindBlacklist && kind != models.KindAllowlist {
return models.DomainList{}, Invalid("List kind must be blacklist or allowlist.")
}
name := strings.TrimSpace(in.Name)
if name == "" {
return models.DomainList{}, Invalid("A list name is required.")
}
enabled := true
if in.Enabled != nil {
enabled = *in.Enabled
}
l := models.DomainList{
Kind: kind,
Name: name,
Description: strings.TrimSpace(in.Description),
Enabled: enabled,
SourceURL: strings.TrimSpace(in.SourceURL),
}
created, err := a.DB.CreateDomainList(ctx, l)
if err != nil {
return models.DomainList{}, translate(err, "List not found.",
fmt.Sprintf("A %s named %q already exists.", kind, name))
}
a.Audit.RecordID(ctx, actor, "list.create", auditlog.ObjectList, created.ID, created.Name,
auditlog.Changes("kind", created.Kind))
a.Runtime.RequestReload()
return created, nil
}
// UpdateDomainList saves list metadata.
func (a *App) UpdateDomainList(ctx context.Context, actor auditlog.Actor, id int64, in ListInput) (models.DomainList, error) {
existing, err := a.DomainList(ctx, id)
if err != nil {
return models.DomainList{}, err
}
if name := strings.TrimSpace(in.Name); name != "" {
existing.Name = name
}
existing.Description = strings.TrimSpace(in.Description)
existing.SourceURL = strings.TrimSpace(in.SourceURL)
if in.Enabled != nil {
existing.Enabled = *in.Enabled
}
updated, err := a.DB.UpdateDomainList(ctx, existing)
if err != nil {
return models.DomainList{}, translate(err, fmt.Sprintf("List %d was not found.", id),
fmt.Sprintf("A list named %q already exists.", existing.Name))
}
a.Audit.RecordID(ctx, actor, "list.update", auditlog.ObjectList, id, updated.Name, "")
a.Runtime.RequestReload()
return updated, nil
}
// SetDomainListEnabled toggles a list.
func (a *App) SetDomainListEnabled(ctx context.Context, actor auditlog.Actor, id int64, enabled bool) error {
l, err := a.DomainList(ctx, id)
if err != nil {
return err
}
if err := a.DB.SetDomainListEnabled(ctx, id, enabled); err != nil {
return translate(err, fmt.Sprintf("List %d was not found.", id), "")
}
action := "list.disable"
if enabled {
action = "list.enable"
}
a.Audit.RecordID(ctx, actor, action, auditlog.ObjectList, id, l.Name, "")
a.Runtime.RequestReload()
return nil
}
// DeleteDomainList removes a list and every domain in it.
func (a *App) DeleteDomainList(ctx context.Context, actor auditlog.Actor, id int64) error {
l, err := a.DomainList(ctx, id)
if err != nil {
return err
}
if err := a.DB.DeleteDomainList(ctx, id); err != nil {
return translate(err, fmt.Sprintf("List %d was not found.", id), "")
}
a.Audit.RecordID(ctx, actor, "list.delete", auditlog.ObjectList, id, l.Name,
auditlog.Changes("domains", fmt.Sprint(l.DomainCount)))
a.Runtime.RequestReload()
return nil
}
// --- Domain entries -----------------------------------------------------
// DomainEntries pages through a list's domains.
func (a *App) DomainEntries(ctx context.Context, listID int64, search string, limit, offset int) ([]models.DomainEntry, int, error) {
entries, total, err := a.DB.DomainEntries(ctx, listID, search, limit, offset)
if err != nil {
return nil, 0, Internal(err, "The domains could not be loaded.")
}
return entries, total, nil
}
// AddDomain adds one domain to a list.
func (a *App) AddDomain(ctx context.Context, actor auditlog.Actor, listID int64, domain string, matchSubdomains bool, comment string) (models.DomainEntry, error) {
l, err := a.DomainList(ctx, listID)
if err != nil {
return models.DomainEntry{}, err
}
d := strings.TrimSpace(domain)
if strings.HasPrefix(d, "*.") {
d = strings.TrimPrefix(d, "*.")
matchSubdomains = true
}
normalised, err := validate.NormaliseDomain(d)
if err != nil {
return models.DomainEntry{}, Invalid("%s", err.Error())
}
entry, err := a.DB.AddDomain(ctx, models.DomainEntry{
ListID: listID,
Domain: normalised,
MatchSubdomains: matchSubdomains,
Enabled: true,
Comment: strings.TrimSpace(comment),
})
if err != nil {
return models.DomainEntry{}, translate(err, "List not found.",
fmt.Sprintf("%s is already in %s.", normalised, l.Name))
}
a.Audit.RecordID(ctx, actor, "domain.add", auditlog.ObjectDomain, entry.ID, normalised,
auditlog.Changes("list", l.Name))
a.Runtime.RequestReload()
return entry, nil
}
// UpdateDomain saves an existing domain entry.
func (a *App) UpdateDomain(ctx context.Context, actor auditlog.Actor, listID, id int64, domain string, matchSubdomains, enabled bool, comment string) error {
normalised, err := validate.NormaliseDomain(strings.TrimPrefix(strings.TrimSpace(domain), "*."))
if err != nil {
return Invalid("%s", err.Error())
}
e := models.DomainEntry{
ID: id, ListID: listID, Domain: normalised,
MatchSubdomains: matchSubdomains, Enabled: enabled, Comment: strings.TrimSpace(comment),
}
if err := a.DB.UpdateDomain(ctx, e); err != nil {
return translate(err, fmt.Sprintf("Domain %d was not found.", id),
fmt.Sprintf("%s is already in this list.", normalised))
}
a.Audit.RecordID(ctx, actor, "domain.update", auditlog.ObjectDomain, id, normalised, "")
a.Runtime.RequestReload()
return nil
}
// DeleteDomain removes one domain from a list.
func (a *App) DeleteDomain(ctx context.Context, actor auditlog.Actor, id int64) error {
if err := a.DB.DeleteDomain(ctx, id); err != nil {
return translate(err, fmt.Sprintf("Domain %d was not found.", id), "")
}
a.Audit.RecordID(ctx, actor, "domain.delete", auditlog.ObjectDomain, id, "", "")
a.Runtime.RequestReload()
return nil
}
// ClearDomains empties a list.
func (a *App) ClearDomains(ctx context.Context, actor auditlog.Actor, listID int64) (int64, error) {
l, err := a.DomainList(ctx, listID)
if err != nil {
return 0, err
}
n, err := a.DB.ClearDomains(ctx, listID)
if err != nil {
return 0, Internal(err, "The list could not be cleared.")
}
a.Audit.RecordID(ctx, actor, "list.clear", auditlog.ObjectList, listID, l.Name,
auditlog.Changes("removed", fmt.Sprint(n)))
a.Runtime.RequestReload()
return n, nil
}
// ImportDomains parses a domain list and stores it.
//
// Parsing happens fully in memory and the insert runs as a single transaction
// with one prepared statement, so a list of several hundred thousand domains
// is one commit rather than one commit per domain.
func (a *App) ImportDomains(ctx context.Context, actor auditlog.Actor, listID int64,
r io.Reader, matchSubdomains bool) (models.ImportSummary, error) {
l, err := a.DomainList(ctx, listID)
if err != nil {
return models.ImportSummary{}, err
}
parsed, summary := blacklist.Parse(r, blacklist.ParseOptions{DefaultMatchSubdomains: matchSubdomains})
if len(parsed) == 0 {
if summary.LinesProcessed == 0 {
return summary, Invalid("The import was empty.")
}
return summary, Invalid("No valid domains were found in %d lines. "+
"Supported formats are a plain domain list, a hosts file, or Adblock-style ||domain^ rules.",
summary.LinesProcessed)
}
rows := make([]database.ImportDomain, 0, len(parsed))
for _, p := range parsed {
rows = append(rows, database.ImportDomain{Domain: p.Domain, MatchSubdomains: p.MatchSubdomains})
}
imported, duplicates, err := a.DB.ImportDomains(ctx, listID, rows)
if err != nil {
return summary, Internal(err, "The domains could not be imported.")
}
// The parser counts duplicates within the file; the database reports
// domains that were already present. The summary shows the total.
summary.Imported = imported
summary.Duplicates += duplicates
a.Audit.RecordID(ctx, actor, "list.import", auditlog.ObjectList, listID, l.Name,
auditlog.Changes(
"imported", fmt.Sprint(summary.Imported),
"duplicates", fmt.Sprint(summary.Duplicates),
"invalid", fmt.Sprint(summary.Invalid),
"lines", fmt.Sprint(summary.LinesProcessed)))
a.Runtime.RequestReload()
return summary, nil
}
// ExportDomains writes a list as a plain domain list.
func (a *App) ExportDomains(ctx context.Context, listID int64, w io.Writer) (models.DomainList, error) {
l, err := a.DomainList(ctx, listID)
if err != nil {
return l, err
}
fmt.Fprintf(w, "# %s\n", l.Name)
if l.Description != "" {
fmt.Fprintf(w, "# %s\n", l.Description)
}
fmt.Fprintf(w, "# %d domains exported by VibeDNS\n", l.DomainCount)
err = a.DB.ExportDomains(ctx, listID, func(domain string, matchSubdomains bool) {
if matchSubdomains {
fmt.Fprintln(w, domain)
return
}
// A domain that must match exactly is written in a form the importer
// will not silently widen.
fmt.Fprintf(w, "%s # exact\n", domain)
})
if err != nil {
return l, Internal(err, "The domains could not be exported.")
}
return l, nil
}
// LookupDomain reports which lists cover a name, for the "why was this
// blocked?" tool.
type LookupHit struct {
ListID int64 `json:"list_id"`
ListName string `json:"list_name"`
Kind string `json:"kind"`
Matched string `json:"matched_domain"`
}
// LookupDomain checks a name against every compiled list.
func (a *App) LookupDomain(ctx context.Context, name string) ([]LookupHit, error) {
domain, err := validate.NormaliseDomain(name)
if err != nil {
return nil, Invalid("%s", err.Error())
}
lists, err := a.DB.DomainLists(ctx, "", "")
if err != nil {
return nil, Internal(err, "The lists could not be loaded.")
}
sets := a.Snapshot().Policy.Sets()
var hits []LookupHit
for _, l := range lists {
set, ok := sets[l.ID]
if !ok {
continue // list is disabled, so it was not compiled
}
if matched, found := set.Match(domain); found {
hits = append(hits, LookupHit{
ListID: l.ID, ListName: l.Name, Kind: l.Kind, Matched: matched,
})
}
}
return hits, nil
}