cmd/cli: ensure DNS goroutines terminated before self-uninstall

Otherwise, these goroutines could mess up with what resetDNS function
do, reverting DHCP DNS settings to ctrld listeners.
This commit is contained in:
Cuong Manh Le
2024-08-16 13:50:11 +07:00
committed by Cuong Manh Le
parent 79476add12
commit 5af3ec4f7b
6 changed files with 88 additions and 75 deletions
+27 -22
View File
@@ -1135,13 +1135,14 @@ func run(appCallback *AppCallback, stopCh chan struct{}) {
} }
waitCh := make(chan struct{}) waitCh := make(chan struct{})
p := &prog{ p := &prog{
waitCh: waitCh, waitCh: waitCh,
stopCh: stopCh, stopCh: stopCh,
reloadCh: make(chan struct{}), reloadCh: make(chan struct{}),
reloadDoneCh: make(chan struct{}), reloadDoneCh: make(chan struct{}),
apiReloadCh: make(chan *ctrld.Config), dnsWatcherStopCh: make(chan struct{}),
cfg: &cfg, apiReloadCh: make(chan *ctrld.Config),
appCallback: appCallback, cfg: &cfg,
appCallback: appCallback,
} }
if homedir == "" { if homedir == "" {
if dir, err := userHomeDir(); err == nil { if dir, err := userHomeDir(); err == nil {
@@ -1232,7 +1233,11 @@ func run(appCallback *AppCallback, stopCh chan struct{}) {
} }
cdLogger := mainLog.Load().With().Str("mode", "cd").Logger() cdLogger := mainLog.Load().With().Str("mode", "cd").Logger()
_ = uninstallIfInvalidCdUID(err, p, cdLogger) // Performs self-uninstallation if the ControlD device does not exist.
var uer *controld.UtilityErrorResponse
if errors.As(err, &uer) && uer.ErrorField.Code == controld.InvalidConfigCode {
_ = uninstallInvalidCdUID(p, cdLogger, false)
}
cdLogger.Fatal().Err(err).Msg("failed to fetch resolver config") cdLogger.Fatal().Err(err).Msg("failed to fetch resolver config")
} }
} }
@@ -2696,23 +2701,23 @@ func doValidateCdRemoteConfig(cdUID string) {
v = oldV v = oldV
} }
// uninstallIfInvalidCdUID performs self-uninstallation if the ControlD device does not exist. // uninstallInvalidCdUID performs self-uninstallation because the ControlD device does not exist.
func uninstallIfInvalidCdUID(err error, p *prog, logger zerolog.Logger) bool { func uninstallInvalidCdUID(p *prog, logger zerolog.Logger, doStop bool) bool {
var uer *controld.UtilityErrorResponse s, err := newService(p, svcConfig)
if errors.As(err, &uer) && uer.ErrorField.Code == controld.InvalidConfigCode { if err != nil {
s, err := newService(p, svcConfig) logger.Warn().Err(err).Msg("failed to create new service")
if err != nil { return false
logger.Warn().Err(err).Msg("failed to create new service") }
return false
}
p.resetDNS() p.resetDNS()
tasks := []task{{s.Uninstall, true}} tasks := []task{{s.Uninstall, true}}
if doTasks(tasks) { if doTasks(tasks) {
logger.Info().Msg("uninstalled service") logger.Info().Msg("uninstalled service")
return true if doStop {
_ = s.Stop()
} }
return true
} }
return false return false
} }
+1 -1
View File
@@ -863,7 +863,7 @@ func (p *prog) doSelfUninstall(answer *dns.Msg) {
p.checkingSelfUninstall = true p.checkingSelfUninstall = true
_, err := controld.FetchResolverConfig(cdUID, rootCmd.Version, cdDev) _, err := controld.FetchResolverConfig(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")
selfUninstall(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")
+30 -13
View File
@@ -22,6 +22,7 @@ import (
"time" "time"
"github.com/kardianos/service" "github.com/kardianos/service"
"github.com/rs/zerolog"
"github.com/spf13/viper" "github.com/spf13/viper"
"tailscale.com/net/interfaces" "tailscale.com/net/interfaces"
"tailscale.com/net/tsaddr" "tailscale.com/net/tsaddr"
@@ -67,18 +68,19 @@ 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
logConn net.Conn logConn net.Conn
cs *controlServer cs *controlServer
csSetDnsDone chan struct{} csSetDnsDone chan struct{}
csSetDnsOk bool csSetDnsOk bool
dnsWatchDogOnce sync.Once dnsWatchDogOnce sync.Once
dnsWg sync.WaitGroup dnsWg sync.WaitGroup
dnsWatcherStopCh chan struct{}
cfg *ctrld.Config cfg *ctrld.Config
localUpstreams []string localUpstreams []string
@@ -261,7 +263,7 @@ func (p *prog) apiConfigReload() {
select { select {
case <-ticker.C: case <-ticker.C:
resolverConfig, err := controld.FetchResolverConfig(cdUID, rootCmd.Version, cdDev) resolverConfig, err := controld.FetchResolverConfig(cdUID, rootCmd.Version, cdDev)
selfUninstall(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")
continue continue
@@ -650,6 +652,8 @@ func (p *prog) dnsWatchdog(iface *net.Interface, nameservers []string, allIfaces
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:
return
case <-p.stopCh: case <-p.stopCh:
mainLog.Load().Debug().Msg("stop dns watchdog") mainLog.Load().Debug().Msg("stop dns watchdog")
return return
@@ -975,3 +979,16 @@ func dnsChanged(iface *net.Interface, nameservers []string) bool {
slices.Sort(curNameservers) slices.Sort(curNameservers)
return !slices.Equal(curNameservers, nameservers) return !slices.Equal(curNameservers, nameservers)
} }
// selfUninstallCheck checks if the error dues to controld.InvalidConfigCode, perform self-uninstall then.
func selfUninstallCheck(uninstallErr error, p *prog, logger zerolog.Logger) {
var uer *controld.UtilityErrorResponse
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.
close(p.dnsWatcherStopCh)
p.dnsWg.Wait()
// Perform self-uninstall now.
selfUninstall(p, logger)
}
}
+2
View File
@@ -34,6 +34,8 @@ func (p *prog) watchResolvConf(iface *net.Interface, ns []netip.Addr, setDnsFn f
for { for {
select { select {
case <-p.dnsWatcherStopCh:
return
case <-p.stopCh: case <-p.stopCh:
mainLog.Load().Debug().Msgf("stopping watcher for %s", resolvConfPath) mainLog.Load().Debug().Msgf("stopping watcher for %s", resolvConfPath)
return return
+2 -2
View File
@@ -8,8 +8,8 @@ import (
"github.com/rs/zerolog" "github.com/rs/zerolog"
) )
func selfUninstall(err error, p *prog, logger zerolog.Logger) { func selfUninstall(p *prog, logger zerolog.Logger) {
if uninstallIfInvalidCdUID(err, p, logger) { if uninstallInvalidCdUID(p, logger, false) {
logger.Warn().Msgf("service was uninstalled because device %q does not exist", cdUID) logger.Warn().Msgf("service was uninstalled because device %q does not exist", cdUID)
os.Exit(0) os.Exit(0)
} }
+26 -37
View File
@@ -3,54 +3,43 @@
package cli package cli
import ( import (
"errors"
"fmt" "fmt"
"os" "os"
"os/exec" "os/exec"
"runtime" "runtime"
"syscall" "syscall"
"github.com/Control-D-Inc/ctrld/internal/controld"
"github.com/rs/zerolog" "github.com/rs/zerolog"
) )
func selfUninstall(uninstallErr error, p *prog, logger zerolog.Logger) { func selfUninstall(p *prog, logger zerolog.Logger) {
var uer *controld.UtilityErrorResponse if runtime.GOOS == "linux" {
if errors.As(uninstallErr, &uer) && uer.ErrorField.Code == controld.InvalidConfigCode { selfUninstallLinux(p, logger)
if runtime.GOOS == "linux" { }
s, err := newService(p, svcConfig)
if err != nil {
logger.Warn().Err(err).Msg("failed to create new service")
} else {
selfUninstallLinux(uninstallErr, p, logger)
_ = s.Stop()
os.Exit(0)
}
}
bin, err := os.Executable() bin, err := os.Executable()
if err != nil { if err != nil {
logger.Fatal().Err(err).Msg("could not determine executable") logger.Fatal().Err(err).Msg("could not determine executable")
} }
args := []string{"uninstall"} args := []string{"uninstall"}
if !deactivationPinNotSet() { if !deactivationPinNotSet() {
args = append(args, fmt.Sprintf("--pin=%d", cdDeactivationPin)) args = append(args, fmt.Sprintf("--pin=%d", cdDeactivationPin))
} }
cmd := exec.Command(bin, args...) cmd := exec.Command(bin, args...)
cmd.SysProcAttr = &syscall.SysProcAttr{Setpgid: true} cmd.SysProcAttr = &syscall.SysProcAttr{Setpgid: true}
if err := cmd.Start(); err != nil { if err := cmd.Start(); err != nil {
logger.Fatal().Err(err).Msg("could not start self uninstall command") logger.Fatal().Err(err).Msg("could not start self uninstall command")
} }
cmd.Stdout = os.Stdout cmd.Stdout = os.Stdout
cmd.Stderr = os.Stderr cmd.Stderr = os.Stderr
logger.Warn().Msgf("service was uninstalled because device %q does not exist", cdUID)
_ = cmd.Wait()
os.Exit(0)
}
func selfUninstallLinux(p *prog, logger zerolog.Logger) {
if uninstallInvalidCdUID(p, logger, true) {
logger.Warn().Msgf("service was uninstalled because device %q does not exist", cdUID) logger.Warn().Msgf("service was uninstalled because device %q does not exist", cdUID)
_ = cmd.Wait()
os.Exit(0) os.Exit(0)
} }
} }
func selfUninstallLinux(err error, p *prog, logger zerolog.Logger) {
if uninstallIfInvalidCdUID(err, p, logger) {
logger.Warn().Msgf("service was uninstalled because device %q does not exist", cdUID)
}
}