From f71d2c7f213244ecc3d4d705538d6343b0d99d0d Mon Sep 17 00:00:00 2001 From: Francisco Obispo Date: Wed, 27 May 2026 08:59:36 -0700 Subject: [PATCH 1/2] feat(dns): add DNS auto-provisioning via RFC 2136 dynamic updates Add port.DNSProvisioner interface and DDNS adapter for automatic DNS record management (TXT, TLSA, SVCB) using RFC 2136 with TSIG auth. When configured, registration issues certs, provisions DNS, and activates in a single call. Revoke auto-deletes records. Domain suffix config auto-qualifies short hostnames. SVCB records provide agent connectivity parameters per the DNS-AID proposal (RFC 9460). Includes the _ans-badge URL fix: badge records point to the TL badge endpoint when tl-client.public-base-url is configured. Signed-off-by: Francisco Obispo Signed-off-by: Francisco Obispo --- cmd/ans-ra/main.go | 30 +- config/ra-local.yaml | 13 + internal/adapter/dns/ddns.go | 241 ++++++++++++++++ internal/adapter/dns/ddns_test.go | 349 +++++++++++++++++++++++ internal/adapter/dns/noop_provisioner.go | 21 ++ internal/config/config.go | 94 +++++- internal/config/config_test.go | 58 ++++ internal/domain/dnsrecords.go | 70 ++++- internal/domain/dnsrecords_test.go | 62 ++++ internal/port/dns.go | 9 + internal/ra/handler/registration.go | 6 +- internal/ra/handler/v1registration.go | 3 +- internal/ra/service/helpers_test.go | 68 +++++ internal/ra/service/lifecycle.go | 25 +- internal/ra/service/registration.go | 199 ++++++++++--- internal/ra/service/registration_test.go | 109 +++++++ 16 files changed, 1298 insertions(+), 59 deletions(-) create mode 100644 internal/adapter/dns/ddns.go create mode 100644 internal/adapter/dns/ddns_test.go create mode 100644 internal/adapter/dns/noop_provisioner.go diff --git a/cmd/ans-ra/main.go b/cmd/ans-ra/main.go index b61a5b7..c02ea20 100644 --- a/cmd/ans-ra/main.go +++ b/cmd/ans-ra/main.go @@ -146,8 +146,14 @@ func run(cfgPath string) error { // remove WithSkipChainVerify in its config factory. validator := cert.NewX509Validator(cert.WithSkipChainVerify()) - // DNS verifier. + // DNS verifier + optional provisioner. var dnsVerifier = selectDNSVerifier(cfg) + var dnsProvisioner = selectDNSProvisioner(cfg) + + logger.Info(). + Str("tlPublicBaseURL", cfg.TLClient.PublicBaseURL). + Str("tlBaseURL", cfg.TLClient.BaseURL). + Msg("transparency log endpoints configured") logger.Info(). Str("tlPublicBaseURL", cfg.TLClient.PublicBaseURL). @@ -171,8 +177,10 @@ func run(cfgPath string) error { KeyID: signerKeyID, RaID: cfg.Signer.RaID, }).WithDNSVerifier(dnsVerifier). + WithDNSProvisioner(dnsProvisioner). WithServerCertificateAuthority(serverCA). - WithTLPublicBaseURL(cfg.TLClient.PublicBaseURL) + WithTLPublicBaseURL(cfg.TLClient.PublicBaseURL). + WithDomainSuffix(cfg.Registration.DomainSuffix) // HTTP. r := chi.NewRouter() @@ -419,3 +427,21 @@ func selectDNSVerifier(cfg *config.RAConfig) port.DNSVerifier { return dns.NewNoopVerifier() } } + +func selectDNSProvisioner(cfg *config.RAConfig) port.DNSProvisioner { + if cfg.DNS.Provisioner == nil || cfg.DNS.Provisioner.Type == "" { + return nil + } + switch cfg.DNS.Provisioner.Type { + case "ddns": + d := cfg.DNS.Provisioner.DDNS + return dns.NewDDNSProvisioner( + dns.WithDDNSServer(d.Server), + dns.WithDDNSZone(d.Zone), + dns.WithTSIG(d.TSIGName, d.TSIGSecret, d.TSIGAlgorithm), + dns.WithDDNSTimeout(d.Timeout), + ) + default: + return nil + } +} diff --git a/config/ra-local.yaml b/config/ra-local.yaml index f4e67b2..b816b8b 100644 --- a/config/ra-local.yaml +++ b/config/ra-local.yaml @@ -24,6 +24,19 @@ dns: # the resulting TL attestation. type: noop # server: "127.0.0.1:15353" + # + # provisioner: auto-create/delete DNS records via RFC 2136 DDNS. + # When configured, VerifyDNS provisions records before verifying + # them, and Revoke cleans them up. + # provisioner: + # type: ddns + # ddns: + # server: "127.0.0.1:53" + # zone: "example.com." + # tsig-name: "ans-updater." + # tsig-secret: "base64-encoded-secret" + # tsig-algorithm: "hmac-sha256" + # timeout: 5s keys: type: file diff --git a/internal/adapter/dns/ddns.go b/internal/adapter/dns/ddns.go new file mode 100644 index 0000000..2154dfb --- /dev/null +++ b/internal/adapter/dns/ddns.go @@ -0,0 +1,241 @@ +package dns + +import ( + "context" + "errors" + "fmt" + "math" + "strconv" + "strings" + "time" + + "github.com/godaddy/ans/internal/domain" + "github.com/miekg/dns" +) + +// DDNSProvisioner implements port.DNSProvisioner via RFC 2136 dynamic +// DNS updates with TSIG authentication. +type DDNSProvisioner struct { + server string + zone string + tsigName string + tsigSecret string + tsigAlgorithm string + timeout time.Duration +} + +type DDNSOption func(*DDNSProvisioner) + +func WithDDNSServer(addr string) DDNSOption { + return func(p *DDNSProvisioner) { p.server = addr } +} + +func WithDDNSZone(zone string) DDNSOption { + return func(p *DDNSProvisioner) { p.zone = dns.Fqdn(zone) } +} + +func WithTSIG(name, secret, algorithm string) DDNSOption { + return func(p *DDNSProvisioner) { + p.tsigName = dns.Fqdn(name) + p.tsigSecret = secret + if algorithm != "" { + p.tsigAlgorithm = dns.Fqdn(algorithm) + } + } +} + +func WithDDNSTimeout(d time.Duration) DDNSOption { + return func(p *DDNSProvisioner) { p.timeout = d } +} + +func NewDDNSProvisioner(opts ...DDNSOption) *DDNSProvisioner { + p := &DDNSProvisioner{ + tsigAlgorithm: dns.HmacSHA256, + timeout: 5 * time.Second, + } + for _, o := range opts { + o(p) + } + return p +} + +func (p *DDNSProvisioner) ProvisionRecords(ctx context.Context, _ string, records []domain.ExpectedDNSRecord) error { + if len(records) == 0 { + return nil + } + + msg := new(dns.Msg) + msg.SetUpdate(p.zone) + + for _, rec := range records { + rrs, err := p.toRR(rec) + if err != nil { + return fmt.Errorf("build RR for %s %s: %w", rec.Name, rec.Type, err) + } + // Remove-then-insert gives replace semantics (idempotent). + msg.RemoveRRset(rrs) + msg.Insert(rrs) + } + + return p.exchange(ctx, msg) +} + +func (p *DDNSProvisioner) DeleteRecords(ctx context.Context, _ string, records []domain.ExpectedDNSRecord) error { + if len(records) == 0 { + return nil + } + + msg := new(dns.Msg) + msg.SetUpdate(p.zone) + + for _, rec := range records { + rrs, err := p.toRR(rec) + if err != nil { + return fmt.Errorf("build RR for %s %s: %w", rec.Name, rec.Type, err) + } + msg.RemoveRRset(rrs) + } + + return p.exchange(ctx, msg) +} + +func (p *DDNSProvisioner) exchange(ctx context.Context, msg *dns.Msg) error { + client := new(dns.Client) + client.Timeout = p.timeout + + if p.tsigName != "" && p.tsigSecret != "" { + client.TsigSecret = map[string]string{p.tsigName: p.tsigSecret} + msg.SetTsig(p.tsigName, p.tsigAlgorithm, 300, time.Now().Unix()) + } + + resp, _, err := client.ExchangeContext(ctx, msg, p.server) + if err != nil { + return fmt.Errorf("dns update exchange: %w", err) + } + if resp.Rcode != dns.RcodeSuccess { + return fmt.Errorf("dns update failed: %s", dns.RcodeToString[resp.Rcode]) + } + return nil +} + +func (p *DDNSProvisioner) toRR(rec domain.ExpectedDNSRecord) ([]dns.RR, error) { + name := dns.Fqdn(rec.Name) + ttl := uint32(3600) + if rec.TTL > 0 && rec.TTL <= math.MaxUint32 { + ttl = uint32(rec.TTL) + } + + switch rec.Type { + case domain.DNSRecordTXT: + rr := &dns.TXT{ + Hdr: dns.RR_Header{Name: name, Rrtype: dns.TypeTXT, Class: dns.ClassINET, Ttl: ttl}, + Txt: splitTXT(rec.Value), + } + return []dns.RR{rr}, nil + + case domain.DNSRecordTLSA: + tlsa, err := parseTLSA(rec.Value) + if err != nil { + return nil, err + } + tlsa.Hdr = dns.RR_Header{Name: name, Rrtype: dns.TypeTLSA, Class: dns.ClassINET, Ttl: ttl} + return []dns.RR{tlsa}, nil + + case domain.DNSRecordSVCB: + svcb, err := parseSVCBValue(rec.Value) + if err != nil { + return nil, err + } + svcb.Hdr = dns.RR_Header{Name: name, Rrtype: dns.TypeSVCB, Class: dns.ClassINET, Ttl: ttl} + return []dns.RR{svcb}, nil + + case domain.DNSRecordHTTPS: + return nil, errors.New("HTTPS record provisioning not supported") + + default: + return nil, fmt.Errorf("unsupported record type: %s", rec.Type) + } +} + +// splitTXT splits a TXT value into 255-byte chunks per RFC 1035 §3.3.14. +func splitTXT(val string) []string { + if len(val) <= 255 { + return []string{val} + } + var chunks []string + for len(val) > 0 { + end := 255 + if end > len(val) { + end = len(val) + } + chunks = append(chunks, val[:end]) + val = val[end:] + } + return chunks +} + +// parseTLSA parses "usage selector mtype hex" into a dns.TLSA RR. +func parseTLSA(val string) (*dns.TLSA, error) { + parts := strings.Fields(val) + if len(parts) != 4 { + return nil, fmt.Errorf("TLSA value must have 4 fields, got %d: %q", len(parts), val) + } + usage, err := strconv.ParseUint(parts[0], 10, 8) + if err != nil { + return nil, fmt.Errorf("TLSA usage: %w", err) + } + selector, err := strconv.ParseUint(parts[1], 10, 8) + if err != nil { + return nil, fmt.Errorf("TLSA selector: %w", err) + } + matchingType, err := strconv.ParseUint(parts[2], 10, 8) + if err != nil { + return nil, fmt.Errorf("TLSA matching type: %w", err) + } + return &dns.TLSA{ + Usage: uint8(usage), + Selector: uint8(selector), + MatchingType: uint8(matchingType), + Certificate: strings.ToLower(parts[3]), + }, nil +} + +// parseSVCBValue parses "priority target key=val ..." into a dns.SVCB RR. +func parseSVCBValue(val string) (*dns.SVCB, error) { + parts := strings.Fields(val) + if len(parts) < 2 { + return nil, fmt.Errorf("SVCB value must have at least priority and target, got: %q", val) + } + + priority, err := strconv.ParseUint(parts[0], 10, 16) + if err != nil { + return nil, fmt.Errorf("SVCB priority: %w", err) + } + + target := parts[1] + + var params []dns.SVCBKeyValue + for _, kv := range parts[2:] { + eqIdx := strings.IndexByte(kv, '=') + if eqIdx < 0 { + continue + } + key, value := kv[:eqIdx], kv[eqIdx+1:] + switch key { + case "alpn": + params = append(params, &dns.SVCBAlpn{Alpn: strings.Split(value, ",")}) + case "port": + p, err := strconv.ParseUint(value, 10, 16) + if err != nil { + return nil, fmt.Errorf("SVCB port: %w", err) + } + params = append(params, &dns.SVCBPort{Port: uint16(p)}) + } + } + + return &dns.SVCB{ + Priority: uint16(priority), + Target: target, + Value: params, + }, nil +} diff --git a/internal/adapter/dns/ddns_test.go b/internal/adapter/dns/ddns_test.go new file mode 100644 index 0000000..9cb798c --- /dev/null +++ b/internal/adapter/dns/ddns_test.go @@ -0,0 +1,349 @@ +package dns + +import ( + "context" + "net" + "sync" + "testing" + "time" + + "github.com/godaddy/ans/internal/domain" + "github.com/miekg/dns" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// ddnsTestServer is a minimal DNS server that accepts RFC 2136 UPDATE +// messages and records them for assertion. +type ddnsTestServer struct { + addr string + updates []*dns.Msg + mu sync.Mutex +} + +func newDDNSTestServer(t *testing.T) *ddnsTestServer { + t.Helper() + + s := &ddnsTestServer{} + + // Use a raw UDP listener and handle DNS manually to avoid + // miekg/dns server's default NOTIMP response for UPDATE opcodes. + pc, err := net.ListenPacket("udp", "127.0.0.1:0") + require.NoError(t, err) + s.addr = pc.LocalAddr().String() + + go func() { + buf := make([]byte, 65535) + for { + n, addr, rerr := pc.ReadFrom(buf) + if rerr != nil { + return + } + req := new(dns.Msg) + if uerr := req.Unpack(buf[:n]); uerr != nil { + continue + } + + s.mu.Lock() + s.updates = append(s.updates, req.Copy()) + s.mu.Unlock() + + resp := new(dns.Msg) + resp.Id = req.Id + resp.Response = true + resp.Opcode = req.Opcode + resp.Rcode = dns.RcodeSuccess + out, _ := resp.Pack() + _, _ = pc.WriteTo(out, addr) + } + }() + t.Cleanup(func() { _ = pc.Close() }) + return s +} + +func (s *ddnsTestServer) lastUpdate() *dns.Msg { + s.mu.Lock() + defer s.mu.Unlock() + if len(s.updates) == 0 { + return nil + } + return s.updates[len(s.updates)-1] +} + +func TestDDNSProvisioner_ProvisionTXTRecords(t *testing.T) { + t.Parallel() + srv := newDDNSTestServer(t) + p := NewDDNSProvisioner( + WithDDNSServer(srv.addr), + WithDDNSZone("example.com."), + ) + + records := []domain.ExpectedDNSRecord{ + {Name: "_ans.agent.example.com", Type: domain.DNSRecordTXT, Value: "v=ans1; version=1.0.0", TTL: 3600}, + {Name: "_ans-badge.agent.example.com", Type: domain.DNSRecordTXT, Value: "v=ans-badge1; version=1.0.0; url=https://tl.example.com/v1/agents/123", TTL: 3600}, + } + + err := p.ProvisionRecords(context.Background(), "agent.example.com", records) + require.NoError(t, err) + + msg := srv.lastUpdate() + require.NotNil(t, msg) + assert.True(t, msg.Opcode == dns.OpcodeUpdate, "message should be an UPDATE") + + // Should have both remove and insert sections for each record. + assert.NotEmpty(t, msg.Ns, "NS (update) section should contain RRs") +} + +func TestDDNSProvisioner_ProvisionTLSARecord(t *testing.T) { + t.Parallel() + srv := newDDNSTestServer(t) + p := NewDDNSProvisioner( + WithDDNSServer(srv.addr), + WithDDNSZone("example.com."), + ) + + records := []domain.ExpectedDNSRecord{ + {Name: "_443._tcp.agent.example.com", Type: domain.DNSRecordTLSA, Value: "3 1 1 abcdef0123456789", TTL: 3600}, + } + + err := p.ProvisionRecords(context.Background(), "agent.example.com", records) + require.NoError(t, err) + + msg := srv.lastUpdate() + require.NotNil(t, msg) + assert.True(t, msg.Opcode == dns.OpcodeUpdate) +} + +func TestDDNSProvisioner_DeleteRecords(t *testing.T) { + t.Parallel() + srv := newDDNSTestServer(t) + p := NewDDNSProvisioner( + WithDDNSServer(srv.addr), + WithDDNSZone("example.com."), + ) + + records := []domain.ExpectedDNSRecord{ + {Name: "_ans.agent.example.com", Type: domain.DNSRecordTXT, Value: "v=ans1; version=1.0.0", TTL: 3600}, + } + + err := p.DeleteRecords(context.Background(), "agent.example.com", records) + require.NoError(t, err) + + msg := srv.lastUpdate() + require.NotNil(t, msg) + assert.True(t, msg.Opcode == dns.OpcodeUpdate) +} + +func TestDDNSProvisioner_EmptyRecordsNoOp(t *testing.T) { + t.Parallel() + srv := newDDNSTestServer(t) + p := NewDDNSProvisioner( + WithDDNSServer(srv.addr), + WithDDNSZone("example.com."), + ) + + err := p.ProvisionRecords(context.Background(), "agent.example.com", nil) + require.NoError(t, err) + assert.Nil(t, srv.lastUpdate(), "no UPDATE should be sent for empty records") + + err = p.DeleteRecords(context.Background(), "agent.example.com", nil) + require.NoError(t, err) +} + +func TestDDNSProvisioner_Idempotent(t *testing.T) { + t.Parallel() + srv := newDDNSTestServer(t) + p := NewDDNSProvisioner( + WithDDNSServer(srv.addr), + WithDDNSZone("example.com."), + ) + + records := []domain.ExpectedDNSRecord{ + {Name: "_ans.agent.example.com", Type: domain.DNSRecordTXT, Value: "v=ans1", TTL: 3600}, + } + + require.NoError(t, p.ProvisionRecords(context.Background(), "agent.example.com", records)) + require.NoError(t, p.ProvisionRecords(context.Background(), "agent.example.com", records)) + + srv.mu.Lock() + count := len(srv.updates) + srv.mu.Unlock() + assert.Equal(t, 2, count, "two UPDATE messages sent (both succeed — server-side idempotent)") +} + +func TestDDNSProvisioner_ServerUnreachable(t *testing.T) { + t.Parallel() + p := NewDDNSProvisioner( + WithDDNSServer("127.0.0.1:1"), + WithDDNSTimeout(100*time.Millisecond), + WithDDNSZone("example.com."), + ) + + records := []domain.ExpectedDNSRecord{ + {Name: "_ans.agent.example.com", Type: domain.DNSRecordTXT, Value: "v=ans1", TTL: 3600}, + } + + err := p.ProvisionRecords(context.Background(), "agent.example.com", records) + assert.Error(t, err) +} + +func TestDDNSProvisioner_WithTSIG(t *testing.T) { + t.Parallel() + srv := newDDNSTestServer(t) + p := NewDDNSProvisioner( + WithDDNSServer(srv.addr), + WithDDNSZone("example.com."), + WithTSIG("ans-key.", "c2VjcmV0", "hmac-sha256"), + ) + + records := []domain.ExpectedDNSRecord{ + {Name: "_ans.agent.example.com", Type: domain.DNSRecordTXT, Value: "v=ans1", TTL: 3600}, + } + + err := p.ProvisionRecords(context.Background(), "agent.example.com", records) + require.NoError(t, err) + + msg := srv.lastUpdate() + require.NotNil(t, msg) + assert.NotEmpty(t, msg.Extra, "TSIG should be in the additional section") +} + +func TestDDNSProvisioner_UnsupportedHTTPS(t *testing.T) { + t.Parallel() + srv := newDDNSTestServer(t) + p := NewDDNSProvisioner( + WithDDNSServer(srv.addr), + WithDDNSZone("example.com."), + ) + + records := []domain.ExpectedDNSRecord{ + {Name: "agent.example.com", Type: domain.DNSRecordHTTPS, Value: "1 . alpn=h2", TTL: 3600}, + } + + err := p.ProvisionRecords(context.Background(), "agent.example.com", records) + assert.Error(t, err) + assert.Contains(t, err.Error(), "HTTPS record provisioning not supported") +} + +func TestDDNSProvisioner_ProvisionSVCBRecord(t *testing.T) { + t.Parallel() + srv := newDDNSTestServer(t) + p := NewDDNSProvisioner( + WithDDNSServer(srv.addr), + WithDDNSZone("example.com."), + ) + + records := []domain.ExpectedDNSRecord{ + {Name: "agent.example.com", Type: domain.DNSRecordSVCB, Value: "1 real-host.example.com. alpn=h2 port=8080", TTL: 3600}, + } + + err := p.ProvisionRecords(context.Background(), "agent.example.com", records) + require.NoError(t, err) + + msg := srv.lastUpdate() + require.NotNil(t, msg) + assert.True(t, msg.Opcode == dns.OpcodeUpdate) +} + +func TestParseSVCBValue(t *testing.T) { + svcb, err := parseSVCBValue("1 host.example.com. alpn=h2 port=8080") + require.NoError(t, err) + assert.Equal(t, uint16(1), svcb.Priority) + assert.Equal(t, "host.example.com.", svcb.Target) + assert.Len(t, svcb.Value, 2) +} + +func TestParseSVCBValue_AlpnOnly(t *testing.T) { + svcb, err := parseSVCBValue("1 host.example.com. alpn=h2") + require.NoError(t, err) + assert.Equal(t, "host.example.com.", svcb.Target) + assert.Len(t, svcb.Value, 1) +} + +func TestParseSVCBValue_Invalid(t *testing.T) { + _, err := parseSVCBValue("1") + assert.Error(t, err) +} + +func TestParseTLSA(t *testing.T) { + tlsa, err := parseTLSA("3 1 1 abcdef0123456789") + require.NoError(t, err) + assert.Equal(t, uint8(3), tlsa.Usage) + assert.Equal(t, uint8(1), tlsa.Selector) + assert.Equal(t, uint8(1), tlsa.MatchingType) + assert.Equal(t, "abcdef0123456789", tlsa.Certificate) +} + +func TestParseTLSA_Invalid(t *testing.T) { + _, err := parseTLSA("3 1") + assert.Error(t, err) +} + +func TestDDNSProvisioner_UnsupportedRecordType(t *testing.T) { + t.Parallel() + srv := newDDNSTestServer(t) + p := NewDDNSProvisioner( + WithDDNSServer(srv.addr), + WithDDNSZone("example.com."), + ) + + records := []domain.ExpectedDNSRecord{ + {Name: "agent.example.com", Type: domain.DNSRecordType("BOGUS"), Value: "x", TTL: 3600}, + } + + err := p.ProvisionRecords(context.Background(), "agent.example.com", records) + assert.Error(t, err) + assert.Contains(t, err.Error(), "unsupported record type") +} + +func TestDDNSProvisioner_ZeroTTLDefaultsTo3600(t *testing.T) { + t.Parallel() + srv := newDDNSTestServer(t) + p := NewDDNSProvisioner( + WithDDNSServer(srv.addr), + WithDDNSZone("example.com."), + ) + + records := []domain.ExpectedDNSRecord{ + {Name: "_ans.agent.example.com", Type: domain.DNSRecordTXT, Value: "v=ans1", TTL: 0}, + } + + err := p.ProvisionRecords(context.Background(), "agent.example.com", records) + require.NoError(t, err) + + msg := srv.lastUpdate() + require.NotNil(t, msg) + for _, rr := range msg.Ns { + if rr.Header().Ttl != 0 { + assert.Equal(t, uint32(3600), rr.Header().Ttl, "zero TTL should default to 3600") + } + } +} + +func TestParseTLSA_NonNumericFields(t *testing.T) { + _, err := parseTLSA("x 1 1 abcdef") + assert.Error(t, err) +} + +func TestParseSVCBValue_UnknownParam(t *testing.T) { + svcb, err := parseSVCBValue("1 host.example.com. alpn=h2 unknown=value") + require.NoError(t, err) + assert.Equal(t, "host.example.com.", svcb.Target) +} + +func TestNoopProvisioner(t *testing.T) { + p := NewNoopProvisioner() + assert.NoError(t, p.ProvisionRecords(context.Background(), "x", nil)) + assert.NoError(t, p.DeleteRecords(context.Background(), "x", nil)) +} + +func TestSplitTXT(t *testing.T) { + short := "hello" + assert.Equal(t, []string{"hello"}, splitTXT(short)) + + long := string(make([]byte, 300)) + chunks := splitTXT(long) + assert.Equal(t, 2, len(chunks)) + assert.Equal(t, 255, len(chunks[0])) + assert.Equal(t, 45, len(chunks[1])) +} diff --git a/internal/adapter/dns/noop_provisioner.go b/internal/adapter/dns/noop_provisioner.go new file mode 100644 index 0000000..2c90eed --- /dev/null +++ b/internal/adapter/dns/noop_provisioner.go @@ -0,0 +1,21 @@ +package dns + +import ( + "context" + + "github.com/godaddy/ans/internal/domain" +) + +// NoopProvisioner is a no-op DNSProvisioner for dev and test +// environments. It accepts any input and always succeeds. +type NoopProvisioner struct{} + +func NewNoopProvisioner() *NoopProvisioner { return &NoopProvisioner{} } + +func (NoopProvisioner) ProvisionRecords(_ context.Context, _ string, _ []domain.ExpectedDNSRecord) error { + return nil +} + +func (NoopProvisioner) DeleteRecords(_ context.Context, _ string, _ []domain.ExpectedDNSRecord) error { + return nil +} diff --git a/internal/config/config.go b/internal/config/config.go index 4a32a87..6360e70 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -92,7 +92,7 @@ type CAServerSelf struct { DataDir string `koanf:"data-dir"` } -// DNS holds DNS verifier configuration. +// DNS holds DNS verifier and optional provisioner configuration. type DNS struct { // Type selects the verifier adapter. "noop" accepts any DNS // state; "lookup" queries a real nameserver via miekg/dns and @@ -106,6 +106,26 @@ type DNS struct { // server (e.g. "127.0.0.1:15353") for self-contained local // testing. Server string `koanf:"server"` + + // Provisioner configures automatic DNS record creation/deletion. + // When nil or empty, operators manage DNS manually (the default). + Provisioner *DNSProvisionerConfig `koanf:"provisioner"` +} + +// DNSProvisionerConfig selects the DNS provisioning adapter. +type DNSProvisionerConfig struct { + Type string `koanf:"type"` // "ddns" + DDNS *DNSDDNSConfig `koanf:"ddns"` +} + +// DNSDDNSConfig holds RFC 2136 dynamic DNS update settings. +type DNSDDNSConfig struct { + Server string `koanf:"server"` // "host:port" of authoritative nameserver + Zone string `koanf:"zone"` // zone name (FQDN, e.g. "obispo.link.") + TSIGName string `koanf:"tsig-name"` // TSIG key name (FQDN) + TSIGSecret string `koanf:"tsig-secret"` // base64-encoded shared secret + TSIGAlgorithm string `koanf:"tsig-algorithm"` // default: hmac-sha256 + Timeout time.Duration `koanf:"timeout"` // default: 5s } // Keys holds key-manager configuration. @@ -188,17 +208,27 @@ type Log struct { Format string `koanf:"format"` // "text" | "json" } +// Registration holds agent registration policy settings. +type Registration struct { + // DomainSuffix is appended to the agentHost submitted by the + // caller. When set, the agent sends a short name (e.g. "my-agent") + // and the RA constructs the full FQDN ("my-agent.agents.example.com"). + // When empty, the caller must provide the full FQDN. + DomainSuffix string `koanf:"domain-suffix"` +} + // RAConfig is the full configuration for ans-ra. type RAConfig struct { - Server Server `koanf:"server"` - Auth Auth `koanf:"auth"` - CA CA `koanf:"ca"` - DNS DNS `koanf:"dns"` - Keys Keys `koanf:"keys"` - Store Store `koanf:"store"` - TLClient TLClient `koanf:"tl-client"` - Signer SignerCfg `koanf:"signer"` - Log Log `koanf:"log"` + Server Server `koanf:"server"` + Auth Auth `koanf:"auth"` + CA CA `koanf:"ca"` + DNS DNS `koanf:"dns"` + Keys Keys `koanf:"keys"` + Store Store `koanf:"store"` + TLClient TLClient `koanf:"tl-client"` + Signer SignerCfg `koanf:"signer"` + Registration Registration `koanf:"registration"` + Log Log `koanf:"log"` } // SignerCfg names the KeyManager-managed key the RA uses to sign @@ -338,6 +368,10 @@ func (c *RAConfig) Validate() error { default: return fmt.Errorf("dns.type %q not supported (expected 'noop' or 'lookup')", c.DNS.Type) } + applyDNSProvisionerDefaults(c.DNS.Provisioner) + if err := validateDNSProvisioner(c.DNS.Provisioner); err != nil { + return err + } if err := validateKeys(&c.Keys); err != nil { return err } @@ -421,6 +455,46 @@ func validateStore(s *Store) error { return nil } +// applyDNSProvisionerDefaults fills in default values for optional +// DDNS provisioner fields. Called before validation so that Validate() +// itself is side-effect-free. +func applyDNSProvisionerDefaults(p *DNSProvisionerConfig) { + if p == nil || p.Type != "ddns" || p.DDNS == nil { + return + } + if p.DDNS.TSIGAlgorithm == "" { + p.DDNS.TSIGAlgorithm = "hmac-sha256" + } + if p.DDNS.Timeout <= 0 { + p.DDNS.Timeout = 5 * time.Second + } +} + +func validateDNSProvisioner(p *DNSProvisionerConfig) error { + if p == nil || p.Type == "" { + return nil + } + switch p.Type { + case "ddns": + if p.DDNS == nil { + return errors.New("dns.provisioner.ddns block required when dns.provisioner.type is 'ddns'") + } + d := p.DDNS + if d.Server == "" { + return errors.New("dns.provisioner.ddns.server is required") + } + if d.Zone == "" { + return errors.New("dns.provisioner.ddns.zone is required") + } + if d.TSIGName == "" || d.TSIGSecret == "" { + return errors.New("dns.provisioner.ddns.tsig-name and tsig-secret are required") + } + default: + return fmt.Errorf("dns.provisioner.type %q not supported (expected 'ddns')", p.Type) + } + return nil +} + func validatePublicBaseURL(raw string) error { if raw == "" { return errors.New("tl-client.public-base-url is required") diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 5face48..60b47e6 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -227,6 +227,36 @@ func TestRAConfig_Validate_Errors(t *testing.T) { {"unsupported store.type", func(c *RAConfig) { c.Store.Type = "postgres" }, "store.type"}, {"missing store.sqlite.path", func(c *RAConfig) { c.Store.SQLite.Path = "" }, "store.sqlite.path"}, {"missing tl-client.base-url", func(c *RAConfig) { c.TLClient.BaseURL = "" }, "tl-client.base-url"}, + {"unsupported dns.provisioner.type", func(c *RAConfig) { + c.DNS.Provisioner = &DNSProvisionerConfig{Type: "bogus"} + }, "dns.provisioner.type"}, + {"ddns missing block", func(c *RAConfig) { + c.DNS.Provisioner = &DNSProvisionerConfig{Type: "ddns"} + }, "dns.provisioner.ddns block"}, + {"ddns missing server", func(c *RAConfig) { + c.DNS.Provisioner = &DNSProvisionerConfig{Type: "ddns", DDNS: &DNSDDNSConfig{Zone: "z.", TSIGName: "k.", TSIGSecret: "s"}} + }, "ddns.server"}, + {"ddns missing zone", func(c *RAConfig) { + c.DNS.Provisioner = &DNSProvisionerConfig{Type: "ddns", DDNS: &DNSDDNSConfig{Server: "s:53", TSIGName: "k.", TSIGSecret: "s"}} + }, "ddns.zone"}, + {"ddns missing tsig", func(c *RAConfig) { + c.DNS.Provisioner = &DNSProvisionerConfig{Type: "ddns", DDNS: &DNSDDNSConfig{Server: "s:53", Zone: "z."}} + }, "tsig-name"}, + {"public-base-url http scheme", func(c *RAConfig) { + c.TLClient.PublicBaseURL = "http://tl.example.com" + }, "https scheme"}, + {"public-base-url with userinfo", func(c *RAConfig) { + c.TLClient.PublicBaseURL = "https://user:pass@tl.example.com" + }, "userinfo not allowed"}, + {"public-base-url with query", func(c *RAConfig) { + c.TLClient.PublicBaseURL = "https://tl.example.com?foo=bar" + }, "query string not allowed"}, + {"public-base-url with fragment", func(c *RAConfig) { + c.TLClient.PublicBaseURL = "https://tl.example.com#frag" + }, "fragment not allowed"}, + {"public-base-url empty", func(c *RAConfig) { + c.TLClient.PublicBaseURL = "" + }, "public-base-url is required"}, } for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { @@ -243,6 +273,34 @@ func TestRAConfig_Validate_Errors(t *testing.T) { } } +func TestRAConfig_Validate_DDNSDefaults(t *testing.T) { + dir := t.TempDir() + c := defaultRAConfig() + c.Auth.Static = &AuthStatic{APIKey: "x"} + c.CA.Self.DataDir = dir + c.Keys.File.Path = dir + c.Store.SQLite.Path = filepath.Join(dir, "db") + c.DNS.Provisioner = &DNSProvisionerConfig{ + Type: "ddns", + DDNS: &DNSDDNSConfig{ + Server: "127.0.0.1:53", + Zone: "example.com.", + TSIGName: "ans-key.", + TSIGSecret: "c2VjcmV0", + }, + } + + if err := c.Validate(); err != nil { + t.Fatalf("validate: %v", err) + } + if c.DNS.Provisioner.DDNS.TSIGAlgorithm != "hmac-sha256" { + t.Errorf("TSIGAlgorithm default not applied: got %q", c.DNS.Provisioner.DDNS.TSIGAlgorithm) + } + if c.DNS.Provisioner.DDNS.Timeout != 5*time.Second { + t.Errorf("Timeout default not applied: got %v", c.DNS.Provisioner.DDNS.Timeout) + } +} + func TestRAConfig_Validate_AppliesDefaultTLClientTimeout(t *testing.T) { dir := t.TempDir() c := defaultRAConfig() diff --git a/internal/domain/dnsrecords.go b/internal/domain/dnsrecords.go index 4f0cb81..a3eecb0 100644 --- a/internal/domain/dnsrecords.go +++ b/internal/domain/dnsrecords.go @@ -3,6 +3,7 @@ package domain import ( "fmt" "net/url" + "strconv" ) // DNSRecordType represents a DNS record type. @@ -12,6 +13,7 @@ const ( DNSRecordTXT DNSRecordType = "TXT" DNSRecordTLSA DNSRecordType = "TLSA" DNSRecordHTTPS DNSRecordType = "HTTPS" + DNSRecordSVCB DNSRecordType = "SVCB" ) // DNSRecordPurpose describes why a DNS record is needed. @@ -22,6 +24,7 @@ const ( PurposeTrust DNSRecordPurpose = "TRUST" PurposeCertificateBinding DNSRecordPurpose = "CERTIFICATE_BINDING" PurposeBadge DNSRecordPurpose = "BADGE" + PurposeConnectivity DNSRecordPurpose = "CONNECTIVITY" ) // ExpectedDNSRecord represents a DNS record the operator must configure. @@ -105,21 +108,68 @@ func ComputeRequiredDNSRecords(reg *AgentRegistration, tlPublicBaseURL string) [ // fingerprint (otherwise an attacker rewrote the record in a // signed zone — the worst failure mode). That post-verify // check lives alongside the verifier, not in the record set. - if reg.ServerCert == nil { - return records + if reg.ServerCert != nil { + records = append(records, ExpectedDNSRecord{ + Name: fmt.Sprintf("_443._tcp.%s", fqdn), + Type: DNSRecordTLSA, + Value: fmt.Sprintf("3 1 1 %s", reg.ServerCert.Fingerprint), + Purpose: PurposeCertificateBinding, + Required: false, + TTL: 3600, + }) + } + + // SVCB record for each endpoint — connectivity parameters. + // Clients query the agent's FQDN for SVCB to get the target + // hostname, port, and ALPN, then resolve the target separately. + for _, ep := range reg.Endpoints { + svcb := buildSVCBValue(ep.AgentURL) + if svcb == "" { + continue + } + records = append(records, ExpectedDNSRecord{ + Name: fqdn, + Type: DNSRecordSVCB, + Value: svcb, + Purpose: PurposeConnectivity, + Required: true, + TTL: 3600, + }) } - records = append(records, ExpectedDNSRecord{ - Name: fmt.Sprintf("_443._tcp.%s", fqdn), - Type: DNSRecordTLSA, - Value: fmt.Sprintf("3 1 1 %s", reg.ServerCert.Fingerprint), - Purpose: PurposeCertificateBinding, - Required: false, - TTL: 3600, - }) return records } +// buildSVCBValue parses an agent endpoint URL and returns the SVCB +// record value in wire-presentation format: "priority target key=val ...". +func buildSVCBValue(agentURL string) string { + u, err := url.Parse(agentURL) + if err != nil || u.Hostname() == "" { + return "" + } + + target := u.Hostname() + alpn := "h2" + if u.Scheme == "http" { + alpn = "http/1.1" + } + + port := u.Port() + defaultPort := "443" + if u.Scheme == "http" { + defaultPort = "80" + } + + if port == "" || port == defaultPort { + return fmt.Sprintf("1 %s. alpn=%s", target, alpn) + } + portNum, err := strconv.Atoi(port) + if err != nil { + return "" + } + return fmt.Sprintf("1 %s. alpn=%s port=%d", target, alpn, portNum) +} + func protocolToANSValue(p Protocol) string { switch p { case ProtocolA2A: diff --git a/internal/domain/dnsrecords_test.go b/internal/domain/dnsrecords_test.go index d5cd812..cf5bb6e 100644 --- a/internal/domain/dnsrecords_test.go +++ b/internal/domain/dnsrecords_test.go @@ -131,6 +131,68 @@ func TestComputeRequiredDNSRecords_BadgeFallbackWithoutTLURL(t *testing.T) { t.Fatal("no badge record found") } +func TestComputeRequiredDNSRecords_SVCBFromEndpoint(t *testing.T) { + ansName, _ := NewAnsName(mustSemVer(1, 0, 0), "agent.example.com") + reg := &AgentRegistration{ + AnsName: ansName, + Endpoints: []AgentEndpoint{ + {Protocol: ProtocolA2A, AgentURL: "https://real-host.example.com:8080/a2a"}, + }, + } + + records := ComputeRequiredDNSRecords(reg, "") + var svcbFound bool + for _, r := range records { + if r.Purpose == PurposeConnectivity { + svcbFound = true + assert.Equal(t, DNSRecordSVCB, r.Type) + assert.Equal(t, "agent.example.com", r.Name) + assert.Contains(t, r.Value, "real-host.example.com.") + assert.Contains(t, r.Value, "alpn=h2") + assert.Contains(t, r.Value, "port=8080") + assert.True(t, r.Required) + } + } + assert.True(t, svcbFound, "SVCB record should be generated") +} + +func TestComputeRequiredDNSRecords_SVCBDefaultPort(t *testing.T) { + ansName, _ := NewAnsName(mustSemVer(1, 0, 0), "agent.example.com") + reg := &AgentRegistration{ + AnsName: ansName, + Endpoints: []AgentEndpoint{ + {Protocol: ProtocolMCP, AgentURL: "https://real-host.example.com/mcp"}, + }, + } + + records := ComputeRequiredDNSRecords(reg, "") + for _, r := range records { + if r.Purpose == PurposeConnectivity { + assert.Contains(t, r.Value, "real-host.example.com.") + assert.Contains(t, r.Value, "alpn=h2") + assert.NotContains(t, r.Value, "port=", "default port 443 should be omitted") + return + } + } + t.Fatal("no SVCB record found") +} + +func TestBuildSVCBValue(t *testing.T) { + tests := []struct { + url string + expected string + }{ + {"https://host.example.com/a2a", "1 host.example.com. alpn=h2"}, + {"https://host.example.com:8080/mcp", "1 host.example.com. alpn=h2 port=8080"}, + {"http://host.example.com:9000/api", "1 host.example.com. alpn=http/1.1 port=9000"}, + {"https://host.example.com:443/default", "1 host.example.com. alpn=h2"}, + {"", ""}, + } + for _, tc := range tests { + assert.Equal(t, tc.expected, buildSVCBValue(tc.url), "url=%s", tc.url) + } +} + func TestProtocolToANSValue(t *testing.T) { assert.Equal(t, "a2a", protocolToANSValue(ProtocolA2A)) assert.Equal(t, "mcp", protocolToANSValue(ProtocolMCP)) diff --git a/internal/port/dns.go b/internal/port/dns.go index 03d96f2..6d0060a 100644 --- a/internal/port/dns.go +++ b/internal/port/dns.go @@ -28,6 +28,15 @@ type VerificationResult struct { Results []RecordVerification } +// DNSProvisioner creates and deletes DNS records on behalf of the RA. +// Implementations must be idempotent: provisioning a record that already +// exists with the correct value is a no-op; deleting a record that does +// not exist succeeds silently. +type DNSProvisioner interface { + ProvisionRecords(ctx context.Context, fqdn string, records []domain.ExpectedDNSRecord) error + DeleteRecords(ctx context.Context, fqdn string, records []domain.ExpectedDNSRecord) error +} + // DNSVerifier checks that the operator's DNS zone contains the records // the domain expects. It does NOT create, update, or delete records — // ANS is verification-only. Operators manage their own DNS. diff --git a/internal/ra/handler/registration.go b/internal/ra/handler/registration.go index 51656bc..15b6978 100644 --- a/internal/ra/handler/registration.go +++ b/internal/ra/handler/registration.go @@ -125,13 +125,15 @@ func (h *RegistrationHandler) Register(w http.ResponseWriter, r *http.Request) { return } - // Parse version + host into an AnsName. + // Parse version + host into an AnsName. The service may append + // a domain suffix to the host (e.g. "my-agent" → "my-agent.agents.example.com"). + qualifiedHost := h.svc.QualifyHost(req.AgentHost) semver, err := domain.ParseSemVer(req.Version) if err != nil { WriteError(w, err) return } - ansName, err := domain.NewAnsName(semver, req.AgentHost) + ansName, err := domain.NewAnsName(semver, qualifiedHost) if err != nil { WriteError(w, err) return diff --git a/internal/ra/handler/v1registration.go b/internal/ra/handler/v1registration.go index ec4baa6..bf7470e 100644 --- a/internal/ra/handler/v1registration.go +++ b/internal/ra/handler/v1registration.go @@ -190,12 +190,13 @@ func (h *V1RegistrationHandler) Register(w http.ResponseWriter, r *http.Request) return } + qualifiedHost := h.svc.QualifyHost(req.AgentHost) semver, err := domain.ParseSemVer(req.Version) if err != nil { WriteError(w, err) return } - ansName, err := domain.NewAnsName(semver, req.AgentHost) + ansName, err := domain.NewAnsName(semver, qualifiedHost) if err != nil { WriteError(w, err) return diff --git a/internal/ra/service/helpers_test.go b/internal/ra/service/helpers_test.go index 7eb0204..995b4a4 100644 --- a/internal/ra/service/helpers_test.go +++ b/internal/ra/service/helpers_test.go @@ -201,3 +201,71 @@ func selfSignedCertPEM(t *testing.T) string { } return string(pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der})) } + +// ----- WithTLPublicBaseURL / TLPublicBaseURL ----- + +func TestTLPublicBaseURL(t *testing.T) { + svc := &RegistrationService{} + if svc.TLPublicBaseURL() != "" { + t.Error("empty by default") + } + svc.WithTLPublicBaseURL("https://tl.example.com") + if svc.TLPublicBaseURL() != "https://tl.example.com" { + t.Errorf("got %q", svc.TLPublicBaseURL()) + } +} + +// ----- WithDNSProvisioner ----- + +func TestWithDNSProvisioner(t *testing.T) { + svc := &RegistrationService{} + svc.WithDNSProvisioner(nil) + if svc.dnsProvisioner != nil { + t.Error("expected nil provisioner") + } +} + +// ----- DomainSuffix ----- + +func TestDomainSuffix(t *testing.T) { + svc := &RegistrationService{} + if svc.DomainSuffix() != "" { + t.Error("empty by default") + } + svc.WithDomainSuffix("agents.example.com") + if svc.DomainSuffix() != "agents.example.com" { + t.Errorf("got %q", svc.DomainSuffix()) + } +} + +// ----- QualifyHost ----- + +func TestQualifyHost(t *testing.T) { + tests := []struct { + name string + suffix string + host string + want string + }{ + {"empty suffix passthrough", "", "my-agent.example.com", "my-agent.example.com"}, + {"suffix applied", "agents.example.com", "my-agent", "my-agent.agents.example.com"}, + {"already qualified", "agents.example.com", "my-agent.agents.example.com", "my-agent.agents.example.com"}, + {"case insensitive match", "AGENTS.EXAMPLE.COM", "my-agent.agents.example.com", "my-agent.agents.example.com"}, + {"leading dot in suffix normalized", ".agents.example.com", "my-agent", "my-agent.agents.example.com"}, + {"trailing dot in suffix normalized", "agents.example.com.", "my-agent", "my-agent.agents.example.com"}, + {"trailing dot on host normalized", "agents.example.com", "my-agent.", "my-agent.agents.example.com"}, + {"host equals suffix", "agents.example.com", "agents.example.com", "agents.example.com"}, + {"uppercase host lowercased", "agents.example.com", "MY-AGENT", "my-agent.agents.example.com"}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + svc := &RegistrationService{} + svc.WithDomainSuffix(tc.suffix) + got := svc.QualifyHost(tc.host) + if got != tc.want { + t.Errorf("QualifyHost(%q) with suffix %q = %q, want %q", + tc.host, tc.suffix, got, tc.want) + } + }) + } +} diff --git a/internal/ra/service/lifecycle.go b/internal/ra/service/lifecycle.go index cf15a65..ecf059f 100644 --- a/internal/ra/service/lifecycle.go +++ b/internal/ra/service/lifecycle.go @@ -537,8 +537,17 @@ func (s *RegistrationService) VerifyDNS(ctx context.Context, agentID string, in expected := domain.ComputeRequiredDNSRecords(reg, s.tlPublicBaseURL) + if s.dnsProvisioner != nil { + if err := s.dnsProvisioner.ProvisionRecords(ctx, reg.FQDN(), expected); err != nil { + return nil, fmt.Errorf("dns provision: %w", err) + } + } + mismatches, perRecord, err := s.verifyDNSRecords(ctx, reg.FQDN(), expected) if err != nil { + if s.dnsProvisioner != nil { + _ = s.dnsProvisioner.DeleteRecords(ctx, reg.FQDN(), expected) + } return nil, fmt.Errorf("dns verify: %w", err) } if len(mismatches) > 0 { @@ -585,6 +594,9 @@ func (s *RegistrationService) VerifyDNS(ctx context.Context, agentID string, in } return s.enqueueTLEvent(txCtx, string(event.TypeAgentRegistered), reg, inner, now) }); err != nil { + if s.dnsProvisioner != nil { + _ = s.dnsProvisioner.DeleteRecords(ctx, reg.FQDN(), expected) + } return nil, err } @@ -781,6 +793,10 @@ type RevokeResult struct { Registration *domain.AgentRegistration RevokedAt time.Time DNSRecordsToRemove []domain.ExpectedDNSRecord + // DNSCleanupErr is non-nil when DNS auto-cleanup was attempted + // but failed. Revocation itself succeeded — the DNS records are + // stale and need manual removal by the operator. + DNSCleanupErr error } // Revoke transitions the registration to REVOKED, marks every active @@ -942,9 +958,16 @@ func (s *RegistrationService) Revoke(ctx context.Context, agentID string, in Rev return nil, err } + dnsToRemove := domain.ComputeRequiredDNSRecords(reg, s.tlPublicBaseURL) + var dnsCleanupErr error + if s.dnsProvisioner != nil && len(dnsToRemove) > 0 { + dnsCleanupErr = s.dnsProvisioner.DeleteRecords(ctx, reg.FQDN(), dnsToRemove) + } + return &RevokeResult{ Registration: reg, RevokedAt: now, - DNSRecordsToRemove: domain.ComputeRequiredDNSRecords(reg, s.tlPublicBaseURL), + DNSRecordsToRemove: dnsToRemove, + DNSCleanupErr: dnsCleanupErr, }, nil } diff --git a/internal/ra/service/registration.go b/internal/ra/service/registration.go index 652ef53..9ce6ddb 100644 --- a/internal/ra/service/registration.go +++ b/internal/ra/service/registration.go @@ -7,6 +7,7 @@ import ( "context" "encoding/json" "fmt" + "strings" "time" "github.com/google/uuid" @@ -123,21 +124,21 @@ type OutboxPayload struct { // sqlx.Tx, cloud adapters can use TransactWriteItems-style atomic // batches. type RegistrationService struct { - agents port.AgentStore - endpoints port.EndpointStore - certs port.CertificateStore - byoc port.ByocCertificateStore - renewals port.RenewalStore - validator port.CertificateValidator - identityCA port.IdentityCertificateAuthority - serverCA port.ServerCertificateAuthority // optional; nil = CSR path rejected - bus port.EventBus - outbox OutboxEnqueuer - uow port.UnitOfWork - dnsVerifier port.DNSVerifier - // tlPublicBaseURL is the externally-reachable Transparency Log URL - // used in _ans-badge DNS records (e.g. "https://tl.example.org"). + agents port.AgentStore + endpoints port.EndpointStore + certs port.CertificateStore + byoc port.ByocCertificateStore + renewals port.RenewalStore + validator port.CertificateValidator + identityCA port.IdentityCertificateAuthority + serverCA port.ServerCertificateAuthority // optional; nil = CSR path rejected + bus port.EventBus + outbox OutboxEnqueuer + uow port.UnitOfWork + dnsVerifier port.DNSVerifier + dnsProvisioner port.DNSProvisioner tlPublicBaseURL string + domainSuffix string // signer is the KeyManager + keyID + raID tuple used to sign // outbox events. When nil, events are still persisted but without // a signature — this is only valid for tests; production configs @@ -209,6 +210,14 @@ func (s *RegistrationService) WithDNSVerifier(v port.DNSVerifier) *RegistrationS return s } +// WithDNSProvisioner wires a DNSProvisioner that auto-creates and +// deletes DNS records during VerifyDNS and Revoke. When nil (or +// never called), operators manage DNS manually. +func (s *RegistrationService) WithDNSProvisioner(p port.DNSProvisioner) *RegistrationService { + s.dnsProvisioner = p + return s +} + // WithTLPublicBaseURL sets the externally-reachable Transparency Log // URL used in _ans-badge DNS TXT records. Without this, badge records // fall back to the agent's own endpoint URL. @@ -222,6 +231,36 @@ func (s *RegistrationService) TLPublicBaseURL() string { return s.tlPublicBaseURL } +// WithDomainSuffix sets the domain suffix appended to agent hostnames +// at registration time. When set, agents submit a short name +// ("my-agent") and the RA constructs the FQDN ("my-agent.example.com"). +// The suffix is normalized: leading/trailing dots are trimmed and the +// value is lowercased. +func (s *RegistrationService) WithDomainSuffix(suffix string) *RegistrationService { + s.domainSuffix = strings.ToLower(strings.Trim(suffix, ".")) + return s +} + +// DomainSuffix returns the configured domain suffix. +func (s *RegistrationService) DomainSuffix() string { + return s.domainSuffix +} + +// QualifyHost appends the domain suffix to a hostname if configured +// and the host doesn't already end with it. Comparison is +// case-insensitive. +func (s *RegistrationService) QualifyHost(host string) string { + if s.domainSuffix == "" { + return host + } + normalized := strings.ToLower(strings.Trim(host, ".")) + suffix := "." + s.domainSuffix + if strings.HasSuffix(normalized, suffix) || normalized == s.domainSuffix { + return normalized + } + return normalized + suffix +} + // RegisterAgent implements the V2 registration flow: // 1. Validate the request shape via domain constructors. // 2. Check ANS name uniqueness. @@ -333,6 +372,14 @@ func (s *RegistrationService) RegisterAgent(ctx context.Context, req RegisterReq } reg.ServerCSR = pendingServerCSR + // When a DNS provisioner is configured, the RA controls the zone + // and domain ownership proof is unnecessary. Issue certs immediately + // and start at PENDING_DNS so the operator only needs to call + // verify-dns (which auto-provisions records). + if s.dnsProvisioner != nil { + return s.registerWithAutoProvision(ctx, reg, byocCert, pendingServerCSR, now) + } + // Generate the ACME DNS-01 challenge token + expiry. The only // DNS action the operator should take before verify-acme. dns01, _, err := generateChallengeTokens() @@ -347,14 +394,6 @@ func (s *RegistrationService) RegisterAgent(ctx context.Context, req RegisterReq } // Persist the aggregate + CSR rows + BYOC cert (if any) atomically. - // Each Save participates in the same transaction via the scoped - // txCtx the UnitOfWork hands fn — partial failure rolls the whole - // batch back so a crash mid-chain can never leave an agent row - // without its endpoints, identity CSR, or server cert. - // - // No signed certs yet: verify-acme signs the identity CSR and - // the server CSR (CSR path); BYOC is already a cert and doesn't - // need signing here. if err := s.uow.Run(ctx, func(txCtx context.Context) error { if err := s.agents.Save(txCtx, reg); err != nil { return err @@ -384,29 +423,123 @@ func (s *RegistrationService) RegisterAgent(ctx context.Context, req RegisterReq return nil, err } - // Publish in-process events AFTER the commit. The bus is - // fire-and-forget for cross-cutting handlers (audit, metrics); - // publishing inside the transaction would mean a downstream - // subscriber that takes a long time, errors, or blocks could - // roll back state that's already durable. for _, ev := range reg.ClearEvents() { if err := s.bus.Publish(ctx, ev); err != nil { return nil, err } } - // Register-time 202 carries NO `dnsRecords[]`. The only DNS - // action the operator can take before verify-acme is installing - // the ACME challenge TXT record, which lives in `challenges[]`. - // Production DNS records (TRUST / BADGE / DISCOVERY / TLSA) - // don't appear until after verify-acme issues certs — the TLSA - // fingerprint can't exist before the server cert does. return &RegisterResponse{ Registration: reg, DNSRecords: nil, }, nil } +// registerWithAutoProvision handles registration when a DNS provisioner +// is configured. Issues certs, provisions DNS records, and activates +// the agent in a single registration call. +func (s *RegistrationService) registerWithAutoProvision( + ctx context.Context, + reg *domain.AgentRegistration, + byocCert *domain.ByocServerCertificate, + pendingServerCSR *domain.AgentCSR, + now time.Time, +) (*RegisterResponse, error) { + // Issue identity certificate. + issuedID, err := s.identityCA.IssueIdentityCertificate(ctx, reg.IdentityCSR.CSRContent, reg.AnsName.String()) + if err != nil { + return nil, domain.NewInternalError("CERT_ISSUE_FAILED", "failed to issue identity cert", err) + } + signedID, err := reg.IdentityCSR.MarkSigned(now) + if err != nil { + return nil, err + } + reg.IdentityCSR = &signedID + storedID := &domain.StoredCertificate{ + CSRID: signedID.CSRID, + CertificateType: domain.CertTypeIdentity, + CertificatePEM: issuedID.CertPEM, + ChainPEM: issuedID.ChainPEM, + Status: domain.CertStatusValid, + IssueTimestamp: issuedID.IssuedAt, + ExpirationTimestamp: issuedID.ExpiresAt, + } + + // Issue server certificate (CSR path). + if pendingServerCSR != nil { + var err error + byocCert, *pendingServerCSR, err = s.signServerCSRForVerifyACME(ctx, reg, pendingServerCSR, now) + if err != nil { + return nil, err + } + reg.ServerCert = byocCert + reg.ServerCSR = pendingServerCSR + } + + // Transition through PENDING_DNS → ACTIVE. + if err := reg.AdvanceToPendingDNS(); err != nil { + return nil, err + } + + // Provision DNS records via DDNS. + expected := domain.ComputeRequiredDNSRecords(reg, s.tlPublicBaseURL) + if err := s.dnsProvisioner.ProvisionRecords(ctx, reg.FQDN(), expected); err != nil { + return nil, fmt.Errorf("dns provision: %w", err) + } + + // Activate. + if err := reg.Activate(now); err != nil { + return nil, err + } + + // Persist agent + certs + TL event atomically. + if err := s.uow.Run(ctx, func(txCtx context.Context) error { + if err := s.agents.Save(txCtx, reg); err != nil { + return err + } + if err := s.endpoints.Save(txCtx, &domain.AgentEndpoints{ + AgentID: reg.AgentID, Endpoints: reg.Endpoints, + }); err != nil { + return err + } + if err := s.certs.SaveCSR(txCtx, reg.AgentID, &signedID); err != nil { + return err + } + if err := s.certs.SaveIdentityCertificate(txCtx, reg.AgentID, storedID); err != nil { + return err + } + if pendingServerCSR != nil { + if err := s.certs.SaveCSR(txCtx, reg.AgentID, pendingServerCSR); err != nil { + return err + } + } + if byocCert != nil { + if err := s.byoc.Save(txCtx, reg.AgentID, byocCert); err != nil { + return err + } + } + // Emit AGENT_REGISTERED event to the TL. + inner, err := s.buildAgentRegisteredEvent(txCtx, reg, expected, nil, now) + if err != nil { + return err + } + return s.enqueueTLEvent(txCtx, string(event.TypeAgentRegistered), reg, inner, now) + }); err != nil { + return nil, err + } + + for _, ev := range reg.ClearEvents() { + if err := s.bus.Publish(ctx, ev); err != nil { + return nil, err + } + } + + return &RegisterResponse{ + Registration: reg, + DNSRecords: expected, + }, nil +} + // baseInnerEvent populates the fields every event carries about its // agent: ansId, ansName, eventType, the agent host/name/version // block, raId (if the RA is configured with a signer), and the diff --git a/internal/ra/service/registration_test.go b/internal/ra/service/registration_test.go index 5eadb75..8951b9b 100644 --- a/internal/ra/service/registration_test.go +++ b/internal/ra/service/registration_test.go @@ -16,6 +16,7 @@ import ( "github.com/rs/zerolog" "github.com/godaddy/ans/internal/adapter/cert" + dnsadapter "github.com/godaddy/ans/internal/adapter/dns" "github.com/godaddy/ans/internal/adapter/eventbus" "github.com/godaddy/ans/internal/adapter/keymanager" "github.com/godaddy/ans/internal/adapter/store/sqlite" @@ -366,6 +367,86 @@ func newRegFixture(t *testing.T) *regFixture { } } +func newRegFixtureWithProvisioner(t *testing.T) *regFixture { + t.Helper() + dir := t.TempDir() + + db, err := sqlite.Open(context.Background(), ":memory:") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = db.Close() }) + + agents := sqlite.NewAgentStore(db) + endpoints := sqlite.NewEndpointStore(db) + certsStore := sqlite.NewCertificateStore(db) + byoc := sqlite.NewByocCertificateStore(db) + renewals := sqlite.NewRenewalStore(db) + outbox := sqlite.NewOutboxStore(db) + + identityCA, err := cert.NewSelfCA(dir+"/ca", "Test CA", 365) + if err != nil { + t.Fatal(err) + } + validator := cert.NewX509Validator(cert.WithSkipChainVerify()) + bus := eventbus.NewInMemoryBus(zerolog.Nop()) + + km, err := keymanager.NewFileKeyManager(dir + "/keys") + if err != nil { + t.Fatal(err) + } + if _, err := km.EnsureKey(context.Background(), "ra-signer", port.AlgorithmECDSAP256); err != nil { + t.Fatal(err) + } + + serverCA, err := cert.NewServerSelfCA(dir+"/server-ca", "Test Server CA", 365) + if err != nil { + t.Fatal(err) + } + + svc := service.NewRegistrationService( + agents, endpoints, certsStore, byoc, renewals, validator, identityCA, bus, outbox, db, + ).WithSigner(service.EventSigner{ + KeyManager: km, + KeyID: "ra-signer", + RaID: "ra-test", + }).WithServerCertificateAuthority(serverCA). + WithDNSProvisioner(dnsadapter.NewNoopProvisioner()) + + semver, _ := domain.ParseSemVer("1.0.0") + ansName, _ := domain.NewAnsName(semver, "agent.example.com") + csrPEM := testCSR(t, ansName.String()) + serverCSR := testServerCSR(t, ansName.FQDN()) + + return ®Fixture{ + svc: svc, + outboxStore: outbox, + uow: db, + agents: agents, + endpoints: endpoints, + certs: certsStore, + byoc: byoc, + renewals: renewals, + validator: validator, + identityCA: identityCA, + serverCA: serverCA, + bus: bus, + req: service.RegisterRequest{ + OwnerID: "owner-1", + AnsName: ansName, + DisplayName: "test-agent", + Description: "a test agent", + Endpoints: []domain.AgentEndpoint{{ + Protocol: domain.Protocol("MCP"), + AgentURL: "https://agent.example.com/mcp", + Transports: []domain.Transport{domain.Transport("SSE")}, + }}, + IdentityCSRPEM: csrPEM, + ServerCsrPEM: serverCSR, + }, + } +} + // testServerCSR builds a server-shaped CSR (DNS SAN matching the // agent FQDN) suitable for the server-cert issuance path. func testServerCSR(t *testing.T, fqdn string) string { @@ -407,6 +488,34 @@ func testCSR(t *testing.T, uri string) string { return string(pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE REQUEST", Bytes: der})) } +// TestRegistration_AutoProvision wires a DNSProvisioner and verifies +// that RegisterAgent activates the agent in a single call (skipping +// the manual verify-acme / verify-dns steps). +func TestRegistration_AutoProvision(t *testing.T) { + t.Parallel() + fx := newRegFixtureWithProvisioner(t) + + semver, _ := domain.ParseSemVer("5.0.0") + an, _ := domain.NewAnsName(semver, "autoprov.example.com") + req := fx.req + req.AnsName = an + req.IdentityCSRPEM = testCSR(t, an.String()) + req.ServerCsrPEM = testServerCSR(t, an.FQDN()) + req.Endpoints = []domain.AgentEndpoint{{ + Protocol: domain.Protocol("MCP"), + AgentURL: "https://autoprov.example.com/mcp", + Transports: []domain.Transport{domain.Transport("SSE")}, + }} + + resp, err := fx.svc.RegisterAgent(context.Background(), req) + if err != nil { + t.Fatalf("RegisterAgent with auto-provision: %v", err) + } + if resp.Registration.Status != domain.StatusActive { + t.Fatalf("status: got %q want ACTIVE", resp.Registration.Status) + } +} + func parseTestURI(t *testing.T, s string) []*url.URL { t.Helper() u, err := url.Parse(s) From 8603286aae5f64e1ddaedc74ec142078418a7a8c Mon Sep 17 00:00:00 2001 From: Francisco Obispo Date: Thu, 16 Jul 2026 10:36:16 -0700 Subject: [PATCH 2/2] feat(ra): add public unauthenticated discovery endpoints MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Add GET /v2/public/agents and GET /v2/public/agents/{agentId} as unauthenticated read-only routes on the RA. These enable agent discovery without API credentials — agent metadata is inherently public (published to DNS and the transparency log). Changes span the full hexagonal stack: - port: ListAll on AgentStore (owner-agnostic query) - sqlite: ListAll implementation with same pagination/filtering - service: ListPublic method on RegistrationService - handler: new PublicHandler with List and Detail methods - router: /v2/public/ registered as anonymous path prefix - dto: parameterized self-link prefix for public vs private routes The public detail endpoint strips the registrationPending block (contains ACME challenge tokens) for security. All 21 new tests pass and the 90% coverage gate is maintained. Signed-off-by: Fernando Obispo Signed-off-by: Francisco Obispo --- cmd/ans-ra/main.go | 9 + internal/adapter/store/sqlite/agent.go | 75 +++++ internal/adapter/store/sqlite/agent_test.go | 196 ++++++++++++ internal/port/store.go | 4 + internal/ra/handler/dto.go | 6 +- internal/ra/handler/public.go | 83 +++++ internal/ra/handler/public_test.go | 333 ++++++++++++++++++++ internal/ra/middleware/ownership_test.go | 3 + internal/ra/service/lifecycle.go | 28 ++ internal/ra/service/lifecycle_test.go | 54 ++++ 10 files changed, 790 insertions(+), 1 deletion(-) create mode 100644 internal/ra/handler/public.go create mode 100644 internal/ra/handler/public_test.go diff --git a/cmd/ans-ra/main.go b/cmd/ans-ra/main.go index c02ea20..a362179 100644 --- a/cmd/ans-ra/main.go +++ b/cmd/ans-ra/main.go @@ -234,6 +234,13 @@ func run(cfgPath string) error { r.With(writeOwnership).Delete("/v2/ans/agents/{agentId}/certificates/server/renewal", lifeH.CancelServerCertRenewal) r.With(writeOwnership).Post("/v2/ans/agents/{agentId}/certificates/server/renewal/verify-acme", lifeH.VerifyRenewalACME) + // Public discovery routes — unauthenticated reads. The /v2/public/ + // prefix is registered as an anonymous path in buildAuth so the auth + // middleware skips it entirely. + pubH := handler.NewPublicHandler(regSvc) + r.Get("/v2/public/agents", pubH.List) + r.Get("/v2/public/agents/{agentId}", pubH.Detail) + // V1 RA surface — byte-for-byte parity with the reference V1 API // spec. Shares the same RegistrationService as the V2 routes; // only the DTO marshalling + TL-emit schema version differ. See @@ -379,6 +386,7 @@ func buildAuth(ctx context.Context, cfg *config.RAConfig) (providerWithAnonymous auth.WithAPISecret(cfg.Auth.Static.APISecret), auth.WithAnonymousPath("/v2/admin/health"), auth.WithAnonymousPath("/v2/admin/ready"), + auth.WithAnonymousPath("/v2/public/"), auth.WithAnonymousPath("/docs"), ), nil case "oidc": @@ -389,6 +397,7 @@ func buildAuth(ctx context.Context, cfg *config.RAConfig) (providerWithAnonymous cfg.Auth.OIDC.ClientID, auth.WithOIDCAnonymousPath("/v2/admin/health"), auth.WithOIDCAnonymousPath("/v2/admin/ready"), + auth.WithOIDCAnonymousPath("/v2/public/"), auth.WithOIDCAnonymousPath("/docs"), // Empty AdminGroups means no OIDC user is admin — // preserves prior behaviour for operators who haven't diff --git a/internal/adapter/store/sqlite/agent.go b/internal/adapter/store/sqlite/agent.go index a06be3f..5962813 100644 --- a/internal/adapter/store/sqlite/agent.go +++ b/internal/adapter/store/sqlite/agent.go @@ -283,6 +283,81 @@ func (s *AgentStore) ListByOwner( }, nil } +// ListAll returns a cursor-paginated list of all agents regardless of owner. +func (s *AgentStore) ListAll( + ctx context.Context, + filter port.ListFilter, +) (*port.CursorPage[*domain.AgentRegistration], error) { + limit := filter.Limit + if limit <= 0 { + limit = 20 + } + if limit > 100 { + limit = 100 + } + + var args []any + var clauses []string + + if filter.AgentHost != "" { + clauses = append(clauses, "agent_host = ?") + args = append(args, filter.AgentHost) + } + + if len(filter.Statuses) > 0 { + placeholders := make([]string, len(filter.Statuses)) + for i, st := range filter.Statuses { + placeholders[i] = "?" + args = append(args, string(st)) + } + clauses = append(clauses, fmt.Sprintf("status IN (%s)", joinStrings(placeholders, ","))) + } + + if filter.Cursor != "" { + id, err := decodeCursor(filter.Cursor) + if err != nil { + return nil, domain.NewValidationError("INVALID_CURSOR", err.Error()) + } + clauses = append(clauses, "id < ?") + args = append(args, id) + } + + where := "1=1" + if len(clauses) > 0 { + where = joinStrings(clauses, " AND ") + } + + q := fmt.Sprintf( + `SELECT * FROM agent_registrations WHERE %s ORDER BY id DESC LIMIT ?`, + where, + ) + args = append(args, limit+1) + + var rows []agentRow + if err := s.db.extx(ctx).SelectContext(ctx, &rows, q, args...); err != nil { + return nil, err + } + + hasMore := len(rows) > limit + if hasMore { + rows = rows[:limit] + } + items, err := rowsToDomain(rows) + if err != nil { + return nil, err + } + nextCursor := "" + if hasMore && len(items) > 0 { + nextCursor = encodeCursor(items[len(items)-1].ID) + } + return &port.CursorPage[*domain.AgentRegistration]{ + Items: items, + NextCursor: nextCursor, + HasMore: hasMore, + ReturnedCount: len(items), + }, nil +} + // Delete removes a registration by ID. func (s *AgentStore) Delete(ctx context.Context, id int64) error { _, err := s.db.extx(ctx).ExecContext(ctx, `DELETE FROM agent_registrations WHERE id = ?`, id) diff --git a/internal/adapter/store/sqlite/agent_test.go b/internal/adapter/store/sqlite/agent_test.go index 5915bbf..5d70ca0 100644 --- a/internal/adapter/store/sqlite/agent_test.go +++ b/internal/adapter/store/sqlite/agent_test.go @@ -342,6 +342,202 @@ func TestAgentStore_ListByOwner_InvalidCursor(t *testing.T) { } } +// ----- ListAll (public discovery — no owner filter) ----- + +func TestAgentStore_ListAll_ReturnsAllOwners(t *testing.T) { + store := NewAgentStore(newTestDB(t)) + ctx := context.Background() + + // Insert agents for two different owners. + for i, owner := range []string{"alice", "bob"} { + ansName, _ := domain.NewAnsName(mustSemVer(t, 1, 0, i), owner+".example.com") + _ = store.Save(ctx, &domain.AgentRegistration{ + AgentID: owner + "-agent", OwnerID: owner, + AnsName: ansName, Status: domain.StatusActive, + Details: domain.RegistrationDetails{RegistrationTimestamp: time.Now()}, + }) + } + + page, err := store.ListAll(ctx, port.ListFilter{Limit: 10}) + if err != nil { + t.Fatal(err) + } + if len(page.Items) != 2 { + t.Errorf("ListAll: got %d, want 2 (both owners)", len(page.Items)) + } +} + +func TestAgentStore_ListAll_StatusFilter(t *testing.T) { + store := NewAgentStore(newTestDB(t)) + ctx := context.Background() + + for _, spec := range []struct { + id string + host string + status domain.RegistrationStatus + }{ + {"active-1", "a.example.com", domain.StatusActive}, + {"revoked-1", "b.example.com", domain.StatusRevoked}, + {"active-2", "c.example.com", domain.StatusActive}, + } { + ansName, _ := domain.NewAnsName(mustSemVer(t, 1, 0, 0), spec.host) + _ = store.Save(ctx, &domain.AgentRegistration{ + AgentID: spec.id, OwnerID: "owner", + AnsName: ansName, Status: spec.status, + Details: domain.RegistrationDetails{RegistrationTimestamp: time.Now()}, + }) + } + + page, err := store.ListAll(ctx, port.ListFilter{ + Statuses: []domain.RegistrationStatus{domain.StatusActive}, + Limit: 10, + }) + if err != nil { + t.Fatal(err) + } + if len(page.Items) != 2 { + t.Errorf("status filter: got %d, want 2 ACTIVE", len(page.Items)) + } + for _, item := range page.Items { + if item.Status != domain.StatusActive { + t.Errorf("leaked status %s", item.Status) + } + } +} + +func TestAgentStore_ListAll_AgentHostFilter(t *testing.T) { + store := NewAgentStore(newTestDB(t)) + ctx := context.Background() + + for i, host := range []string{"target.example.com", "other.example.com"} { + ansName, _ := domain.NewAnsName(mustSemVer(t, 1, 0, i), host) + _ = store.Save(ctx, &domain.AgentRegistration{ + AgentID: "id-" + host, OwnerID: "owner", + AnsName: ansName, Status: domain.StatusActive, + Details: domain.RegistrationDetails{RegistrationTimestamp: time.Now()}, + }) + } + + page, err := store.ListAll(ctx, port.ListFilter{ + AgentHost: "target.example.com", + Limit: 10, + }) + if err != nil { + t.Fatal(err) + } + if len(page.Items) != 1 { + t.Errorf("host filter: got %d, want 1", len(page.Items)) + } + if page.Items[0].AnsName.FQDN() != "target.example.com" { + t.Errorf("wrong host: %s", page.Items[0].AnsName.FQDN()) + } +} + +func TestAgentStore_ListAll_Pagination(t *testing.T) { + store := NewAgentStore(newTestDB(t)) + ctx := context.Background() + + for i := 1; i <= 5; i++ { + ansName, _ := domain.NewAnsName(mustSemVer(t, 1, 0, i), "p.example.com") + _ = store.Save(ctx, &domain.AgentRegistration{ + AgentID: "all-paged-" + ansName.String(), OwnerID: "alice", + AnsName: ansName, Status: domain.StatusActive, + Details: domain.RegistrationDetails{RegistrationTimestamp: time.Now()}, + }) + } + + page1, err := store.ListAll(ctx, port.ListFilter{Limit: 2}) + if err != nil { + t.Fatal(err) + } + if len(page1.Items) != 2 { + t.Errorf("page1: got %d, want 2", len(page1.Items)) + } + if !page1.HasMore || page1.NextCursor == "" { + t.Error("HasMore + NextCursor expected") + } + + page2, err := store.ListAll(ctx, port.ListFilter{Limit: 2, Cursor: page1.NextCursor}) + if err != nil { + t.Fatal(err) + } + if len(page2.Items) != 2 { + t.Errorf("page2: got %d, want 2", len(page2.Items)) + } + if page1.Items[len(page1.Items)-1].ID <= page2.Items[0].ID { + t.Errorf("cursor did not advance: page1 last=%d page2 first=%d", + page1.Items[len(page1.Items)-1].ID, page2.Items[0].ID) + } + + // Last page. + page3, err := store.ListAll(ctx, port.ListFilter{Limit: 2, Cursor: page2.NextCursor}) + if err != nil { + t.Fatal(err) + } + if len(page3.Items) != 1 { + t.Errorf("page3: got %d, want 1", len(page3.Items)) + } + if page3.HasMore { + t.Error("last page should not have more") + } +} + +func TestAgentStore_ListAll_DefaultLimit(t *testing.T) { + store := NewAgentStore(newTestDB(t)) + ctx := context.Background() + + // Insert 3 agents — less than default 20. + for i := 1; i <= 3; i++ { + ansName, _ := domain.NewAnsName(mustSemVer(t, 1, 0, i), "def.example.com") + _ = store.Save(ctx, &domain.AgentRegistration{ + AgentID: "def-" + ansName.String(), OwnerID: "owner", + AnsName: ansName, Status: domain.StatusActive, + Details: domain.RegistrationDetails{RegistrationTimestamp: time.Now()}, + }) + } + + // Limit 0 → default 20. + page, err := store.ListAll(ctx, port.ListFilter{}) + if err != nil { + t.Fatal(err) + } + if len(page.Items) != 3 { + t.Errorf("default limit: got %d, want 3", len(page.Items)) + } + if page.HasMore { + t.Error("should not have more with 3 < 20") + } +} + +func TestAgentStore_ListAll_Empty(t *testing.T) { + store := NewAgentStore(newTestDB(t)) + page, err := store.ListAll(context.Background(), port.ListFilter{Limit: 10}) + if err != nil { + t.Fatal(err) + } + if len(page.Items) != 0 { + t.Errorf("empty: got %d, want 0", len(page.Items)) + } + if page.HasMore { + t.Error("empty DB should not have more") + } + if page.NextCursor != "" { + t.Errorf("empty DB should have empty cursor, got %q", page.NextCursor) + } +} + +func TestAgentStore_ListAll_InvalidCursor(t *testing.T) { + store := NewAgentStore(newTestDB(t)) + _, err := store.ListAll(context.Background(), port.ListFilter{Cursor: "!!!"}) + if err == nil { + t.Fatal("expected cursor decode error") + } + var de *domain.Error + if !errors.As(err, &de) || de.Code != "INVALID_CURSOR" { + t.Errorf("want INVALID_CURSOR domain error, got %v", err) + } +} + // ----- Delete ----- func TestAgentStore_Delete(t *testing.T) { diff --git a/internal/port/store.go b/internal/port/store.go index 38c79d9..befe6b9 100644 --- a/internal/port/store.go +++ b/internal/port/store.go @@ -60,6 +60,10 @@ type AgentStore interface { filter ListFilter, ) (*CursorPage[*domain.AgentRegistration], error) + // ListAll returns a cursor page of agents across all owners. + // Used by the public discovery endpoint. + ListAll(ctx context.Context, filter ListFilter) (*CursorPage[*domain.AgentRegistration], error) + // Delete removes the registration with the given ID. Used only for // administrative cleanup; normal lifecycle uses Revoke. Delete(ctx context.Context, id int64) error diff --git a/internal/ra/handler/dto.go b/internal/ra/handler/dto.go index d2d3588..b2eb4b7 100644 --- a/internal/ra/handler/dto.go +++ b/internal/ra/handler/dto.go @@ -34,6 +34,10 @@ type listItem struct { } func mapListResponse(res *service.ListResult) listResponse { + return mapListResponseWithPrefix(res, "/v2/ans/agents/") +} + +func mapListResponseWithPrefix(res *service.ListResult, linkPrefix string) listResponse { items := make([]listItem, 0, len(res.Items)) for _, reg := range res.Items { epSlice := []domain.AgentEndpoint{} @@ -52,7 +56,7 @@ func mapListResponse(res *service.ListResult) listResponse { RegistrationTimestamp: reg.Details.RegistrationTimestamp.Format("2006-01-02T15:04:05Z07:00"), Endpoints: mapEndpointsToDTO(epSlice), Links: []linkDTO{ - {Rel: "self", Href: "/v2/ans/agents/" + reg.AgentID}, + {Rel: "self", Href: linkPrefix + reg.AgentID}, }, }) } diff --git a/internal/ra/handler/public.go b/internal/ra/handler/public.go new file mode 100644 index 0000000..233e30c --- /dev/null +++ b/internal/ra/handler/public.go @@ -0,0 +1,83 @@ +package handler + +import ( + "net/http" + "strconv" + + "github.com/go-chi/chi/v5" + + "github.com/godaddy/ans/internal/domain" + "github.com/godaddy/ans/internal/port" + "github.com/godaddy/ans/internal/ra/service" +) + +// PublicHandler groups unauthenticated read-only discovery endpoints. +// Routes served by this handler bypass the auth middleware via anonymous +// path registration. +type PublicHandler struct { + svc *service.RegistrationService +} + +// NewPublicHandler constructs a PublicHandler. +func NewPublicHandler(svc *service.RegistrationService) *PublicHandler { + return &PublicHandler{svc: svc} +} + +const publicAgentsPrefix = "/v2/public/agents/" + +// List handles GET /v2/public/agents. +func (h *PublicHandler) List(w http.ResponseWriter, r *http.Request) { + q := r.URL.Query() + filter := port.ListFilter{ + AgentHost: q.Get("agentHost"), + Cursor: q.Get("cursor"), + } + + if statuses, ok := q["status"]; ok && len(statuses) > 0 { + for _, s := range statuses { + if s == "ALL" { + filter.Statuses = nil + break + } + filter.Statuses = append(filter.Statuses, domain.RegistrationStatus(s)) + } + } else { + filter.Statuses = []domain.RegistrationStatus{domain.StatusActive} + } + + if lv := q.Get("limit"); lv != "" { + n, err := strconv.Atoi(lv) + if err != nil || n < 1 || n > 100 { + WriteError(w, domain.NewValidationError( + "INVALID_LIMIT", "limit must be between 1 and 100", + )) + return + } + filter.Limit = n + } else { + filter.Limit = 20 + } + + res, err := h.svc.ListPublic(r.Context(), filter) + if err != nil { + WriteError(w, err) + return + } + WriteJSON(w, http.StatusOK, mapListResponseWithPrefix(res, publicAgentsPrefix)) +} + +// Detail handles GET /v2/public/agents/{agentId}. +func (h *PublicHandler) Detail(w http.ResponseWriter, r *http.Request) { + agentID := chi.URLParam(r, "agentId") + res, err := h.svc.GetByAgentID(r.Context(), agentID) + if err != nil { + WriteError(w, err) + return + } + d := mapAgentDetails(res, r, h.svc.TLPublicBaseURL()) + d.RegistrationPending = nil + d.Links = []linkDTO{ + {Rel: "self", Href: publicAgentsPrefix + res.Registration.AgentID}, + } + WriteJSON(w, http.StatusOK, d) +} diff --git a/internal/ra/handler/public_test.go b/internal/ra/handler/public_test.go new file mode 100644 index 0000000..cf2e8c5 --- /dev/null +++ b/internal/ra/handler/public_test.go @@ -0,0 +1,333 @@ +package handler_test + +import ( + "encoding/json" + "fmt" + "net/http" + "testing" + + "github.com/godaddy/ans/internal/ra/handler" +) + +// Public-endpoint integration tests. Reuses the same handlerFixture as +// the lifecycle tests; the difference is the /v2/public/agents routes +// are exercised WITHOUT injecting an identity (no asOwner tweak). + +func TestPublicList_ReturnsActiveAgents(t *testing.T) { + t.Parallel() + fx := newPublicFixture(t) + + agentID := fx.registerAgent(t, "alice", "pub.example.com", "1.0.0") + fx.activateAgent(t, "alice", agentID) + + rec := fx.request(t, http.MethodGet, "/v2/public/agents", nil, nil) + if rec.Code != http.StatusOK { + t.Fatalf("status=%d body=%s", rec.Code, rec.Body) + } + var resp struct { + Items []struct { + AgentID string `json:"agentId"` + AnsName string `json:"ansName"` + AgentHost string `json:"agentHost"` + } `json:"items"` + ReturnedCount int `json:"returnedCount"` + HasMore bool `json:"hasMore"` + } + _ = json.Unmarshal(rec.Body.Bytes(), &resp) + if resp.ReturnedCount != 1 { + t.Fatalf("want 1 item, got %d", resp.ReturnedCount) + } + if resp.Items[0].AgentID != agentID { + t.Errorf("agentId: got %q want %q", resp.Items[0].AgentID, agentID) + } +} + +func TestPublicList_CrossOwnerVisibility(t *testing.T) { + t.Parallel() + fx := newPublicFixture(t) + + a1 := fx.registerAgent(t, "alice", "alice.example.com", "1.0.0") + fx.activateAgent(t, "alice", a1) + a2 := fx.registerAgent(t, "bob", "bob.example.com", "1.0.0") + fx.activateAgent(t, "bob", a2) + + rec := fx.request(t, http.MethodGet, "/v2/public/agents", nil, nil) + if rec.Code != http.StatusOK { + t.Fatalf("status=%d body=%s", rec.Code, rec.Body) + } + var resp struct { + Items []struct { + AgentID string `json:"agentId"` + } `json:"items"` + } + _ = json.Unmarshal(rec.Body.Bytes(), &resp) + if len(resp.Items) != 2 { + t.Fatalf("public list should see both owners' agents; got %d", len(resp.Items)) + } + seen := map[string]bool{} + for _, it := range resp.Items { + seen[it.AgentID] = true + } + if !seen[a1] || !seen[a2] { + t.Errorf("missing agent(s): %v", seen) + } +} + +func TestPublicList_DefaultStatusACTIVE(t *testing.T) { + t.Parallel() + fx := newPublicFixture(t) + + // Register without activating → PENDING_VALIDATION. + fx.registerAgent(t, "alice", "pending.example.com", "1.0.0") + + rec := fx.request(t, http.MethodGet, "/v2/public/agents", nil, nil) + if rec.Code != http.StatusOK { + t.Fatalf("status=%d", rec.Code) + } + var resp struct { + Items []any `json:"items"` + } + _ = json.Unmarshal(rec.Body.Bytes(), &resp) + if len(resp.Items) != 0 { + t.Errorf("default ACTIVE filter should hide pending agents; got %d items", len(resp.Items)) + } +} + +func TestPublicList_InvalidLimit(t *testing.T) { + t.Parallel() + fx := newPublicFixture(t) + + rec := fx.request(t, http.MethodGet, "/v2/public/agents?limit=999", nil, nil) + if rec.Code != http.StatusUnprocessableEntity { + t.Fatalf("want 422, got %d", rec.Code) + } +} + +func TestPublicList_StatusFilterALL(t *testing.T) { + t.Parallel() + fx := newPublicFixture(t) + + a1 := fx.registerAgent(t, "alice", "all1.example.com", "1.0.0") + fx.activateAgent(t, "alice", a1) + // a2 stays in PENDING_VALIDATION. + fx.registerAgent(t, "alice", "all2.example.com", "1.0.0") + + rec := fx.request(t, http.MethodGet, "/v2/public/agents?status=ALL", nil, nil) + if rec.Code != http.StatusOK { + t.Fatalf("status=%d body=%s", rec.Code, rec.Body) + } + var resp struct { + ReturnedCount int `json:"returnedCount"` + } + _ = json.Unmarshal(rec.Body.Bytes(), &resp) + if resp.ReturnedCount != 2 { + t.Errorf("status=ALL should return both agents; got %d", resp.ReturnedCount) + } +} + +func TestPublicList_ExplicitStatusFilter(t *testing.T) { + t.Parallel() + fx := newPublicFixture(t) + + a1 := fx.registerAgent(t, "alice", "exp1.example.com", "1.0.0") + fx.activateAgent(t, "alice", a1) + // a2 stays in PENDING_VALIDATION. + fx.registerAgent(t, "alice", "exp2.example.com", "1.0.0") + + rec := fx.request(t, http.MethodGet, "/v2/public/agents?status=PENDING_VALIDATION", nil, nil) + if rec.Code != http.StatusOK { + t.Fatalf("status=%d body=%s", rec.Code, rec.Body) + } + var resp struct { + ReturnedCount int `json:"returnedCount"` + } + _ = json.Unmarshal(rec.Body.Bytes(), &resp) + if resp.ReturnedCount != 1 { + t.Errorf("status=PENDING_VALIDATION should return 1; got %d", resp.ReturnedCount) + } +} + +func TestPublicList_ValidLimit(t *testing.T) { + t.Parallel() + fx := newPublicFixture(t) + + for i := range 3 { + id := fx.registerAgent(t, "alice", fmt.Sprintf("lim%d.example.com", i), "1.0.0") + fx.activateAgent(t, "alice", id) + } + + rec := fx.request(t, http.MethodGet, "/v2/public/agents?limit=2", nil, nil) + if rec.Code != http.StatusOK { + t.Fatalf("status=%d body=%s", rec.Code, rec.Body) + } + var resp struct { + ReturnedCount int `json:"returnedCount"` + Limit int `json:"limit"` + HasMore bool `json:"hasMore"` + } + _ = json.Unmarshal(rec.Body.Bytes(), &resp) + if resp.ReturnedCount != 2 { + t.Errorf("want 2 items, got %d", resp.ReturnedCount) + } + if resp.Limit != 2 { + t.Errorf("limit: got %d want 2", resp.Limit) + } + if !resp.HasMore { + t.Error("hasMore should be true with 3 agents and limit=2") + } +} + +func TestPublicList_HostFilter(t *testing.T) { + t.Parallel() + fx := newPublicFixture(t) + + a1 := fx.registerAgent(t, "alice", "alpha.example.com", "1.0.0") + fx.activateAgent(t, "alice", a1) + a2 := fx.registerAgent(t, "bob", "beta.example.com", "1.0.0") + fx.activateAgent(t, "bob", a2) + + rec := fx.request(t, http.MethodGet, "/v2/public/agents?agentHost=alpha.example.com", nil, nil) + if rec.Code != http.StatusOK { + t.Fatalf("status=%d body=%s", rec.Code, rec.Body) + } + var resp struct { + Items []struct { + AgentHost string `json:"agentHost"` + } `json:"items"` + } + _ = json.Unmarshal(rec.Body.Bytes(), &resp) + if len(resp.Items) != 1 { + t.Fatalf("want 1 filtered result, got %d", len(resp.Items)) + } + if resp.Items[0].AgentHost != "alpha.example.com" { + t.Errorf("host: got %q", resp.Items[0].AgentHost) + } +} + +func TestPublicList_SelfLinkPrefix(t *testing.T) { + t.Parallel() + fx := newPublicFixture(t) + + agentID := fx.registerAgent(t, "alice", "link.example.com", "1.0.0") + fx.activateAgent(t, "alice", agentID) + + rec := fx.request(t, http.MethodGet, "/v2/public/agents", nil, nil) + if rec.Code != http.StatusOK { + t.Fatalf("status=%d", rec.Code) + } + var resp struct { + Items []struct { + Links []struct { + Rel string `json:"rel"` + Href string `json:"href"` + } `json:"links"` + } `json:"items"` + } + _ = json.Unmarshal(rec.Body.Bytes(), &resp) + if len(resp.Items) != 1 || len(resp.Items[0].Links) == 0 { + t.Fatalf("expected one item with links; got %d items", len(resp.Items)) + } + want := "/v2/public/agents/" + agentID + got := resp.Items[0].Links[0].Href + if got != want { + t.Errorf("self link: got %q want %q", got, want) + } +} + +func TestPublicDetail_ReturnsAgent(t *testing.T) { + t.Parallel() + fx := newPublicFixture(t) + + agentID := fx.registerAgent(t, "alice", "detail.example.com", "1.0.0") + fx.activateAgent(t, "alice", agentID) + + rec := fx.request(t, http.MethodGet, "/v2/public/agents/"+agentID, nil, nil) + if rec.Code != http.StatusOK { + t.Fatalf("status=%d body=%s", rec.Code, rec.Body) + } + var resp struct { + AgentID string `json:"agentId"` + AgentStatus string `json:"agentStatus"` + AnsName string `json:"ansName"` + } + _ = json.Unmarshal(rec.Body.Bytes(), &resp) + if resp.AgentID != agentID { + t.Errorf("agentId: got %q want %q", resp.AgentID, agentID) + } + if resp.AgentStatus != "ACTIVE" { + t.Errorf("status: got %q want ACTIVE", resp.AgentStatus) + } +} + +func TestPublicDetail_NoPendingBlock(t *testing.T) { + t.Parallel() + fx := newPublicFixture(t) + + // Leave agent in PENDING_VALIDATION — it has a non-nil + // registrationPending block internally. The public endpoint + // must strip it. + agentID := fx.registerAgent(t, "alice", "pend.example.com", "1.0.0") + + rec := fx.request(t, http.MethodGet, "/v2/public/agents/"+agentID, nil, nil) + if rec.Code != http.StatusOK { + t.Fatalf("status=%d body=%s", rec.Code, rec.Body) + } + var raw map[string]any + _ = json.Unmarshal(rec.Body.Bytes(), &raw) + if rp, ok := raw["registrationPending"]; ok && rp != nil { + t.Errorf("registrationPending must be null/absent on public detail; got %v", rp) + } +} + +func TestPublicDetail_SelfLinkPrefix(t *testing.T) { + t.Parallel() + fx := newPublicFixture(t) + + agentID := fx.registerAgent(t, "alice", "selflink.example.com", "1.0.0") + + rec := fx.request(t, http.MethodGet, "/v2/public/agents/"+agentID, nil, nil) + if rec.Code != http.StatusOK { + t.Fatalf("status=%d", rec.Code) + } + var resp struct { + Links []struct { + Rel string `json:"rel"` + Href string `json:"href"` + } `json:"links"` + } + _ = json.Unmarshal(rec.Body.Bytes(), &resp) + want := "/v2/public/agents/" + agentID + found := false + for _, l := range resp.Links { + if l.Rel == "self" && l.Href == want { + found = true + } + } + if !found { + t.Errorf("self link not found or wrong; got %+v", resp.Links) + } +} + +func TestPublicDetail_NotFound(t *testing.T) { + t.Parallel() + fx := newPublicFixture(t) + + rec := fx.request(t, http.MethodGet, "/v2/public/agents/00000000-0000-0000-0000-000000000000", nil, nil) + if rec.Code != http.StatusNotFound { + t.Fatalf("want 404, got %d body=%s", rec.Code, rec.Body) + } +} + +// ----- fixture ----- + +// newPublicFixture reuses the handler fixture but also mounts the +// public discovery routes. Authenticated routes are still available +// (for registerAgent / activateAgent which need an owner identity). +func newPublicFixture(t *testing.T) *handlerFixture { + t.Helper() + fx := newHandlerFixture(t) + pubH := handler.NewPublicHandler(fx.svc) + fx.router.Get("/v2/public/agents", pubH.List) + fx.router.Get("/v2/public/agents/{agentId}", pubH.Detail) + return fx +} diff --git a/internal/ra/middleware/ownership_test.go b/internal/ra/middleware/ownership_test.go index cf4b8e3..8c789d0 100644 --- a/internal/ra/middleware/ownership_test.go +++ b/internal/ra/middleware/ownership_test.go @@ -279,6 +279,9 @@ func (f *fakeAgentStore) FindExistingByFQDN(_ context.Context, _ string) ([]*dom func (f *fakeAgentStore) ListByOwner(_ context.Context, _ string, _ port.ListFilter) (*port.CursorPage[*domain.AgentRegistration], error) { return nil, nil } +func (f *fakeAgentStore) ListAll(_ context.Context, _ port.ListFilter) (*port.CursorPage[*domain.AgentRegistration], error) { + return nil, nil +} func (f *fakeAgentStore) Delete(_ context.Context, _ int64) error { return nil } // Silence "unused" on timing imports the testbed may shed. diff --git a/internal/ra/service/lifecycle.go b/internal/ra/service/lifecycle.go index ecf059f..83dcaf1 100644 --- a/internal/ra/service/lifecycle.go +++ b/internal/ra/service/lifecycle.go @@ -61,6 +61,34 @@ func (s *RegistrationService) List(ctx context.Context, ownerID string, filter p }, nil } +// ListPublic returns agents matching the filter across all owners. +func (s *RegistrationService) ListPublic(ctx context.Context, filter port.ListFilter) (*ListResult, error) { + page, err := s.agents.ListAll(ctx, filter) + if err != nil { + return nil, err + } + + endpointsByAgent := map[string]*domain.AgentEndpoints{} + if len(page.Items) > 0 { + ids := make([]string, 0, len(page.Items)) + for _, a := range page.Items { + ids = append(ids, a.AgentID) + } + endpointsByAgent, err = s.endpoints.FindByAgentIDs(ctx, ids) + if err != nil { + return nil, err + } + } + + return &ListResult{ + Items: page.Items, + Endpoints: endpointsByAgent, + NextCursor: page.NextCursor, + HasMore: page.HasMore, + Limit: filter.Limit, + }, nil +} + // DetailResult carries everything the detail handler needs to build // an AgentDetails response. type DetailResult struct { diff --git a/internal/ra/service/lifecycle_test.go b/internal/ra/service/lifecycle_test.go index c05bcbb..3be661d 100644 --- a/internal/ra/service/lifecycle_test.go +++ b/internal/ra/service/lifecycle_test.go @@ -7,6 +7,7 @@ import ( "testing" "github.com/godaddy/ans/internal/domain" + "github.com/godaddy/ans/internal/port" "github.com/godaddy/ans/internal/ra/service" event "github.com/godaddy/ans/internal/tl/event" eventv1 "github.com/godaddy/ans/internal/tl/event/v1" @@ -198,3 +199,56 @@ func TestSubmitIdentityCSR_NoIdentityCSR_Rejected(t *testing.T) { t.Fatalf("expected code IDENTITY_CSR_NOT_PERMITTED; got %q", de.Code) } } + +// TestListPublic_ReturnsAllOwners verifies that ListPublic returns +// agents across all owners, unlike List which is ownership-scoped. +func TestListPublic_ReturnsAllOwners(t *testing.T) { + t.Parallel() + fx := newRegFixture(t) + ctx := context.Background() + + // Register agent for owner-1 (the fixture default). + resp1, err := fx.svc.RegisterAgent(ctx, fx.req) + if err != nil { + t.Fatalf("register owner-1: %v", err) + } + + // Register a second agent for owner-2. + semver2, _ := domain.ParseSemVer("2.0.0") + ansName2, _ := domain.NewAnsName(semver2, "other.example.com") + req2 := fx.req + req2.OwnerID = "owner-2" + req2.AnsName = ansName2 + req2.Endpoints = []domain.AgentEndpoint{{ + Protocol: domain.Protocol("MCP"), + AgentURL: "https://other.example.com/mcp", + Transports: []domain.Transport{domain.Transport("SSE")}, + }} + req2.IdentityCSRPEM = testCSR(t, ansName2.String()) + req2.ServerCsrPEM = testServerCSR(t, ansName2.FQDN()) + resp2, err := fx.svc.RegisterAgent(ctx, req2) + if err != nil { + t.Fatalf("register owner-2: %v", err) + } + + // ListPublic with status=ALL should return both. + result, err := fx.svc.ListPublic(ctx, port.ListFilter{ + Statuses: []domain.RegistrationStatus{ + domain.StatusPendingValidation, + domain.StatusActive, + }, + }) + if err != nil { + t.Fatalf("ListPublic: %v", err) + } + if len(result.Items) != 2 { + t.Fatalf("want 2 agents across both owners; got %d", len(result.Items)) + } + seen := map[string]bool{} + for _, item := range result.Items { + seen[item.AgentID] = true + } + if !seen[resp1.Registration.AgentID] || !seen[resp2.Registration.AgentID] { + t.Errorf("expected both agents in result; seen=%v", seen) + } +}