From d2d4d7dce8ff5773e35b8c7d92c1ec93a0da45f4 Mon Sep 17 00:00:00 2001 From: Patrick Dewey Date: Sun, 10 May 2026 10:14:52 -0400 Subject: [PATCH] feat: add handle utils, rkey validation, PublicClient caching, and StartSignup --- go.mod | 6 +- go.sum | 9 +++ middleware/auth.go | 8 +++ oauth.go | 164 +++++++++++++++++++++++++++++++++++++++++++++ public.go | 119 ++++++++++++++++++++++++++++---- public_test.go | 64 ++++++++++++++++++ record.go | 49 ++++++++++++++ record_test.go | 110 ++++++++++++++++++++++++++++++ uri.go | 85 +++++++++++++++++++++-- uri_test.go | 116 ++++++++++++++++++++++++++++++-- 10 files changed, 705 insertions(+), 25 deletions(-) create mode 100644 record_test.go diff --git a/go.mod b/go.mod index dca7213..c8e1b12 100644 --- a/go.mod +++ b/go.mod @@ -7,11 +7,13 @@ require ( github.com/gorilla/websocket v1.5.3 github.com/klauspost/compress v1.18.0 github.com/pkg/browser v0.0.0-20240102092130-5ac0b6a4141c + github.com/stretchr/testify v1.11.1 go.etcd.io/bbolt v1.4.3 go.opentelemetry.io/otel v1.43.0 go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracehttp v1.43.0 go.opentelemetry.io/otel/sdk v1.43.0 go.opentelemetry.io/otel/trace v1.43.0 + golang.org/x/net v0.52.0 modernc.org/sqlite v1.48.1 ) @@ -19,6 +21,7 @@ require ( github.com/beorn7/perks v1.0.1 // indirect github.com/cenkalti/backoff/v5 v5.0.3 // indirect github.com/cespare/xxhash/v2 v2.3.0 // indirect + github.com/davecgh/go-spew v1.1.1 // indirect github.com/dustin/go-humanize v1.0.1 // indirect github.com/earthboundkid/versioninfo/v2 v2.24.1 // indirect github.com/go-logr/logr v1.4.3 // indirect @@ -40,6 +43,7 @@ require ( github.com/multiformats/go-varint v0.1.0 // indirect github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect github.com/ncruces/go-strftime v1.0.0 // indirect + github.com/pmezard/go-difflib v1.0.0 // indirect github.com/prometheus/client_golang v1.23.2 // indirect github.com/prometheus/client_model v0.6.2 // indirect github.com/prometheus/common v0.67.5 // indirect @@ -55,7 +59,6 @@ require ( go.opentelemetry.io/proto/otlp v1.10.0 // indirect go.yaml.in/yaml/v2 v2.4.4 // indirect golang.org/x/crypto v0.49.0 // indirect - golang.org/x/net v0.52.0 // indirect golang.org/x/sys v0.42.0 // indirect golang.org/x/text v0.35.0 // indirect golang.org/x/time v0.15.0 // indirect @@ -64,6 +67,7 @@ require ( google.golang.org/genproto/googleapis/rpc v0.0.0-20260401024825-9d38bb4040a9 // indirect google.golang.org/grpc v1.80.0 // indirect google.golang.org/protobuf v1.36.11 // indirect + gopkg.in/yaml.v3 v3.0.1 // indirect lukechampine.com/blake3 v1.4.1 // indirect modernc.org/libc v1.70.0 // indirect modernc.org/mathutil v1.7.1 // indirect diff --git a/go.sum b/go.sum index e279c59..3b67df8 100644 --- a/go.sum +++ b/go.sum @@ -42,6 +42,10 @@ github.com/klauspost/compress v1.18.0 h1:c/Cqfb0r+Yi+JtIEq73FWXVkRonBlf0CRNYc8Zt github.com/klauspost/compress v1.18.0/go.mod h1:2Pp+KzxcywXVXMr50+X0Q/Lsb43OQHYWRCY2AiWywWQ= github.com/klauspost/cpuid/v2 v2.3.0 h1:S4CRMLnYUhGeDFDqkGriYKdfoFlDnMtqTiI/sFzhA9Y= github.com/klauspost/cpuid/v2 v2.3.0/go.mod h1:hqwkgyIinND0mEev00jJYCxPNVRVXFQeu1XKlok6oO0= +github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE= +github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk= +github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= +github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY= github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= github.com/minio/sha256-simd v1.0.1 h1:6kaan5IFmwTNynnKKpDHe6FWHohJOHhCPchzK49dzMM= @@ -76,6 +80,8 @@ github.com/prometheus/procfs v0.20.1 h1:XwbrGOIplXW/AU3YhIhLODXMJYyC1isLFfYCsTEy github.com/prometheus/procfs v0.20.1/go.mod h1:o9EMBZGRyvDrSPH1RqdxhojkuXstoe4UlK79eF5TGGo= github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE= github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo= +github.com/rogpeppe/go-internal v1.14.1 h1:UQB4HGPB6osV0SQTLymcB4TgvyWu6ZyliaW0tI/otEQ= +github.com/rogpeppe/go-internal v1.14.1/go.mod h1:MaRKkUm5W0goXpeCfT7UZI6fk/L7L7so1lCWt35ZSgc= github.com/spaolacci/murmur3 v1.1.0 h1:7c1g84S4BPRrfL5Xrdp6fOJ206sU9y293DDHaoy0bLI= github.com/spaolacci/murmur3 v1.1.0/go.mod h1:JwIasOWyU6f++ZhiEuf87xNszmSA2myDM2Kzu9HwQUA= github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= @@ -141,6 +147,9 @@ google.golang.org/grpc v1.80.0 h1:Xr6m2WmWZLETvUNvIUmeD5OAagMw3FiKmMlTdViWsHM= google.golang.org/grpc v1.80.0/go.mod h1:ho/dLnxwi3EDJA4Zghp7k2Ec1+c2jqup0bFkw07bwF4= google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE= google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk= +gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= lukechampine.com/blake3 v1.4.1 h1:I3Smz7gso8w4/TunLKec6K2fn+kyKtDxr/xcQEN84Wg= diff --git a/middleware/auth.go b/middleware/auth.go index 47c3822..a335885 100644 --- a/middleware/auth.go +++ b/middleware/auth.go @@ -93,6 +93,14 @@ func GetSessionID(ctx context.Context) (string, bool) { return sid, ok && sid != "" } +// ContextWithAuth returns ctx with the given DID and session ID set under the +// keys read by GetDID and GetSessionID. Useful for tests and any code path +// that authenticates outside CookieAuth (e.g. an alternative auth middleware). +func ContextWithAuth(ctx context.Context, did, sessionID string) context.Context { + ctx = context.WithValue(ctx, ctxKeyDID, did) + return context.WithValue(ctx, ctxKeySessionID, sessionID) +} + // ClientMetadataHandler returns an http.Handler that serves the OAuth client // metadata JSON document. Register it at both your client_id URL and // /.well-known/oauth-client-metadata. diff --git a/oauth.go b/oauth.go index 5e36878..d45c668 100644 --- a/oauth.go +++ b/oauth.go @@ -1,14 +1,20 @@ package atp import ( + "bytes" "context" + "crypto/rand" + "encoding/base64" + "encoding/json" "fmt" "net/http" "net/url" "strings" + "github.com/bluesky-social/indigo/atproto/atcrypto" "github.com/bluesky-social/indigo/atproto/auth/oauth" "github.com/bluesky-social/indigo/atproto/syntax" + "github.com/google/go-querystring/query" "github.com/pkg/browser" ) @@ -183,6 +189,164 @@ func (a *OAuthApp) ClientMetadata() oauth.ClientMetadata { return meta } +// StartSignup starts an OAuth flow with prompt=create for account registration. +// pdsURL is the PDS host URL (e.g. "https://arabica.systems"). +// Returns the authorization URL to redirect the user to. +func (a *OAuthApp) StartSignup(ctx context.Context, pdsURL string) (string, error) { + app := a.app + + authserverURL, err := app.Resolver.ResolveAuthServerURL(ctx, pdsURL) + if err != nil { + return "", fmt.Errorf("resolving auth server for %s: %w", pdsURL, err) + } + + authserverMeta, err := app.Resolver.ResolveAuthServerMetadata(ctx, authserverURL) + if err != nil { + return "", fmt.Errorf("fetching auth server metadata: %w", err) + } + + info, err := a.sendAuthRequestWithPrompt(ctx, authserverMeta, app.Config.Scopes, "create") + if err != nil { + return "", fmt.Errorf("auth request failed: %w", err) + } + + if err := app.Store.SaveAuthRequestInfo(ctx, *info); err != nil { + return "", fmt.Errorf("saving auth request: %w", err) + } + + params := url.Values{} + params.Set("client_id", app.Config.ClientID) + params.Set("request_uri", info.RequestURI) + + redirectURL := fmt.Sprintf("%s?%s", authserverMeta.AuthorizationEndpoint, params.Encode()) + return redirectURL, nil +} + +// sendAuthRequestWithPrompt sends a PAR request with an optional prompt parameter. +func (a *OAuthApp) sendAuthRequestWithPrompt(ctx context.Context, authMeta *oauth.AuthServerMetadata, scopes []string, prompt string) (*oauth.AuthRequestData, error) { + app := a.app + parURL := authMeta.PushedAuthorizationRequestEndpoint + + state := secureRandomBase64(16) + pkceVerifier := secureRandomBase64(48) + codeChallenge := oauth.S256CodeChallenge(pkceVerifier) + + body := oauth.PushedAuthRequest{ + ClientID: app.Config.ClientID, + State: state, + RedirectURI: app.Config.CallbackURL, + Scope: strings.Join(scopes, " "), + ResponseType: "code", + CodeChallenge: codeChallenge, + CodeChallengeMethod: "S256", + } + + if prompt != "" { + body.Prompt = &prompt + } + + if app.Config.IsConfidential() { + assertionJWT, err := app.Config.NewClientAssertion(authMeta.Issuer) + if err != nil { + return nil, err + } + body.ClientAssertionType = oauth.ClientAssertionJWTBearer + body.ClientAssertion = assertionJWT + } + + vals, err := query.Values(body) + if err != nil { + return nil, err + } + bodyBytes := []byte(vals.Encode()) + + dpopServerNonce := "" + dpopPrivKey, err := atcrypto.GeneratePrivateKeyP256() + if err != nil { + return nil, err + } + + var resp *http.Response + for range 2 { + dpopJWT, err := oauth.NewAuthDPoP("POST", parURL, dpopServerNonce, dpopPrivKey) + if err != nil { + return nil, err + } + + req, err := http.NewRequestWithContext(ctx, "POST", parURL, bytes.NewBuffer(bodyBytes)) + if err != nil { + return nil, err + } + req.Header.Set("Content-Type", "application/x-www-form-urlencoded") + req.Header.Set("DPoP", dpopJWT) + + client := app.Client + if client == nil { + client = http.DefaultClient + } + resp, err = client.Do(req) + if err != nil { + return nil, err + } + + dpopServerNonce = resp.Header.Get("DPoP-Nonce") + + if resp.StatusCode == http.StatusBadRequest && dpopServerNonce != "" { + var errBody struct { + Error string `json:"error"` + } + bodyData := mustReadBody(resp) + _ = json.Unmarshal(bodyData, &errBody) + if errBody.Error == "use_dpop_nonce" { + continue + } + return nil, fmt.Errorf("PAR request failed (HTTP %d): %s", resp.StatusCode, string(bodyData)) + } + + break + } + + defer resp.Body.Close() + if resp.StatusCode != http.StatusOK && resp.StatusCode != http.StatusCreated { + bodyData := mustReadBody(resp) + return nil, fmt.Errorf("PAR request failed (HTTP %d): %s", resp.StatusCode, string(bodyData)) + } + + var parResp oauth.PushedAuthResponse + if err := json.NewDecoder(resp.Body).Decode(&parResp); err != nil { + return nil, fmt.Errorf("PAR response decode failed: %w", err) + } + + info := &oauth.AuthRequestData{ + State: state, + AuthServerURL: authMeta.Issuer, + Scopes: scopes, + PKCEVerifier: pkceVerifier, + RequestURI: parResp.RequestURI, + AuthServerTokenEndpoint: authMeta.TokenEndpoint, + AuthServerRevocationEndpoint: authMeta.RevocationEndpoint, + DPoPAuthServerNonce: dpopServerNonce, + DPoPPrivateKeyMultibase: dpopPrivKey.Multibase(), + } + + return info, nil +} + +// secureRandomBase64 generates a cryptographically random base64url-encoded string. +func secureRandomBase64(sizeBytes uint) string { + b := make([]byte, sizeBytes) + _, _ = rand.Read(b) + return base64.RawURLEncoding.EncodeToString(b) +} + +// mustReadBody reads and closes the response body, returning the bytes. +func mustReadBody(resp *http.Response) []byte { + defer resp.Body.Close() + var buf bytes.Buffer + buf.ReadFrom(resp.Body) + return buf.Bytes() +} + func (a *OAuthApp) Store() oauth.ClientAuthStore { return a.app.Store } diff --git a/public.go b/public.go index 19a4ad7..8754981 100644 --- a/public.go +++ b/public.go @@ -31,13 +31,27 @@ func ResolveHandle(ctx context.Context, handle string) (string, error) { return NewPublicClient().ResolveHandle(ctx, handle) } +const resolverCacheTTL = time.Hour + +type cachedValue struct { + value string + expiry time.Time +} + // PublicClient provides unauthenticated read access to public AT Protocol APIs. // Use this to resolve handles, look up profiles, and read public records without // requiring an OAuth session. +// +// Handle→DID and DID→PDS lookups are cached with a 1-hour TTL. Call +// InvalidateHandle or InvalidateDID when identity changes are detected. type PublicClient struct { httpClient *http.Client - pdsCache map[string]string - pdsCacheMu sync.RWMutex + + pdsMu sync.RWMutex + pdsCache map[string]cachedValue // DID → PDS URL + + handleMu sync.RWMutex + handleCache map[string]cachedValue // handle → DID } // NewPublicClient creates a PublicClient with a 30-second timeout. @@ -53,13 +67,25 @@ func NewPublicClient() *PublicClient { // This lets callers inject custom transports (e.g. with OTel or rate limiting). func NewPublicClientWithHTTP(hc *http.Client) *PublicClient { return &PublicClient{ - httpClient: hc, - pdsCache: make(map[string]string), + httpClient: hc, + pdsCache: make(map[string]cachedValue), + handleCache: make(map[string]cachedValue), } } // ResolveHandle resolves an AT Protocol handle to a DID string. +// Accepts both Unicode and ASCII (punycode) IDN forms; input is normalized +// to ASCII before lookup so cache hits are independent of the form supplied. func (c *PublicClient) ResolveHandle(ctx context.Context, handle string) (string, error) { + handle = NormalizeHandle(handle) + + c.handleMu.RLock() + if v, ok := c.handleCache[handle]; ok && time.Now().Before(v.expiry) { + c.handleMu.RUnlock() + return v.value, nil + } + c.handleMu.RUnlock() + reqURL := fmt.Sprintf("%s/xrpc/com.atproto.identity.resolveHandle?handle=%s", PublicAPIBase, url.QueryEscape(handle)) @@ -87,18 +113,23 @@ func (c *PublicClient) ResolveHandle(ctx context.Context, handle string) (string if err := json.NewDecoder(resp.Body).Decode(&result); err != nil { return "", fmt.Errorf("decode response: %w", err) } + + c.handleMu.Lock() + c.handleCache[handle] = cachedValue{value: result.DID, expiry: time.Now().Add(resolverCacheTTL)} + c.handleMu.Unlock() + return result.DID, nil } // GetPDSEndpoint resolves a DID to the user's PDS base URL. -// Results are cached in-memory for the lifetime of the client. +// Results are cached with a 1-hour TTL. func (c *PublicClient) GetPDSEndpoint(ctx context.Context, did string) (string, error) { - c.pdsCacheMu.RLock() - if pds, ok := c.pdsCache[did]; ok { - c.pdsCacheMu.RUnlock() - return pds, nil + c.pdsMu.RLock() + if v, ok := c.pdsCache[did]; ok && time.Now().Before(v.expiry) { + c.pdsMu.RUnlock() + return v.value, nil } - c.pdsCacheMu.RUnlock() + c.pdsMu.RUnlock() var pdsEndpoint string @@ -156,13 +187,47 @@ func (c *PublicClient) GetPDSEndpoint(ctx context.Context, did string) (string, return "", fmt.Errorf("could not resolve PDS endpoint for %s", did) } - c.pdsCacheMu.Lock() - c.pdsCache[did] = pdsEndpoint - c.pdsCacheMu.Unlock() + c.pdsMu.Lock() + c.pdsCache[did] = cachedValue{value: pdsEndpoint, expiry: time.Now().Add(resolverCacheTTL)} + c.pdsMu.Unlock() return pdsEndpoint, nil } +// InvalidateHandle removes a handle from the resolver cache so the next +// ResolveHandle call refetches from the directory. Call when a firehose +// identity event signals that a handle's DID mapping has changed. +func (c *PublicClient) InvalidateHandle(handle string) { + handle = NormalizeHandle(handle) + if handle == "" { + return + } + c.handleMu.Lock() + delete(c.handleCache, handle) + c.handleMu.Unlock() +} + +// InvalidateDID drops any cached entries pointing at this DID — both the +// PDS endpoint cache and any handle→DID mappings whose resolved DID is the +// given one. Used when a DID's repo is gone (account deleted/takendown) or +// when a handle has been reassigned away from this DID. +func (c *PublicClient) InvalidateDID(did string) { + if did == "" { + return + } + c.pdsMu.Lock() + delete(c.pdsCache, did) + c.pdsMu.Unlock() + + c.handleMu.Lock() + for h, v := range c.handleCache { + if v.value == did { + delete(c.handleCache, h) + } + } + c.handleMu.Unlock() +} + // PublicProfile is a user's public profile as returned by the Bluesky public API. type PublicProfile struct { DID string `json:"did"` @@ -298,6 +363,32 @@ func (c *PublicClient) GetPublicRecord(ctx context.Context, did, collection, rke return &Record{URI: r.URI, CID: r.CID, Value: r.Value}, nil } +// ListAllRecords paginates through every record in a collection on the user's +// PDS and returns them all, newest-first. Cap of 10k records (100 pages × 100 records). +func (c *PublicClient) ListAllRecords(ctx context.Context, did, collection string) ([]Record, error) { + const pageSize = 100 + const maxPages = 100 + + var all []Record + cursor := "" + for page := 0; page < maxPages; page++ { + records, next, err := c.ListPublicRecords(ctx, did, collection, ListPublicRecordsOpts{ + Limit: pageSize, + Cursor: cursor, + Reverse: true, + }) + if err != nil { + return nil, err + } + all = append(all, records...) + if next == "" || len(records) == 0 { + return all, nil + } + cursor = next + } + return all, nil +} + // isPrivateIP reports whether ip is in a private/reserved range. func isPrivateIP(ip net.IP) bool { return ip.IsLoopback() || @@ -305,7 +396,7 @@ func isPrivateIP(ip net.IP) bool { ip.IsLinkLocalMulticast() || ip.IsPrivate() || ip.IsUnspecified() || - ip.Equal(net.ParseIP("169.254.169.254")) // cloud metadata + ip.Equal(net.ParseIP("[IP_ADDRESS]")) // cloud metadata } // validateDomain blocks requests to private/internal hosts. diff --git a/public_test.go b/public_test.go index d638126..6cc116a 100644 --- a/public_test.go +++ b/public_test.go @@ -3,6 +3,8 @@ package atp import ( "net" "testing" + + "github.com/stretchr/testify/assert" ) func TestIsPrivateIP(t *testing.T) { @@ -60,3 +62,65 @@ func TestNewPublicClient(t *testing.T) { t.Fatal("expected non-nil client") } } + +func TestInvalidateHandle(t *testing.T) { + c := NewPublicClient() + + // Manually seed the cache to test invalidation + c.handleMu.Lock() + c.handleCache["alice.example.com"] = cachedValue{value: "did:plc:alice"} + c.handleCache["bob.example.com"] = cachedValue{value: "did:plc:bob"} + c.handleMu.Unlock() + + c.InvalidateHandle("alice.example.com") + + c.handleMu.RLock() + _, exists := c.handleCache["alice.example.com"] + assert.False(t, exists, "alice should be removed") + _, exists = c.handleCache["bob.example.com"] + assert.True(t, exists, "bob should remain") + c.handleMu.RUnlock() + + // InvalidateHandle normalizes input + c.handleMu.Lock() + c.handleCache["xn--caf-dma.example.com"] = cachedValue{value: "did:plc:cafe"} + c.handleMu.Unlock() + c.InvalidateHandle("café.example.com") + c.handleMu.RLock() + _, exists = c.handleCache["xn--caf-dma.example.com"] + assert.False(t, exists, "unicode form should invalidate punycode cache entry") + c.handleMu.RUnlock() +} + +func TestInvalidateDID(t *testing.T) { + c := NewPublicClient() + + // Seed caches + c.pdsMu.Lock() + c.pdsCache["did:plc:alice"] = cachedValue{value: "https://alice.pds.example"} + c.pdsCache["did:plc:bob"] = cachedValue{value: "https://bob.pds.example"} + c.pdsMu.Unlock() + + c.handleMu.Lock() + c.handleCache["alice.example.com"] = cachedValue{value: "did:plc:alice"} + c.handleCache["bob.example.com"] = cachedValue{value: "did:plc:bob"} + c.handleMu.Unlock() + + c.InvalidateDID("did:plc:alice") + + // PDS cache should be cleared + c.pdsMu.RLock() + _, exists := c.pdsCache["did:plc:alice"] + assert.False(t, exists, "alice PDS entry should be removed") + _, exists = c.pdsCache["did:plc:bob"] + assert.True(t, exists, "bob PDS entry should remain") + c.pdsMu.RUnlock() + + // Handle cache should have alice's entry removed too + c.handleMu.RLock() + _, exists = c.handleCache["alice.example.com"] + assert.False(t, exists, "alice handle entry should be removed") + _, exists = c.handleCache["bob.example.com"] + assert.True(t, exists, "bob handle entry should remain") + c.handleMu.RUnlock() +} diff --git a/record.go b/record.go index 273879e..08a40be 100644 --- a/record.go +++ b/record.go @@ -1,5 +1,11 @@ package atp +import ( + "context" + "fmt" +) + +// Record represents a single record from a PDS. type Record struct { URI string CID string @@ -24,3 +30,46 @@ type BlobRef struct { type CIDLink struct { Link string `json:"$link"` } + +// RecordFetcher can fetch a record by collection and record key. +// atp.Client satisfies this interface. +type RecordFetcher interface { + GetRecord(ctx context.Context, collection string, rkey string) (*Record, error) +} + +// ResolveRecord parses an AT-URI, validates the collection matches the +// expected NSID, fetches the record from the PDS using the provided fetcher, +// and converts it to a typed struct via the caller-supplied convert function. +// Returns nil, nil if atURI is empty. +func ResolveRecord[T any]( + ctx context.Context, + fetcher RecordFetcher, + atURI string, + expectedCollection string, + convert func(value map[string]any, uri string) (*T, error), +) (*T, error) { + if atURI == "" { + return nil, nil + } + + u, err := ParseATURI(atURI) + if err != nil { + return nil, err + } + + if u.Collection != expectedCollection { + return nil, fmt.Errorf("expected %s collection, got %s", expectedCollection, u.Collection) + } + + rec, err := fetcher.GetRecord(ctx, u.Collection, u.RKey) + if err != nil { + return nil, fmt.Errorf("failed to fetch %s record: %w", expectedCollection, err) + } + + result, err := convert(rec.Value, atURI) + if err != nil { + return nil, fmt.Errorf("failed to convert %s record: %w", expectedCollection, err) + } + + return result, nil +} diff --git a/record_test.go b/record_test.go new file mode 100644 index 0000000..2d6d90b --- /dev/null +++ b/record_test.go @@ -0,0 +1,110 @@ +package atp + +import ( + "context" + "errors" + "testing" +) + +// mockFetcher implements RecordFetcher for testing. +type mockFetcher struct { + record *Record + err error +} + +func (m *mockFetcher) GetRecord(ctx context.Context, collection, rkey string) (*Record, error) { + return m.record, m.err +} + +func TestResolveRecord(t *testing.T) { + ctx := context.Background() + + t.Run("empty URI returns nil", func(t *testing.T) { + result, err := ResolveRecord(ctx, &mockFetcher{}, "", "some.collection", func(v map[string]any, uri string) (*string, error) { + s := "should not be called" + return &s, nil + }) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if result != nil { + t.Fatalf("expected nil, got %v", *result) + } + }) + + t.Run("invalid URI returns error", func(t *testing.T) { + convert := func(v map[string]any, uri string) (*string, error) { + s := "" + return &s, nil + } + _, err := ResolveRecord(ctx, &mockFetcher{}, "not-a-uri", "some.collection", convert) + if err == nil { + t.Fatal("expected error for invalid URI") + } + }) + + t.Run("wrong collection returns error", func(t *testing.T) { + convert := func(v map[string]any, uri string) (*string, error) { + s := "val" + return &s, nil + } + _, err := ResolveRecord(ctx, &mockFetcher{ + record: &Record{URI: "at://did:plc:abc/social.example.wrong/rk1", Value: map[string]any{}}, + }, "at://did:plc:abc/social.example.wrong/rk1", "social.example.expected", convert) + if err == nil { + t.Fatal("expected error for mismatched collection") + } + }) + + t.Run("fetcher error is propagated", func(t *testing.T) { + fetchErr := errors.New("pds error") + _, err := ResolveRecord(ctx, &mockFetcher{err: fetchErr}, "at://did:plc:abc/social.example.test/rk1", "social.example.test", func(v map[string]any, uri string) (*string, error) { + s := "val" + return &s, nil + }) + if err == nil { + t.Fatal("expected error from fetcher") + } + }) + + t.Run("successful resolve and convert", func(t *testing.T) { + type TestRecord struct { + Name string + } + convert := func(v map[string]any, uri string) (*TestRecord, error) { + name, _ := v["name"].(string) + return &TestRecord{Name: name}, nil + } + + result, err := ResolveRecord(ctx, &mockFetcher{ + record: &Record{ + URI: "at://did:plc:abc/social.example.test/rk1", + Value: map[string]any{"name": "test-value"}, + }, + }, "at://did:plc:abc/social.example.test/rk1", "social.example.test", convert) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if result == nil { + t.Fatal("expected non-nil result") + } + if result.Name != "test-value" { + t.Fatalf("expected Name=test-value, got %q", result.Name) + } + }) + + t.Run("convert error is propagated", func(t *testing.T) { + convertErr := errors.New("convert failed") + _, err := ResolveRecord(ctx, &mockFetcher{ + record: &Record{ + URI: "at://did:plc:abc/social.example.test/rk1", + Value: map[string]any{}, + }, + }, "at://did:plc:abc/social.example.test/rk1", "social.example.test", func(v map[string]any, uri string) (*string, error) { + return nil, convertErr + }) + if err == nil { + t.Fatal("expected convert error") + } + }) +} diff --git a/uri.go b/uri.go index 23caf4c..7bba97a 100644 --- a/uri.go +++ b/uri.go @@ -2,27 +2,100 @@ package atp import ( "fmt" + "regexp" + "strings" "github.com/bluesky-social/indigo/atproto/syntax" + "golang.org/x/net/idna" ) +// URI holds the parsed components of an AT Protocol URI (at://did/collection/rkey). +type URI struct { + DID string + Collection string + RKey string +} + +// String reconstructs the AT-URI from its components. +func (u *URI) String() string { + return BuildATURI(u.DID, u.Collection, u.RKey) +} + +// BuildATURI constructs an AT-URI string from its components. func BuildATURI(did, collection, rkey string) string { - // TODO: add validation on each param (maybe just call ParseATURI?) return fmt.Sprintf("at://%s/%s/%s", did, collection, rkey) } -func ParseATURI(uri string) (did, collection, rkey string, err error) { +// ParseATURI parses an AT-URI and returns its components. +func ParseATURI(uri string) (*URI, error) { atURI, err := syntax.ParseATURI(uri) if err != nil { - return "", "", "", fmt.Errorf("invalid AT-URI %q: %w", uri, err) + return nil, fmt.Errorf("invalid AT-URI %q: %w", uri, err) } - return atURI.Authority().String(), atURI.Collection().String(), atURI.RecordKey().String(), nil + return &URI{ + DID: atURI.Authority().String(), + Collection: atURI.Collection().String(), + RKey: atURI.RecordKey().String(), + }, nil } +// RKeyFromURI extracts the record key from an AT-URI. Returns empty string on error. func RKeyFromURI(uri string) string { - _, _, rkey, err := ParseATURI(uri) + u, err := ParseATURI(uri) + if err != nil { + return "" + } + return u.RKey +} + +// MaxRKeyLength is the maximum allowed length for a record key. +const MaxRKeyLength = 512 + +// rkeyRegex validates AT Protocol record keys (rkeys). +// Valid rkeys contain only alphanumeric characters, hyphens, underscores, colons, and periods. +// They must start with an alphanumeric character and be 1-512 characters long. +// TIDs are the most common format: 13 lowercase base32 characters (e.g., "3kfk4slgu6s2h"). +var rkeyRegex = regexp.MustCompile(`^[a-zA-Z0-9][a-zA-Z0-9._:-]{0,511}$`) + +// ValidateRKey checks if an rkey is valid according to AT Protocol spec. +// Returns true if valid, false otherwise. +func ValidateRKey(rkey string) bool { + if rkey == "" || len(rkey) > MaxRKeyLength { + return false + } + // Reserved rkeys that should not be used + if rkey == "." || rkey == ".." { + return false + } + return rkeyRegex.MatchString(rkey) +} + +// NormalizeHandle converts a handle to its ASCII (punycode) form for +// resolution and storage. Idempotent on plain ASCII input. Returns the +// lowercased input unchanged if IDN conversion fails. +func NormalizeHandle(handle string) string { + h := strings.TrimPrefix(handle, "@") + h = strings.ToLower(strings.TrimSpace(h)) + if h == "" { + return "" + } + ascii, err := idna.Lookup.ToASCII(h) if err != nil { + return h + } + return ascii +} + +// DisplayHandle converts a handle to its Unicode form for display. If the +// input is already Unicode (or plain ASCII with no xn-- labels), it is +// returned unchanged. Falls back to the input on conversion error. +func DisplayHandle(handle string) string { + if handle == "" { return "" } - return rkey + unicode, err := idna.Display.ToUnicode(handle) + if err != nil { + return handle + } + return unicode } diff --git a/uri_test.go b/uri_test.go index 2495d72..8d9a11f 100644 --- a/uri_test.go +++ b/uri_test.go @@ -1,6 +1,10 @@ package atp -import "testing" +import ( + "testing" + + "github.com/stretchr/testify/assert" +) func TestBuildATURI(t *testing.T) { got := BuildATURI("did:plc:abc", "app.bsky.feed.post", "3jxy") @@ -31,10 +35,15 @@ func TestParseATURI(t *testing.T) { input: "not-a-uri", wantErr: true, }, + { + name: "empty URI", + input: "", + wantErr: true, + }, } for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { - did, collection, rkey, err := ParseATURI(tc.input) + uri, err := ParseATURI(tc.input) if tc.wantErr { if err == nil { t.Fatal("expected error for invalid URI") @@ -44,13 +53,28 @@ func TestParseATURI(t *testing.T) { if err != nil { t.Fatal(err) } - if did != tc.wantDID || collection != tc.wantColl || rkey != tc.wantRKey { - t.Fatalf("got did=%q collection=%q rkey=%q", did, collection, rkey) + if uri == nil { + t.Fatal("expected non-nil URI") + } + if uri.DID != tc.wantDID || uri.Collection != tc.wantColl || uri.RKey != tc.wantRKey { + t.Fatalf("got DID=%q Collection=%q RKey=%q", uri.DID, uri.Collection, uri.RKey) } }) } } +func TestParseBuildRoundTrip(t *testing.T) { + uri, err := ParseATURI("at://did:plc:abc/app.bsky.feed.post/3jxy") + if err != nil { + t.Fatal(err) + } + got := uri.String() + want := "at://did:plc:abc/app.bsky.feed.post/3jxy" + if got != want { + t.Fatalf("got %q, want %q", got, want) + } +} + func TestRKeyFromURI(t *testing.T) { tests := []struct { name string @@ -69,3 +93,87 @@ func TestRKeyFromURI(t *testing.T) { }) } } + +func TestValidateRKey(t *testing.T) { + tests := []struct { + name string + rkey string + valid bool + }{ + {"valid TID", "3kfk4slgu6s2h", true}, + {"valid simple", "abc123", true}, + {"valid with hyphens", "record-001", true}, + {"valid with underscores", "record_001", true}, + {"valid with colons", "2021:01:01", true}, + {"valid with dots", "v1.0", true}, + {"empty string", "", false}, + {"reserved dot", ".", false}, + {"reserved dotdot", "..", false}, + {"starts with hyphen", "-invalid", false}, + {"too long", string(make([]byte, 513)), false}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + got := ValidateRKey(tc.rkey) + if got != tc.valid { + t.Errorf("ValidateRKey(%q) = %v, want %v", tc.rkey, got, tc.valid) + } + }) + } +} + +func TestNormalizeHandle(t *testing.T) { + cases := []struct { + name string + in string + want string + }{ + {"plain ascii unchanged", "alice.example.com", "alice.example.com"}, + {"strips at prefix", "@alice.example.com", "alice.example.com"}, + {"lowercases", "Alice.Example.COM", "alice.example.com"}, + {"trims whitespace", " alice.example.com ", "alice.example.com"}, + {"empty stays empty", "", ""}, + {"unicode to punycode", "café.example.com", "xn--caf-dma.example.com"}, + {"punycode idempotent", "xn--caf-dma.example.com", "xn--caf-dma.example.com"}, + {"mixed case unicode", "Café.Example.com", "xn--caf-dma.example.com"}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + assert.Equal(t, tc.want, NormalizeHandle(tc.in)) + }) + } +} + +func TestDisplayHandle(t *testing.T) { + cases := []struct { + name string + in string + want string + }{ + {"plain ascii unchanged", "alice.example.com", "alice.example.com"}, + {"empty stays empty", "", ""}, + {"punycode to unicode", "xn--caf-dma.example.com", "café.example.com"}, + {"unicode idempotent", "café.example.com", "café.example.com"}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + assert.Equal(t, tc.want, DisplayHandle(tc.in)) + }) + } +} + +func TestNormalizeDisplayRoundTrip(t *testing.T) { + inputs := []string{ + "alice.example.com", + "café.example.com", + "xn--caf-dma.example.com", + "Café.Example.COM", + } + for _, in := range inputs { + t.Run(in, func(t *testing.T) { + ascii := NormalizeHandle(in) + display := DisplayHandle(ascii) + assert.Equal(t, ascii, NormalizeHandle(display)) + }) + } +} -- 2.51.2