cmd/cli: decouple reset DNS task from ctrld status

So it can be run regardless of ctrld current status. This prevents a
racy behavior when reset DNS task restores DNS settings of the system,
but current running ctrld process may revert it immediately.
This commit is contained in:
Cuong Manh Le authored and Cuong Manh Le committed 2024-09-30 18:17:31 +07:00
1 parent 8c661c4401
commit 5a88a7c22c
5 files changed
+106 -74

No files matched your search

+32 -18
View File
@@ -194,11 +194,15 @@ NOTE: running "ctrld start" without any arguments will start already installed c
isCtrldRunning := status == service.StatusRunning isCtrldRunning := status == service.StatusRunning
isCtrldInstalled := !errors.Is(err, service.ErrNotInstalled) isCtrldInstalled := !errors.Is(err, service.ErrNotInstalled)
// Get current running iface, if any.
var currentIface string
// If pin code was set, do not allow running start command. // If pin code was set, do not allow running start command.
if isCtrldRunning { if isCtrldRunning {
if err := checkDeactivationPin(s, nil); isCheckDeactivationPinErr(err) { if err := checkDeactivationPin(s, nil); isCheckDeactivationPinErr(err) {
os.Exit(deactivationPinInvalidExitCode) os.Exit(deactivationPinInvalidExitCode)
} }
currentIface = runningIface(s)
} }
if !startOnly { if !startOnly {
@@ -213,12 +217,15 @@ NOTE: running "ctrld start" without any arguments will start already installed c
initLogging() initLogging()
tasks := []task{ tasks := []task{
resetDnsTask(p, s),
{s.Stop, false}, {s.Stop, false},
resetDnsTask(p, s, isCtrldInstalled, currentIface),
{func() error { {func() error {
// Save current DNS so we can restore later. // Save current DNS so we can restore later.
withEachPhysicalInterfaces("", "save DNS settings", func(i *net.Interface) error { withEachPhysicalInterfaces("", "", func(i *net.Interface) error {
return saveCurrentStaticDNS(i) if err := saveCurrentStaticDNS(i); !errors.Is(err, errSaveCurrentStaticDNSNotSupported) && err != nil {
return err
}
return nil
}) })
return nil return nil
}, false}, }, false},
@@ -334,14 +341,17 @@ NOTE: running "ctrld start" without any arguments will start already installed c
} }
tasks := []task{ tasks := []task{
resetDnsTask(p, s),
{s.Stop, false}, {s.Stop, false},
{func() error { return doGenerateNextDNSConfig(nextdns) }, true}, {func() error { return doGenerateNextDNSConfig(nextdns) }, true},
{func() error { return ensureUninstall(s) }, false}, {func() error { return ensureUninstall(s) }, false},
resetDnsTask(p, s, isCtrldInstalled, currentIface),
{func() error { {func() error {
// Save current DNS so we can restore later. // Save current DNS so we can restore later.
withEachPhysicalInterfaces("", "save DNS settings", func(i *net.Interface) error { withEachPhysicalInterfaces("", "", func(i *net.Interface) error {
return saveCurrentStaticDNS(i) if err := saveCurrentStaticDNS(i); !errors.Is(err, errSaveCurrentStaticDNSNotSupported) && err != nil {
return err
}
return nil
}) })
return nil return nil
}, false}, }, false},
@@ -1340,9 +1350,7 @@ func run(appCallback *AppCallback, stopCh chan struct{}) {
close(waitCh) close(waitCh)
<-stopCh <-stopCh
// Wait goroutines which watches/manipulates DNS settings terminated, p.stopDnsWatchers()
// ensuring that changes to DNS since here won't be reverted.
p.dnsWg.Wait()
for _, f := range p.onStopped { for _, f := range p.onStopped {
f() f()
} }
@@ -2642,17 +2650,20 @@ func runningIface(s service.Service) string {
// resetDnsNoLog performs resetting DNS with logging disable. // resetDnsNoLog performs resetting DNS with logging disable.
func resetDnsNoLog(p *prog) { func resetDnsNoLog(p *prog) {
lvl := zerolog.GlobalLevel() // Normally, disable log to prevent annoying users.
zerolog.SetGlobalLevel(zerolog.Disabled) if verbose < 3 {
lvl := zerolog.GlobalLevel()
zerolog.SetGlobalLevel(zerolog.Disabled)
p.resetDNS()
zerolog.SetGlobalLevel(lvl)
return
}
// For debugging purpose, still emit log.
p.resetDNS() p.resetDNS()
zerolog.SetGlobalLevel(lvl)
} }
// resetDnsTask returns a task which perform reset DNS operation. // resetDnsTask returns a task which perform reset DNS operation.
func resetDnsTask(p *prog, s service.Service) task { func resetDnsTask(p *prog, s service.Service, isCtrldInstalled bool, currentRunningIface string) task {
status, err := s.Status()
isCtrldInstalled := !errors.Is(err, service.ErrNotInstalled)
isCtrldRunning := status == service.StatusRunning
return task{func() error { return task{func() error {
if iface == "" { if iface == "" {
return nil return nil
@@ -2662,11 +2673,14 @@ func resetDnsTask(p *prog, s service.Service) task {
// process to reset what setDNS has done properly. // process to reset what setDNS has done properly.
oldIface := iface oldIface := iface
iface = "auto" iface = "auto"
if isCtrldRunning { if currentRunningIface != "" {
iface = runningIface(s) iface = currentRunningIface
} }
if isCtrldInstalled { if isCtrldInstalled {
mainLog.Load().Debug().Msg("restore system DNS settings") mainLog.Load().Debug().Msg("restore system DNS settings")
if status, _ := s.Status(); status == service.StatusRunning {
mainLog.Load().Fatal().Msg("reset DNS while ctrld still running is not safe")
}
resetDnsNoLog(p) resetDnsNoLog(p)
} }
iface = oldIface iface = oldIface
+2
View File
@@ -915,6 +915,8 @@ func (p *prog) performCaptivePortalDetection() {
if found { if found {
resetDnsOnce.Do(func() { resetDnsOnce.Do(func() {
mainLog.Load().Warn().Msg("found captive portal, leaking query to OS resolver") mainLog.Load().Warn().Msg("found captive portal, leaking query to OS resolver")
// Store the result once here, so changes made below won't be reverted by DNS watchers.
p.captivePortalDetected.Store(found)
p.resetDNS() p.resetDNS()
}) })
} }
+1
View File
@@ -119,6 +119,7 @@ func resetDNS(iface *net.Interface) error {
if len(ns) == 0 { if len(ns) == 0 {
continue continue
} }
mainLog.Load().Debug().Msgf("setting static DNS for interface %q", iface.Name)
if err := setDNS(iface, ns); err != nil { if err := setDNS(iface, ns); err != nil {
return err return err
} }
+68 -56
View File
@@ -69,21 +69,21 @@ var svcConfig = &service.Config{
var useSystemdResolved = false var useSystemdResolved = false
type prog struct { type prog struct {
mu sync.Mutex mu sync.Mutex
waitCh chan struct{} waitCh chan struct{}
stopCh chan struct{} stopCh chan struct{}
reloadCh chan struct{} // For Windows. reloadCh chan struct{} // For Windows.
reloadDoneCh chan struct{} reloadDoneCh chan struct{}
apiReloadCh chan *ctrld.Config apiReloadCh chan *ctrld.Config
apiForceReloadCh chan struct{} apiForceReloadCh chan struct{}
apiForceReloadGroup singleflight.Group apiForceReloadGroup singleflight.Group
logConn net.Conn logConn net.Conn
cs *controlServer cs *controlServer
csSetDnsDone chan struct{} csSetDnsDone chan struct{}
csSetDnsOk bool csSetDnsOk bool
dnsWatchDogOnce sync.Once dnsWg sync.WaitGroup
dnsWg sync.WaitGroup dnsWatcherClosedOnce sync.Once
dnsWatcherStopCh chan struct{} dnsWatcherStopCh chan struct{}
cfg *ctrld.Config cfg *ctrld.Config
localUpstreams []string localUpstreams []string
@@ -512,6 +512,8 @@ func (p *prog) metricsEnabled() bool {
} }
func (p *prog) Stop(s service.Service) error { func (p *prog) Stop(s service.Service) error {
p.stopDnsWatchers()
mainLog.Load().Debug().Msg("dns watchers stopped")
mainLog.Load().Info().Msg("Service stopped") mainLog.Load().Info().Msg("Service stopped")
close(p.stopCh) close(p.stopCh)
if err := p.deAllocateIP(); err != nil { if err := p.deAllocateIP(); err != nil {
@@ -521,6 +523,15 @@ func (p *prog) Stop(s service.Service) error {
return nil return nil
} }
func (p *prog) stopDnsWatchers() {
// Ensure all DNS watchers goroutine are terminated,
// so it won't mess up with other DNS changes.
p.dnsWatcherClosedOnce.Do(func() {
close(p.dnsWatcherStopCh)
})
p.dnsWg.Wait()
}
func (p *prog) allocateIP(ip string) error { func (p *prog) allocateIP(ip string) error {
p.mu.Lock() p.mu.Lock()
defer p.mu.Unlock() defer p.mu.Unlock()
@@ -611,6 +622,11 @@ func (p *prog) setDNS() {
} }
setDnsOK = true setDnsOK = true
logger.Debug().Msg("setting DNS successfully") logger.Debug().Msg("setting DNS successfully")
if allIfaces {
withEachPhysicalInterfaces(netIface.Name, "set DNS", func(i *net.Interface) error {
return setDnsIgnoreUnusableInterface(i, nameservers)
})
}
if shouldWatchResolvconf() { if shouldWatchResolvconf() {
servers := make([]netip.Addr, len(nameservers)) servers := make([]netip.Addr, len(nameservers))
for i := range nameservers { for i := range nameservers {
@@ -622,11 +638,6 @@ func (p *prog) setDNS() {
p.watchResolvConf(netIface, servers, setResolvConf) p.watchResolvConf(netIface, servers, setResolvConf)
}() }()
} }
if allIfaces {
withEachPhysicalInterfaces(netIface.Name, "set DNS", func(i *net.Interface) error {
return setDnsIgnoreUnusableInterface(i, nameservers)
})
}
if p.dnsWatchdogEnabled() { if p.dnsWatchdogEnabled() {
p.dnsWg.Add(1) p.dnsWg.Add(1)
go func() { go func() {
@@ -661,41 +672,42 @@ func (p *prog) dnsWatchdog(iface *net.Interface, nameservers []string, allIfaces
return return
} }
p.dnsWatchDogOnce.Do(func() { mainLog.Load().Debug().Msg("start DNS settings watchdog")
mainLog.Load().Debug().Msg("start DNS settings watchdog") ns := nameservers
ns := nameservers slices.Sort(ns)
slices.Sort(ns) ticker := time.NewTicker(p.dnsWatchdogDuration())
ticker := time.NewTicker(p.dnsWatchdogDuration()) logger := mainLog.Load().With().Str("iface", iface.Name).Logger()
logger := mainLog.Load().With().Str("iface", iface.Name).Logger() for {
for { select {
select { case <-p.dnsWatcherStopCh:
case <-p.dnsWatcherStopCh: return
case <-p.stopCh:
mainLog.Load().Debug().Msg("stop dns watchdog")
return
case <-ticker.C:
if p.captivePortalDetected.Load() {
return return
case <-p.stopCh: }
mainLog.Load().Debug().Msg("stop dns watchdog") if dnsChanged(iface, ns) {
return logger.Debug().Msg("DNS settings were changed, re-applying settings")
case <-ticker.C: if err := setDNS(iface, ns); err != nil {
if dnsChanged(iface, ns) { mainLog.Load().Error().Err(err).Str("iface", iface.Name).Msgf("could not re-apply DNS settings")
logger.Debug().Msg("DNS settings were changed, re-applying settings")
if err := setDNS(iface, ns); err != nil {
mainLog.Load().Error().Err(err).Str("iface", iface.Name).Msgf("could not re-apply DNS settings")
}
}
if allIfaces {
withEachPhysicalInterfaces(iface.Name, "", func(i *net.Interface) error {
if dnsChanged(i, ns) {
if err := setDnsIgnoreUnusableInterface(i, nameservers); err != nil {
mainLog.Load().Error().Err(err).Str("iface", i.Name).Msgf("could not re-apply DNS settings")
} else {
mainLog.Load().Debug().Msgf("re-applying DNS for interface %q successfully", i.Name)
}
}
return nil
})
} }
} }
if allIfaces {
withEachPhysicalInterfaces(iface.Name, "", func(i *net.Interface) error {
if dnsChanged(i, ns) {
if err := setDnsIgnoreUnusableInterface(i, nameservers); err != nil {
mainLog.Load().Error().Err(err).Str("iface", i.Name).Msgf("could not re-apply DNS settings")
} else {
mainLog.Load().Debug().Msgf("re-applying DNS for interface %q successfully", i.Name)
}
}
return nil
})
}
} }
}) }
} }
func (p *prog) resetDNS() { func (p *prog) resetDNS() {
@@ -965,11 +977,13 @@ func saveCurrentStaticDNS(iface *net.Interface) error {
if err := os.Remove(file); err != nil && !errors.Is(err, fs.ErrNotExist) { if err := os.Remove(file); err != nil && !errors.Is(err, fs.ErrNotExist) {
mainLog.Load().Warn().Err(err).Msg("could not remove old static DNS settings file") mainLog.Load().Warn().Err(err).Msg("could not remove old static DNS settings file")
} }
mainLog.Load().Debug().Msgf("DNS settings for %s is static, saving ...", iface.Name) nss := strings.Join(ns, ",")
if err := os.WriteFile(file, []byte(strings.Join(ns, ",")), 0600); err != nil { mainLog.Load().Debug().Msgf("DNS settings for %q is static: %v, saving ...", iface.Name, nss)
if err := os.WriteFile(file, []byte(nss), 0600); err != nil {
mainLog.Load().Err(err).Msgf("could not save DNS settings for iface: %s", iface.Name) mainLog.Load().Err(err).Msgf("could not save DNS settings for iface: %s", iface.Name)
return err return err
} }
mainLog.Load().Debug().Msgf("save DNS settings for interface %q successfully", iface.Name)
return nil return nil
} }
@@ -1005,9 +1019,7 @@ func dnsChanged(iface *net.Interface, nameservers []string) bool {
func selfUninstallCheck(uninstallErr error, p *prog, logger zerolog.Logger) { func selfUninstallCheck(uninstallErr error, p *prog, logger zerolog.Logger) {
var uer *controld.UtilityErrorResponse var uer *controld.UtilityErrorResponse
if errors.As(uninstallErr, &uer) && uer.ErrorField.Code == controld.InvalidConfigCode { if errors.As(uninstallErr, &uer) && uer.ErrorField.Code == controld.InvalidConfigCode {
// Ensure all DNS watchers goroutine are terminated, so it won't mess up with self-uninstall. p.stopDnsWatchers()
close(p.dnsWatcherStopCh)
p.dnsWg.Wait()
// Perform self-uninstall now. // Perform self-uninstall now.
selfUninstall(p, logger) selfUninstall(p, logger)
+3
View File
@@ -40,6 +40,9 @@ func (p *prog) watchResolvConf(iface *net.Interface, ns []netip.Addr, setDnsFn f
mainLog.Load().Debug().Msgf("stopping watcher for %s", resolvConfPath) mainLog.Load().Debug().Msgf("stopping watcher for %s", resolvConfPath)
return return
case event, ok := <-watcher.Events: case event, ok := <-watcher.Events:
if p.captivePortalDetected.Load() {
return
}
if !ok { if !ok {
return return
} }