package agent import ( "context" "fmt" "log/slog" "net" "time" ) // Resolver resolves dns-set FQDNs to host CIDRs on-device, honoring the // device's configured resolver (falling back to the system resolver). type Resolver struct { // Servers are resolver addresses (host or host:port); empty uses the system // resolver. The literal "system" is treated the same as empty. Servers []string } // NewResolver builds a Resolver for the given server list. func NewResolver(servers []string) *Resolver { if len(servers) == 1 && servers[0] == "system" { servers = nil } return &Resolver{Servers: servers} } func (r *Resolver) netResolver() *net.Resolver { if len(r.Servers) == 0 { return net.DefaultResolver } servers := r.Servers dialer := &net.Dialer{Timeout: 5 * time.Second} var idx int return &net.Resolver{ PreferGo: true, Dial: func(ctx context.Context, network, _ string) (net.Conn, error) { // Round-robin across configured servers for resilience. addr := servers[idx%len(servers)] idx++ if _, _, err := net.SplitHostPort(addr); err != nil { addr = net.JoinHostPort(addr, "53") } return dialer.DialContext(ctx, network, addr) }, } } // Resolve returns host CIDRs (/32 or /128) for a FQDN's A and AAAA records. func (r *Resolver) Resolve(ctx context.Context, fqdn string) ([]string, error) { ips, err := r.netResolver().LookupIP(ctx, "ip", fqdn) if err != nil { return nil, err } out := make([]string, 0, len(ips)) for _, ip := range ips { if ip4 := ip.To4(); ip4 != nil { out = append(out, ip4.String()+"/32") } else { out = append(out, ip.String()+"/128") } } return out, nil } // ExpandDNSSets resolves every dns set's FQDNs and populates its Members in // place. Resolution failures are logged and leave the prior Members untouched // (fail-safe): a resolver outage must never empty a set. func (r *Resolver) ExpandDNSSets(ctx context.Context, cfg *RenderedConfig) { for i := range cfg.Sets { set := &cfg.Sets[i] if set.Kind != "dns" { continue } var members []string var anyErr bool for _, fqdn := range set.FQDNs { cidrs, err := r.Resolve(ctx, fqdn) if err != nil { slog.Warn("agent: dns resolution failed, keeping last-good", "set", set.Name, "fqdn", fqdn, "err", err) anyErr = true continue } members = append(members, cidrs...) } // Only replace membership when we resolved something; never empty a set // on total failure. if len(members) > 0 { set.Members = dedup(members) } else if anyErr { slog.Warn("agent: dns set kept last-good members", "set", set.Name) } } } func dedup(in []string) []string { seen := make(map[string]struct{}, len(in)) out := in[:0] for _, s := range in { if _, ok := seen[s]; ok { continue } seen[s] = struct{}{} out = append(out, s) } return out } // validateCIDR is a small guard used by translation to skip malformed members. func validateCIDR(s string) error { if _, _, err := net.ParseCIDR(s); err != nil { return fmt.Errorf("invalid CIDR %q: %w", s, err) } return nil }