diff --git a/go.mod b/go.mod index 8190ff71b..f7956827e 100644 --- a/go.mod +++ b/go.mod @@ -3,6 +3,7 @@ module go.probo.inc/probo go 1.25.5 require ( + codeberg.org/miekg/dns v0.6.2 github.com/99designs/gqlgen v0.17.83 github.com/aws/aws-sdk-go-v2 v1.40.0 github.com/aws/aws-sdk-go-v2/credentials v1.19.0 diff --git a/go.sum b/go.sum index b9de83e78..b7c68fcb4 100644 --- a/go.sum +++ b/go.sum @@ -1,3 +1,5 @@ +codeberg.org/miekg/dns v0.6.2 h1:gkuKad3tKCs4On/3ZlydASeXPc1uQaZNuDSxEXR7Mt0= +codeberg.org/miekg/dns v0.6.2/go.mod h1:IKSpRNHVdUyxHC457VnDQ9Dv05UhXXsALZAywCM7D54= github.com/99designs/gqlgen v0.17.83 h1:LZOd4Of2snK5V22/ZWfBAPa3WoAZkBO70dKXM0ODHQk= github.com/99designs/gqlgen v0.17.83/go.mod h1:q6Lb64wknFqNFSbSUGzKRKupklvY/xgNr62g0GGWPB8= github.com/PuerkitoBio/goquery v1.10.3 h1:pFYcNSqHxBD06Fpj/KsbStFRsgRATgnf3LeXiUkhzPo= diff --git a/pkg/certmanager/provisioner.go b/pkg/certmanager/provisioner.go index aabc9cc00..53ef33c85 100644 --- a/pkg/certmanager/provisioner.go +++ b/pkg/certmanager/provisioner.go @@ -18,7 +18,6 @@ import ( "context" "errors" "fmt" - "net" "strings" "time" @@ -29,6 +28,8 @@ import ( "go.probo.inc/probo/pkg/coredata" "go.probo.inc/probo/pkg/crypto/cipher" "go.probo.inc/probo/pkg/gid" + + "codeberg.org/miekg/dns" ) type ( @@ -38,6 +39,7 @@ type ( encryptionKey cipher.EncryptionKey cnameTarget string interval time.Duration + resolverAddr string logger *log.Logger } ) @@ -52,6 +54,7 @@ func NewProvisioner( encryptionKey cipher.EncryptionKey, cnameTarget string, interval time.Duration, + resolverAddr string, logger *log.Logger, ) *Provisioner { return &Provisioner{ @@ -60,6 +63,7 @@ func NewProvisioner( encryptionKey: encryptionKey, cnameTarget: cnameTarget, interval: interval, + resolverAddr: resolverAddr, logger: logger.Named("certmanager.provisioner"), } } @@ -88,19 +92,45 @@ func (p *Provisioner) Run(ctx context.Context) error { } func (p *Provisioner) checkDNSConfiguration(domain string) error { - cnameRecords, err := net.LookupCNAME(domain) - if err != nil { - return fmt.Errorf("cannot lookup cname for domain %q: %w", domain, err) + customerFQDN := domain + if !strings.HasSuffix(customerFQDN, ".") { + customerFQDN = customerFQDN + "." } - expectedTarget := strings.TrimSuffix(p.cnameTarget, ".") - actualTarget := strings.TrimSuffix(cnameRecords, ".") - if !strings.EqualFold(actualTarget, expectedTarget) { + expectedFQDN := p.cnameTarget + if !strings.HasSuffix(expectedFQDN, ".") { + expectedFQDN = expectedFQDN + "." + } + + msg := &dns.Msg{MsgHeader: dns.MsgHeader{ID: dns.ID(), RecursionDesired: true}} + msg.Question = []dns.RR{&dns.CNAME{Hdr: dns.Header{Name: customerFQDN, Class: dns.ClassINET}}} + + client := dns.NewClient() + + resp, _, err := client.Exchange(context.Background(), msg, "udp", p.resolverAddr) + if err != nil { + return fmt.Errorf("cannot exchange dns message: %w", err) + } + + if len(resp.Answer) == 0 { + return fmt.Errorf("no cname records found for domain %q", domain) + } + + if len(resp.Answer) > 1 { + return fmt.Errorf("multiple cname records found for domain %q", domain) + } + + resolvedRecord, ok := resp.Answer[0].(*dns.CNAME) + if !ok { + return fmt.Errorf("first answer is not a cname record for domain %q", domain) + } + + if !strings.EqualFold(expectedFQDN, resolvedRecord.Target) { return fmt.Errorf( "cname target mismatch: domain %q resolves to %q, expected %q", domain, - actualTarget, - p.cnameTarget, + resolvedRecord.Target, + expectedFQDN, ) } diff --git a/pkg/probod/custom_domains_config.go b/pkg/probod/custom_domains_config.go index 5c44724ac..17b0e5f1b 100644 --- a/pkg/probod/custom_domains_config.go +++ b/pkg/probod/custom_domains_config.go @@ -17,6 +17,7 @@ package probod type customDomainsConfig struct { RenewalInterval int `json:"renewal-interval"` ProvisionInterval int `json:"provision-interval"` + ResolverAddr string `json:"resolver-addr"` CnameTarget string `json:"cname-target"` ACME acmeConfig `json:"acme"` } diff --git a/pkg/probod/probod.go b/pkg/probod/probod.go index 5339d1738..9a2b14f3b 100644 --- a/pkg/probod/probod.go +++ b/pkg/probod/probod.go @@ -159,6 +159,7 @@ func New() *Implm { CustomDomains: customDomainsConfig{ RenewalInterval: 3600, ProvisionInterval: 30, + ResolverAddr: "8.8.8.8:53", ACME: acmeConfig{ Directory: "https://acme-v02.api.letsencrypt.org/directory", Email: "admin@getprobo.com", @@ -681,7 +682,7 @@ func (impl *Implm) runTrustCenterServer( if certProvisioningInterval == 0 { certProvisioningInterval = 30 * time.Second } - certProvisioner := certmanager.NewProvisioner(pgClient, acmeService, impl.cfg.EncryptionKey, impl.cfg.CustomDomains.CnameTarget, certProvisioningInterval, l) + certProvisioner := certmanager.NewProvisioner(pgClient, acmeService, impl.cfg.EncryptionKey, impl.cfg.CustomDomains.CnameTarget, certProvisioningInterval, impl.cfg.CustomDomains.ResolverAddr, l) g, ctx := errgroup.WithContext(ctx)