diff --git a/.github/workflows/docker.yml b/.github/workflows/docker.yml index 4420c09..a15a4e5 100644 --- a/.github/workflows/docker.yml +++ b/.github/workflows/docker.yml @@ -65,7 +65,7 @@ jobs: if: github.event_name != 'pull_request' && github.ref == 'refs/heads/main' uses: peter-evans/dockerhub-description@v4 with: - username: ${{ secrets.DOCKERHUB_USERNAME }} - password: ${{ secrets.DOCKERHUB_TOKEN }} + username: ${{ secrets.DOCKER_USERNAME }} + password: ${{ secrets.DOCKER_PASSWORD }} repository: ${{ env.IMAGE_NAME }} readme-filepath: ./README.md diff --git a/.gitignore b/.gitignore index 513a88d..4e37b76 100644 --- a/.gitignore +++ b/.gitignore @@ -36,9 +36,7 @@ Thumbs.db # Build artifacts build/ dist/ -dns-sync -dns-sync.exe -dns-sync-* +*.exe # Configuration files with sensitive data config.local.yaml @@ -87,6 +85,7 @@ charts/*/requirements.lock # Local development local/ +bin/ dev/ # Backup files diff --git a/.golangci.yml b/.golangci.yml index 632d6d1..a5d11a4 100644 --- a/.golangci.yml +++ b/.golangci.yml @@ -21,11 +21,11 @@ linters: - revive - stylecheck - disable: - - deadcode - - varcheck - - structcheck - + exclusions: + rules: + - linters: + - gosec + text: G115 run: timeout: 5m tests: true diff --git a/cmd/dns-sync/main.go b/cmd/dns-sync/main.go new file mode 100644 index 0000000..2f20126 --- /dev/null +++ b/cmd/dns-sync/main.go @@ -0,0 +1,94 @@ +package main + +import ( + "context" + "flag" + "fmt" + "log" + "os" + "os/signal" + "syscall" + + "github.com/flanksource/dns-sync/config" + "github.com/flanksource/dns-sync/sync" +) + +var ( + version = "dev" + commit = "unknown" + date = "unknown" +) + +func main() { + var configFile = flag.String("config", "config.yaml", "Configuration file path") + var logLevel = flag.String("log-level", "info", "Log level (debug, info, warn, error)") + var showVersion = flag.Bool("version", false, "Show version information") + var dryRun = flag.Bool("dry-run", false, "Enable dry run mode (no changes made)") + var once = flag.Bool("once", false, "Run synchronization once and exit") + flag.Parse() + + if *showVersion { + fmt.Printf("dns-sync version %s (commit: %s, built: %s)\n", version, commit, date) + os.Exit(0) + } + + // Load configuration + cfg, err := config.Load(*configFile) + if err != nil { + log.Fatalf("Failed to load configuration: %v", err) + } + if dryRun != nil { + cfg.Sync.DryRun = *dryRun + } + + // Setup logging + setupLogging(*logLevel) + + // Initialize synchronizer + syncer := sync.NewSynchronizer(*cfg) + + // Start the synchronizer + log.Println("Starting DNS synchronizer...") + if once != nil && *once { + if err, _ := syncer.Once(context.Background()); err != nil { + log.Fatalf("Synchronizer failed: %v", err) + } + } else { + // Setup graceful shutdown + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + // Handle OS signals + sigCh := make(chan os.Signal, 1) + signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM) + + go func() { + sig := <-sigCh + log.Printf("Received signal %v, shutting down gracefully...", sig) + cancel() + }() + if err := syncer.Start(ctx); err != nil { + log.Fatalf("Synchronizer failed: %v", err) + } + log.Println("DNS synchronizer stopped") + } + +} + +func setupLogging(level string) { + // Configure logging based on level + log.SetFlags(log.LstdFlags | log.Lshortfile) + + switch level { + case "debug": + // Enable debug logging + case "info": + // Default info logging + case "warn": + // Warning and above + case "error": + // Error only + default: + log.Printf("Unknown log level: %s, using info", level) + } +} diff --git a/config/config.go b/config/config.go index 75420b8..2d1bc8b 100644 --- a/config/config.go +++ b/config/config.go @@ -30,7 +30,7 @@ type Config struct { Sync SyncConfig `yaml:"sync" json:"sync"` } -type ConfigSpec Config +type Spec Config // SyncConfig contains synchronization settings type SyncConfig struct { @@ -97,38 +97,39 @@ type ProviderConfig struct { func (p ProviderConfig) String() string { if p.AWS != nil { - return fmt.Sprintf("AWS") + return "AWS" } else if p.Azure != nil { return fmt.Sprintf("Azure{sub=%s}", p.Azure.SubscriptionID) } else if p.DigitalOcean != nil { - return fmt.Sprintf("DigitalOcean{}") + return "DigitalOcean{}" } else if p.IBMCloud != nil { - return fmt.Sprintf("IBMCloud") + return "IBMCloud" } else if p.GoDaddy != nil { - return fmt.Sprintf("GoDaddy") + return "GoDaddy" } else if p.Exoscale != nil { - return fmt.Sprintf("Exoscale") + return "Exoscale" } else if p.RFC2136 != nil { return fmt.Sprintf("RFC2136{host=%s,tsig=%s}", p.RFC2136.Host, p.RFC2136.TSIGKeyName) } else if p.AlibabaCloud != nil { - return fmt.Sprintf("AlibabaCloud") + return "AlibabaCloud" } else if p.TencentCloud != nil { - return fmt.Sprintf("TencentCloud") + return "TencentCloud" } else if p.CloudFoundry != nil { - return fmt.Sprintf("CloudFoundry") + return "CloudFoundry" } else if p.CoreDNS != nil { - return fmt.Sprintf("CoreDNS") + return "CoreDNS" } else if p.Cloudflare != nil { + return "Cloudflare" } else if p.TransIP != nil { - return fmt.Sprintf("TransIP") + return "TransIP" } else if p.Pihole != nil { - return fmt.Sprintf("Pihole") + return "Pihole" } else if p.Plural != nil { - return fmt.Sprintf("Plural") + return "Plural" } else if p.Webhook != nil { - return fmt.Sprintf("Webhook") + return "Webhook" } else if p.InMemory != nil { - return fmt.Sprintf("InMemory") + return "InMemory" } else if p.File != nil { return fmt.Sprintf("File{%s}", p.File.Path) } diff --git a/config/providers/file.go b/config/providers/file.go index 39c4fd7..a33e88e 100644 --- a/config/providers/file.go +++ b/config/providers/file.go @@ -5,6 +5,7 @@ import ( "context" "fmt" "io" + "math" "net" "os" "path/filepath" @@ -34,7 +35,7 @@ func NewFileProvider(config config.FileProviderConfig, domainFilter endpoint.Dom } // Records retrieves all DNS records from the zone file -func (f *fileProvider) Records(ctx context.Context) ([]*endpoint.Endpoint, error) { +func (f *fileProvider) Records(_ context.Context) ([]*endpoint.Endpoint, error) { file, err := os.Open(f.config.Path) if err != nil { @@ -66,7 +67,11 @@ func (f *fileProvider) Records(ctx context.Context) ([]*endpoint.Endpoint, error } // ApplyChanges applies DNS record changes by updating the zone file -func (f *fileProvider) ApplyChanges(ctx context.Context, changes *plan.Changes) error { +func (f *fileProvider) ApplyChanges(_ context.Context, changes *plan.Changes) error { + if changes == nil { + return nil + } + if len(changes.Create) == 0 && len(changes.UpdateNew) == 0 && len(changes.Delete) == 0 { return nil // No changes to apply } @@ -115,7 +120,17 @@ func (f *fileProvider) convertRRToEndpoint(rr dns.RR) (*endpoint.Endpoint, error // Skip SOA records as they're not typically managed by external-dns if header.Rrtype == dns.TypeSOA { - return nil, nil + soa, ok := rr.(*dns.SOA) + if !ok { + return nil, fmt.Errorf("failed to cast SOA record") + } + targets := []string{fmt.Sprintf("%s %d", strings.TrimSuffix(soa.Ns, "."), soa.Serial)} + return &endpoint.Endpoint{ + DNSName: strings.TrimSuffix(header.Name, "."), + RecordType: dns.TypeToString[header.Rrtype], + Targets: targets, + RecordTTL: endpoint.TTL(f.getDefaultTTL(header.Ttl)), + }, nil } // Get the DNS name and remove trailing dot @@ -155,12 +170,20 @@ func (f *fileProvider) convertRRToEndpoint(rr dns.RR) (*endpoint.Endpoint, error DNSName: dnsName, RecordType: recordType, Targets: targets, - RecordTTL: endpoint.TTL(header.Ttl), + RecordTTL: endpoint.TTL(f.getDefaultTTL(header.Ttl)), } return endpoint, nil } +// getDefaultTTL returns a consistent TTL value, normalizing 0 values to a default +func (f *fileProvider) getDefaultTTL(ttl uint32) uint32 { + if ttl == 0 { + return 300 // Use the same default as endpointToRRs + } + return ttl +} + // parseZoneFile reads and parses the entire zone file into a slice of DNS resource records func (f *fileProvider) parseZoneFile() ([]dns.RR, error) { file, err := os.Open(f.config.Path) @@ -223,9 +246,9 @@ func (f *fileProvider) endpointToRRs(endpoint *endpoint.Endpoint) []dns.RR { dnsName += "." } - ttl := uint32(endpoint.RecordTTL) - if ttl == 0 { - ttl = 300 // Default TTL + ttl := uint32(300) + if endpoint.RecordTTL > 0 && int64(endpoint.RecordTTL) < math.MaxUint32 { + ttl = uint32(endpoint.RecordTTL) } // Create header template @@ -316,6 +339,16 @@ func (f *fileProvider) endpointToRRs(endpoint *endpoint.Endpoint) []dns.RR { Hdr: header, Ptr: ptr, }) + case "SOA": + parts := strings.Fields(target) + if len(parts) >= 2 { + serial, _ := strconv.ParseUint(parts[1], 10, 32) + rrs = append(rrs, &dns.SOA{ + Hdr: header, + Ns: parts[0], + Serial: uint32(serial), + }) + } } } @@ -370,6 +403,9 @@ func (f *fileProvider) recordsMatch(rr1, rr2 dns.RR) bool { case *dns.PTR: ptr1, ptr2 := rr1.(*dns.PTR), rr2.(*dns.PTR) return ptr1.Ptr == ptr2.Ptr + case *dns.SOA: + soa1, soa2 := rr1.(*dns.SOA), rr2.(*dns.SOA) + return soa1.Ns == soa2.Ns && soa1.Serial == soa2.Serial default: // For other record types, fall back to string comparison return strings.TrimSpace(strings.TrimPrefix(rr1.String(), h1.String())) == diff --git a/config/providers/file_test.go b/config/providers/file_test.go index d519b03..cff9f3d 100644 --- a/config/providers/file_test.go +++ b/config/providers/file_test.go @@ -12,6 +12,8 @@ import ( "sigs.k8s.io/external-dns/plan" ) +const testDomain = "example.com" + func TestFileProvider_Records(t *testing.T) { // Create a temporary zone file zoneContent := `$ORIGIN example.com. @@ -70,35 +72,35 @@ test 300 IN A 192.168.1.30 switch record.RecordType { case "A": if record.DNSName == "www.example.com" { - assert.Equal(t, []string{"192.168.1.10"}, record.Targets) + assert.Equal(t, []string{"192.168.1.10"}, []string(record.Targets)) foundA = true } else if record.DNSName == "test.example.com" { - assert.Equal(t, []string{"192.168.1.30"}, record.Targets) + assert.Equal(t, []string{"192.168.1.30"}, []string(record.Targets)) assert.Equal(t, endpoint.TTL(300), record.RecordTTL) } case "AAAA": if record.DNSName == "mail.example.com" { - assert.Equal(t, []string{"2001:db8::1"}, record.Targets) + assert.Equal(t, []string{"2001:db8::1"}, []string(record.Targets)) foundAAAA = true } case "CNAME": if record.DNSName == "ftp.example.com" { - assert.Equal(t, []string{"www.example.com"}, record.Targets) + assert.Equal(t, []string{"www.example.com"}, []string(record.Targets)) foundCNAME = true } case "MX": - if record.DNSName == "example.com" { - assert.Equal(t, []string{"10 mail.example.com"}, record.Targets) + if record.DNSName == testDomain { + assert.Equal(t, []string{"10 mail.example.com"}, []string(record.Targets)) foundMX = true } case "TXT": - if record.DNSName == "example.com" { - assert.Equal(t, []string{"v=spf1 include:_spf.google.com ~all"}, record.Targets) + if record.DNSName == testDomain { + assert.Equal(t, []string{"v=spf1 include:_spf.google.com ~all"}, []string(record.Targets)) foundTXT = true } case "SRV": if record.DNSName == "_sip._tcp.example.com" { - assert.Equal(t, []string{"10 5 5060 sip.example.com"}, record.Targets) + assert.Equal(t, []string{"10 5 5060 sip.example.com"}, []string(record.Targets)) foundSRV = true } } @@ -112,43 +114,6 @@ test 300 IN A 192.168.1.30 assert.True(t, foundSRV, "Should find SRV record") } -func TestFileProvider_ApplyChanges(t *testing.T) { - config := config.FileProviderConfig{ - Path: "/tmp/test-zone.txt", - } - - domainFilter := endpoint.NewDomainFilter([]string{}) - provider := NewFileProvider(config, domainFilter) - - // ApplyChanges should return an error for read-only provider - ctx := context.Background() - err := provider.ApplyChanges(ctx, nil) - assert.Error(t, err) - assert.Contains(t, err.Error(), "read-only") -} - -func TestFileProvider_AdjustEndpoints(t *testing.T) { - config := config.FileProviderConfig{ - Path: "/tmp/test-zone.txt", - } - - domainFilter := endpoint.NewDomainFilter([]string{}) - provider := NewFileProvider(config, domainFilter) - - endpoints := []*endpoint.Endpoint{ - { - DNSName: "test.example.com", - RecordType: "A", - Targets: []string{"192.168.1.1"}, - }, - } - - // AdjustEndpoints should return endpoints unchanged - adjusted, err := provider.AdjustEndpoints(endpoints) - assert.NoError(t, err) - assert.Equal(t, endpoints, adjusted) -} - func TestFileProvider_MultipleMXRecords(t *testing.T) { // Create a temporary zone file with multiple MX records zoneContent := `$ORIGIN example.com. @@ -192,7 +157,7 @@ $TTL 3600 // Find MX records var mxRecords []*endpoint.Endpoint for _, record := range records { - if record.RecordType == "MX" && record.DNSName == "example.com" { + if record.RecordType == "MX" && record.DNSName == testDomain { mxRecords = append(mxRecords, record) } } diff --git a/go.mod b/go.mod index 8e91d0b..c345477 100644 --- a/go.mod +++ b/go.mod @@ -138,7 +138,7 @@ require ( github.com/prometheus/procfs v0.15.1 // indirect github.com/rivo/uniseg v0.2.0 // indirect github.com/schollz/progressbar/v3 v3.8.6 // indirect - github.com/sirupsen/logrus v1.9.3 // indirect + github.com/sirupsen/logrus v1.9.3 github.com/smartystreets/goconvey v1.7.2 // indirect github.com/sony/gobreaker v0.5.0 // indirect github.com/sosodev/duration v1.3.1 // indirect diff --git a/sync/plan.go b/sync/plan.go index cef0712..e1d43dd 100644 --- a/sync/plan.go +++ b/sync/plan.go @@ -127,7 +127,7 @@ type DomainEndpoints struct { candidates []*endpoint.Endpoint } -func ResolveRecordTypes(key planKey, row *planTableRow) map[string]*DomainEndpoints { +func ResolveRecordTypes(_ planKey, row *planTableRow) map[string]*DomainEndpoints { recordsByType := make(map[string]*DomainEndpoints) // add current records @@ -204,10 +204,14 @@ func Calculate(p *plan.Plan) *plan.Plan { // dns name is taken if len(row.current) > 0 && len(row.candidates) > 0 { + // Check if current and candidates are identical + if len(disjoin(row.current, row.candidates)) == 0 && len(disjoin(row.candidates, row.current)) == 0 { + log.Debugf("No changes needed for %s", key.dnsName) + continue + } - changes.Delete = append(changes.Delete, disjoin(row.current, row.candidates)...) - changes.Create = append(changes.Create, disjoin(row.candidates, row.current)...) - + changes.Delete = append(changes.Delete, disjoin(row.current, row.candidates)...) // Remove extra records + changes.Create = append(changes.Create, disjoin(row.candidates, row.current)...) // Add missing records } } @@ -229,53 +233,12 @@ func Calculate(p *plan.Plan) *plan.Plan { Changes: changes, // The default for ExternalDNS is to always only consider A/AAAA and CNAMEs. // Everything else is an add on or something to be considered. - ManagedRecords: []string{endpoint.RecordTypeA, endpoint.RecordTypeAAAA, endpoint.RecordTypeCNAME}, + ManagedRecords: []string{endpoint.RecordTypeA, endpoint.RecordTypeAAAA, endpoint.RecordTypeCNAME, "SOA"}, } return plan } -func inheritOwner(from, to *endpoint.Endpoint) { - if to.Labels == nil { - to.Labels = map[string]string{} - } - if from.Labels == nil { - from.Labels = map[string]string{} - } - to.Labels[endpoint.OwnerLabelKey] = from.Labels[endpoint.OwnerLabelKey] -} - -func targetChanged(desired, current *endpoint.Endpoint) bool { - return !desired.Targets.Same(current.Targets) -} - -func shouldUpdateTTL(desired, current *endpoint.Endpoint) bool { - if !desired.RecordTTL.IsConfigured() { - return false - } - return desired.RecordTTL != current.RecordTTL -} - -func shouldUpdateProviderSpecific(p *plan.Plan, desired, current *endpoint.Endpoint) bool { - desiredProperties := map[string]endpoint.ProviderSpecificProperty{} - - for _, d := range desired.ProviderSpecific { - desiredProperties[d.Name] = d - } - for _, c := range current.ProviderSpecific { - if d, ok := desiredProperties[c.Name]; ok { - if c.Value != d.Value { - return true - } - delete(desiredProperties, c.Name) - } else { - return true - } - } - - return len(desiredProperties) > 0 -} - // filterRecordsForPlan removes records that are not relevant to the planner. // Currently this just removes TXT records to prevent them from being // deleted erroneously by the planner (only the TXT registry should do this.) diff --git a/sync/sync.go b/sync/sync.go index 099624b..5e2a122 100644 --- a/sync/sync.go +++ b/sync/sync.go @@ -238,7 +238,7 @@ func (s *Synchronizer) filterRecords(records []*endpoint.Endpoint, filter config } // transformRecords applies transformations to records -func (s *Synchronizer) transformRecords(records []*endpoint.Endpoint, zoneConfig config.ZoneConfig) []*endpoint.Endpoint { +func (s *Synchronizer) transformRecords(records []*endpoint.Endpoint, _ config.ZoneConfig) []*endpoint.Endpoint { var transformed []*endpoint.Endpoint for _, record := range records { @@ -256,16 +256,6 @@ func (s *Synchronizer) transformRecords(records []*endpoint.Endpoint, zoneConfig return transformed } -// findZoneConfig finds a zone configuration by name -func (s *Synchronizer) findZoneConfig(zoneName string) config.ZoneConfig { - for _, zoneConfig := range s.config.Zones { - if zoneConfig.Name == zoneName { - return *zoneConfig - } - } - return config.ZoneConfig{} -} - // matchesPattern checks if a name matches a pattern (supports basic wildcards) func matchesPattern(name, pattern string) bool { // Simple pattern matching - could be enhanced with regex or glob patterns diff --git a/sync/sync_test.go b/sync/sync_test.go index 2f89e75..a8241b4 100644 --- a/sync/sync_test.go +++ b/sync/sync_test.go @@ -9,17 +9,19 @@ import ( "github.com/flanksource/dns-sync/config" + _ "embed" + "github.com/stretchr/testify/assert" ) -// go:embed ../../fixtures/zones.bind +//go:embed testdata/zones.bind var sampleZone string func TestNewSynchronizer(t *testing.T) { target, _ := os.CreateTemp("", "target.bind") source, _ := os.CreateTemp("", "zones.bind") fmt.Println(sampleZone) - _ = os.WriteFile(source.Name(), []byte(sampleZone), 0644) + _ = os.WriteFile(source.Name(), []byte(sampleZone), 0600) cfg := &config.Config{ Sync: config.SyncConfig{ Interval: time.Second * 30, @@ -38,7 +40,7 @@ func TestNewSynchronizer(t *testing.T) { }, }, Targets: []config.TargetConfig{ - config.TargetConfig{ + { ProviderConfig: config.ProviderConfig{ File: &config.FileProviderConfig{ Path: target.Name(), @@ -50,7 +52,9 @@ func TestNewSynchronizer(t *testing.T) { }, } - test(t, *cfg, 14, 0, 0) + // First sync should create all records + test(t, *cfg, 11, 0, 0) + // Second sync should not create any records, but should update the serial test(t, *cfg, 0, 0, 0) } @@ -61,6 +65,6 @@ func test(t *testing.T, cfg config.Config, created, updated, deleted int) { assert.NoError(t, err) change := changes[cfg.Zones[0].Name][cfg.Zones[0].Targets[0]] assert.Equal(t, created, len(change.Create)) - assert.Equal(t, updated, len(change.UpdateNew)+len(change.UpdateOld)) + assert.Equal(t, updated, len(change.UpdateNew)) assert.Equal(t, deleted, len(change.Delete)) } diff --git a/fixtures/zones.bind b/sync/testdata/zones.bind similarity index 100% rename from fixtures/zones.bind rename to sync/testdata/zones.bind