cmd/ctrld: optimizing set/reset DNS

Currently, when reset DNS, ctrld always find the net.Interface by
interface name. This may produce unexpected error because the interface
table may be cleared at the time ctrld is being stopped.

Instead, we can get the net interface only once, and use that interface
for restoring the DNS before shutting down.

While at it, also making logging message clearer.
This commit is contained in:
Cuong Manh Le
2023-01-24 16:57:16 +07:00
committed by Cuong Manh Le
parent b0dc96aa01
commit c82a0e2562
4 changed files with 66 additions and 39 deletions
+24 -17
View File
@@ -421,32 +421,28 @@ func initCLI() {
} }
} }
func writeConfigFile() { func writeConfigFile() error {
if cfu := v.ConfigFileUsed(); cfu != "" { if cfu := v.ConfigFileUsed(); cfu != "" {
defaultConfigFile = cfu defaultConfigFile = cfu
} }
f, err := os.OpenFile(defaultConfigFile, os.O_CREATE|os.O_TRUNC|os.O_WRONLY, os.FileMode(0o644)) f, err := os.OpenFile(defaultConfigFile, os.O_CREATE|os.O_TRUNC|os.O_WRONLY, os.FileMode(0o644))
if err != nil { if err != nil {
log.Printf("failed to open config file: %v\n", err) return err
os.Exit(1)
} }
defer f.Close() defer f.Close()
if cdUID != "" { if cdUID != "" {
if _, err := f.WriteString("# AUTO-GENERATED VIA CD FLAG - DO NOT MODIFY\n\n"); err != nil { if _, err := f.WriteString("# AUTO-GENERATED VIA CD FLAG - DO NOT MODIFY\n\n"); err != nil {
log.Printf("failed to write header to config file: %v\n", err) return err
os.Exit(1)
} }
} }
enc := toml.NewEncoder(f).SetIndentTables(true) enc := toml.NewEncoder(f).SetIndentTables(true)
if err := enc.Encode(v.AllSettings()); err != nil { if err := enc.Encode(v.AllSettings()); err != nil {
log.Printf("failed to encode config file: %v\n", err) return err
os.Exit(1)
} }
if err := f.Close(); err != nil { if err := f.Close(); err != nil {
log.Printf("failed to write config file: %v\n", err) return err
os.Exit(1)
} }
fmt.Println("writing config file to:", defaultConfigFile) return nil
} }
func readConfigFile(writeDefaultConfig bool) bool { func readConfigFile(writeDefaultConfig bool) bool {
@@ -464,7 +460,11 @@ func readConfigFile(writeDefaultConfig bool) bool {
// If error is viper.ConfigFileNotFoundError, write default config. // If error is viper.ConfigFileNotFoundError, write default config.
if _, ok := err.(viper.ConfigFileNotFoundError); ok { if _, ok := err.(viper.ConfigFileNotFoundError); ok {
writeConfigFile() if err := writeConfigFile(); err != nil {
log.Fatalf("failed to write default config file: %v", err)
} else {
fmt.Println("writing default config file to: " + defaultConfigFile)
}
defaultConfigWritten = true defaultConfigWritten = true
return false return false
} }
@@ -526,6 +526,7 @@ func processCDFlags() {
iface = "auto" iface = "auto"
} }
logger := mainLog.With().Str("mode", "cd").Logger() logger := mainLog.With().Str("mode", "cd").Logger()
logger.Info().Msg("fetching Controld-D configuration")
resolverConfig, err := controld.FetchResolverConfig(cdUID) resolverConfig, err := controld.FetchResolverConfig(cdUID)
if uer, ok := err.(*controld.UtilityErrorResponse); ok && uer.ErrorField.Code == controld.InvalidConfigCode { if uer, ok := err.(*controld.UtilityErrorResponse); ok && uer.ErrorField.Code == controld.InvalidConfigCode {
s, err := service.New(&prog{}, svcConfig) s, err := service.New(&prog{}, svcConfig)
@@ -533,10 +534,8 @@ func processCDFlags() {
logger.Warn().Err(err).Msg("failed to create new service") logger.Warn().Err(err).Msg("failed to create new service")
return return
} }
if iface == "auto" {
iface = defaultIfaceName() if netIface, _ := netInterface(iface); netIface != nil {
}
if netIface, _ := netIfaceFromName(iface); netIface != nil {
if err := resetDNS(netIface); err != nil { if err := resetDNS(netIface); err != nil {
logger.Warn().Err(err).Msg("something went wrong while restoring DNS") logger.Warn().Err(err).Msg("something went wrong while restoring DNS")
} }
@@ -552,6 +551,7 @@ func processCDFlags() {
return return
} }
logger.Info().Msg("generating ctrld config from Controld-D configuration")
cfg = ctrld.Config{} cfg = ctrld.Config{}
cfg.Network = make(map[string]*ctrld.NetworkConfig) cfg.Network = make(map[string]*ctrld.NetworkConfig)
cfg.Network["0"] = &ctrld.NetworkConfig{ cfg.Network["0"] = &ctrld.NetworkConfig{
@@ -583,7 +583,11 @@ func processCDFlags() {
v.Set("upstream", cfg.Upstream) v.Set("upstream", cfg.Upstream)
v.Set("listener", cfg.Listener) v.Set("listener", cfg.Listener)
processLogAndCacheFlags() processLogAndCacheFlags()
writeConfigFile() if err := writeConfigFile(); err != nil {
logger.Fatal().Err(err).Msg("failed to write config file")
} else {
logger.Info().Msg("writing config file to: " + defaultConfigFile)
}
} }
func processListenFlag() { func processListenFlag() {
@@ -620,7 +624,10 @@ func processLogAndCacheFlags() {
v.Set("service", cfg.Service) v.Set("service", cfg.Service)
} }
func netIfaceFromName(ifaceName string) (*net.Interface, error) { func netInterface(ifaceName string) (*net.Interface, error) {
if ifaceName == "auto" {
ifaceName = defaultIfaceName()
}
var iface *net.Interface var iface *net.Interface
err := interfaces.ForeachInterface(func(i interfaces.Interface, prefixes []netip.Prefix) { err := interfaces.ForeachInterface(func(i interfaces.Interface, prefixes []netip.Prefix) {
if i.Name == ifaceName { if i.Name == ifaceName {
+2
View File
@@ -3,6 +3,7 @@ package main
import ( import (
"fmt" "fmt"
"io" "io"
"net"
"os" "os"
"path/filepath" "path/filepath"
"time" "time"
@@ -35,6 +36,7 @@ var (
cdUID string cdUID string
iface string iface string
netIface *net.Interface
ifaceStartStop string ifaceStartStop string
) )
+1 -1
View File
@@ -81,7 +81,7 @@ func resetDNS(iface *net.Interface) error {
c := client6.NewClient() c := client6.NewClient()
conversation, err := c.Exchange(iface.Name) conversation, err := c.Exchange(iface.Name)
if err != nil { if err != nil {
mainLog.Warn().Err(err).Msg("failed to exchange DHCPv6") mainLog.Warn().Err(err).Msg("could not exchange DHCPv6")
} }
for _, packet := range conversation { for _, packet := range conversation {
if packet.Type() == dhcpv6.MessageTypeReply { if packet.Type() == dhcpv6.MessageTypeReply {
+39 -21
View File
@@ -34,17 +34,7 @@ func (p *prog) Start(s service.Service) error {
} }
func (p *prog) run() { func (p *prog) run() {
if iface != "" { p.setDNS()
netIface, err := netIfaceFromName(iface)
if err != nil {
mainLog.Error().Err(err).Msg("could not get interface")
} else {
if err := setDNS(netIface, []string{cfg.Listener["0"].IP}); err != nil {
mainLog.Error().Err(err).Str("iface", iface).Msgf("could not set DNS for interface")
}
}
}
if p.cfg.Service.CacheEnable { if p.cfg.Service.CacheEnable {
cacher, err := dnscache.NewLRUCache(p.cfg.Service.CacheSize) cacher, err := dnscache.NewLRUCache(p.cfg.Service.CacheSize)
if err != nil { if err != nil {
@@ -150,7 +140,11 @@ func (p *prog) run() {
Port: port, Port: port,
}, },
}) })
writeConfigFile() if err := writeConfigFile(); err != nil {
proxyLog.Fatal().Err(err).Msg("failed to write config file")
} else {
mainLog.Info().Msg("writing config file to: " + defaultConfigFile)
}
mainLog.Info().Msgf("Starting DNS server on listener.%s: %s", listenerNum, pc.LocalAddr()) mainLog.Info().Msgf("Starting DNS server on listener.%s: %s", listenerNum, pc.LocalAddr())
// There can be a race between closing the listener and start our own UDP server, but it's // There can be a race between closing the listener and start our own UDP server, but it's
// rare, and we only do this once, so let conservative here. // rare, and we only do this once, so let conservative here.
@@ -175,15 +169,7 @@ func (p *prog) Stop(s service.Service) error {
mainLog.Error().Err(err).Msg("de-allocate ip failed") mainLog.Error().Err(err).Msg("de-allocate ip failed")
return err return err
} }
if iface != "" { p.resetDNS()
if netIface, err := netIfaceFromName(iface); err == nil {
if err := resetDNS(netIface); err != nil {
mainLog.Error().Err(err).Str("iface", iface).Msgf("could not reset DNS")
}
} else {
mainLog.Error().Err(err).Msg("could not get interface")
}
}
return nil return nil
} }
@@ -205,3 +191,35 @@ func (p *prog) deAllocateIP() error {
} }
return nil return nil
} }
func (p *prog) setDNS() {
if iface == "" {
return
}
logger := mainLog.With().Str("iface", iface).Logger()
var err error
netIface, err = netInterface(iface)
if err != nil {
logger.Error().Err(err).Msg("could not get interface")
return
}
logger.Debug().Msg("setting DNS for interface")
if err := setDNS(netIface, []string{p.cfg.Listener["0"].IP}); err != nil {
logger.Error().Err(err).Msgf("could not set DNS for interface")
return
}
logger.Debug().Msg("setting DNS successfully")
}
func (p *prog) resetDNS() {
if netIface == nil {
return
}
logger := mainLog.With().Str("iface", iface).Logger()
logger.Debug().Msg("Restoring DNS for interface")
if err := resetDNS(netIface); err != nil {
logger.Error().Err(err).Msgf("could not reset DNS")
return
}
logger.Debug().Msg("Restoring DNS successfully")
}