mirror of
https://github.com/Control-D-Inc/ctrld.git
synced 2026-09-29 13:51:51 +02:00
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:
1 parent
8c661c4401
commit
5a88a7c22c
5 files changed
+106
-74
No files matched your search
+32
-18
@@ -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
|
||||||
|
|||||||
@@ -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()
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in new issue
Block a user