all: eliminate usage of global ProxyLogger

So setting up logging for ctrld binary and ctrld packages could be done
more easily, decouple the required setup for interactive vs daemon
running.

This is the first step toward replacing rs/zerolog libary with a
different logging library.
This commit is contained in:
Cuong Manh Le
2025-04-03 21:17:02 +07:00
committed by Cuong Manh Le
parent 47c04bf0f6
commit 0e66697247
39 changed files with 425 additions and 420 deletions
+16 -16
View File
@@ -349,7 +349,7 @@ func run(appCallback *AppCallback, stopCh chan struct{}) {
if newLogPath := cfg.Service.LogPath; newLogPath != "" && oldLogPath != newLogPath { if newLogPath := cfg.Service.LogPath; newLogPath != "" && oldLogPath != newLogPath {
// After processCDFlags, log config may change, so reset mainLog and re-init logging. // After processCDFlags, log config may change, so reset mainLog and re-init logging.
l := zerolog.New(io.Discard) l := zerolog.New(io.Discard)
mainLog.Store(&l) mainLog.Store(&ctrld.Logger{Logger: &l})
// Copy logs written so far to new log file if possible. // Copy logs written so far to new log file if possible.
if buf, err := os.ReadFile(oldLogPath); err == nil { if buf, err := os.ReadFile(oldLogPath); err == nil {
@@ -502,8 +502,7 @@ func readConfigFile(writeDefaultConfig, notice bool) bool {
if err := v.Unmarshal(&cfg); err != nil { if err := v.Unmarshal(&cfg); err != nil {
mainLog.Load().Fatal().Msgf("failed to unmarshal default config: %v", err) mainLog.Load().Fatal().Msgf("failed to unmarshal default config: %v", err)
} }
nop := zerolog.Nop() _, _ = tryUpdateListenerConfig(&cfg, func() {}, true)
_, _ = tryUpdateListenerConfig(&cfg, &nop, func() {}, true)
addExtraSplitDnsRule(&cfg) addExtraSplitDnsRule(&cfg)
if err := writeConfigFile(&cfg); err != nil { if err := writeConfigFile(&cfg); err != nil {
mainLog.Load().Fatal().Msgf("failed to write default config file: %v", err) mainLog.Load().Fatal().Msgf("failed to write default config file: %v", err)
@@ -591,7 +590,8 @@ func processNoConfigFlags(noConfigStart bool) {
Type: pType, Type: pType,
Timeout: 5000, Timeout: 5000,
} }
puc.Init() loggerCtx := ctrld.LoggerCtx(context.Background(), mainLog.Load())
puc.Init(loggerCtx)
upstream := map[string]*ctrld.UpstreamConfig{"0": puc} upstream := map[string]*ctrld.UpstreamConfig{"0": puc}
if secondaryUpstream != "" { if secondaryUpstream != "" {
sEndpoint, sType := endpointAndTyp(secondaryUpstream) sEndpoint, sType := endpointAndTyp(secondaryUpstream)
@@ -601,7 +601,7 @@ func processNoConfigFlags(noConfigStart bool) {
Type: sType, Type: sType,
Timeout: 5000, Timeout: 5000,
} }
suc.Init() suc.Init(loggerCtx)
upstream["1"] = suc upstream["1"] = suc
rules := make([]ctrld.Rule, 0, len(domains)) rules := make([]ctrld.Rule, 0, len(domains))
for _, domain := range domains { for _, domain := range domains {
@@ -634,13 +634,13 @@ func processCDFlags(cfg *ctrld.Config) (*controld.ResolverConfig, error) {
logger.Info().Msgf("fetching Controld D configuration from API: %s", cdUID) logger.Info().Msgf("fetching Controld D configuration from API: %s", cdUID)
bo := backoff.NewBackoff("processCDFlags", logf, 30*time.Second) bo := backoff.NewBackoff("processCDFlags", logf, 30*time.Second)
bo.LogLongerThan = 30 * time.Second bo.LogLongerThan = 30 * time.Second
ctx := context.Background() ctx := ctrld.LoggerCtx(context.Background(), mainLog.Load())
resolverConfig, err := controld.FetchResolverConfig(cdUID, rootCmd.Version, cdDev) resolverConfig, err := controld.FetchResolverConfig(ctx, cdUID, rootCmd.Version, cdDev)
for { for {
if errUrlNetworkError(err) { if errUrlNetworkError(err) {
bo.BackOff(ctx, err) bo.BackOff(ctx, err)
logger.Warn().Msg("could not fetch resolver using bootstrap DNS, retrying...") logger.Warn().Msg("could not fetch resolver using bootstrap DNS, retrying...")
resolverConfig, err = controld.FetchResolverConfig(cdUID, rootCmd.Version, cdDev) resolverConfig, err = controld.FetchResolverConfig(ctx, cdUID, rootCmd.Version, cdDev)
continue continue
} }
break break
@@ -938,9 +938,10 @@ func selfCheckResolveDomain(ctx context.Context, addr, scope string, domain stri
bo.BackOff(ctx, fmt.Errorf("ExchangeContext: %w", exErr)) bo.BackOff(ctx, fmt.Errorf("ExchangeContext: %w", exErr))
} }
mainLog.Load().Debug().Msgf("self-check against %q failed", domain) mainLog.Load().Debug().Msgf("self-check against %q failed", domain)
loggerCtx := ctrld.LoggerCtx(ctx, mainLog.Load())
// Ping all upstreams to provide better error message to users. // Ping all upstreams to provide better error message to users.
for name, uc := range cfg.Upstream { for name, uc := range cfg.Upstream {
if err := uc.ErrorPing(); err != nil { if err := uc.ErrorPing(loggerCtx); err != nil {
mainLog.Load().Err(err).Msgf("failed to connect to upstream.%s, endpoint: %s", name, uc.Endpoint) mainLog.Load().Err(err).Msgf("failed to connect to upstream.%s, endpoint: %s", name, uc.Endpoint)
} }
} }
@@ -1181,7 +1182,7 @@ func mobileListenerIp() string {
// or defined but invalid to be used, e.g: using loopback address other // or defined but invalid to be used, e.g: using loopback address other
// than 127.0.0.1 with systemd-resolved. // than 127.0.0.1 with systemd-resolved.
func updateListenerConfig(cfg *ctrld.Config, notifyToLogServerFunc func()) bool { func updateListenerConfig(cfg *ctrld.Config, notifyToLogServerFunc func()) bool {
updated, _ := tryUpdateListenerConfig(cfg, nil, notifyToLogServerFunc, true) updated, _ := tryUpdateListenerConfig(cfg, notifyToLogServerFunc, true)
if addExtraSplitDnsRule(cfg) { if addExtraSplitDnsRule(cfg) {
updated = true updated = true
} }
@@ -1191,7 +1192,7 @@ func updateListenerConfig(cfg *ctrld.Config, notifyToLogServerFunc func()) bool
// tryUpdateListenerConfig tries updating listener config with a working one. // tryUpdateListenerConfig tries updating listener config with a working one.
// If fatal is true, and there's listen address conflicted, the function do // If fatal is true, and there's listen address conflicted, the function do
// fatal error. // fatal error.
func tryUpdateListenerConfig(cfg *ctrld.Config, infoLogger *zerolog.Logger, notifyFunc func(), fatal bool) (updated, ok bool) { func tryUpdateListenerConfig(cfg *ctrld.Config, notifyFunc func(), fatal bool) (updated, ok bool) {
ok = true ok = true
lcc := make(map[string]*listenerConfigCheck) lcc := make(map[string]*listenerConfigCheck)
cdMode := cdUID != "" cdMode := cdUID != ""
@@ -1235,9 +1236,6 @@ func tryUpdateListenerConfig(cfg *ctrld.Config, infoLogger *zerolog.Logger, noti
} }
il := mainLog.Load() il := mainLog.Load()
if infoLogger != nil {
il = infoLogger
}
if isMobile() { if isMobile() {
// On Mobile, only use first listener, ignore others. // On Mobile, only use first listener, ignore others.
firstLn := cfg.FirstListener() firstLn := cfg.FirstListener()
@@ -1492,7 +1490,8 @@ func cdUIDFromProvToken() string {
} }
req := &controld.UtilityOrgRequest{ProvToken: cdOrg, Hostname: customHostname} req := &controld.UtilityOrgRequest{ProvToken: cdOrg, Hostname: customHostname}
// Process provision token if provided. // Process provision token if provided.
resolverConfig, err := controld.FetchResolverUID(req, rootCmd.Version, cdDev) loggerCtx := ctrld.LoggerCtx(context.Background(), mainLog.Load())
resolverConfig, err := controld.FetchResolverUID(loggerCtx, req, rootCmd.Version, cdDev)
if err != nil { if err != nil {
mainLog.Load().Fatal().Err(err).Msgf("failed to fetch resolver uid with provision token: %s", cdOrg) mainLog.Load().Fatal().Err(err).Msgf("failed to fetch resolver uid with provision token: %s", cdOrg)
} }
@@ -1819,7 +1818,8 @@ func runningIface(s service.Service) *ifaceResponse {
// doValidateCdRemoteConfig fetches and validates custom config for cdUID. // doValidateCdRemoteConfig fetches and validates custom config for cdUID.
func doValidateCdRemoteConfig(cdUID string, fatal bool) error { func doValidateCdRemoteConfig(cdUID string, fatal bool) error {
rc, err := controld.FetchResolverConfig(cdUID, rootCmd.Version, cdDev) loggerCtx := ctrld.LoggerCtx(context.Background(), mainLog.Load())
rc, err := controld.FetchResolverConfig(loggerCtx, cdUID, rootCmd.Version, cdDev)
if err != nil { if err != nil {
logger := mainLog.Load().Fatal() logger := mainLog.Load().Fatal()
if !fatal { if !fatal {
+4 -2
View File
@@ -216,8 +216,9 @@ func (p *prog) registerControlServerHandler() {
return return
} }
loggerCtx := ctrld.LoggerCtx(context.Background(), mainLog.Load())
// Re-fetch pin code from API. // Re-fetch pin code from API.
if rc, err := controld.FetchResolverConfig(cdUID, rootCmd.Version, cdDev); rc != nil { if rc, err := controld.FetchResolverConfig(loggerCtx, cdUID, rootCmd.Version, cdDev); rc != nil {
if rc.DeactivationPin != nil { if rc.DeactivationPin != nil {
cdDeactivationPin.Store(*rc.DeactivationPin) cdDeactivationPin.Store(*rc.DeactivationPin)
} else { } else {
@@ -321,7 +322,8 @@ func (p *prog) registerControlServerHandler() {
} }
mainLog.Load().Debug().Msg("sending log file to ControlD server") mainLog.Load().Debug().Msg("sending log file to ControlD server")
resp := logSentResponse{Size: r.size} resp := logSentResponse{Size: r.size}
if err := controld.SendLogs(req, cdDev); err != nil { loggerCtx := ctrld.LoggerCtx(context.Background(), mainLog.Load())
if err := controld.SendLogs(loggerCtx, req, cdDev); err != nil {
mainLog.Load().Error().Msgf("could not send log file to ControlD server: %v", err) mainLog.Load().Error().Msgf("could not send log file to ControlD server: %v", err)
resp.Error = err.Error() resp.Error = err.Error()
w.WriteHeader(http.StatusInternalServerError) w.WriteHeader(http.StatusInternalServerError)
+17 -14
View File
@@ -110,6 +110,7 @@ func (p *prog) serveDNS(mainCtx context.Context, listenerNum string) error {
listenerConfig := p.cfg.Listener[listenerNum] listenerConfig := p.cfg.Listener[listenerNum]
reqId := requestID() reqId := requestID()
ctx := context.WithValue(context.Background(), ctrld.ReqIdCtxKey{}, reqId) ctx := context.WithValue(context.Background(), ctrld.ReqIdCtxKey{}, reqId)
ctx = ctrld.LoggerCtx(ctx, mainLog.Load())
if !listenerConfig.AllowWanClients && isWanClient(w.RemoteAddr()) { if !listenerConfig.AllowWanClients && isWanClient(w.RemoteAddr()) {
ctrld.Log(ctx, mainLog.Load().Debug(), "query refused, listener does not allow WAN clients: %s", w.RemoteAddr().String()) ctrld.Log(ctx, mainLog.Load().Debug(), "query refused, listener does not allow WAN clients: %s", w.RemoteAddr().String())
answer := new(dns.Msg) answer := new(dns.Msg)
@@ -514,7 +515,7 @@ func (p *prog) proxy(ctx context.Context, req *proxyRequest) *proxyResponse {
} }
resolve1 := func(upstream string, upstreamConfig *ctrld.UpstreamConfig, msg *dns.Msg) (*dns.Msg, error) { resolve1 := func(upstream string, upstreamConfig *ctrld.UpstreamConfig, msg *dns.Msg) (*dns.Msg, error) {
ctrld.Log(ctx, mainLog.Load().Debug(), "sending query to %s: %s", upstream, upstreamConfig.Name) ctrld.Log(ctx, mainLog.Load().Debug(), "sending query to %s: %s", upstream, upstreamConfig.Name)
dnsResolver, err := ctrld.NewResolver(upstreamConfig) dnsResolver, err := ctrld.NewResolver(ctx, upstreamConfig)
if err != nil { if err != nil {
ctrld.Log(ctx, mainLog.Load().Error().Err(err), "failed to create resolver") ctrld.Log(ctx, mainLog.Load().Error().Err(err), "failed to create resolver")
return nil, err return nil, err
@@ -549,11 +550,11 @@ func (p *prog) proxy(ctx context.Context, req *proxyRequest) *proxyResponse {
// For timeout error (i.e: context deadline exceed), force re-bootstrapping. // For timeout error (i.e: context deadline exceed), force re-bootstrapping.
var e net.Error var e net.Error
if errors.As(err, &e) && e.Timeout() { if errors.As(err, &e) && e.Timeout() {
upstreamConfig.ReBootstrap() upstreamConfig.ReBootstrap(ctx)
} }
// For network error, turn ipv6 off if enabled. // For network error, turn ipv6 off if enabled.
if ctrld.HasIPv6() && (errUrlNetworkError(err) || errNetworkError(err)) { if ctrld.HasIPv6(ctx) && (errUrlNetworkError(err) || errNetworkError(err)) {
ctrld.DisableIPv6() ctrld.DisableIPv6(ctx)
} }
} }
@@ -960,7 +961,8 @@ func (p *prog) doSelfUninstall(answer *dns.Msg) {
logger := mainLog.Load().With().Str("mode", "self-uninstall").Logger() logger := mainLog.Load().With().Str("mode", "self-uninstall").Logger()
if p.refusedQueryCount > selfUninstallMaxQueries { if p.refusedQueryCount > selfUninstallMaxQueries {
p.checkingSelfUninstall = true p.checkingSelfUninstall = true
_, err := controld.FetchResolverConfig(cdUID, rootCmd.Version, cdDev) loggerCtx := ctrld.LoggerCtx(context.Background(), mainLog.Load())
_, err := controld.FetchResolverConfig(loggerCtx, cdUID, rootCmd.Version, cdDev)
logger.Debug().Msg("maximum number of refused queries reached, checking device status") logger.Debug().Msg("maximum number of refused queries reached, checking device status")
selfUninstallCheck(err, p, logger) selfUninstallCheck(err, p, logger)
@@ -1326,13 +1328,13 @@ func (p *prog) monitorNetworkChanges(ctx context.Context) error {
// Only set the IPv4 default if selfIP is a valid IPv4 address. // Only set the IPv4 default if selfIP is a valid IPv4 address.
if ip := net.ParseIP(selfIP); ip != nil && ip.To4() != nil { if ip := net.ParseIP(selfIP); ip != nil && ip.To4() != nil {
ctrld.SetDefaultLocalIPv4(ip) ctrld.SetDefaultLocalIPv4(ctrld.LoggerCtx(ctx, mainLog.Load()), ip)
if !isMobile() && p.ciTable != nil { if !isMobile() && p.ciTable != nil {
p.ciTable.SetSelfIP(selfIP) p.ciTable.SetSelfIP(selfIP)
} }
} }
if ip := net.ParseIP(ipv6); ip != nil { if ip := net.ParseIP(ipv6); ip != nil {
ctrld.SetDefaultLocalIPv6(ip) ctrld.SetDefaultLocalIPv6(ctrld.LoggerCtx(ctx, mainLog.Load()), ip)
} }
mainLog.Load().Debug().Msgf("Set default local IPv4: %s, IPv6: %s", selfIP, ipv6) mainLog.Load().Debug().Msgf("Set default local IPv4: %s, IPv6: %s", selfIP, ipv6)
@@ -1400,7 +1402,7 @@ func interfaceIPsEqual(a, b []netip.Prefix) bool {
func (p *prog) checkUpstreamOnce(upstream string, uc *ctrld.UpstreamConfig) error { func (p *prog) checkUpstreamOnce(upstream string, uc *ctrld.UpstreamConfig) error {
mainLog.Load().Debug().Msgf("Starting check for upstream: %s", upstream) mainLog.Load().Debug().Msgf("Starting check for upstream: %s", upstream)
resolver, err := ctrld.NewResolver(uc) resolver, err := ctrld.NewResolver(ctrld.LoggerCtx(context.Background(), mainLog.Load()), uc)
if err != nil { if err != nil {
mainLog.Load().Error().Err(err).Msgf("Failed to create resolver for upstream %s", upstream) mainLog.Load().Error().Err(err).Msgf("Failed to create resolver for upstream %s", upstream)
return err return err
@@ -1418,7 +1420,7 @@ func (p *prog) checkUpstreamOnce(upstream string, uc *ctrld.UpstreamConfig) erro
ctx, cancel := context.WithTimeout(context.Background(), timeout) ctx, cancel := context.WithTimeout(context.Background(), timeout)
defer cancel() defer cancel()
uc.ReBootstrap() uc.ReBootstrap(ctrld.LoggerCtx(ctx, mainLog.Load()))
mainLog.Load().Debug().Msgf("Rebootstrapping resolver for upstream: %s", upstream) mainLog.Load().Debug().Msgf("Rebootstrapping resolver for upstream: %s", upstream)
start := time.Now() start := time.Now()
@@ -1474,10 +1476,11 @@ func (p *prog) handleRecovery(reason RecoveryReason) {
// will be appended to nameservers from the saved interface values // will be appended to nameservers from the saved interface values
p.resetDNS(false, false) p.resetDNS(false, false)
loggerCtx := ctrld.LoggerCtx(context.Background(), mainLog.Load())
// For an OS failure, reinitialize OS resolver nameservers immediately. // For an OS failure, reinitialize OS resolver nameservers immediately.
if reason == RecoveryReasonOSFailure { if reason == RecoveryReasonOSFailure {
mainLog.Load().Debug().Msg("OS resolver failure detected; reinitializing OS resolver nameservers") mainLog.Load().Debug().Msg("OS resolver failure detected; reinitializing OS resolver nameservers")
ns := ctrld.InitializeOsResolver(true) ns := ctrld.InitializeOsResolver(loggerCtx, true)
if len(ns) == 0 { if len(ns) == 0 {
mainLog.Load().Warn().Msg("No nameservers found for OS resolver; using existing values") mainLog.Load().Warn().Msg("No nameservers found for OS resolver; using existing values")
} else { } else {
@@ -1504,7 +1507,7 @@ func (p *prog) handleRecovery(reason RecoveryReason) {
// For network changes we also reinitialize the OS resolver. // For network changes we also reinitialize the OS resolver.
if reason == RecoveryReasonNetworkChange { if reason == RecoveryReasonNetworkChange {
ns := ctrld.InitializeOsResolver(true) ns := ctrld.InitializeOsResolver(loggerCtx, true)
if len(ns) == 0 { if len(ns) == 0 {
mainLog.Load().Warn().Msg("No nameservers found for OS resolver during network-change recovery; using existing values") mainLog.Load().Warn().Msg("No nameservers found for OS resolver during network-change recovery; using existing values")
} else { } else {
@@ -1564,7 +1567,7 @@ func (p *prog) waitForUpstreamRecovery(ctx context.Context, upstreams map[string
// we should try to reinit the OS resolver to ensure we can recover // we should try to reinit the OS resolver to ensure we can recover
if name == upstreamOS && attempts%3 == 0 { if name == upstreamOS && attempts%3 == 0 {
mainLog.Load().Debug().Msgf("UpstreamOS check failed on attempt %d, reinitializing OS resolver", attempts) mainLog.Load().Debug().Msgf("UpstreamOS check failed on attempt %d, reinitializing OS resolver", attempts)
ns := ctrld.InitializeOsResolver(true) ns := ctrld.InitializeOsResolver(ctrld.LoggerCtx(ctx, mainLog.Load()), true)
if len(ns) == 0 { if len(ns) == 0 {
mainLog.Load().Warn().Msg("No nameservers found for OS resolver; using existing values") mainLog.Load().Warn().Msg("No nameservers found for OS resolver; using existing values")
} else { } else {
@@ -1624,12 +1627,12 @@ func ValidateDefaultLocalIPsFromDelta(newState *netmon.State) {
// Check if the default IPv4 is still active. // Check if the default IPv4 is still active.
if currentIPv4 != nil && !activeIPs[currentIPv4.String()] { if currentIPv4 != nil && !activeIPs[currentIPv4.String()] {
mainLog.Load().Debug().Msgf("DefaultLocalIPv4 %s is no longer active in the new state. Resetting.", currentIPv4) mainLog.Load().Debug().Msgf("DefaultLocalIPv4 %s is no longer active in the new state. Resetting.", currentIPv4)
ctrld.SetDefaultLocalIPv4(nil) ctrld.SetDefaultLocalIPv4(ctrld.LoggerCtx(context.Background(), mainLog.Load()), nil)
} }
// Check if the default IPv6 is still active. // Check if the default IPv6 is still active.
if currentIPv6 != nil && !activeIPs[currentIPv6.String()] { if currentIPv6 != nil && !activeIPs[currentIPv6.String()] {
mainLog.Load().Debug().Msgf("DefaultLocalIPv6 %s is no longer active in the new state. Resetting.", currentIPv6) mainLog.Load().Debug().Msgf("DefaultLocalIPv6 %s is no longer active in the new state. Resetting.", currentIPv6)
ctrld.SetDefaultLocalIPv6(nil) ctrld.SetDefaultLocalIPv6(ctrld.LoggerCtx(context.Background(), mainLog.Load()), nil)
} }
} }
+1 -2
View File
@@ -137,8 +137,7 @@ func (p *prog) initInternalLogging(writers []io.Writer) {
}) })
multi := zerolog.MultiLevelWriter(writers...) multi := zerolog.MultiLevelWriter(writers...)
l := mainLog.Load().Output(multi).With().Logger() l := mainLog.Load().Output(multi).With().Logger()
mainLog.Store(&l) mainLog.Store(&ctrld.Logger{Logger: &l})
ctrld.ProxyLogger.Store(&l)
} }
// needInternalLogging reports whether prog needs to run internal logging. // needInternalLogging reports whether prog needs to run internal logging.
+2 -1
View File
@@ -102,6 +102,7 @@ func (p *prog) checkDnsLoop() {
} }
p.loopMu.Unlock() p.loopMu.Unlock()
loggerCtx := ctrld.LoggerCtx(context.Background(), mainLog.Load())
for uid := range p.loop { for uid := range p.loop {
msg := loopTestMsg(uid) msg := loopTestMsg(uid)
uc := upstream[uid] uc := upstream[uid]
@@ -109,7 +110,7 @@ func (p *prog) checkDnsLoop() {
if uc == nil { if uc == nil {
continue continue
} }
resolver, err := ctrld.NewResolver(uc) resolver, err := ctrld.NewResolver(loggerCtx, uc)
if err != nil { if err != nil {
mainLog.Load().Warn().Err(err).Msgf("could not perform loop check for upstream: %q, endpoint: %q", uc.Name, uc.Endpoint) mainLog.Load().Warn().Err(err).Msgf("could not perform loop check for upstream: %q, endpoint: %q", uc.Name, uc.Endpoint)
continue continue
+4 -10
View File
@@ -40,7 +40,7 @@ var (
cleanup bool cleanup bool
startOnly bool startOnly bool
mainLog atomic.Pointer[zerolog.Logger] mainLog atomic.Pointer[ctrld.Logger]
consoleWriter zerolog.ConsoleWriter consoleWriter zerolog.ConsoleWriter
noConfigStart bool noConfigStart bool
) )
@@ -54,7 +54,7 @@ const (
func init() { func init() {
l := zerolog.New(io.Discard) l := zerolog.New(io.Discard)
mainLog.Store(&l) mainLog.Store(&ctrld.Logger{Logger: &l})
} }
func Main() { func Main() {
@@ -87,16 +87,14 @@ func initConsoleLogging() {
}) })
multi := zerolog.MultiLevelWriter(consoleWriter) multi := zerolog.MultiLevelWriter(consoleWriter)
l := mainLog.Load().Output(multi).With().Timestamp().Logger() l := mainLog.Load().Output(multi).With().Timestamp().Logger()
mainLog.Store(&l) mainLog.Store(&ctrld.Logger{Logger: &l})
switch { switch {
case silent: case silent:
zerolog.SetGlobalLevel(zerolog.NoLevel) zerolog.SetGlobalLevel(zerolog.NoLevel)
case verbose == 1: case verbose == 1:
ctrld.ProxyLogger.Store(&l)
zerolog.SetGlobalLevel(zerolog.InfoLevel) zerolog.SetGlobalLevel(zerolog.InfoLevel)
case verbose > 1: case verbose > 1:
ctrld.ProxyLogger.Store(&l)
zerolog.SetGlobalLevel(zerolog.DebugLevel) zerolog.SetGlobalLevel(zerolog.DebugLevel)
default: default:
zerolog.SetGlobalLevel(zerolog.NoticeLevel) zerolog.SetGlobalLevel(zerolog.NoticeLevel)
@@ -113,8 +111,6 @@ func initInteractiveLogging() {
zerolog.TimeFieldFormat = time.RFC3339 + ".000" zerolog.TimeFieldFormat = time.RFC3339 + ".000"
initLoggingWithBackup(false) initLoggingWithBackup(false)
cfg.Service.LogPath = old cfg.Service.LogPath = old
l := zerolog.New(io.Discard)
ctrld.ProxyLogger.Store(&l)
} }
// initLoggingWithBackup initializes log setup base on current config. // initLoggingWithBackup initializes log setup base on current config.
@@ -153,9 +149,7 @@ func initLoggingWithBackup(doBackup bool) []io.Writer {
writers = append(writers, consoleWriter) writers = append(writers, consoleWriter)
multi := zerolog.MultiLevelWriter(writers...) multi := zerolog.MultiLevelWriter(writers...)
l := mainLog.Load().Output(multi).With().Logger() l := mainLog.Load().Output(multi).With().Logger()
mainLog.Store(&l) mainLog.Store(&ctrld.Logger{Logger: &l})
// TODO: find a better way.
ctrld.ProxyLogger.Store(&l)
zerolog.SetGlobalLevel(zerolog.NoticeLevel) zerolog.SetGlobalLevel(zerolog.NoticeLevel)
logLevel := cfg.Service.LogLevel logLevel := cfg.Service.LogLevel
+3 -1
View File
@@ -6,12 +6,14 @@ import (
"testing" "testing"
"github.com/rs/zerolog" "github.com/rs/zerolog"
"github.com/Control-D-Inc/ctrld"
) )
var logOutput strings.Builder var logOutput strings.Builder
func TestMain(m *testing.M) { func TestMain(m *testing.M) {
l := zerolog.New(&logOutput) l := zerolog.New(&logOutput)
mainLog.Store(&l) mainLog.Store(&ctrld.Logger{Logger: &l})
os.Exit(m.Run()) os.Exit(m.Run())
} }
+3 -1
View File
@@ -5,6 +5,8 @@ import (
"github.com/vishvananda/netlink" "github.com/vishvananda/netlink"
"golang.org/x/sys/unix" "golang.org/x/sys/unix"
"github.com/Control-D-Inc/ctrld"
) )
func (p *prog) watchLinkState(ctx context.Context) { func (p *prog) watchLinkState(ctx context.Context) {
@@ -26,7 +28,7 @@ func (p *prog) watchLinkState(ctx context.Context) {
if lu.Change&unix.IFF_UP != 0 { if lu.Change&unix.IFF_UP != 0 {
mainLog.Load().Debug().Msgf("link state changed, re-bootstrapping") mainLog.Load().Debug().Msgf("link state changed, re-bootstrapping")
for _, uc := range p.cfg.Upstream { for _, uc := range p.cfg.Upstream {
uc.ReBootstrap() uc.ReBootstrap(ctrld.LoggerCtx(ctx, mainLog.Load()))
} }
} }
} }
+9 -7
View File
@@ -286,7 +286,7 @@ func (p *prog) postRun() {
mainLog.Load().Debug().Msgf("running on domain controller: %t, role: %d", p.runningOnDomainController, roleInt) mainLog.Load().Debug().Msgf("running on domain controller: %t, role: %d", p.runningOnDomainController, roleInt)
} }
p.resetDNS(false, false) p.resetDNS(false, false)
ns := ctrld.InitializeOsResolver(false) ns := ctrld.InitializeOsResolver(ctrld.LoggerCtx(context.Background(), mainLog.Load()), false)
mainLog.Load().Debug().Msgf("initialized OS resolver with nameservers: %v", ns) mainLog.Load().Debug().Msgf("initialized OS resolver with nameservers: %v", ns)
p.setDNS() p.setDNS()
p.csSetDnsDone <- struct{}{} p.csSetDnsDone <- struct{}{}
@@ -319,7 +319,8 @@ func (p *prog) apiConfigReload() {
} }
doReloadApiConfig := func(forced bool, logger zerolog.Logger) { doReloadApiConfig := func(forced bool, logger zerolog.Logger) {
resolverConfig, err := controld.FetchResolverConfig(cdUID, rootCmd.Version, cdDev) loggerCtx := ctrld.LoggerCtx(context.Background(), mainLog.Load())
resolverConfig, err := controld.FetchResolverConfig(loggerCtx, cdUID, rootCmd.Version, cdDev)
selfUninstallCheck(err, p, logger) selfUninstallCheck(err, p, logger)
if err != nil { if err != nil {
logger.Warn().Err(err).Msg("could not fetch resolver config") logger.Warn().Err(err).Msg("could not fetch resolver config")
@@ -377,7 +378,7 @@ func (p *prog) apiConfigReload() {
} }
if cfgErr != nil { if cfgErr != nil {
logger.Warn().Err(err).Msg("skipping invalid custom config") logger.Warn().Err(err).Msg("skipping invalid custom config")
if _, err := controld.UpdateCustomLastFailed(cdUID, rootCmd.Version, cdDev, true); err != nil { if _, err := controld.UpdateCustomLastFailed(loggerCtx, cdUID, rootCmd.Version, cdDev, true); err != nil {
logger.Error().Err(err).Msg("could not mark custom last update failed") logger.Error().Err(err).Msg("could not mark custom last update failed")
} }
return return
@@ -404,22 +405,23 @@ func (p *prog) setupUpstream(cfg *ctrld.Config) {
localUpstreams := make([]string, 0, len(cfg.Upstream)) localUpstreams := make([]string, 0, len(cfg.Upstream))
ptrNameservers := make([]string, 0, len(cfg.Upstream)) ptrNameservers := make([]string, 0, len(cfg.Upstream))
isControlDUpstream := false isControlDUpstream := false
loggerCtx := ctrld.LoggerCtx(context.Background(), mainLog.Load())
for n := range cfg.Upstream { for n := range cfg.Upstream {
uc := cfg.Upstream[n] uc := cfg.Upstream[n]
sdns := uc.Type == ctrld.ResolverTypeSDNS sdns := uc.Type == ctrld.ResolverTypeSDNS
uc.Init() uc.Init(loggerCtx)
if sdns { if sdns {
mainLog.Load().Debug().Msgf("initialized DNS Stamps with endpoint: %s, type: %s", uc.Endpoint, uc.Type) mainLog.Load().Debug().Msgf("initialized DNS Stamps with endpoint: %s, type: %s", uc.Endpoint, uc.Type)
} }
isControlDUpstream = isControlDUpstream || uc.IsControlD() isControlDUpstream = isControlDUpstream || uc.IsControlD()
if uc.BootstrapIP == "" { if uc.BootstrapIP == "" {
uc.SetupBootstrapIP() uc.SetupBootstrapIP(ctrld.LoggerCtx(context.Background(), mainLog.Load()))
mainLog.Load().Info().Msgf("bootstrap IPs for upstream.%s: %q", n, uc.BootstrapIPs()) mainLog.Load().Info().Msgf("bootstrap IPs for upstream.%s: %q", n, uc.BootstrapIPs())
} else { } else {
mainLog.Load().Info().Str("bootstrap_ip", uc.BootstrapIP).Msgf("using bootstrap IP for upstream.%s", n) mainLog.Load().Info().Str("bootstrap_ip", uc.BootstrapIP).Msgf("using bootstrap IP for upstream.%s", n)
} }
uc.SetCertPool(rootCertPool) uc.SetCertPool(rootCertPool)
go uc.Ping() go uc.Ping(loggerCtx)
if canBeLocalUpstream(uc.Domain) { if canBeLocalUpstream(uc.Domain) {
localUpstreams = append(localUpstreams, upstreamPrefix+n) localUpstreams = append(localUpstreams, upstreamPrefix+n)
@@ -601,7 +603,7 @@ func (p *prog) run(reload bool, reloadCh chan struct{}) {
// setupClientInfoDiscover performs necessary works for running client info discover. // setupClientInfoDiscover performs necessary works for running client info discover.
func (p *prog) setupClientInfoDiscover(selfIP string) { func (p *prog) setupClientInfoDiscover(selfIP string) {
p.ciTable = clientinfo.NewTable(&cfg, selfIP, cdUID, p.ptrNameservers) p.ciTable = clientinfo.NewTable(&cfg, selfIP, cdUID, p.ptrNameservers, mainLog.Load())
if leaseFile := p.cfg.Service.DHCPLeaseFile; leaseFile != "" { if leaseFile := p.cfg.Service.DHCPLeaseFile; leaseFile != "" {
mainLog.Load().Debug().Msgf("watching custom lease file: %s", leaseFile) mainLog.Load().Debug().Msgf("watching custom lease file: %s", leaseFile)
format := ctrld.LeaseFileFormat(p.cfg.Service.DHCPLeaseFileFormat) format := ctrld.LeaseFileFormat(p.cfg.Service.DHCPLeaseFileFormat)
+55 -48
View File
@@ -325,12 +325,13 @@ type ListenerPolicyConfig struct {
type Rule map[string][]string type Rule map[string][]string
// Init initialized necessary values for an UpstreamConfig. // Init initialized necessary values for an UpstreamConfig.
func (uc *UpstreamConfig) Init() { func (uc *UpstreamConfig) Init(ctx context.Context) {
logger := LoggerFromCtx(ctx)
if err := uc.initDnsStamps(); err != nil { if err := uc.initDnsStamps(); err != nil {
ProxyLogger.Load().Fatal().Err(err).Msg("invalid DNS Stamps") logger.Fatal().Err(err).Msg("invalid DNS Stamps")
} }
uc.initDoHScheme() uc.initDoHScheme()
uc.uid = upstreamUID() uc.uid = upstreamUID(ctx)
if u, err := url.Parse(uc.Endpoint); err == nil { if u, err := url.Parse(uc.Endpoint); err == nil {
uc.Domain = u.Hostname() uc.Domain = u.Hostname()
switch uc.Type { switch uc.Type {
@@ -434,12 +435,13 @@ func (uc *UpstreamConfig) UID() string {
// - ControlD Bootstrap DNS 76.76.2.22 // - ControlD Bootstrap DNS 76.76.2.22
// //
// The setup process will block until there's usable IPs found. // The setup process will block until there's usable IPs found.
func (uc *UpstreamConfig) SetupBootstrapIP() { func (uc *UpstreamConfig) SetupBootstrapIP(ctx context.Context) {
b := backoff.NewBackoff("setupBootstrapIP", func(format string, args ...any) {}, 10*time.Second) b := backoff.NewBackoff("setupBootstrapIP", func(format string, args ...any) {}, 10*time.Second)
isControlD := uc.IsControlD() isControlD := uc.IsControlD()
nss := initDefaultOsResolver() logger := LoggerFromCtx(ctx)
nss := initDefaultOsResolver(ctx)
for { for {
uc.bootstrapIPs = lookupIP(uc.Domain, uc.Timeout, nss) uc.bootstrapIPs = lookupIP(ctx, uc.Domain, uc.Timeout, nss)
// For ControlD upstream, the bootstrap IPs could not be RFC 1918 addresses, // For ControlD upstream, the bootstrap IPs could not be RFC 1918 addresses,
// filtering them out here to prevent weird behavior. // filtering them out here to prevent weird behavior.
if isControlD { if isControlD {
@@ -454,18 +456,18 @@ func (uc *UpstreamConfig) SetupBootstrapIP() {
uc.bootstrapIPs = uc.bootstrapIPs[:n] uc.bootstrapIPs = uc.bootstrapIPs[:n]
if len(uc.bootstrapIPs) == 0 { if len(uc.bootstrapIPs) == 0 {
uc.bootstrapIPs = bootstrapIPsFromControlDDomain(uc.Domain) uc.bootstrapIPs = bootstrapIPsFromControlDDomain(uc.Domain)
ProxyLogger.Load().Warn().Msgf("no record found for %q, lookup from direct IP table", uc.Domain) logger.Warn().Msgf("no record found for %q, lookup from direct IP table", uc.Domain)
} }
} }
if len(uc.bootstrapIPs) == 0 { if len(uc.bootstrapIPs) == 0 {
ProxyLogger.Load().Warn().Msgf("no record found for %q, using bootstrap server: %s", uc.Domain, PremiumDNSBoostrapIP) logger.Warn().Msgf("no record found for %q, using bootstrap server: %s", uc.Domain, PremiumDNSBoostrapIP)
uc.bootstrapIPs = lookupIP(uc.Domain, uc.Timeout, []string{net.JoinHostPort(PremiumDNSBoostrapIP, "53")}) uc.bootstrapIPs = lookupIP(ctx, uc.Domain, uc.Timeout, []string{net.JoinHostPort(PremiumDNSBoostrapIP, "53")})
} }
if len(uc.bootstrapIPs) > 0 { if len(uc.bootstrapIPs) > 0 {
break break
} }
ProxyLogger.Load().Warn().Msg("could not resolve bootstrap IPs, retrying...") logger.Warn().Msg("could not resolve bootstrap IPs, retrying...")
b.BackOff(context.Background(), errors.New("no bootstrap IPs")) b.BackOff(context.Background(), errors.New("no bootstrap IPs"))
} }
for _, ip := range uc.bootstrapIPs { for _, ip := range uc.bootstrapIPs {
@@ -475,11 +477,11 @@ func (uc *UpstreamConfig) SetupBootstrapIP() {
uc.bootstrapIPs4 = append(uc.bootstrapIPs4, ip) uc.bootstrapIPs4 = append(uc.bootstrapIPs4, ip)
} }
} }
ProxyLogger.Load().Debug().Msgf("bootstrap IPs: %v", uc.bootstrapIPs) logger.Debug().Msgf("bootstrap IPs: %v", uc.bootstrapIPs)
} }
// ReBootstrap re-setup the bootstrap IP and the transport. // ReBootstrap re-setup the bootstrap IP and the transport.
func (uc *UpstreamConfig) ReBootstrap() { func (uc *UpstreamConfig) ReBootstrap(ctx context.Context) {
switch uc.Type { switch uc.Type {
case ResolverTypeDOH, ResolverTypeDOH3: case ResolverTypeDOH, ResolverTypeDOH3:
default: default:
@@ -487,7 +489,8 @@ func (uc *UpstreamConfig) ReBootstrap() {
} }
_, _, _ = uc.g.Do("ReBootstrap", func() (any, error) { _, _, _ = uc.g.Do("ReBootstrap", func() (any, error) {
if uc.rebootstrap.CompareAndSwap(false, true) { if uc.rebootstrap.CompareAndSwap(false, true) {
ProxyLogger.Load().Debug().Msgf("re-bootstrapping upstream ip for %v", uc) logger := LoggerFromCtx(ctx)
logger.Debug().Msgf("re-bootstrapping upstream ip for %v", uc)
} }
return true, nil return true, nil
}) })
@@ -495,35 +498,35 @@ func (uc *UpstreamConfig) ReBootstrap() {
// SetupTransport initializes the network transport used to connect to upstream server. // SetupTransport initializes the network transport used to connect to upstream server.
// For now, only DoH upstream is supported. // For now, only DoH upstream is supported.
func (uc *UpstreamConfig) SetupTransport() { func (uc *UpstreamConfig) SetupTransport(ctx context.Context) {
switch uc.Type { switch uc.Type {
case ResolverTypeDOH: case ResolverTypeDOH:
uc.setupDOHTransport() uc.setupDOHTransport(ctx)
case ResolverTypeDOH3: case ResolverTypeDOH3:
uc.setupDOH3Transport() uc.setupDOH3Transport(ctx)
} }
} }
func (uc *UpstreamConfig) setupDOHTransport() { func (uc *UpstreamConfig) setupDOHTransport(ctx context.Context) {
switch uc.IPStack { switch uc.IPStack {
case IpStackBoth, "": case IpStackBoth, "":
uc.transport = uc.newDOHTransport(uc.bootstrapIPs) uc.transport = uc.newDOHTransport(ctx, uc.bootstrapIPs)
case IpStackV4: case IpStackV4:
uc.transport = uc.newDOHTransport(uc.bootstrapIPs4) uc.transport = uc.newDOHTransport(ctx, uc.bootstrapIPs4)
case IpStackV6: case IpStackV6:
uc.transport = uc.newDOHTransport(uc.bootstrapIPs6) uc.transport = uc.newDOHTransport(ctx, uc.bootstrapIPs6)
case IpStackSplit: case IpStackSplit:
uc.transport4 = uc.newDOHTransport(uc.bootstrapIPs4) uc.transport4 = uc.newDOHTransport(ctx, uc.bootstrapIPs4)
if HasIPv6() { if HasIPv6(ctx) {
uc.transport6 = uc.newDOHTransport(uc.bootstrapIPs6) uc.transport6 = uc.newDOHTransport(ctx, uc.bootstrapIPs6)
} else { } else {
uc.transport6 = uc.transport4 uc.transport6 = uc.transport4
} }
uc.transport = uc.newDOHTransport(uc.bootstrapIPs) uc.transport = uc.newDOHTransport(ctx, uc.bootstrapIPs)
} }
} }
func (uc *UpstreamConfig) newDOHTransport(addrs []string) *http.Transport { func (uc *UpstreamConfig) newDOHTransport(ctx context.Context, addrs []string) *http.Transport {
transport := http.DefaultTransport.(*http.Transport).Clone() transport := http.DefaultTransport.(*http.Transport).Clone()
transport.MaxIdleConnsPerHost = 100 transport.MaxIdleConnsPerHost = 100
transport.TLSClientConfig = &tls.Config{ transport.TLSClientConfig = &tls.Config{
@@ -543,12 +546,13 @@ func (uc *UpstreamConfig) newDOHTransport(addrs []string) *http.Transport {
dialerTimeoutMs = uc.Timeout dialerTimeoutMs = uc.Timeout
} }
dialerTimeout := time.Duration(dialerTimeoutMs) * time.Millisecond dialerTimeout := time.Duration(dialerTimeoutMs) * time.Millisecond
logger := LoggerFromCtx(ctx)
transport.DialContext = func(ctx context.Context, network, addr string) (net.Conn, error) { transport.DialContext = func(ctx context.Context, network, addr string) (net.Conn, error) {
_, port, _ := net.SplitHostPort(addr) _, port, _ := net.SplitHostPort(addr)
if uc.BootstrapIP != "" { if uc.BootstrapIP != "" {
dialer := net.Dialer{Timeout: dialerTimeout, KeepAlive: dialerTimeout} dialer := net.Dialer{Timeout: dialerTimeout, KeepAlive: dialerTimeout}
addr := net.JoinHostPort(uc.BootstrapIP, port) addr := net.JoinHostPort(uc.BootstrapIP, port)
Log(ctx, ProxyLogger.Load().Debug(), "sending doh request to: %s", addr) logger.Debug().Msgf("sending doh request to: %s", addr)
return dialer.DialContext(ctx, network, addr) return dialer.DialContext(ctx, network, addr)
} }
pd := &ctrldnet.ParallelDialer{} pd := &ctrldnet.ParallelDialer{}
@@ -558,11 +562,11 @@ func (uc *UpstreamConfig) newDOHTransport(addrs []string) *http.Transport {
for i := range addrs { for i := range addrs {
dialAddrs[i] = net.JoinHostPort(addrs[i], port) dialAddrs[i] = net.JoinHostPort(addrs[i], port)
} }
conn, err := pd.DialContext(ctx, network, dialAddrs, ProxyLogger.Load()) conn, err := pd.DialContext(ctx, network, dialAddrs, logger.Logger)
if err != nil { if err != nil {
return nil, err return nil, err
} }
Log(ctx, ProxyLogger.Load().Debug(), "sending doh request to: %s", conn.RemoteAddr()) logger.Debug().Msgf("sending doh request to: %s", conn.RemoteAddr())
return conn, nil return conn, nil
} }
runtime.SetFinalizer(transport, func(transport *http.Transport) { runtime.SetFinalizer(transport, func(transport *http.Transport) {
@@ -572,19 +576,20 @@ func (uc *UpstreamConfig) newDOHTransport(addrs []string) *http.Transport {
} }
// Ping warms up the connection to DoH/DoH3 upstream. // Ping warms up the connection to DoH/DoH3 upstream.
func (uc *UpstreamConfig) Ping() { func (uc *UpstreamConfig) Ping(ctx context.Context) {
if err := uc.ping(); err != nil { if err := uc.ping(ctx); err != nil {
ProxyLogger.Load().Debug().Err(err).Msgf("upstream ping failed: %s", uc.Endpoint) logger := LoggerFromCtx(ctx)
_ = uc.FallbackToDirectIP() logger.Debug().Err(err).Msgf("upstream ping failed: %s", uc.Endpoint)
_ = uc.FallbackToDirectIP(ctx)
} }
} }
// ErrorPing is like Ping, but return an error if any. // ErrorPing is like Ping, but return an error if any.
func (uc *UpstreamConfig) ErrorPing() error { func (uc *UpstreamConfig) ErrorPing(ctx context.Context) error {
return uc.ping() return uc.ping(ctx)
} }
func (uc *UpstreamConfig) ping() error { func (uc *UpstreamConfig) ping(ctx context.Context) error {
switch uc.Type { switch uc.Type {
case ResolverTypeDOH, ResolverTypeDOH3: case ResolverTypeDOH, ResolverTypeDOH3:
default: default:
@@ -613,11 +618,11 @@ func (uc *UpstreamConfig) ping() error {
for _, typ := range []uint16{dns.TypeA, dns.TypeAAAA} { for _, typ := range []uint16{dns.TypeA, dns.TypeAAAA} {
switch uc.Type { switch uc.Type {
case ResolverTypeDOH: case ResolverTypeDOH:
if err := ping(uc.dohTransport(typ)); err != nil { if err := ping(uc.dohTransport(ctx, typ)); err != nil {
return err return err
} }
case ResolverTypeDOH3: case ResolverTypeDOH3:
if err := ping(uc.doh3Transport(typ)); err != nil { if err := ping(uc.doh3Transport(ctx, typ)); err != nil {
return err return err
} }
} }
@@ -652,12 +657,12 @@ func (uc *UpstreamConfig) isNextDNS() bool {
return domain == "dns.nextdns.io" return domain == "dns.nextdns.io"
} }
func (uc *UpstreamConfig) dohTransport(dnsType uint16) http.RoundTripper { func (uc *UpstreamConfig) dohTransport(ctx context.Context, dnsType uint16) http.RoundTripper {
uc.transportOnce.Do(func() { uc.transportOnce.Do(func() {
uc.SetupTransport() uc.SetupTransport(ctx)
}) })
if uc.rebootstrap.CompareAndSwap(true, false) { if uc.rebootstrap.CompareAndSwap(true, false) {
uc.SetupTransport() uc.SetupTransport(ctx)
} }
switch uc.IPStack { switch uc.IPStack {
case IpStackBoth, IpStackV4, IpStackV6: case IpStackBoth, IpStackV4, IpStackV6:
@@ -673,7 +678,7 @@ func (uc *UpstreamConfig) dohTransport(dnsType uint16) http.RoundTripper {
return uc.transport return uc.transport
} }
func (uc *UpstreamConfig) bootstrapIPForDNSType(dnsType uint16) string { func (uc *UpstreamConfig) bootstrapIPForDNSType(ctx context.Context, dnsType uint16) string {
switch uc.IPStack { switch uc.IPStack {
case IpStackBoth: case IpStackBoth:
return pick(uc.bootstrapIPs) return pick(uc.bootstrapIPs)
@@ -686,7 +691,7 @@ func (uc *UpstreamConfig) bootstrapIPForDNSType(dnsType uint16) string {
case dns.TypeA: case dns.TypeA:
return pick(uc.bootstrapIPs4) return pick(uc.bootstrapIPs4)
default: default:
if HasIPv6() { if HasIPv6(ctx) {
return pick(uc.bootstrapIPs6) return pick(uc.bootstrapIPs6)
} }
return pick(uc.bootstrapIPs4) return pick(uc.bootstrapIPs4)
@@ -695,7 +700,7 @@ func (uc *UpstreamConfig) bootstrapIPForDNSType(dnsType uint16) string {
return pick(uc.bootstrapIPs) return pick(uc.bootstrapIPs)
} }
func (uc *UpstreamConfig) netForDNSType(dnsType uint16) (string, string) { func (uc *UpstreamConfig) netForDNSType(ctx context.Context, dnsType uint16) (string, string) {
switch uc.IPStack { switch uc.IPStack {
case IpStackBoth: case IpStackBoth:
return "tcp-tls", "udp" return "tcp-tls", "udp"
@@ -708,7 +713,7 @@ func (uc *UpstreamConfig) netForDNSType(dnsType uint16) (string, string) {
case dns.TypeA: case dns.TypeA:
return "tcp4-tls", "udp4" return "tcp4-tls", "udp4"
default: default:
if HasIPv6() { if HasIPv6(ctx) {
return "tcp6-tls", "udp6" return "tcp6-tls", "udp6"
} }
return "tcp4-tls", "udp4" return "tcp4-tls", "udp4"
@@ -789,7 +794,7 @@ func (uc *UpstreamConfig) Context(ctx context.Context) (context.Context, context
} }
// FallbackToDirectIP changes ControlD upstream endpoint to use direct IP instead of domain. // FallbackToDirectIP changes ControlD upstream endpoint to use direct IP instead of domain.
func (uc *UpstreamConfig) FallbackToDirectIP() bool { func (uc *UpstreamConfig) FallbackToDirectIP(ctx context.Context) bool {
if !uc.IsControlD() { if !uc.IsControlD() {
return false return false
} }
@@ -808,7 +813,8 @@ func (uc *UpstreamConfig) FallbackToDirectIP() bool {
default: default:
return return
} }
ProxyLogger.Load().Warn().Msgf("using direct IP for %q: %s", uc.Endpoint, ip) logger := LoggerFromCtx(ctx)
logger.Warn().Msgf("using direct IP for %q: %s", uc.Endpoint, ip)
uc.u.Host = ip uc.u.Host = ip
done = true done = true
}) })
@@ -942,11 +948,12 @@ func pick(s []string) string {
} }
// upstreamUID generates an unique identifier for an upstream. // upstreamUID generates an unique identifier for an upstream.
func upstreamUID() string { func upstreamUID(ctx context.Context) string {
logger := LoggerFromCtx(ctx)
b := make([]byte, 4) b := make([]byte, 4)
for { for {
if _, err := crand.Read(b); err != nil { if _, err := crand.Read(b); err != nil {
ProxyLogger.Load().Warn().Err(err).Msg("could not generate uid for upstream, retrying...") logger.Warn().Err(err).Msg("could not generate uid for upstream, retrying...")
continue continue
} }
return hex.EncodeToString(b) return hex.EncodeToString(b)
+6 -5
View File
@@ -1,6 +1,7 @@
package ctrld package ctrld
import ( import (
"context"
"net/url" "net/url"
"testing" "testing"
@@ -36,10 +37,10 @@ func TestUpstreamConfig_SetupBootstrapIP(t *testing.T) {
t.Run(tc.name, func(t *testing.T) { t.Run(tc.name, func(t *testing.T) {
// Enable parallel tests once https://github.com/microsoft/wmi/issues/165 fixed. // Enable parallel tests once https://github.com/microsoft/wmi/issues/165 fixed.
// t.Parallel() // t.Parallel()
tc.uc.Init() tc.uc.Init(context.Background())
tc.uc.SetupBootstrapIP() tc.uc.SetupBootstrapIP(context.Background())
if len(tc.uc.bootstrapIPs) == 0 { if len(tc.uc.bootstrapIPs) == 0 {
t.Log(defaultNameservers()) t.Log(defaultNameservers(context.Background()))
t.Fatalf("could not bootstrap ip: %s", tc.uc.String()) t.Fatalf("could not bootstrap ip: %s", tc.uc.String())
} }
}) })
@@ -355,7 +356,7 @@ func TestUpstreamConfig_Init(t *testing.T) {
tc := tc tc := tc
t.Run(tc.name, func(t *testing.T) { t.Run(tc.name, func(t *testing.T) {
t.Parallel() t.Parallel()
tc.uc.Init() tc.uc.Init(context.Background())
tc.uc.uid = "" // we don't care about the uid. tc.uc.uid = "" // we don't care about the uid.
assert.Equal(t, tc.expected, tc.uc) assert.Equal(t, tc.expected, tc.uc)
}) })
@@ -497,7 +498,7 @@ func TestUpstreamConfig_IsDiscoverable(t *testing.T) {
tc := tc tc := tc
t.Run(tc.name, func(t *testing.T) { t.Run(tc.name, func(t *testing.T) {
t.Parallel() t.Parallel()
tc.uc.Init() tc.uc.Init(context.Background())
if got := tc.uc.IsDiscoverable(); got != tc.discoverable { if got := tc.uc.IsDiscoverable(); got != tc.discoverable {
t.Errorf("unexpected result, want: %v, got: %v", tc.discoverable, got) t.Errorf("unexpected result, want: %v, got: %v", tc.discoverable, got)
} }
+15 -14
View File
@@ -14,34 +14,35 @@ import (
"github.com/quic-go/quic-go/http3" "github.com/quic-go/quic-go/http3"
) )
func (uc *UpstreamConfig) setupDOH3Transport() { func (uc *UpstreamConfig) setupDOH3Transport(ctx context.Context) {
switch uc.IPStack { switch uc.IPStack {
case IpStackBoth, "": case IpStackBoth, "":
uc.http3RoundTripper = uc.newDOH3Transport(uc.bootstrapIPs) uc.http3RoundTripper = uc.newDOH3Transport(ctx, uc.bootstrapIPs)
case IpStackV4: case IpStackV4:
uc.http3RoundTripper = uc.newDOH3Transport(uc.bootstrapIPs4) uc.http3RoundTripper = uc.newDOH3Transport(ctx, uc.bootstrapIPs4)
case IpStackV6: case IpStackV6:
uc.http3RoundTripper = uc.newDOH3Transport(uc.bootstrapIPs6) uc.http3RoundTripper = uc.newDOH3Transport(ctx, uc.bootstrapIPs6)
case IpStackSplit: case IpStackSplit:
uc.http3RoundTripper4 = uc.newDOH3Transport(uc.bootstrapIPs4) uc.http3RoundTripper4 = uc.newDOH3Transport(ctx, uc.bootstrapIPs4)
if HasIPv6() { if HasIPv6(ctx) {
uc.http3RoundTripper6 = uc.newDOH3Transport(uc.bootstrapIPs6) uc.http3RoundTripper6 = uc.newDOH3Transport(ctx, uc.bootstrapIPs6)
} else { } else {
uc.http3RoundTripper6 = uc.http3RoundTripper4 uc.http3RoundTripper6 = uc.http3RoundTripper4
} }
uc.http3RoundTripper = uc.newDOH3Transport(uc.bootstrapIPs) uc.http3RoundTripper = uc.newDOH3Transport(ctx, uc.bootstrapIPs)
} }
} }
func (uc *UpstreamConfig) newDOH3Transport(addrs []string) http.RoundTripper { func (uc *UpstreamConfig) newDOH3Transport(ctx context.Context, addrs []string) http.RoundTripper {
rt := &http3.Transport{} rt := &http3.Transport{}
rt.TLSClientConfig = &tls.Config{RootCAs: uc.certPool} rt.TLSClientConfig = &tls.Config{RootCAs: uc.certPool}
logger := LoggerFromCtx(ctx)
rt.Dial = func(ctx context.Context, addr string, tlsCfg *tls.Config, cfg *quic.Config) (quic.EarlyConnection, error) { rt.Dial = func(ctx context.Context, addr string, tlsCfg *tls.Config, cfg *quic.Config) (quic.EarlyConnection, error) {
_, port, _ := net.SplitHostPort(addr) _, port, _ := net.SplitHostPort(addr)
// if we have a bootstrap ip set, use it to avoid DNS lookup // if we have a bootstrap ip set, use it to avoid DNS lookup
if uc.BootstrapIP != "" { if uc.BootstrapIP != "" {
addr = net.JoinHostPort(uc.BootstrapIP, port) addr = net.JoinHostPort(uc.BootstrapIP, port)
ProxyLogger.Load().Debug().Msgf("sending doh3 request to: %s", addr) logger.Debug().Msgf("sending doh3 request to: %s", addr)
udpConn, err := net.ListenUDP("udp", nil) udpConn, err := net.ListenUDP("udp", nil)
if err != nil { if err != nil {
return nil, err return nil, err
@@ -61,7 +62,7 @@ func (uc *UpstreamConfig) newDOH3Transport(addrs []string) http.RoundTripper {
if err != nil { if err != nil {
return nil, err return nil, err
} }
ProxyLogger.Load().Debug().Msgf("sending doh3 request to: %s", conn.RemoteAddr()) logger.Debug().Msgf("sending doh3 request to: %s", conn.RemoteAddr())
return conn, err return conn, err
} }
runtime.SetFinalizer(rt, func(rt *http3.Transport) { runtime.SetFinalizer(rt, func(rt *http3.Transport) {
@@ -70,12 +71,12 @@ func (uc *UpstreamConfig) newDOH3Transport(addrs []string) http.RoundTripper {
return rt return rt
} }
func (uc *UpstreamConfig) doh3Transport(dnsType uint16) http.RoundTripper { func (uc *UpstreamConfig) doh3Transport(ctx context.Context, dnsType uint16) http.RoundTripper {
uc.transportOnce.Do(func() { uc.transportOnce.Do(func() {
uc.SetupTransport() uc.SetupTransport(ctx)
}) })
if uc.rebootstrap.CompareAndSwap(true, false) { if uc.rebootstrap.CompareAndSwap(true, false) {
uc.SetupTransport() uc.SetupTransport(ctx)
} }
switch uc.IPStack { switch uc.IPStack {
case IpStackBoth, IpStackV4, IpStackV6: case IpStackBoth, IpStackV4, IpStackV6:
+7 -5
View File
@@ -105,19 +105,20 @@ func (r *dohResolver) Resolve(ctx context.Context, msg *dns.Msg) (*dns.Msg, erro
if len(msg.Question) > 0 { if len(msg.Question) > 0 {
dnsTyp = msg.Question[0].Qtype dnsTyp = msg.Question[0].Qtype
} }
c := http.Client{Transport: r.uc.dohTransport(dnsTyp)} c := http.Client{Transport: r.uc.dohTransport(ctx, dnsTyp)}
if r.isDoH3 { if r.isDoH3 {
transport := r.uc.doh3Transport(dnsTyp) transport := r.uc.doh3Transport(ctx, dnsTyp)
if transport == nil { if transport == nil {
return nil, errors.New("DoH3 is not supported") return nil, errors.New("DoH3 is not supported")
} }
c.Transport = transport c.Transport = transport
} }
resp, err := c.Do(req) resp, err := c.Do(req)
if err != nil && r.uc.FallbackToDirectIP() { if err != nil && r.uc.FallbackToDirectIP(ctx) {
retryCtx, cancel := r.uc.Context(context.WithoutCancel(ctx)) retryCtx, cancel := r.uc.Context(context.WithoutCancel(ctx))
defer cancel() defer cancel()
Log(ctx, ProxyLogger.Load().Warn().Err(err), "retrying request after fallback to direct ip") logger := LoggerFromCtx(ctx)
logger.Warn().Err(err).Msg("retrying request after fallback to direct ip")
resp, err = c.Do(req.Clone(retryCtx)) resp, err = c.Do(req.Clone(retryCtx))
} }
if err != nil { if err != nil {
@@ -163,7 +164,8 @@ func addHeader(ctx context.Context, req *http.Request, uc *UpstreamConfig) {
} }
} }
if printed { if printed {
Log(ctx, ProxyLogger.Load().Debug(), "sending request header: %v", dohHeader) logger := LoggerFromCtx(ctx)
logger.Debug().Msgf("sending request header: %v", dohHeader)
} }
dohHeader.Set("Content-Type", headerApplicationDNS) dohHeader.Set("Content-Type", headerApplicationDNS)
dohHeader.Set("Accept", headerApplicationDNS) dohHeader.Set("Accept", headerApplicationDNS)
+5 -4
View File
@@ -157,20 +157,21 @@ func Test_ClientCertificateVerificationError(t *testing.T) {
}, },
} }
ctx := context.Background()
for _, tc := range tests { for _, tc := range tests {
tc := tc tc := tc
t.Run(tc.name, func(t *testing.T) { t.Run(tc.name, func(t *testing.T) {
t.Parallel() t.Parallel()
tc.uc.Init() tc.uc.Init(ctx)
tc.uc.SetupBootstrapIP() tc.uc.SetupBootstrapIP(ctx)
r, err := NewResolver(tc.uc) r, err := NewResolver(ctx, tc.uc)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
msg := new(dns.Msg) msg := new(dns.Msg)
msg.SetQuestion("verify.controld.com.", dns.TypeA) msg.SetQuestion("verify.controld.com.", dns.TypeA)
msg.RecursionDesired = true msg.RecursionDesired = true
_, err = r.Resolve(context.Background(), msg) _, err = r.Resolve(ctx, msg)
// Verify the error contains the expected certificate information // Verify the error contains the expected certificate information
if err == nil { if err == nil {
t.Fatal("expected certificate verification error, got nil") t.Fatal("expected certificate verification error, got nil")
+1 -1
View File
@@ -26,7 +26,7 @@ func (r *doqResolver) Resolve(ctx context.Context, msg *dns.Msg) (*dns.Msg, erro
if msg != nil && len(msg.Question) > 0 { if msg != nil && len(msg.Question) > 0 {
dnsTyp = msg.Question[0].Qtype dnsTyp = msg.Question[0].Qtype
} }
ip = r.uc.bootstrapIPForDNSType(dnsTyp) ip = r.uc.bootstrapIPForDNSType(ctx, dnsTyp)
} }
tlsConfig.ServerName = r.uc.Domain tlsConfig.ServerName = r.uc.Domain
_, port, _ := net.SplitHostPort(endpoint) _, port, _ := net.SplitHostPort(endpoint)
+1 -1
View File
@@ -23,7 +23,7 @@ func (r *dotResolver) Resolve(ctx context.Context, msg *dns.Msg) (*dns.Msg, erro
if msg != nil && len(msg.Question) > 0 { if msg != nil && len(msg.Question) > 0 {
dnsTyp = msg.Question[0].Qtype dnsTyp = msg.Question[0].Qtype
} }
tcpNet, _ := r.uc.netForDNSType(dnsTyp) tcpNet, _ := r.uc.netForDNSType(ctx, dnsTyp)
dnsClient := &dns.Client{ dnsClient := &dns.Client{
Net: tcpNet, Net: tcpNet,
Dialer: dialer, Dialer: dialer,
+31 -23
View File
@@ -79,6 +79,7 @@ type Table struct {
initOnce sync.Once initOnce sync.Once
stopOnce sync.Once stopOnce sync.Once
refreshInterval int refreshInterval int
logger *ctrld.Logger
dhcp *dhcp dhcp *dhcp
merlin *merlinDiscover merlin *merlinDiscover
@@ -98,11 +99,14 @@ type Table struct {
ptrNameservers []string ptrNameservers []string
} }
func NewTable(cfg *ctrld.Config, selfIP, cdUID string, ns []string) *Table { func NewTable(cfg *ctrld.Config, selfIP, cdUID string, ns []string, logger *ctrld.Logger) *Table {
refreshInterval := cfg.Service.DiscoverRefreshInterval refreshInterval := cfg.Service.DiscoverRefreshInterval
if refreshInterval <= 0 { if refreshInterval <= 0 {
refreshInterval = 2 * 60 // 2 minutes refreshInterval = 2 * 60 // 2 minutes
} }
if logger == nil {
logger = ctrld.NopLogger
}
return &Table{ return &Table{
svcCfg: cfg.Service, svcCfg: cfg.Service,
quitCh: make(chan struct{}), quitCh: make(chan struct{}),
@@ -111,6 +115,7 @@ func NewTable(cfg *ctrld.Config, selfIP, cdUID string, ns []string) *Table {
cdUID: cdUID, cdUID: cdUID,
ptrNameservers: ns, ptrNameservers: ns,
refreshInterval: refreshInterval, refreshInterval: refreshInterval,
logger: logger,
} }
} }
@@ -179,7 +184,7 @@ func (t *Table) SetSelfIP(ip string) {
// initSelfDiscover initializes necessary client metadata for self query. // initSelfDiscover initializes necessary client metadata for self query.
func (t *Table) initSelfDiscover() { func (t *Table) initSelfDiscover() {
t.dhcp = &dhcp{selfIP: t.selfIP} t.dhcp = &dhcp{selfIP: t.selfIP, logger: t.logger}
t.dhcp.addSelf() t.dhcp.addSelf()
t.ipResolvers = append(t.ipResolvers, t.dhcp) t.ipResolvers = append(t.ipResolvers, t.dhcp)
t.macResolvers = append(t.macResolvers, t.dhcp) t.macResolvers = append(t.macResolvers, t.dhcp)
@@ -189,14 +194,14 @@ func (t *Table) initSelfDiscover() {
func (t *Table) init() { func (t *Table) init() {
// Custom client ID presents, use it as the only source. // Custom client ID presents, use it as the only source.
if _, clientID := controld.ParseRawUID(t.cdUID); clientID != "" { if _, clientID := controld.ParseRawUID(t.cdUID); clientID != "" {
ctrld.ProxyLogger.Load().Debug().Msg("start self discovery with custom client id") t.logger.Debug().Msg("start self discovery with custom client id")
t.initSelfDiscover() t.initSelfDiscover()
return return
} }
// If we are running on platforms that should only do self discover, use it as the only source, too. // If we are running on platforms that should only do self discover, use it as the only source, too.
if ctrld.SelfDiscover() { if ctrld.SelfDiscover() {
ctrld.ProxyLogger.Load().Debug().Msg("start self discovery on desktop platforms") t.logger.Debug().Msg("start self discovery on desktop platforms")
t.initSelfDiscover() t.initSelfDiscover()
return return
} }
@@ -208,7 +213,7 @@ func (t *Table) init() {
// - Merlin // - Merlin
// - Ubios // - Ubios
if t.discoverDHCP() || t.discoverARP() { if t.discoverDHCP() || t.discoverARP() {
t.merlin = &merlinDiscover{} t.merlin = &merlinDiscover{logger: t.logger}
t.ubios = &ubiosDiscover{} t.ubios = &ubiosDiscover{}
discovers := map[string]interface { discovers := map[string]interface {
refresher refresher
@@ -219,7 +224,7 @@ func (t *Table) init() {
} }
for platform, discover := range discovers { for platform, discover := range discovers {
if err := discover.refresh(); err != nil { if err := discover.refresh(); err != nil {
ctrld.ProxyLogger.Load().Warn().Err(err).Msgf("failed to init %s discover", platform) t.logger.Warn().Err(err).Msgf("failed to init %s discover", platform)
} }
t.hostnameResolvers = append(t.hostnameResolvers, discover) t.hostnameResolvers = append(t.hostnameResolvers, discover)
t.refreshers = append(t.refreshers, discover) t.refreshers = append(t.refreshers, discover)
@@ -227,10 +232,10 @@ func (t *Table) init() {
} }
// Hosts file mapping. // Hosts file mapping.
if t.discoverHosts() { if t.discoverHosts() {
t.hf = &hostsFile{} t.hf = &hostsFile{logger: t.logger}
ctrld.ProxyLogger.Load().Debug().Msg("start hosts file discovery") t.logger.Debug().Msg("start hosts file discovery")
if err := t.hf.init(); err != nil { if err := t.hf.init(); err != nil {
ctrld.ProxyLogger.Load().Error().Err(err).Msg("could not init hosts file discover") t.logger.Error().Err(err).Msg("could not init hosts file discover")
} else { } else {
t.hostnameResolvers = append(t.hostnameResolvers, t.hf) t.hostnameResolvers = append(t.hostnameResolvers, t.hf)
t.refreshers = append(t.refreshers, t.hf) t.refreshers = append(t.refreshers, t.hf)
@@ -239,10 +244,10 @@ func (t *Table) init() {
} }
// DHCP lease files. // DHCP lease files.
if t.discoverDHCP() { if t.discoverDHCP() {
t.dhcp = &dhcp{selfIP: t.selfIP} t.dhcp = &dhcp{selfIP: t.selfIP, logger: t.logger}
ctrld.ProxyLogger.Load().Debug().Msg("start dhcp discovery") t.logger.Debug().Msg("start dhcp discovery")
if err := t.dhcp.init(); err != nil { if err := t.dhcp.init(); err != nil {
ctrld.ProxyLogger.Load().Error().Err(err).Msg("could not init DHCP discover") t.logger.Error().Err(err).Msg("could not init DHCP discover")
} else { } else {
t.ipResolvers = append(t.ipResolvers, t.dhcp) t.ipResolvers = append(t.ipResolvers, t.dhcp)
t.macResolvers = append(t.macResolvers, t.dhcp) t.macResolvers = append(t.macResolvers, t.dhcp)
@@ -253,8 +258,8 @@ func (t *Table) init() {
// ARP/NDP table. // ARP/NDP table.
if t.discoverARP() { if t.discoverARP() {
t.arp = &arpDiscover{} t.arp = &arpDiscover{}
t.ndp = &ndpDiscover{} t.ndp = &ndpDiscover{logger: t.logger}
ctrld.ProxyLogger.Load().Debug().Msg("start arp discovery") t.logger.Debug().Msg("start arp discovery")
discovers := map[string]interface { discovers := map[string]interface {
refresher refresher
IpResolver IpResolver
@@ -266,7 +271,7 @@ func (t *Table) init() {
for protocol, discover := range discovers { for protocol, discover := range discovers {
if err := discover.refresh(); err != nil { if err := discover.refresh(); err != nil {
ctrld.ProxyLogger.Load().Error().Err(err).Msgf("could not init %s discover", protocol) t.logger.Error().Err(err).Msgf("could not init %s discover", protocol)
} else { } else {
t.ipResolvers = append(t.ipResolvers, discover) t.ipResolvers = append(t.ipResolvers, discover)
t.macResolvers = append(t.macResolvers, discover) t.macResolvers = append(t.macResolvers, discover)
@@ -283,7 +288,10 @@ func (t *Table) init() {
} }
// PTR lookup. // PTR lookup.
if t.discoverPTR() { if t.discoverPTR() {
t.ptr = &ptrDiscover{resolver: ctrld.NewPrivateResolver()} t.ptr = &ptrDiscover{
resolver: ctrld.NewPrivateResolver(context.Background()),
logger: t.logger,
}
if len(t.ptrNameservers) > 0 { if len(t.ptrNameservers) > 0 {
nss := make([]string, 0, len(t.ptrNameservers)) nss := make([]string, 0, len(t.ptrNameservers))
for _, ns := range t.ptrNameservers { for _, ns := range t.ptrNameservers {
@@ -295,18 +303,18 @@ func (t *Table) init() {
if _, portErr := strconv.Atoi(port); portErr == nil && port != "0" && net.ParseIP(host) != nil { if _, portErr := strconv.Atoi(port); portErr == nil && port != "0" && net.ParseIP(host) != nil {
nss = append(nss, net.JoinHostPort(host, port)) nss = append(nss, net.JoinHostPort(host, port))
} else { } else {
ctrld.ProxyLogger.Load().Warn().Msgf("ignoring invalid nameserver for ptr discover: %q", ns) t.logger.Warn().Msgf("ignoring invalid nameserver for ptr discover: %q", ns)
} }
} }
if len(nss) > 0 { if len(nss) > 0 {
t.ptr.resolver = ctrld.NewResolverWithNameserver(nss) t.ptr.resolver = ctrld.NewResolverWithNameserver(nss)
ctrld.ProxyLogger.Load().Debug().Msgf("using nameservers %v for ptr discovery", nss) t.logger.Debug().Msgf("using nameservers %v for ptr discovery", nss)
} }
} }
ctrld.ProxyLogger.Load().Debug().Msg("start ptr discovery") t.logger.Debug().Msg("start ptr discovery")
if err := t.ptr.refresh(); err != nil { if err := t.ptr.refresh(); err != nil {
ctrld.ProxyLogger.Load().Error().Err(err).Msg("could not init PTR discover") t.logger.Error().Err(err).Msg("could not init PTR discover")
} else { } else {
t.hostnameResolvers = append(t.hostnameResolvers, t.ptr) t.hostnameResolvers = append(t.hostnameResolvers, t.ptr)
t.refreshers = append(t.refreshers, t.ptr) t.refreshers = append(t.refreshers, t.ptr)
@@ -314,10 +322,10 @@ func (t *Table) init() {
} }
// mdns. // mdns.
if t.discoverMDNS() { if t.discoverMDNS() {
t.mdns = &mdns{} t.mdns = &mdns{logger: t.logger}
ctrld.ProxyLogger.Load().Debug().Msg("start mdns discovery") t.logger.Debug().Msg("start mdns discovery")
if err := t.mdns.init(t.quitCh); err != nil { if err := t.mdns.init(t.quitCh); err != nil {
ctrld.ProxyLogger.Load().Error().Err(err).Msg("could not init mDNS discover") t.logger.Error().Err(err).Msg("could not init mDNS discover")
} else { } else {
t.hostnameResolvers = append(t.hostnameResolvers, t.mdns) t.hostnameResolvers = append(t.hostnameResolvers, t.mdns)
} }
+5 -2
View File
@@ -2,6 +2,8 @@ package clientinfo
import ( import (
"testing" "testing"
"github.com/Control-D-Inc/ctrld"
) )
func Test_normalizeIP(t *testing.T) { func Test_normalizeIP(t *testing.T) {
@@ -28,8 +30,9 @@ func Test_normalizeIP(t *testing.T) {
func TestTable_LookupRFC1918IPv4(t *testing.T) { func TestTable_LookupRFC1918IPv4(t *testing.T) {
table := &Table{ table := &Table{
dhcp: &dhcp{}, dhcp: &dhcp{},
arp: &arpDiscover{}, arp: &arpDiscover{},
logger: ctrld.NopLogger,
} }
table.ipResolvers = append(table.ipResolvers, table.dhcp) table.ipResolvers = append(table.ipResolvers, table.dhcp)
+10 -10
View File
@@ -13,9 +13,8 @@ import (
"strings" "strings"
"sync" "sync"
"tailscale.com/net/netmon"
"github.com/fsnotify/fsnotify" "github.com/fsnotify/fsnotify"
"tailscale.com/net/netmon"
"tailscale.com/util/lineread" "tailscale.com/util/lineread"
"github.com/Control-D-Inc/ctrld" "github.com/Control-D-Inc/ctrld"
@@ -30,6 +29,7 @@ type dhcp struct {
watcher *fsnotify.Watcher watcher *fsnotify.Watcher
selfIP string selfIP string
logger *ctrld.Logger
} }
func (d *dhcp) init() error { func (d *dhcp) init() error {
@@ -52,7 +52,7 @@ func (d *dhcp) watchChanges() {
} }
if dir := router.LeaseFilesDir(); dir != "" { if dir := router.LeaseFilesDir(); dir != "" {
if err := d.watcher.Add(dir); err != nil { if err := d.watcher.Add(dir); err != nil {
ctrld.ProxyLogger.Load().Err(err).Str("dir", dir).Msg("could not watch lease dir") d.logger.Err(err).Str("dir", dir).Msg("could not watch lease dir")
} }
} }
for { for {
@@ -64,7 +64,7 @@ func (d *dhcp) watchChanges() {
if event.Has(fsnotify.Create) { if event.Has(fsnotify.Create) {
if format, ok := clientInfoFiles[event.Name]; ok { if format, ok := clientInfoFiles[event.Name]; ok {
if err := d.addLeaseFile(event.Name, format); err != nil { if err := d.addLeaseFile(event.Name, format); err != nil {
ctrld.ProxyLogger.Load().Err(err).Str("file", event.Name).Msg("could not add lease file") d.logger.Err(err).Str("file", event.Name).Msg("could not add lease file")
} }
} }
continue continue
@@ -72,14 +72,14 @@ func (d *dhcp) watchChanges() {
if event.Has(fsnotify.Write) || event.Has(fsnotify.Rename) || event.Has(fsnotify.Chmod) || event.Has(fsnotify.Remove) { if event.Has(fsnotify.Write) || event.Has(fsnotify.Rename) || event.Has(fsnotify.Chmod) || event.Has(fsnotify.Remove) {
format := clientInfoFiles[event.Name] format := clientInfoFiles[event.Name]
if err := d.readLeaseFile(event.Name, format); err != nil && !os.IsNotExist(err) { if err := d.readLeaseFile(event.Name, format); err != nil && !os.IsNotExist(err) {
ctrld.ProxyLogger.Load().Err(err).Str("file", event.Name).Msg("leases file changed but failed to update client info") d.logger.Err(err).Str("file", event.Name).Msg("leases file changed but failed to update client info")
} }
} }
case err, ok := <-d.watcher.Errors: case err, ok := <-d.watcher.Errors:
if !ok { if !ok {
return return
} }
ctrld.ProxyLogger.Load().Err(err).Msg("could not watch client info file") d.logger.Err(err).Msg("could not watch client info file")
} }
} }
@@ -222,7 +222,7 @@ func (d *dhcp) dnsmasqReadClientInfoReader(reader io.Reader) error {
} }
ip := normalizeIP(string(fields[2])) ip := normalizeIP(string(fields[2]))
if net.ParseIP(ip) == nil { if net.ParseIP(ip) == nil {
ctrld.ProxyLogger.Load().Warn().Msgf("invalid ip address entry: %q", ip) d.logger.Warn().Msgf("invalid ip address entry: %q", ip)
ip = "" ip = ""
} }
@@ -275,7 +275,7 @@ func (d *dhcp) iscDHCPReadClientInfoReader(reader io.Reader) error {
case "lease": case "lease":
ip = normalizeIP(strings.ToLower(fields[1])) ip = normalizeIP(strings.ToLower(fields[1]))
if net.ParseIP(ip) == nil { if net.ParseIP(ip) == nil {
ctrld.ProxyLogger.Load().Warn().Msgf("invalid ip address entry: %q", ip) d.logger.Warn().Msgf("invalid ip address entry: %q", ip)
ip = "" ip = ""
} }
case "hardware": case "hardware":
@@ -328,7 +328,7 @@ func (d *dhcp) keaDhcp4ReadClientInfoReader(r io.Reader) error {
} }
ip := normalizeIP(record[0]) ip := normalizeIP(record[0])
if net.ParseIP(ip) == nil { if net.ParseIP(ip) == nil {
ctrld.ProxyLogger.Load().Warn().Msgf("invalid ip address entry: %q", ip) d.logger.Warn().Msgf("invalid ip address entry: %q", ip)
ip = "" ip = ""
} }
@@ -350,7 +350,7 @@ func (d *dhcp) keaDhcp4ReadClientInfoReader(r io.Reader) error {
func (d *dhcp) addSelf() { func (d *dhcp) addSelf() {
hostname, err := os.Hostname() hostname, err := os.Hostname()
if err != nil { if err != nil {
ctrld.ProxyLogger.Load().Err(err).Msg("could not get hostname") d.logger.Err(err).Msg("could not get hostname")
return return
} }
hostname = normalizeHostname(hostname) hostname = normalizeHostname(hostname)
+4 -3
View File
@@ -27,6 +27,7 @@ type hostsFile struct {
watcher *fsnotify.Watcher watcher *fsnotify.Watcher
mu sync.Mutex mu sync.Mutex
m map[string][]string m map[string][]string
logger *ctrld.Logger
} }
// init performs initialization works, which is necessary before hostsFile can be fully operated. // init performs initialization works, which is necessary before hostsFile can be fully operated.
@@ -55,7 +56,7 @@ func (hf *hostsFile) refresh() error {
// override hosts file with host_entries.conf content if present. // override hosts file with host_entries.conf content if present.
hem, err := parseHostEntriesConf(hostEntriesConfPath) hem, err := parseHostEntriesConf(hostEntriesConfPath)
if err != nil && !os.IsNotExist(err) { if err != nil && !os.IsNotExist(err) {
ctrld.ProxyLogger.Load().Debug().Err(err).Msg("could not read host_entries.conf file") hf.logger.Debug().Err(err).Msg("could not read host_entries.conf file")
} }
for k, v := range hem { for k, v := range hem {
hf.m[k] = v hf.m[k] = v
@@ -77,14 +78,14 @@ func (hf *hostsFile) watchChanges() {
} }
if event.Has(fsnotify.Write) || event.Has(fsnotify.Rename) || event.Has(fsnotify.Chmod) || event.Has(fsnotify.Remove) { if event.Has(fsnotify.Write) || event.Has(fsnotify.Rename) || event.Has(fsnotify.Chmod) || event.Has(fsnotify.Remove) {
if err := hf.refresh(); err != nil && !os.IsNotExist(err) { if err := hf.refresh(); err != nil && !os.IsNotExist(err) {
ctrld.ProxyLogger.Load().Err(err).Msg("hosts file changed but failed to update client info") hf.logger.Err(err).Msg("hosts file changed but failed to update client info")
} }
} }
case err, ok := <-hf.watcher.Errors: case err, ok := <-hf.watcher.Errors:
if !ok { if !ok {
return return
} }
ctrld.ProxyLogger.Load().Err(err).Msg("could not watch client info file") hf.logger.Err(err).Msg("could not watch client info file")
} }
} }
+12 -11
View File
@@ -34,7 +34,8 @@ var (
) )
type mdns struct { type mdns struct {
name sync.Map // ip => hostname name sync.Map // ip => hostname
logger *ctrld.Logger
} }
func (m *mdns) LookupHostnameByIP(ip string) string { func (m *mdns) LookupHostnameByIP(ip string) string {
@@ -93,9 +94,9 @@ func (m *mdns) init(quitCh chan struct{}) error {
} }
// Check if IPv6 is available once and use the result for the rest of the function. // Check if IPv6 is available once and use the result for the rest of the function.
ctrld.ProxyLogger.Load().Debug().Msgf("checking for IPv6 availability in mdns init") m.logger.Debug().Msgf("checking for IPv6 availability in mdns init")
ipv6 := ctrldnet.IPv6Available(context.Background()) ipv6 := ctrldnet.IPv6Available(context.Background())
ctrld.ProxyLogger.Load().Debug().Msgf("IPv6 is %v in mdns init", ipv6) m.logger.Debug().Msgf("IPv6 is %v in mdns init", ipv6)
v4ConnList := make([]*net.UDPConn, 0, len(ifaces)) v4ConnList := make([]*net.UDPConn, 0, len(ifaces))
v6ConnList := make([]*net.UDPConn, 0, len(ifaces)) v6ConnList := make([]*net.UDPConn, 0, len(ifaces))
@@ -129,11 +130,11 @@ func (m *mdns) probeLoop(conns []*net.UDPConn, remoteAddr net.Addr, quitCh chan
for { for {
err := m.probe(conns, remoteAddr) err := m.probe(conns, remoteAddr)
if shouldStopProbing(err) { if shouldStopProbing(err) {
ctrld.ProxyLogger.Load().Warn().Msgf("stop probing %q: %v", remoteAddr, err) m.logger.Warn().Msgf("stop probing %q: %v", remoteAddr, err)
break break
} }
if err != nil { if err != nil {
ctrld.ProxyLogger.Load().Warn().Err(err).Msg("error while probing mdns") m.logger.Warn().Err(err).Msg("error while probing mdns")
bo.BackOff(context.Background(), errors.New("mdns probe backoff")) bo.BackOff(context.Background(), errors.New("mdns probe backoff"))
continue continue
} }
@@ -161,7 +162,7 @@ func (m *mdns) readLoop(conn *net.UDPConn) {
if errors.Is(err, net.ErrClosed) { if errors.Is(err, net.ErrClosed) {
return return
} }
ctrld.ProxyLogger.Load().Debug().Err(err).Msg("mdns readLoop error") m.logger.Debug().Err(err).Msg("mdns readLoop error")
return return
} }
@@ -184,11 +185,11 @@ func (m *mdns) readLoop(conn *net.UDPConn) {
if ip != "" && name != "" { if ip != "" && name != "" {
name = normalizeHostname(name) name = normalizeHostname(name)
if val, loaded := m.name.LoadOrStore(ip, name); !loaded { if val, loaded := m.name.LoadOrStore(ip, name); !loaded {
ctrld.ProxyLogger.Load().Debug().Msgf("found hostname: %q, ip: %q via mdns", name, ip) m.logger.Debug().Msgf("found hostname: %q, ip: %q via mdns", name, ip)
} else { } else {
old := val.(string) old := val.(string)
if old != name { if old != name {
ctrld.ProxyLogger.Load().Debug().Msgf("update hostname: %q, ip: %q, old: %q via mdns", name, ip, old) m.logger.Debug().Msgf("update hostname: %q, ip: %q, old: %q via mdns", name, ip, old)
m.name.Store(ip, name) m.name.Store(ip, name)
} }
} }
@@ -227,7 +228,7 @@ func (m *mdns) probe(conns []*net.UDPConn, remoteAddr net.Addr) error {
// getDataFromAvahiDaemonCache reads entries from avahi-daemon cache to update mdns data. // getDataFromAvahiDaemonCache reads entries from avahi-daemon cache to update mdns data.
func (m *mdns) getDataFromAvahiDaemonCache() { func (m *mdns) getDataFromAvahiDaemonCache() {
if _, err := exec.LookPath("avahi-browse"); err != nil { if _, err := exec.LookPath("avahi-browse"); err != nil {
ctrld.ProxyLogger.Load().Debug().Err(err).Msg("could not find avahi-browse binary, skipping.") m.logger.Debug().Err(err).Msg("could not find avahi-browse binary, skipping.")
return return
} }
// Run avahi-browse to discover services from cache: // Run avahi-browse to discover services from cache:
@@ -237,7 +238,7 @@ func (m *mdns) getDataFromAvahiDaemonCache() {
// - "-c" -> read from cache. // - "-c" -> read from cache.
out, err := exec.Command("avahi-browse", "-a", "-r", "-p", "-c").Output() out, err := exec.Command("avahi-browse", "-a", "-r", "-p", "-c").Output()
if err != nil { if err != nil {
ctrld.ProxyLogger.Load().Debug().Err(err).Msg("could not browse services from avahi cache") m.logger.Debug().Err(err).Msg("could not browse services from avahi cache")
return return
} }
m.storeDataFromAvahiBrowseOutput(bytes.NewReader(out)) m.storeDataFromAvahiBrowseOutput(bytes.NewReader(out))
@@ -257,7 +258,7 @@ func (m *mdns) storeDataFromAvahiBrowseOutput(r io.Reader) {
name := normalizeHostname(fields[6]) name := normalizeHostname(fields[6])
// Only using cache value if we don't have existed one. // Only using cache value if we don't have existed one.
if _, loaded := m.name.LoadOrStore(ip, name); !loaded { if _, loaded := m.name.LoadOrStore(ip, name); !loaded {
ctrld.ProxyLogger.Load().Debug().Msgf("found hostname: %q, ip: %q via avahi cache", name, ip) m.logger.Debug().Msgf("found hostname: %q, ip: %q via avahi cache", name, ip)
} }
} }
} }
+3 -1
View File
@@ -3,6 +3,8 @@ package clientinfo
import ( import (
"strings" "strings"
"testing" "testing"
"github.com/Control-D-Inc/ctrld"
) )
func Test_mdns_storeDataFromAvahiBrowseOutput(t *testing.T) { func Test_mdns_storeDataFromAvahiBrowseOutput(t *testing.T) {
@@ -11,7 +13,7 @@ func Test_mdns_storeDataFromAvahiBrowseOutput(t *testing.T) {
=;wlp0s20f3;IPv6;Foo\032\0402\041;_companion-link._tcp;local;Foo-2.local;192.168.1.123;64842;"rpBA=00:00:00:00:00:01" "rpHI=e6ae2cbbca0e" "rpAD=36566f4d850f" "rpVr=510.71.1" "rpHA=0ddc20fdddc8" "rpFl=0x30000" "rpHN=1d4a03afdefa" "rpMac=0" =;wlp0s20f3;IPv6;Foo\032\0402\041;_companion-link._tcp;local;Foo-2.local;192.168.1.123;64842;"rpBA=00:00:00:00:00:01" "rpHI=e6ae2cbbca0e" "rpAD=36566f4d850f" "rpVr=510.71.1" "rpHA=0ddc20fdddc8" "rpFl=0x30000" "rpHN=1d4a03afdefa" "rpMac=0"
=;wlp0s20f3;IPv4;Foo\032\0402\041;_companion-link._tcp;local;Foo-2.local;192.168.1.123;64842;"rpBA=00:00:00:00:00:01" "rpHI=e6ae2cbbca0e" "rpAD=36566f4d850f" "rpVr=510.71.1" "rpHA=0ddc20fdddc8" "rpFl=0x30000" "rpHN=1d4a03afdefa" "rpMac=0" =;wlp0s20f3;IPv4;Foo\032\0402\041;_companion-link._tcp;local;Foo-2.local;192.168.1.123;64842;"rpBA=00:00:00:00:00:01" "rpHI=e6ae2cbbca0e" "rpAD=36566f4d850f" "rpVr=510.71.1" "rpHA=0ddc20fdddc8" "rpFl=0x30000" "rpHN=1d4a03afdefa" "rpMac=0"
` `
m := &mdns{} m := &mdns{logger: ctrld.NopLogger}
m.storeDataFromAvahiBrowseOutput(strings.NewReader(content)) m.storeDataFromAvahiBrowseOutput(strings.NewReader(content))
ip := "192.168.1.123" ip := "192.168.1.123"
val, loaded := m.name.LoadOrStore(ip, "") val, loaded := m.name.LoadOrStore(ip, "")
+2 -1
View File
@@ -15,6 +15,7 @@ const merlinNvramCustomClientListKey = "custom_clientlist"
type merlinDiscover struct { type merlinDiscover struct {
hostname sync.Map // mac => hostname hostname sync.Map // mac => hostname
logger *ctrld.Logger
} }
func (m *merlinDiscover) refresh() error { func (m *merlinDiscover) refresh() error {
@@ -25,7 +26,7 @@ func (m *merlinDiscover) refresh() error {
if err != nil { if err != nil {
return err return err
} }
ctrld.ProxyLogger.Load().Debug().Msg("reading Merlin custom client list") m.logger.Debug().Msg("reading Merlin custom client list")
m.parseMerlinCustomClientList(out) m.parseMerlinCustomClientList(out)
return nil return nil
} }
+7 -6
View File
@@ -20,8 +20,9 @@ import (
// ndpDiscover provides client discovery functionality using NDP protocol. // ndpDiscover provides client discovery functionality using NDP protocol.
type ndpDiscover struct { type ndpDiscover struct {
mac sync.Map // ip => mac mac sync.Map // ip => mac
ip sync.Map // mac => ip ip sync.Map // mac => ip
logger *ctrld.Logger
} }
// refresh re-scans the NDP table. // refresh re-scans the NDP table.
@@ -97,7 +98,7 @@ func (nd *ndpDiscover) saveInfo(ip, mac string) {
func (nd *ndpDiscover) listen(ctx context.Context) { func (nd *ndpDiscover) listen(ctx context.Context) {
ifis, err := allInterfacesWithV6LinkLocal() ifis, err := allInterfacesWithV6LinkLocal()
if err != nil { if err != nil {
ctrld.ProxyLogger.Load().Debug().Err(err).Msg("failed to find valid ipv6 interfaces") nd.logger.Debug().Err(err).Msg("failed to find valid ipv6 interfaces")
return return
} }
for _, ifi := range ifis { for _, ifi := range ifis {
@@ -110,11 +111,11 @@ func (nd *ndpDiscover) listen(ctx context.Context) {
func (nd *ndpDiscover) listenOnInterface(ctx context.Context, ifi *net.Interface) { func (nd *ndpDiscover) listenOnInterface(ctx context.Context, ifi *net.Interface) {
c, ip, err := ndp.Listen(ifi, ndp.Unspecified) c, ip, err := ndp.Listen(ifi, ndp.Unspecified)
if err != nil { if err != nil {
ctrld.ProxyLogger.Load().Debug().Err(err).Msg("ndp listen failed") nd.logger.Debug().Err(err).Msg("ndp listen failed")
return return
} }
defer c.Close() defer c.Close()
ctrld.ProxyLogger.Load().Debug().Msgf("listening ndp on: %s", ip.String()) nd.logger.Debug().Msgf("listening ndp on: %s", ip.String())
for { for {
select { select {
case <-ctx.Done(): case <-ctx.Done():
@@ -128,7 +129,7 @@ func (nd *ndpDiscover) listenOnInterface(ctx context.Context, ifi *net.Interface
if errors.As(readErr, &opErr) && (opErr.Timeout() || opErr.Temporary()) { if errors.As(readErr, &opErr) && (opErr.Timeout() || opErr.Temporary()) {
continue continue
} }
ctrld.ProxyLogger.Load().Debug().Err(readErr).Msg("ndp read loop error") nd.logger.Debug().Err(readErr).Msg("ndp read loop error")
return return
} }
+4 -6
View File
@@ -5,15 +5,13 @@ import (
"github.com/vishvananda/netlink" "github.com/vishvananda/netlink"
"golang.org/x/sys/unix" "golang.org/x/sys/unix"
"github.com/Control-D-Inc/ctrld"
) )
// scan populates NDP table using information from system mappings. // scan populates NDP table using information from system mappings.
func (nd *ndpDiscover) scan() { func (nd *ndpDiscover) scan() {
neighs, err := netlink.NeighList(0, netlink.FAMILY_V6) neighs, err := netlink.NeighList(0, netlink.FAMILY_V6)
if err != nil { if err != nil {
ctrld.ProxyLogger.Load().Warn().Err(err).Msg("could not get neigh list") nd.logger.Warn().Err(err).Msg("could not get neigh list")
return return
} }
@@ -34,7 +32,7 @@ func (nd *ndpDiscover) subscribe(ctx context.Context) {
done := make(chan struct{}) done := make(chan struct{})
defer close(done) defer close(done)
if err := netlink.NeighSubscribe(ch, done); err != nil { if err := netlink.NeighSubscribe(ch, done); err != nil {
ctrld.ProxyLogger.Load().Err(err).Msg("could not perform neighbor subscribing") nd.logger.Err(err).Msg("could not perform neighbor subscribing")
return return
} }
for { for {
@@ -47,7 +45,7 @@ func (nd *ndpDiscover) subscribe(ctx context.Context) {
} }
ip := normalizeIP(nu.IP.String()) ip := normalizeIP(nu.IP.String())
if nu.Type == unix.RTM_DELNEIGH { if nu.Type == unix.RTM_DELNEIGH {
ctrld.ProxyLogger.Load().Debug().Msgf("removing NDP neighbor: %s", ip) nd.logger.Debug().Msgf("removing NDP neighbor: %s", ip)
nd.mac.Delete(ip) nd.mac.Delete(ip)
continue continue
} }
@@ -56,7 +54,7 @@ func (nd *ndpDiscover) subscribe(ctx context.Context) {
case netlink.NUD_REACHABLE: case netlink.NUD_REACHABLE:
nd.saveInfo(ip, mac) nd.saveInfo(ip, mac)
case netlink.NUD_FAILED: case netlink.NUD_FAILED:
ctrld.ProxyLogger.Load().Debug().Msgf("removing NDP neighbor with failed state: %s", ip) nd.logger.Debug().Msgf("removing NDP neighbor with failed state: %s", ip)
nd.mac.Delete(ip) nd.mac.Delete(ip)
} }
} }
+2 -4
View File
@@ -7,8 +7,6 @@ import (
"context" "context"
"os/exec" "os/exec"
"runtime" "runtime"
"github.com/Control-D-Inc/ctrld"
) )
// scan populates NDP table using information from system mappings. // scan populates NDP table using information from system mappings.
@@ -17,14 +15,14 @@ func (nd *ndpDiscover) scan() {
case "windows": case "windows":
data, err := exec.Command("netsh", "interface", "ipv6", "show", "neighbors").Output() data, err := exec.Command("netsh", "interface", "ipv6", "show", "neighbors").Output()
if err != nil { if err != nil {
ctrld.ProxyLogger.Load().Warn().Err(err).Msg("could not query ndp table") nd.logger.Warn().Err(err).Msg("could not query ndp table")
return return
} }
nd.scanWindows(bytes.NewReader(data)) nd.scanWindows(bytes.NewReader(data))
default: default:
data, err := exec.Command("ndp", "-an").Output() data, err := exec.Command("ndp", "-an").Output()
if err != nil { if err != nil {
ctrld.ProxyLogger.Load().Warn().Err(err).Msg("could not query ndp table") nd.logger.Warn().Err(err).Msg("could not query ndp table")
return return
} }
nd.scanUnix(bytes.NewReader(data)) nd.scanUnix(bytes.NewReader(data))
+3 -2
View File
@@ -17,6 +17,7 @@ type ptrDiscover struct {
hostname sync.Map // ip => hostname hostname sync.Map // ip => hostname
resolver ctrld.Resolver resolver ctrld.Resolver
serverDown atomic.Bool serverDown atomic.Bool
logger *ctrld.Logger
} }
func (p *ptrDiscover) refresh() error { func (p *ptrDiscover) refresh() error {
@@ -73,14 +74,14 @@ func (p *ptrDiscover) lookupHostname(ip string) string {
msg := new(dns.Msg) msg := new(dns.Msg)
addr, err := dns.ReverseAddr(ip) addr, err := dns.ReverseAddr(ip)
if err != nil { if err != nil {
ctrld.ProxyLogger.Load().Info().Str("discovery", "ptr").Err(err).Msg("invalid ip address") p.logger.Info().Str("discovery", "ptr").Err(err).Msg("invalid ip address")
return "" return ""
} }
msg.SetQuestion(addr, dns.TypePTR) msg.SetQuestion(addr, dns.TypePTR)
ans, err := p.resolver.Resolve(ctx, msg) ans, err := p.resolver.Resolve(ctx, msg)
if err != nil { if err != nil {
if p.serverDown.CompareAndSwap(false, true) { if p.serverDown.CompareAndSwap(false, true) {
ctrld.ProxyLogger.Load().Info().Str("discovery", "ptr").Err(err).Msg("could not perform PTR lookup") p.logger.Info().Str("discovery", "ptr").Err(err).Msg("could not perform PTR lookup")
go p.checkServer() go p.checkServer()
} }
return "" return ""
+21 -18
View File
@@ -88,18 +88,18 @@ type LogsRequest struct {
} }
// FetchResolverConfig fetch Control D config for given uid. // FetchResolverConfig fetch Control D config for given uid.
func FetchResolverConfig(rawUID, version string, cdDev bool) (*ResolverConfig, error) { func FetchResolverConfig(ctx context.Context, rawUID, version string, cdDev bool) (*ResolverConfig, error) {
uid, clientID := ParseRawUID(rawUID) uid, clientID := ParseRawUID(rawUID)
req := utilityRequest{UID: uid} req := utilityRequest{UID: uid}
if clientID != "" { if clientID != "" {
req.ClientID = clientID req.ClientID = clientID
} }
body, _ := json.Marshal(req) body, _ := json.Marshal(req)
return postUtilityAPI(version, cdDev, false, bytes.NewReader(body)) return postUtilityAPI(ctx, version, cdDev, false, bytes.NewReader(body))
} }
// FetchResolverUID fetch resolver uid from provision token. // FetchResolverUID fetch resolver uid from provision token.
func FetchResolverUID(req *UtilityOrgRequest, version string, cdDev bool) (*ResolverConfig, error) { func FetchResolverUID(ctx context.Context, req *UtilityOrgRequest, version string, cdDev bool) (*ResolverConfig, error) {
if req == nil { if req == nil {
return nil, errors.New("invalid request") return nil, errors.New("invalid request")
} }
@@ -108,21 +108,21 @@ func FetchResolverUID(req *UtilityOrgRequest, version string, cdDev bool) (*Reso
hostname, _ = os.Hostname() hostname, _ = os.Hostname()
} }
body, _ := json.Marshal(UtilityOrgRequest{ProvToken: req.ProvToken, Hostname: hostname}) body, _ := json.Marshal(UtilityOrgRequest{ProvToken: req.ProvToken, Hostname: hostname})
return postUtilityAPI(version, cdDev, false, bytes.NewReader(body)) return postUtilityAPI(ctx, version, cdDev, false, bytes.NewReader(body))
} }
// UpdateCustomLastFailed calls API to mark custom config is bad. // UpdateCustomLastFailed calls API to mark custom config is bad.
func UpdateCustomLastFailed(rawUID, version string, cdDev, lastUpdatedFailed bool) (*ResolverConfig, error) { func UpdateCustomLastFailed(ctx context.Context, rawUID, version string, cdDev, lastUpdatedFailed bool) (*ResolverConfig, error) {
uid, clientID := ParseRawUID(rawUID) uid, clientID := ParseRawUID(rawUID)
req := utilityRequest{UID: uid} req := utilityRequest{UID: uid}
if clientID != "" { if clientID != "" {
req.ClientID = clientID req.ClientID = clientID
} }
body, _ := json.Marshal(req) body, _ := json.Marshal(req)
return postUtilityAPI(version, cdDev, true, bytes.NewReader(body)) return postUtilityAPI(ctx, version, cdDev, true, bytes.NewReader(body))
} }
func postUtilityAPI(version string, cdDev, lastUpdatedFailed bool, body io.Reader) (*ResolverConfig, error) { func postUtilityAPI(ctx context.Context, version string, cdDev, lastUpdatedFailed bool, body io.Reader) (*ResolverConfig, error) {
apiUrl := resolverDataURLCom apiUrl := resolverDataURLCom
if cdDev { if cdDev {
apiUrl = resolverDataURLDev apiUrl = resolverDataURLDev
@@ -139,12 +139,12 @@ func postUtilityAPI(version string, cdDev, lastUpdatedFailed bool, body io.Reade
} }
req.URL.RawQuery = q.Encode() req.URL.RawQuery = q.Encode()
req.Header.Add("Content-Type", "application/json") req.Header.Add("Content-Type", "application/json")
transport := apiTransport(cdDev) transport := apiTransport(ctx, cdDev)
client := &http.Client{ client := &http.Client{
Timeout: defaultTimeout, Timeout: defaultTimeout,
Transport: transport, Transport: transport,
} }
resp, err := doWithFallback(client, req, apiServerIP(cdDev)) resp, err := doWithFallback(ctx, client, req, apiServerIP(cdDev))
if err != nil { if err != nil {
return nil, fmt.Errorf("postUtilityAPI client.Do: %w", err) return nil, fmt.Errorf("postUtilityAPI client.Do: %w", err)
} }
@@ -166,7 +166,7 @@ func postUtilityAPI(version string, cdDev, lastUpdatedFailed bool, body io.Reade
} }
// SendLogs sends runtime log to ControlD API. // SendLogs sends runtime log to ControlD API.
func SendLogs(lr *LogsRequest, cdDev bool) error { func SendLogs(ctx context.Context, lr *LogsRequest, cdDev bool) error {
defer lr.Data.Close() defer lr.Data.Close()
apiUrl := logURLCom apiUrl := logURLCom
if cdDev { if cdDev {
@@ -180,12 +180,12 @@ func SendLogs(lr *LogsRequest, cdDev bool) error {
q.Set("uid", lr.UID) q.Set("uid", lr.UID)
req.URL.RawQuery = q.Encode() req.URL.RawQuery = q.Encode()
req.Header.Add("Content-Type", "application/x-www-form-urlencoded") req.Header.Add("Content-Type", "application/x-www-form-urlencoded")
transport := apiTransport(cdDev) transport := apiTransport(ctx, cdDev)
client := &http.Client{ client := &http.Client{
Timeout: sendLogTimeout, Timeout: sendLogTimeout,
Transport: transport, Transport: transport,
} }
resp, err := doWithFallback(client, req, apiServerIP(cdDev)) resp, err := doWithFallback(ctx, client, req, apiServerIP(cdDev))
if err != nil { if err != nil {
return fmt.Errorf("SendLogs client.Do: %w", err) return fmt.Errorf("SendLogs client.Do: %w", err)
} }
@@ -213,7 +213,7 @@ func ParseRawUID(rawUID string) (string, string) {
} }
// apiTransport returns an HTTP transport for connecting to ControlD API endpoint. // apiTransport returns an HTTP transport for connecting to ControlD API endpoint.
func apiTransport(cdDev bool) *http.Transport { func apiTransport(loggerCtx context.Context, cdDev bool) *http.Transport {
transport := http.DefaultTransport.(*http.Transport).Clone() transport := http.DefaultTransport.(*http.Transport).Clone()
transport.DialContext = func(ctx context.Context, network, addr string) (net.Conn, error) { transport.DialContext = func(ctx context.Context, network, addr string) (net.Conn, error) {
apiDomain := apiDomainCom apiDomain := apiDomainCom
@@ -227,9 +227,10 @@ func apiTransport(cdDev bool) *http.Transport {
apiIPs = []string{apiDomainDevIPv4} apiIPs = []string{apiDomainDevIPv4}
} }
ips := ctrld.LookupIP(apiDomain) ips := ctrld.LookupIP(loggerCtx, apiDomain)
if len(ips) == 0 { if len(ips) == 0 {
ctrld.ProxyLogger.Load().Warn().Msgf("No IPs found for %s, use direct IPs: %v", apiDomain, apiIPs) logger := ctrld.LoggerFromCtx(loggerCtx)
logger.Warn().Msgf("No IPs found for %s, use direct IPs: %v", apiDomain, apiIPs)
ips = apiIPs ips = apiIPs
} }
@@ -245,7 +246,8 @@ func apiTransport(cdDev bool) *http.Transport {
dial := func(ctx context.Context, network string, addrs []string) (net.Conn, error) { dial := func(ctx context.Context, network string, addrs []string) (net.Conn, error) {
d := &ctrldnet.ParallelDialer{} d := &ctrldnet.ParallelDialer{}
return d.DialContext(ctx, network, addrs, ctrld.ProxyLogger.Load()) logger := ctrld.LoggerFromCtx(loggerCtx)
return d.DialContext(ctx, network, addrs, logger.Logger)
} }
_, port, _ := net.SplitHostPort(addr) _, port, _ := net.SplitHostPort(addr)
@@ -283,10 +285,11 @@ func addrsFromPort(ips []string, port string) []string {
return addrs return addrs
} }
func doWithFallback(client *http.Client, req *http.Request, apiIp string) (*http.Response, error) { func doWithFallback(ctx context.Context, client *http.Client, req *http.Request, apiIp string) (*http.Response, error) {
resp, err := client.Do(req) resp, err := client.Do(req)
if err != nil { if err != nil {
ctrld.ProxyLogger.Load().Warn().Err(err).Msgf("failed to send request, fallback to direct IP: %s", apiIp) logger := ctrld.LoggerFromCtx(ctx)
logger.Warn().Err(err).Msgf("failed to send request, fallback to direct IP: %s", apiIp)
ipReq := req.Clone(req.Context()) ipReq := req.Clone(req.Context())
ipReq.Host = apiIp ipReq.Host = apiIp
ipReq.URL.Host = apiIp ipReq.URL.Host = apiIp
+26 -8
View File
@@ -3,19 +3,37 @@ package ctrld
import ( import (
"context" "context"
"fmt" "fmt"
"io"
"sync/atomic"
"github.com/rs/zerolog" "github.com/rs/zerolog"
) )
// ProxyLog emits the log record for proxy operations. // LoggerCtxKey is the context.Context key for a logger.
// The caller should set it only once. type LoggerCtxKey struct{}
// DEPRECATED: use ProxyLogger instead.
var ProxyLog = zerolog.New(io.Discard)
// ProxyLogger emits the log record for proxy operations. // LoggerCtx returns a context.Context with LoggerCtxKey set.
var ProxyLogger atomic.Pointer[zerolog.Logger] func LoggerCtx(ctx context.Context, l *Logger) context.Context {
return context.WithValue(ctx, LoggerCtxKey{}, l)
}
// A Logger provides fast, leveled, structured logging.
type Logger struct {
*zerolog.Logger
}
var noOpZeroLogger = zerolog.Nop()
// NopLogger returns a logger which all operation are no-op.
var NopLogger = &Logger{&noOpZeroLogger}
// LoggerFromCtx returns the logger associated with given ctx.
//
// If there's no logger, a no-op logger will be returned.
func LoggerFromCtx(ctx context.Context) *Logger {
if logger, ok := ctx.Value(LoggerCtxKey{}).(*Logger); ok && logger != nil {
return logger
}
return NopLogger
}
// ReqIdCtxKey is the context.Context key for a request id. // ReqIdCtxKey is the context.Context key for a request id.
type ReqIdCtxKey struct{} type ReqIdCtxKey struct{}
+5 -3
View File
@@ -1,9 +1,11 @@
package ctrld package ctrld
type dnsFn func() []string import "context"
type dnsFn func(ctx context.Context) []string
// nameservers returns DNS nameservers from system settings. // nameservers returns DNS nameservers from system settings.
func nameservers() []string { func nameservers(ctx context.Context) []string {
var dns []string var dns []string
seen := make(map[string]bool) seen := make(map[string]bool)
ch := make(chan []string) ch := make(chan []string)
@@ -11,7 +13,7 @@ func nameservers() []string {
for _, fn := range fns { for _, fn := range fns {
go func(fn dnsFn) { go func(fn dnsFn) {
ch <- fn() ch <- fn(ctx)
}(fn) }(fn)
} }
for range fns { for range fns {
+2 -1
View File
@@ -3,6 +3,7 @@
package ctrld package ctrld
import ( import (
"context"
"net" "net"
"syscall" "syscall"
@@ -13,7 +14,7 @@ func dnsFns() []dnsFn {
return []dnsFn{dnsFromResolvConf, dnsFromRIB} return []dnsFn{dnsFromResolvConf, dnsFromRIB}
} }
func dnsFromRIB() []string { func dnsFromRIB(_ context.Context) []string {
var dns []string var dns []string
rib, err := route.FetchRIB(syscall.AF_UNSPEC, route.RIBTypeRoute, 0) rib, err := route.FetchRIB(syscall.AF_UNSPEC, route.RIBTypeRoute, 0)
if err != nil { if err != nil {
+4 -4
View File
@@ -22,8 +22,8 @@ func dnsFns() []dnsFn {
return []dnsFn{dnsFromResolvConf, getDNSFromScutil, getAllDHCPNameservers} return []dnsFn{dnsFromResolvConf, getDNSFromScutil, getAllDHCPNameservers}
} }
func getDNSFromScutil() []string { func getDNSFromScutil(ctx context.Context) []string {
logger := *ProxyLogger.Load() logger := LoggerFromCtx(ctx)
const ( const (
maxRetries = 10 maxRetries = 10
@@ -109,8 +109,8 @@ func getDHCPNameservers(iface string) ([]string, error) {
return nameservers, nil return nameservers, nil
} }
func getAllDHCPNameservers() []string { func getAllDHCPNameservers(ctx context.Context) []string {
logger := *ProxyLogger.Load() logger := LoggerFromCtx(ctx)
interfaces, err := net.Interfaces() interfaces, err := net.Interfaces()
if err != nil { if err != nil {
+4 -3
View File
@@ -3,6 +3,7 @@ package ctrld
import ( import (
"bufio" "bufio"
"bytes" "bytes"
"context"
"encoding/hex" "encoding/hex"
"net" "net"
"os" "os"
@@ -20,7 +21,7 @@ func dnsFns() []dnsFn {
return []dnsFn{dnsFromResolvConf, dns4, dns6, dnsFromSystemdResolver} return []dnsFn{dnsFromResolvConf, dns4, dns6, dnsFromSystemdResolver}
} }
func dns4() []string { func dns4(_ context.Context) []string {
f, err := os.Open(v4RouteFile) f, err := os.Open(v4RouteFile)
if err != nil { if err != nil {
return nil return nil
@@ -60,7 +61,7 @@ func dns4() []string {
return dns return dns
} }
func dns6() []string { func dns6(_ context.Context) []string {
f, err := os.Open(v6RouteFile) f, err := os.Open(v6RouteFile)
if err != nil { if err != nil {
return nil return nil
@@ -94,7 +95,7 @@ func dns6() []string {
return dns return dns
} }
func dnsFromSystemdResolver() []string { func dnsFromSystemdResolver(_ context.Context) []string {
c, err := resolvconffile.ParseFile("/run/systemd/resolve/resolv.conf") c, err := resolvconffile.ParseFile("/run/systemd/resolve/resolv.conf")
if err != nil { if err != nil {
return nil return nil
+5 -2
View File
@@ -1,9 +1,12 @@
package ctrld package ctrld
import "testing" import (
"context"
"testing"
)
func TestNameservers(t *testing.T) { func TestNameservers(t *testing.T) {
ns := nameservers() ns := nameservers(context.Background())
if len(ns) == 0 { if len(ns) == 0 {
t.Fatal("failed to get nameservers") t.Fatal("failed to get nameservers")
} }
+2 -1
View File
@@ -3,6 +3,7 @@
package ctrld package ctrld
import ( import (
"context"
"net" "net"
"slices" "slices"
"time" "time"
@@ -20,7 +21,7 @@ func currentNameserversFromResolvconf() []string {
// dnsFromResolvConf reads usable nameservers from /etc/resolv.conf file. // dnsFromResolvConf reads usable nameservers from /etc/resolv.conf file.
// A nameserver is usable if it's not one of current machine's IP addresses // A nameserver is usable if it's not one of current machine's IP addresses
// and loopback IP addresses. // and loopback IP addresses.
func dnsFromResolvConf() []string { func dnsFromResolvConf(_ context.Context) []string {
const ( const (
maxRetries = 10 maxRetries = 10
retryInterval = 100 * time.Millisecond retryInterval = 100 * time.Millisecond
+55 -93
View File
@@ -55,28 +55,25 @@ func dnsFns() []dnsFn {
return []dnsFn{dnsFromAdapter} return []dnsFn{dnsFromAdapter}
} }
func dnsFromAdapter() []string { func dnsFromAdapter(ctx context.Context) []string {
ctx, cancel := context.WithTimeout(context.Background(), defaultDNSAdapterTimeout) ctx, cancel := context.WithTimeout(context.Background(), defaultDNSAdapterTimeout)
defer cancel() defer cancel()
var ns []string var ns []string
var err error var err error
logger := *ProxyLogger.Load() logger := LoggerFromCtx(ctx)
for i := 0; i < maxDNSAdapterRetries; i++ { for i := 0; i < maxDNSAdapterRetries; i++ {
if ctx.Err() != nil { if ctx.Err() != nil {
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("dnsFromAdapter lookup cancelled or timed out, attempt %d", i)
"dnsFromAdapter lookup cancelled or timed out, attempt %d", i)
return nil return nil
} }
ns, err = getDNSServers(ctx) ns, err = getDNSServers(ctx)
if err == nil && len(ns) >= minDNSServers { if err == nil && len(ns) >= minDNSServers {
if i > 0 { if i > 0 {
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("Successfully got DNS servers after %d attempts, found %d servers", i+1, len(ns))
"Successfully got DNS servers after %d attempts, found %d servers",
i+1, len(ns))
} }
return ns return ns
} }
@@ -88,11 +85,9 @@ func dnsFromAdapter() []string {
} }
if err != nil { if err != nil {
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("Failed to get DNS servers, attempt %d: %v", i+1, err)
"Failed to get DNS servers, attempt %d: %v", i+1, err)
} else { } else {
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("Got insufficient DNS servers, retrying, found %d servers", len(ns))
"Got insufficient DNS servers, retrying, found %d servers", len(ns))
} }
select { select {
@@ -102,14 +97,12 @@ func dnsFromAdapter() []string {
} }
} }
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("Failed to get sufficient DNS servers after all attempts, max_retries=%d", maxDNSAdapterRetries)
"Failed to get sufficient DNS servers after all attempts, max_retries=%d", maxDNSAdapterRetries)
return ns return ns
} }
func getDNSServers(ctx context.Context) ([]string, error) { func getDNSServers(ctx context.Context) ([]string, error) {
logger := *ProxyLogger.Load()
// Check context before making the call // Check context before making the call
if ctx.Err() != nil { if ctx.Err() != nil {
return nil, ctx.Err() return nil, ctx.Err()
@@ -124,17 +117,16 @@ func getDNSServers(ctx context.Context) ([]string, error) {
return nil, fmt.Errorf("getting adapters: %w", err) return nil, fmt.Errorf("getting adapters: %w", err)
} }
Log(context.Background(), logger.Debug(), logger := LoggerFromCtx(ctx)
"Found network adapters, count=%d", len(aas)) logger.Debug().Msgf("Found network adapters, count=%d", len(aas))
// Try to get domain controller info if domain-joined // Try to get domain controller info if domain-joined
var dcServers []string var dcServers []string
isDomain := checkDomainJoined() isDomain := checkDomainJoined(ctx)
if isDomain { if isDomain {
domainName, err := getLocalADDomain() domainName, err := getLocalADDomain()
if err != nil { if err != nil {
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("Failed to get local AD domain: %v", err)
"Failed to get local AD domain: %v", err)
} else { } else {
// Load netapi32.dll // Load netapi32.dll
netapi32 := windows.NewLazySystemDLL("netapi32.dll") netapi32 := windows.NewLazySystemDLL("netapi32.dll")
@@ -145,11 +137,9 @@ func getDNSServers(ctx context.Context) ([]string, error) {
domainUTF16, err := windows.UTF16PtrFromString(domainName) domainUTF16, err := windows.UTF16PtrFromString(domainName)
if err != nil { if err != nil {
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("Failed to convert domain name to UTF16: %v", err)
"Failed to convert domain name to UTF16: %v", err)
} else { } else {
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("Attempting to get DC for domain: %s with flags: 0x%x", domainName, flags)
"Attempting to get DC for domain: %s with flags: 0x%x", domainName, flags)
// Call DsGetDcNameW with domain name // Call DsGetDcNameW with domain name
ret, _, err := dsDcName.Call( ret, _, err := dsDcName.Call(
@@ -163,20 +153,15 @@ func getDNSServers(ctx context.Context) ([]string, error) {
if ret != 0 { if ret != 0 {
switch ret { switch ret {
case 1355: // ERROR_NO_SUCH_DOMAIN case 1355: // ERROR_NO_SUCH_DOMAIN
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("Domain not found: %s (%d)", domainName, ret)
"Domain not found: %s (%d)", domainName, ret)
case 1311: // ERROR_NO_LOGON_SERVERS case 1311: // ERROR_NO_LOGON_SERVERS
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("No logon servers available for domain: %s (%d)", domainName, ret)
"No logon servers available for domain: %s (%d)", domainName, ret)
case 1004: // ERROR_DC_NOT_FOUND case 1004: // ERROR_DC_NOT_FOUND
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("Domain controller not found for domain: %s (%d)", domainName, ret)
"Domain controller not found for domain: %s (%d)", domainName, ret)
case 1722: // RPC_S_SERVER_UNAVAILABLE case 1722: // RPC_S_SERVER_UNAVAILABLE
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("RPC server unavailable for domain: %s (%d)", domainName, ret)
"RPC server unavailable for domain: %s (%d)", domainName, ret)
default: default:
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("Failed to get domain controller info for domain %s: %d, %v", domainName, ret, err)
"Failed to get domain controller info for domain %s: %d, %v", domainName, ret, err)
} }
} else if info != nil { } else if info != nil {
defer windows.NetApiBufferFree((*byte)(unsafe.Pointer(info))) defer windows.NetApiBufferFree((*byte)(unsafe.Pointer(info)))
@@ -184,17 +169,13 @@ func getDNSServers(ctx context.Context) ([]string, error) {
if info.DomainControllerAddress != nil { if info.DomainControllerAddress != nil {
dcAddr := windows.UTF16PtrToString(info.DomainControllerAddress) dcAddr := windows.UTF16PtrToString(info.DomainControllerAddress)
dcAddr = strings.TrimPrefix(dcAddr, "\\\\") dcAddr = strings.TrimPrefix(dcAddr, "\\\\")
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("Found domain controller address: %s", dcAddr)
"Found domain controller address: %s", dcAddr)
if ip := net.ParseIP(dcAddr); ip != nil { if ip := net.ParseIP(dcAddr); ip != nil {
dcServers = append(dcServers, ip.String()) dcServers = append(dcServers, ip.String())
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("Added domain controller DNS servers: %v", dcServers)
"Added domain controller DNS servers: %v", dcServers)
} }
} else { } else {
Log(context.Background(), logger.Debug(), logger.Debug().Msg("No domain controller address found")
"No domain controller address found")
} }
} }
} }
@@ -209,31 +190,27 @@ func getDNSServers(ctx context.Context) ([]string, error) {
// Collect all local IPs // Collect all local IPs
for _, aa := range aas { for _, aa := range aas {
if aa.OperStatus != winipcfg.IfOperStatusUp { if aa.OperStatus != winipcfg.IfOperStatusUp {
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("Skipping adapter %s - not up, status: %d", aa.FriendlyName(), aa.OperStatus)
"Skipping adapter %s - not up, status: %d", aa.FriendlyName(), aa.OperStatus)
continue continue
} }
// Skip if software loopback or other non-physical types // Skip if software loopback or other non-physical types
// This is to avoid the "Loopback Pseudo-Interface 1" issue we see on windows // This is to avoid the "Loopback Pseudo-Interface 1" issue we see on windows
if aa.IfType == winipcfg.IfTypeSoftwareLoopback { if aa.IfType == winipcfg.IfTypeSoftwareLoopback {
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("Skipping %s (software loopback)", aa.FriendlyName())
"Skipping %s (software loopback)", aa.FriendlyName())
continue continue
} }
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("Processing adapter %s", aa.FriendlyName())
"Processing adapter %s", aa.FriendlyName())
for a := aa.FirstUnicastAddress; a != nil; a = a.Next { for a := aa.FirstUnicastAddress; a != nil; a = a.Next {
ip := a.Address.IP().String() ip := a.Address.IP().String()
addressMap[ip] = struct{}{} addressMap[ip] = struct{}{}
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("Added local IP %s from adapter %s", ip, aa.FriendlyName())
"Added local IP %s from adapter %s", ip, aa.FriendlyName())
} }
} }
validInterfacesMap := validInterfaces() validInterfacesMap := validInterfaces(ctx)
// Collect DNS servers // Collect DNS servers
for _, aa := range aas { for _, aa := range aas {
@@ -244,23 +221,20 @@ func getDNSServers(ctx context.Context) ([]string, error) {
// Skip if software loopback or other non-physical types // Skip if software loopback or other non-physical types
// This is to avoid the "Loopback Pseudo-Interface 1" issue we see on windows // This is to avoid the "Loopback Pseudo-Interface 1" issue we see on windows
if aa.IfType == winipcfg.IfTypeSoftwareLoopback { if aa.IfType == winipcfg.IfTypeSoftwareLoopback {
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("Skipping %s (software loopback)", aa.FriendlyName())
"Skipping %s (software loopback)", aa.FriendlyName())
continue continue
} }
// if not in the validInterfacesMap, skip // if not in the validInterfacesMap, skip
if _, ok := validInterfacesMap[aa.FriendlyName()]; !ok { if _, ok := validInterfacesMap[aa.FriendlyName()]; !ok {
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("Skipping %s (not in validInterfacesMap)", aa.FriendlyName())
"Skipping %s (not in validInterfacesMap)", aa.FriendlyName())
continue continue
} }
for dns := aa.FirstDNSServerAddress; dns != nil; dns = dns.Next { for dns := aa.FirstDNSServerAddress; dns != nil; dns = dns.Next {
ip := dns.Address.IP() ip := dns.Address.IP()
if ip == nil { if ip == nil {
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("Skipping nil IP from adapter %s", aa.FriendlyName())
"Skipping nil IP from adapter %s", aa.FriendlyName())
continue continue
} }
@@ -293,28 +267,23 @@ func getDNSServers(ctx context.Context) ([]string, error) {
if !seen[dcServer] { if !seen[dcServer] {
seen[dcServer] = true seen[dcServer] = true
ns = append(ns, dcServer) ns = append(ns, dcServer)
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("Added additional domain controller DNS server: %s", dcServer)
"Added additional domain controller DNS server: %s", dcServer)
} }
} }
// if we have static DNS servers saved for the current default route, we should add them to the list // if we have static DNS servers saved for the current default route, we should add them to the list
drIfaceName, err := netmon.DefaultRouteInterface() drIfaceName, err := netmon.DefaultRouteInterface()
if err != nil { if err != nil {
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("Failed to get default route interface: %v", err)
"Failed to get default route interface: %v", err)
} else { } else {
drIface, err := net.InterfaceByName(drIfaceName) drIface, err := net.InterfaceByName(drIfaceName)
if err != nil { if err != nil {
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("Failed to get interface by name %s: %v", drIfaceName, err)
"Failed to get interface by name %s: %v", drIfaceName, err)
} else { } else {
staticNs, file := SavedStaticNameserversAndPath(drIface) staticNs, file := SavedStaticNameserversAndPath(drIface)
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("static dns servers from %s: %v", file, staticNs)
"static dns servers from %s: %v", file, staticNs)
if len(staticNs) > 0 { if len(staticNs) > 0 {
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("Adding static DNS servers from %s: %v", drIfaceName, staticNs)
"Adding static DNS servers from %s: %v", drIfaceName, staticNs)
ns = append(ns, staticNs...) ns = append(ns, staticNs...)
} }
} }
@@ -324,9 +293,7 @@ func getDNSServers(ctx context.Context) ([]string, error) {
return nil, fmt.Errorf("no valid DNS servers found") return nil, fmt.Errorf("no valid DNS servers found")
} }
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("DNS server discovery completed, count=%d, servers=%v (including %d DC servers)", len(ns), ns, len(dcServers))
"DNS server discovery completed, count=%d, servers=%v (including %d DC servers)",
len(ns), ns, len(dcServers))
return ns, nil return ns, nil
} }
@@ -337,33 +304,35 @@ func currentNameserversFromResolvconf() []string {
// checkDomainJoined checks if the machine is joined to an Active Directory domain // checkDomainJoined checks if the machine is joined to an Active Directory domain
// Returns whether it's domain joined and the domain name if available // Returns whether it's domain joined and the domain name if available
func checkDomainJoined() bool { func checkDomainJoined(ctx context.Context) bool {
logger := *ProxyLogger.Load() logger := LoggerFromCtx(ctx)
var domain *uint16 var domain *uint16
var status uint32 var status uint32
err := windows.NetGetJoinInformation(nil, &domain, &status) err := windows.NetGetJoinInformation(nil, &domain, &status)
if err != nil { if err != nil {
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("Failed to get domain join status: %v", err)
"Failed to get domain join status: %v", err)
return false return false
} }
defer windows.NetApiBufferFree((*byte)(unsafe.Pointer(domain))) defer windows.NetApiBufferFree((*byte)(unsafe.Pointer(domain)))
domainName := windows.UTF16PtrToString(domain) domainName := windows.UTF16PtrToString(domain)
Log(context.Background(), logger.Debug(), logger.Debug().Msgf(
"Domain join status: domain=%s status=%d (Unknown=0, Workgroup=1, Domain=2, CloudDomain=3)", "Domain join status: domain=%s status=%d (Unknown=0, Workgroup=1, Domain=2, CloudDomain=3)",
domainName, status) domainName,
status,
)
// Consider domain or cloud domain as domain-joined // Consider domain or cloud domain as domain-joined
isDomain := status == NetSetupDomain || status == NetSetupCloudDomain isDomain := status == NetSetupDomain || status == NetSetupCloudDomain
Log(context.Background(), logger.Debug(), logger.Debug().Msgf(
"Is domain joined? status=%d, traditional=%v, cloud=%v, result=%v", "Is domain joined? status=%d, traditional=%v, cloud=%v, result=%v",
status, status,
status == NetSetupDomain, status == NetSetupDomain,
status == NetSetupCloudDomain, status == NetSetupCloudDomain,
isDomain) isDomain,
)
return isDomain return isDomain
} }
@@ -411,12 +380,12 @@ func getLocalADDomain() (string, error) {
// validInterfaces returns a list of all physical interfaces. // validInterfaces returns a list of all physical interfaces.
// this is a duplicate of what is in net_windows.go, we should // this is a duplicate of what is in net_windows.go, we should
// clean this up so there is only one version // clean this up so there is only one version
func validInterfaces() map[string]struct{} { func validInterfaces(ctx context.Context) map[string]struct{} {
log.SetOutput(io.Discard) log.SetOutput(io.Discard)
defer log.SetOutput(os.Stderr) defer log.SetOutput(os.Stderr)
//load the logger //load the logger
logger := *ProxyLogger.Load() logger := LoggerFromCtx(ctx)
whost := host.NewWmiLocalHost() whost := host.NewWmiLocalHost()
q := query.NewWmiQuery("MSFT_NetAdapter") q := query.NewWmiQuery("MSFT_NetAdapter")
@@ -425,23 +394,20 @@ func validInterfaces() map[string]struct{} {
defer instances.Close() defer instances.Close()
} }
if err != nil { if err != nil {
Log(context.Background(), logger.Warn(), logger.Warn().Msgf("failed to get wmi network adapter: %v", err)
"failed to get wmi network adapter: %v", err)
return nil return nil
} }
var adapters []string var adapters []string
for _, i := range instances { for _, i := range instances {
adapter, err := netadapter.NewNetworkAdapter(i) adapter, err := netadapter.NewNetworkAdapter(i)
if err != nil { if err != nil {
Log(context.Background(), logger.Warn(), logger.Warn().Msgf("failed to get network adapter: %v", err)
"failed to get network adapter: %v", err)
continue continue
} }
name, err := adapter.GetPropertyName() name, err := adapter.GetPropertyName()
if err != nil { if err != nil {
Log(context.Background(), logger.Warn(), logger.Warn().Msgf("failed to get interface name: %v", err)
"failed to get interface name: %v", err)
continue continue
} }
@@ -451,13 +417,11 @@ func validInterfaces() map[string]struct{} {
// if this is a physical adapter or FALSE if this is not a physical adapter." // if this is a physical adapter or FALSE if this is not a physical adapter."
physical, err := adapter.GetPropertyConnectorPresent() physical, err := adapter.GetPropertyConnectorPresent()
if err != nil { if err != nil {
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("failed to get network adapter connector present property: %v", err)
"failed to get network adapter connector present property: %v", err)
continue continue
} }
if !physical { if !physical {
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("skipping non-physical adapter: %s", name)
"skipping non-physical adapter: %s", name)
continue continue
} }
@@ -465,13 +429,11 @@ func validInterfaces() map[string]struct{} {
// because some interfaces are not physical but have a connector. // because some interfaces are not physical but have a connector.
hardware, err := adapter.GetPropertyHardwareInterface() hardware, err := adapter.GetPropertyHardwareInterface()
if err != nil { if err != nil {
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("failed to get network adapter hardware interface property: %v", err)
"failed to get network adapter hardware interface property: %v", err)
continue continue
} }
if !hardware { if !hardware {
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("skipping non-hardware interface: %s", name)
"skipping non-hardware interface: %s", name)
continue continue
} }
+10 -8
View File
@@ -17,26 +17,27 @@ var (
) )
// HasIPv6 reports whether the current network stack has IPv6 available. // HasIPv6 reports whether the current network stack has IPv6 available.
func HasIPv6() bool { func HasIPv6(ctx context.Context) bool {
hasIPv6Once.Do(func() { hasIPv6Once.Do(func() {
ProxyLogger.Load().Debug().Msg("checking for IPv6 availability once") logger := LoggerFromCtx(ctx)
logger.Debug().Msg("checking for IPv6 availability once")
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
defer cancel() defer cancel()
val := ctrldnet.IPv6Available(ctx) val := ctrldnet.IPv6Available(ctx)
ipv6Available.Store(val) ipv6Available.Store(val)
ProxyLogger.Load().Debug().Msgf("ipv6 availability: %v", val) logger.Debug().Msgf("ipv6 availability: %v", val)
mon, err := netmon.New(func(format string, args ...any) {}) mon, err := netmon.New(func(format string, args ...any) {})
if err != nil { if err != nil {
ProxyLogger.Load().Debug().Err(err).Msg("failed to monitor IPv6 state") logger.Debug().Err(err).Msg("failed to monitor IPv6 state")
return return
} }
mon.RegisterChangeCallback(func(delta *netmon.ChangeDelta) { mon.RegisterChangeCallback(func(delta *netmon.ChangeDelta) {
old := ipv6Available.Load() old := ipv6Available.Load()
cur := delta.Monitor.InterfaceState().HaveV6 cur := delta.Monitor.InterfaceState().HaveV6
if old != cur { if old != cur {
ProxyLogger.Load().Warn().Msgf("ipv6 availability changed, old: %v, new: %v", old, cur) logger.Warn().Msgf("ipv6 availability changed, old: %v, new: %v", old, cur)
} else { } else {
ProxyLogger.Load().Debug().Msg("ipv6 availability does not changed") logger.Debug().Msg("ipv6 availability does not changed")
} }
ipv6Available.Store(cur) ipv6Available.Store(cur)
}) })
@@ -46,8 +47,9 @@ func HasIPv6() bool {
} }
// DisableIPv6 marks IPv6 as unavailable if enabled. // DisableIPv6 marks IPv6 as unavailable if enabled.
func DisableIPv6() { func DisableIPv6(ctx context.Context) {
if ipv6Available.CompareAndSwap(true, false) { if ipv6Available.CompareAndSwap(true, false) {
ProxyLogger.Load().Debug().Msg("turned off IPv6 availability") logger := LoggerFromCtx(ctx)
logger.Debug().Msg("turned off IPv6 availability")
} }
} }
+58 -77
View File
@@ -4,7 +4,6 @@ import (
"context" "context"
"errors" "errors"
"fmt" "fmt"
"io"
"net" "net"
"net/netip" "net/netip"
"runtime" "runtime"
@@ -15,7 +14,6 @@ import (
"time" "time"
"github.com/miekg/dns" "github.com/miekg/dns"
"github.com/rs/zerolog"
"golang.org/x/sync/singleflight" "golang.org/x/sync/singleflight"
"tailscale.com/net/netmon" "tailscale.com/net/netmon"
"tailscale.com/net/tsaddr" "tailscale.com/net/tsaddr"
@@ -50,10 +48,6 @@ var controldPublicDnsWithPort = net.JoinHostPort(controldPublicDns, "53")
var localResolver Resolver var localResolver Resolver
func init() { func init() {
// Initializing ProxyLogger here, so other places don't have to do nil check.
l := zerolog.New(io.Discard)
ProxyLogger.Store(&l)
localResolver = newLocalResolver() localResolver = newLocalResolver()
} }
@@ -81,8 +75,8 @@ func LanQueryCtx(ctx context.Context) context.Context {
} }
// defaultNameservers is like nameservers with each element formed "ip:53". // defaultNameservers is like nameservers with each element formed "ip:53".
func defaultNameservers() []string { func defaultNameservers(ctx context.Context) []string {
ns := nameservers() ns := nameservers(ctx)
nss := make([]string, len(ns)) nss := make([]string, len(ns))
for i := range ns { for i := range ns {
nss[i] = net.JoinHostPort(ns[i], "53") nss[i] = net.JoinHostPort(ns[i], "53")
@@ -91,42 +85,36 @@ func defaultNameservers() []string {
} }
// availableNameservers returns list of current available DNS servers of the system. // availableNameservers returns list of current available DNS servers of the system.
func availableNameservers() []string { func availableNameservers(ctx context.Context) []string {
var nss []string var nss []string
// Ignore local addresses to prevent loop. // Ignore local addresses to prevent loop.
regularIPs, loopbackIPs, _ := netmon.LocalAddresses() regularIPs, loopbackIPs, _ := netmon.LocalAddresses()
machineIPsMap := make(map[string]struct{}, len(regularIPs)) machineIPsMap := make(map[string]struct{}, len(regularIPs))
//load the logger // Load the logger.
logger := *ProxyLogger.Load() logger := LoggerFromCtx(ctx)
logger.Debug().Msgf("Got local addresses - regular IPs: %v, loopback IPs: %v", regularIPs, loopbackIPs)
Log(context.Background(), logger.Debug(),
"Got local addresses - regular IPs: %v, loopback IPs: %v", regularIPs, loopbackIPs)
for _, v := range slices.Concat(regularIPs, loopbackIPs) { for _, v := range slices.Concat(regularIPs, loopbackIPs) {
ipStr := v.String() ipStr := v.String()
machineIPsMap[ipStr] = struct{}{} machineIPsMap[ipStr] = struct{}{}
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("Added local IP to OS resolverexclusion map: %s", ipStr)
"Added local IP to OS resolverexclusion map: %s", ipStr)
} }
systemNameservers := nameservers() systemNameservers := nameservers(ctx)
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("Got system nameservers: %v", systemNameservers)
"Got system nameservers: %v", systemNameservers)
for _, ns := range systemNameservers { for _, ns := range systemNameservers {
if _, ok := machineIPsMap[ns]; ok { if _, ok := machineIPsMap[ns]; ok {
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("Skipping local nameserver: %s", ns)
"Skipping local nameserver: %s", ns)
continue continue
} }
nss = append(nss, ns) nss = append(nss, ns)
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("Added non-local nameserver: %s", ns)
"Added non-local nameserver: %s", ns)
} }
Log(context.Background(), logger.Debug(), logger.Debug().Msgf("Final available nameservers: %v", nss)
"Final available nameservers: %v", nss)
return nss return nss
} }
@@ -135,8 +123,8 @@ func availableNameservers() []string {
// //
// It's the caller's responsibility to ensure the system DNS is in a clean state before // It's the caller's responsibility to ensure the system DNS is in a clean state before
// calling this function. // calling this function.
func InitializeOsResolver(guardAgainstNoNameservers bool) []string { func InitializeOsResolver(ctx context.Context, guardAgainstNoNameservers bool) []string {
nameservers := availableNameservers() nameservers := availableNameservers(ctx)
// if no nameservers, return empty slice so we dont remove all nameservers // if no nameservers, return empty slice so we dont remove all nameservers
if len(nameservers) == 0 && guardAgainstNoNameservers { if len(nameservers) == 0 && guardAgainstNoNameservers {
return []string{} return []string{}
@@ -188,7 +176,7 @@ type Resolver interface {
var errUnknownResolver = errors.New("unknown resolver") var errUnknownResolver = errors.New("unknown resolver")
// NewResolver creates a Resolver based on the given upstream config. // NewResolver creates a Resolver based on the given upstream config.
func NewResolver(uc *UpstreamConfig) (Resolver, error) { func NewResolver(ctx context.Context, uc *UpstreamConfig) (Resolver, error) {
typ := uc.Type typ := uc.Type
switch typ { switch typ {
case ResolverTypeDOH, ResolverTypeDOH3: case ResolverTypeDOH, ResolverTypeDOH3:
@@ -200,15 +188,16 @@ func NewResolver(uc *UpstreamConfig) (Resolver, error) {
case ResolverTypeOS: case ResolverTypeOS:
resolverMutex.Lock() resolverMutex.Lock()
if or == nil { if or == nil {
ProxyLogger.Load().Debug().Msgf("Initialize new OS resolver") logger := LoggerFromCtx(ctx)
or = newResolverWithNameserver(defaultNameservers()) logger.Debug().Msgf("Initialize new OS resolver")
or = newResolverWithNameserver(defaultNameservers(ctx))
} }
resolverMutex.Unlock() resolverMutex.Unlock()
return or, nil return or, nil
case ResolverTypeLegacy: case ResolverTypeLegacy:
return &legacyResolver{uc: uc}, nil return &legacyResolver{uc: uc}, nil
case ResolverTypePrivate: case ResolverTypePrivate:
return NewPrivateResolver(), nil return NewPrivateResolver(ctx), nil
case ResolverTypeLocal: case ResolverTypeLocal:
return localResolver, nil return localResolver, nil
} }
@@ -235,14 +224,16 @@ type publicResponse struct {
} }
// SetDefaultLocalIPv4 updates the stored local IPv4. // SetDefaultLocalIPv4 updates the stored local IPv4.
func SetDefaultLocalIPv4(ip net.IP) { func SetDefaultLocalIPv4(ctx context.Context, ip net.IP) {
Log(context.Background(), ProxyLogger.Load().Debug(), "SetDefaultLocalIPv4: %s", ip) logger := LoggerFromCtx(ctx)
logger.Debug().Msgf("SetDefaultLocalIPv4: %s", ip)
defaultLocalIPv4.Store(ip) defaultLocalIPv4.Store(ip)
} }
// SetDefaultLocalIPv6 updates the stored local IPv6. // SetDefaultLocalIPv6 updates the stored local IPv6.
func SetDefaultLocalIPv6(ip net.IP) { func SetDefaultLocalIPv6(ctx context.Context, ip net.IP) {
Log(context.Background(), ProxyLogger.Load().Debug(), "SetDefaultLocalIPv6: %s", ip) logger := LoggerFromCtx(ctx)
logger.Debug().Msgf("SetDefaultLocalIPv6: %s", ip)
defaultLocalIPv6.Store(ip) defaultLocalIPv6.Store(ip)
} }
@@ -300,10 +291,11 @@ func (o *osResolver) Resolve(ctx context.Context, msg *dns.Msg) (*dns.Msg, error
// Unique key for the singleflight group. // Unique key for the singleflight group.
key := fmt.Sprintf("%s:%d:", domain, qtype) key := fmt.Sprintf("%s:%d:", domain, qtype)
logger := LoggerFromCtx(ctx)
// Checking the cache first. // Checking the cache first.
if val, ok := o.cache.Load(key); ok { if val, ok := o.cache.Load(key); ok {
if val, ok := val.(*dns.Msg); ok { if val, ok := val.(*dns.Msg); ok {
Log(ctx, ProxyLogger.Load().Debug(), "hit hot cached result: %s - %s", domain, dns.TypeToString[qtype]) Log(ctx, logger.Debug(), "hit hot cached result: %s - %s", domain, dns.TypeToString[qtype])
res := val.Copy() res := val.Copy()
SetCacheReply(res, msg, val.Rcode) SetCacheReply(res, msg, val.Rcode)
return res, nil return res, nil
@@ -338,7 +330,7 @@ func (o *osResolver) Resolve(ctx context.Context, msg *dns.Msg) (*dns.Msg, error
res := sharedMsg.Copy() res := sharedMsg.Copy()
SetCacheReply(res, msg, sharedMsg.Rcode) SetCacheReply(res, msg, sharedMsg.Rcode)
if shared { if shared {
Log(ctx, ProxyLogger.Load().Debug(), "shared result: %s - %s", domain, dns.TypeToString[qtype]) Log(ctx, logger.Debug(), "shared result: %s - %s", domain, dns.TypeToString[qtype])
} }
return res, nil return res, nil
@@ -368,7 +360,8 @@ func (o *osResolver) resolve(ctx context.Context, msg *dns.Msg) (*dns.Msg, error
if msg != nil && len(msg.Question) > 0 { if msg != nil && len(msg.Question) > 0 {
question = msg.Question[0].Name question = msg.Question[0].Name
} }
Log(ctx, ProxyLogger.Load().Debug(), "os resolver query for %s with nameservers: %v public: %v", question, nss, publicServers) logger := LoggerFromCtx(ctx)
Log(ctx, logger.Debug(), "os resolver query for %s with nameservers: %v public: %v", question, nss, publicServers)
// New check: If no resolvers are available, return an error. // New check: If no resolvers are available, return an error.
if numServers == 0 { if numServers == 0 {
@@ -417,7 +410,7 @@ func (o *osResolver) resolve(ctx context.Context, msg *dns.Msg) (*dns.Msg, error
// If splitting fails, fallback to the original server string // If splitting fails, fallback to the original server string
host = server host = server
} }
Log(ctx, ProxyLogger.Load().Debug(), "got answer from nameserver: %s", host) Log(ctx, logger.Debug(), "got answer from nameserver: %s", host)
} }
// try local nameservers // try local nameservers
@@ -444,7 +437,7 @@ func (o *osResolver) resolve(ctx context.Context, msg *dns.Msg) (*dns.Msg, error
switch { switch {
case res.lan: case res.lan:
// Always prefer LAN responses immediately // Always prefer LAN responses immediately
Log(ctx, ProxyLogger.Load().Debug(), "using LAN answer from: %s", res.server) Log(ctx, logger.Debug(), "using LAN answer from: %s", res.server)
cancel() cancel()
logAnswer(res.server) logAnswer(res.server)
return res.answer, nil return res.answer, nil
@@ -454,7 +447,7 @@ func (o *osResolver) resolve(ctx context.Context, msg *dns.Msg) (*dns.Msg, error
// if there are no LAN nameservers, we should not wait // if there are no LAN nameservers, we should not wait
// just use the first response // just use the first response
if len(nss) == 0 { if len(nss) == 0 {
Log(ctx, ProxyLogger.Load().Debug(), "using public answer from: %s", res.server) Log(ctx, logger.Debug(), "using public answer from: %s", res.server)
cancel() cancel()
logAnswer(res.server) logAnswer(res.server)
return res.answer, nil return res.answer, nil
@@ -465,12 +458,12 @@ func (o *osResolver) resolve(ctx context.Context, msg *dns.Msg) (*dns.Msg, error
}) })
} }
case res.answer != nil: case res.answer != nil:
Log(ctx, ProxyLogger.Load().Debug(), "got non-success answer from: %s with code: %d", Log(ctx, logger.Debug(), "got non-success answer from: %s with code: %d",
res.server, res.answer.Rcode) res.server, res.answer.Rcode)
// When there are no LAN nameservers, we should not wait // When there are no LAN nameservers, we should not wait
// for other nameservers to respond. // for other nameservers to respond.
if len(nss) == 0 { if len(nss) == 0 {
Log(ctx, ProxyLogger.Load().Debug(), "no lan nameservers using public non success answer") Log(ctx, logger.Debug(), "no lan nameservers using public non success answer")
cancel() cancel()
logAnswer(res.server) logAnswer(res.server)
return res.answer, nil return res.answer, nil
@@ -483,17 +476,17 @@ func (o *osResolver) resolve(ctx context.Context, msg *dns.Msg) (*dns.Msg, error
if len(publicResponses) > 0 { if len(publicResponses) > 0 {
resp := publicResponses[0] resp := publicResponses[0]
Log(ctx, ProxyLogger.Load().Debug(), "using public answer from: %s", resp.server) Log(ctx, logger.Debug(), "using public answer from: %s", resp.server)
logAnswer(resp.server) logAnswer(resp.server)
return resp.answer, nil return resp.answer, nil
} }
if controldSuccessAnswer != nil { if controldSuccessAnswer != nil {
Log(ctx, ProxyLogger.Load().Debug(), "using ControlD answer from: %s", controldPublicDnsWithPort) Log(ctx, logger.Debug(), "using ControlD answer from: %s", controldPublicDnsWithPort)
logAnswer(controldPublicDnsWithPort) logAnswer(controldPublicDnsWithPort)
return controldSuccessAnswer, nil return controldSuccessAnswer, nil
} }
if nonSuccessAnswer != nil { if nonSuccessAnswer != nil {
Log(ctx, ProxyLogger.Load().Debug(), "using non-success answer from: %s", nonSuccessServer) Log(ctx, logger.Debug(), "using non-success answer from: %s", nonSuccessServer)
logAnswer(nonSuccessServer) logAnswer(nonSuccessServer)
return nonSuccessAnswer, nil return nonSuccessAnswer, nil
} }
@@ -515,7 +508,7 @@ func (r *legacyResolver) Resolve(ctx context.Context, msg *dns.Msg) (*dns.Msg, e
if msg != nil && len(msg.Question) > 0 { if msg != nil && len(msg.Question) > 0 {
dnsTyp = msg.Question[0].Qtype dnsTyp = msg.Question[0].Qtype
} }
_, udpNet := r.uc.netForDNSType(dnsTyp) _, udpNet := r.uc.netForDNSType(ctx, dnsTyp)
dnsClient := &dns.Client{ dnsClient := &dns.Client{
Net: udpNet, Net: udpNet,
Dialer: dialer, Dialer: dialer,
@@ -541,39 +534,43 @@ func (d dummyResolver) Resolve(ctx context.Context, msg *dns.Msg) (*dns.Msg, err
// LookupIP looks up domain using current system nameservers settings. // LookupIP looks up domain using current system nameservers settings.
// It returns a slice of that host's IPv4 and IPv6 addresses. // It returns a slice of that host's IPv4 and IPv6 addresses.
func LookupIP(domain string) []string { func LookupIP(ctx context.Context, domain string) []string {
nss := initDefaultOsResolver() nss := initDefaultOsResolver(ctx)
return lookupIP(domain, -1, nss) return lookupIP(ctx, domain, -1, nss)
} }
// initDefaultOsResolver initializes the default OS resolver with system's default nameservers if it hasn't been initialized yet. // initDefaultOsResolver initializes the default OS resolver with system's default nameservers if it hasn't been initialized yet.
// It returns the combined list of LAN and public nameservers currently held by the resolver. // It returns the combined list of LAN and public nameservers currently held by the resolver.
func initDefaultOsResolver() []string { func initDefaultOsResolver(ctx context.Context) []string {
logger := LoggerFromCtx(ctx)
resolverMutex.Lock() resolverMutex.Lock()
defer resolverMutex.Unlock() defer resolverMutex.Unlock()
if or == nil { if or == nil {
ProxyLogger.Load().Debug().Msgf("Initialize new OS resolver with default nameservers") logger.Debug().Msgf("Initialize new OS resolver with default nameservers")
or = newResolverWithNameserver(defaultNameservers()) or = newResolverWithNameserver(defaultNameservers(ctx))
} }
nss := *or.lanServers.Load() nss := *or.lanServers.Load()
nss = append(nss, *or.publicServers.Load()...) nss = append(nss, *or.publicServers.Load()...)
return nss return nss
} }
// lookupIP looks up domain with given timeout and bootstrapDNS. // lookupIP looks up domain with given timeout and bootstrapDNS.
// If the timeout is negative, default timeout 2000 ms will be used. // If the timeout is negative, default timeout 2000 ms will be used.
// It returns nil if bootstrapDNS is nil or empty. // It returns nil if bootstrapDNS is nil or empty.
func lookupIP(domain string, timeout int, bootstrapDNS []string) (ips []string) { func lookupIP(ctx context.Context, domain string, timeout int, bootstrapDNS []string) (ips []string) {
if net.ParseIP(domain) != nil { if net.ParseIP(domain) != nil {
return []string{domain} return []string{domain}
} }
logger := LoggerFromCtx(ctx)
if bootstrapDNS == nil { if bootstrapDNS == nil {
ProxyLogger.Load().Debug().Msgf("empty bootstrap DNS") logger.Debug().Msgf("empty bootstrap DNS")
return nil return nil
} }
resolver := newResolverWithNameserver(bootstrapDNS) resolver := newResolverWithNameserver(bootstrapDNS)
ProxyLogger.Load().Debug().Msgf("resolving %q using bootstrap DNS %q", domain, bootstrapDNS) logger.Debug().Msgf("resolving %q using bootstrap DNS %q", domain, bootstrapDNS)
timeoutMs := 2000 timeoutMs := 2000
if timeout > 0 && timeout < timeoutMs { if timeout > 0 && timeout < timeoutMs {
timeoutMs = timeout timeoutMs = timeout
@@ -616,15 +613,15 @@ func lookupIP(domain string, timeout int, bootstrapDNS []string) (ips []string)
r, err := resolver.Resolve(ctx, m) r, err := resolver.Resolve(ctx, m)
if err != nil { if err != nil {
ProxyLogger.Load().Error().Err(err).Msgf("could not lookup %q record for domain %q", dns.TypeToString[dnsType], domain) logger.Error().Err(err).Msgf("could not lookup %q record for domain %q", dns.TypeToString[dnsType], domain)
return return
} }
if r.Rcode != dns.RcodeSuccess { if r.Rcode != dns.RcodeSuccess {
ProxyLogger.Load().Error().Msgf("could not resolve domain %q, return code: %s", domain, dns.RcodeToString[r.Rcode]) logger.Error().Msgf("could not resolve domain %q, return code: %s", domain, dns.RcodeToString[r.Rcode])
return return
} }
if len(r.Answer) == 0 { if len(r.Answer) == 0 {
ProxyLogger.Load().Error().Msg("no answer from OS resolver") logger.Error().Msg("no answer from OS resolver")
return return
} }
target := targetDomain(r.Answer) target := targetDomain(r.Answer)
@@ -641,22 +638,6 @@ func lookupIP(domain string, timeout int, bootstrapDNS []string) (ips []string)
return ips return ips
} }
// NewBootstrapResolver returns an OS resolver, which use following nameservers:
//
// - Gateway IP address (depends on OS).
// - Input servers.
func NewBootstrapResolver(servers ...string) Resolver {
logger := *ProxyLogger.Load()
Log(context.Background(), logger.Debug(), "NewBootstrapResolver called with servers: %v", servers)
nss := defaultNameservers()
nss = append([]string{controldPublicDnsWithPort}, nss...)
for _, ns := range servers {
nss = append([]string{net.JoinHostPort(ns, "53")}, nss...)
}
return NewResolverWithNameserver(nss)
}
// NewPrivateResolver returns an OS resolver, which includes only private DNS servers, // NewPrivateResolver returns an OS resolver, which includes only private DNS servers,
// excluding: // excluding:
// //
@@ -664,8 +645,8 @@ func NewBootstrapResolver(servers ...string) Resolver {
// - Nameservers which is local RFC1918 addresses. // - Nameservers which is local RFC1918 addresses.
// //
// This is useful for doing PTR lookup in LAN network. // This is useful for doing PTR lookup in LAN network.
func NewPrivateResolver() Resolver { func NewPrivateResolver(ctx context.Context) Resolver {
nss := initDefaultOsResolver() nss := initDefaultOsResolver(ctx)
resolveConfNss := currentNameserversFromResolvconf() resolveConfNss := currentNameserversFromResolvconf()
localRfc1918Addrs := Rfc1918Addresses() localRfc1918Addrs := Rfc1918Addresses()
n := 0 n := 0
+1 -1
View File
@@ -132,7 +132,7 @@ func Test_osResolver_InitializationRace(t *testing.T) {
for range n { for range n {
go func() { go func() {
defer wg.Done() defer wg.Done()
InitializeOsResolver(false) InitializeOsResolver(LoggerCtx(context.Background(), nil), false)
}() }()
} }
wg.Wait() wg.Wait()