diff --git a/client.go b/client.go index eddc9ab..7956930 100644 --- a/client.go +++ b/client.go @@ -29,14 +29,12 @@ func NewClient(api *atclient.APIClient, did syntax.DID) *Client { return &Client{api: api, did: did} } -// DID returns the authenticated user's DID. func (c *Client) DID() syntax.DID { return c.did } // APIClient returns the underlying indigo APIClient for advanced usage // (custom XRPC calls, service proxying, etc.). func (c *Client) APIClient() *atclient.APIClient { return c.api } -// CreateRecord creates a new record with an auto-generated TID key. func (c *Client) CreateRecord(ctx context.Context, collection string, record any) (uri, cid string, err error) { body := map[string]any{ "repo": c.did.String(), @@ -54,7 +52,6 @@ func (c *Client) CreateRecord(ctx context.Context, collection string, record any return result.URI, result.CID, nil } -// CreateRecordWithRKey creates a new record with a specific record key. func (c *Client) CreateRecordWithRKey(ctx context.Context, collection, rkey string, record any) (uri, cid string, err error) { body := map[string]any{ "repo": c.did.String(), @@ -92,8 +89,8 @@ func (c *Client) GetRecord(ctx context.Context, collection, rkey string) (*Recor return &Record{URI: result.URI, CID: result.CID, Value: result.Value}, nil } -// ListRecords retrieves a single page of records from a collection. -// Pass limit <= 0 for the server default (usually 50). Pass empty cursor for the first page. +// Pass limit <= 0 for the server default (usually 50). +// Pass empty cursor for the first page. func (c *Client) ListRecords(ctx context.Context, collection string, limit int, cursor string) (*ListResult, error) { params := map[string]any{ "repo": c.did.String(), @@ -130,8 +127,6 @@ func (c *Client) ListRecords(ctx context.Context, collection string, limit int, return out, nil } -// ListAllRecords fetches every record in a collection, handling cursor -// pagination automatically. Returns all records at once. func (c *Client) ListAllRecords(ctx context.Context, collection string) ([]Record, error) { var all []Record cursor := "" @@ -157,7 +152,6 @@ func (c *Client) ListAllRecords(ctx context.Context, collection string) ([]Recor return all, nil } -// PutRecord creates or updates a record at a specific record key. func (c *Client) PutRecord(ctx context.Context, collection, rkey string, record any) (uri, cid string, err error) { body := map[string]any{ "repo": c.did.String(), @@ -176,7 +170,6 @@ func (c *Client) PutRecord(ctx context.Context, collection, rkey string, record return result.URI, result.CID, nil } -// DeleteRecord removes a record from the user's repository. func (c *Client) DeleteRecord(ctx context.Context, collection, rkey string) error { body := map[string]any{ "repo": c.did.String(), @@ -191,7 +184,6 @@ func (c *Client) DeleteRecord(ctx context.Context, collection, rkey string) erro return nil } -// UploadBlob uploads a blob to the user's PDS. // Data must be at most 1 MB (MaxBlobSize). The mimeType should match the blob content. func (c *Client) UploadBlob(ctx context.Context, data []byte, mimeType string) (*BlobRef, error) { if len(data) > MaxBlobSize { @@ -235,7 +227,6 @@ func (c *Client) UploadBlob(ctx context.Context, data []byte, mimeType string) ( }, nil } -// GetBlob downloads a blob from the user's PDS by its CID. func (c *Client) GetBlob(ctx context.Context, cid string) ([]byte, error) { data, err := atproto.SyncGetBlob(ctx, c.api, cid, c.did.String()) if err != nil { diff --git a/errors.go b/errors.go index d734c30..462aef1 100644 --- a/errors.go +++ b/errors.go @@ -12,6 +12,7 @@ var ErrSessionExpired = errors.New("oauth session expired") // WrapPDSError inspects an XRPC error for signals that the OAuth grant is no // longer valid and, if so, wraps it with ErrSessionExpired. +// TODO: handle other common error types func WrapPDSError(err error) error { if err == nil { return nil diff --git a/go.mod b/go.mod index 4ce39e6..dd49aab 100644 --- a/go.mod +++ b/go.mod @@ -4,6 +4,8 @@ go 1.25.5 require ( github.com/bluesky-social/indigo v0.0.0-20260318212431-cbaa83aee9dd + github.com/gorilla/websocket v1.5.3 + github.com/klauspost/compress v1.17.3 github.com/pkg/browser v0.0.0-20240102092130-5ac0b6a4141c go.etcd.io/bbolt v1.4.3 go.opentelemetry.io/otel v1.43.0 diff --git a/go.sum b/go.sum index 08f87b9..5aeda22 100644 --- a/go.sum +++ b/go.sum @@ -30,12 +30,16 @@ github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17k github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e/go.mod h1:boTsfXsheKC2y+lKOCMpSfarhxDeIzfZG1jqGcPl3cA= github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg= +github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE= github.com/grpc-ecosystem/grpc-gateway/v2 v2.28.0 h1:HWRh5R2+9EifMyIHV7ZV+MIZqgz+PMpZ14Jynv3O2Zs= github.com/grpc-ecosystem/grpc-gateway/v2 v2.28.0/go.mod h1:JfhWUomR1baixubs02l85lZYYOm7LV6om4ceouMv45c= github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k= github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM= github.com/ipfs/go-cid v0.4.1 h1:A/T3qGvxi4kpKWWcPC/PgbvDA2bjVLO7n4UeVwnbs/s= github.com/ipfs/go-cid v0.4.1/go.mod h1:uQHwDeX4c6CtyrFwdqyhpNcxVewur1M7l7fNU7LKwZk= +github.com/klauspost/compress v1.17.3 h1:qkRjuerhUU1EmXLYGkSH6EZL+vPSxIrYjLNAK4slzwA= +github.com/klauspost/compress v1.17.3/go.mod h1:/dCuZOvVtNoHsyb+cuJD3itjs3NbnF6KH9zAO4BDxPM= github.com/klauspost/cpuid/v2 v2.2.7 h1:ZWSB3igEs+d0qvnxR/ZBzXVmxkgt8DdzP6m9pfuVLDM= github.com/klauspost/cpuid/v2 v2.2.7/go.mod h1:Lcz8mBdAVJIBVzewtcLocK12l3Y+JytZYpaMropDUws= github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY= diff --git a/jetstream/jetstream.go b/jetstream/jetstream.go new file mode 100644 index 0000000..95052ec --- /dev/null +++ b/jetstream/jetstream.go @@ -0,0 +1,360 @@ +// Package jetstream consumes real-time AT Protocol events from a Jetstream relay. +// +// Jetstream is a WebSocket-based relay that delivers a filtered stream of AT +// Protocol repository events (commits, identity changes, account changes). +// This package handles connection management, reconnection with backoff, +// endpoint rotation, cursor tracking, and optional zstd decompression. +// +// Basic usage: +// +// consumer := jetstream.New(&jetstream.Config{ +// WantedCollections: []string{"app.bsky.feed.post"}, +// }, func(ctx context.Context, evt *jetstream.Event) error { +// fmt.Printf("new post from %s\n", evt.DID) +// return nil +// }) +// consumer.Start(ctx) +// defer consumer.Stop() +package jetstream + +import ( + "context" + "encoding/json" + "fmt" + "net/url" + "sync" + "sync/atomic" + "time" + + "github.com/gorilla/websocket" + "github.com/klauspost/compress/zstd" +) + +// DefaultEndpoints are the public Jetstream relay endpoints. +var DefaultEndpoints = []string{ + "wss://jetstream1.us-east.bsky.network/subscribe", + "wss://jetstream2.us-east.bsky.network/subscribe", + "wss://jetstream1.us-west.bsky.network/subscribe", + "wss://jetstream2.us-west.bsky.network/subscribe", +} + +// Event is a single event from the Jetstream relay. +type Event struct { + DID string `json:"did"` + TimeUS int64 `json:"time_us"` + Kind string `json:"kind"` // "commit", "identity", "account" + Commit *Commit `json:"commit,omitempty"` +} + +// Commit is the commit payload within an Event. +type Commit struct { + Rev string `json:"rev"` + Operation string `json:"operation"` // "create", "update", "delete" + Collection string `json:"collection"` + RKey string `json:"rkey"` + Record json.RawMessage `json:"record,omitempty"` + CID string `json:"cid"` +} + +// Handler is called for each event received from Jetstream. +// Returning an error logs a warning but does not stop the consumer. +type Handler func(ctx context.Context, event *Event) error + +// CursorStore persists the Jetstream cursor across restarts. +// If nil, the cursor is tracked in memory only (replay from live on restart). +type CursorStore interface { + GetCursor(ctx context.Context) (int64, error) + SetCursor(ctx context.Context, cursor int64) error +} + +// Config configures a Jetstream consumer. +type Config struct { + // Endpoints is the list of Jetstream WebSocket URLs. Defaults to DefaultEndpoints. + Endpoints []string + + // WantedCollections filters events to specific NSIDs. + // Empty means all collections (high volume). + WantedCollections []string + + // Compress enables zstd compression. Disabled by default because Jetstream + // uses a custom dictionary incompatible with the standard zstd decoder. + Compress bool + + // CursorStore persists the cursor for resume after restart. + // If nil, the consumer starts from live on each restart. + CursorStore CursorStore + + // CursorPersistEvery controls how often the cursor is flushed to CursorStore. + // Defaults to every 1000 events. + CursorPersistEvery int64 + + // OnConnect is called each time a WebSocket connection is established. + OnConnect func() + + // OnDisconnect is called each time a connection is lost. + OnDisconnect func() + + // OnError is called when the handler returns an error. + // If nil, errors are silently dropped (caller should log in the handler). + OnError func(err error, event *Event) +} + +func (c *Config) endpoints() []string { + if len(c.Endpoints) > 0 { + return c.Endpoints + } + return DefaultEndpoints +} + +func (c *Config) cursorPersistEvery() int64 { + if c.CursorPersistEvery > 0 { + return c.CursorPersistEvery + } + return 1000 +} + +// Consumer consumes events from a Jetstream relay. +type Consumer struct { + cfg *Config + handler Handler + + conn *websocket.Conn + connMu sync.Mutex + currentEndpointIdx int + + zstdDecoder *zstd.Decoder + + cursor atomic.Int64 + eventsReceived atomic.Int64 + bytesReceived atomic.Int64 + connected atomic.Bool + + stopCh chan struct{} + wg sync.WaitGroup +} + +// New creates a new Consumer. Call Start to begin consuming events. +func New(cfg *Config, handler Handler) *Consumer { + decoder, err := zstd.NewReader(nil, zstd.WithDecoderConcurrency(1)) + if err != nil { + // zstd.NewReader with nil src only fails on bad options + panic(fmt.Sprintf("jetstream: create zstd decoder: %v", err)) + } + + c := &Consumer{ + cfg: cfg, + handler: handler, + stopCh: make(chan struct{}), + zstdDecoder: decoder, + } + + if cfg.CursorStore != nil { + if cursor, err := cfg.CursorStore.GetCursor(context.Background()); err == nil && cursor > 0 { + c.cursor.Store(cursor) + } + } + + return c +} + +// Start begins consuming events in a background goroutine. +func (c *Consumer) Start(ctx context.Context) { + c.wg.Add(1) + go func() { + defer c.wg.Done() + c.run(ctx) + }() +} + +// Stop gracefully shuts down the consumer and waits for it to finish. +func (c *Consumer) Stop() { + close(c.stopCh) + c.connMu.Lock() + if c.conn != nil { + c.conn.Close() + } + c.connMu.Unlock() + c.wg.Wait() + c.zstdDecoder.Close() +} + +// IsConnected reports whether the consumer is currently connected. +func (c *Consumer) IsConnected() bool { + return c.connected.Load() +} + +// Stats returns cumulative event and byte counts since Start was called. +func (c *Consumer) Stats() (eventsReceived, bytesReceived int64) { + return c.eventsReceived.Load(), c.bytesReceived.Load() +} + +func (c *Consumer) run(ctx context.Context) { + backoff := time.Second + const maxBackoff = 30 * time.Second + + for { + select { + case <-ctx.Done(): + return + case <-c.stopCh: + return + default: + } + + endpoints := c.cfg.endpoints() + endpoint := endpoints[c.currentEndpointIdx] + + if err := c.connectAndConsume(ctx, endpoint); err != nil { + c.connected.Store(false) + if c.cfg.OnDisconnect != nil { + c.cfg.OnDisconnect() + } + + // Rotate to next endpoint + c.currentEndpointIdx = (c.currentEndpointIdx + 1) % len(endpoints) + + select { + case <-ctx.Done(): + return + case <-c.stopCh: + return + case <-time.After(backoff): + } + + backoff *= 2 + if backoff > maxBackoff { + backoff = maxBackoff + } + } else { + backoff = time.Second + } + } +} + +func (c *Consumer) connectAndConsume(ctx context.Context, endpoint string) error { + wsURL, err := c.buildURL(endpoint) + if err != nil { + return fmt.Errorf("build URL: %w", err) + } + + dialer := websocket.Dialer{HandshakeTimeout: 10 * time.Second} + conn, _, err := dialer.DialContext(ctx, wsURL, nil) + if err != nil { + return fmt.Errorf("dial: %w", err) + } + + c.connMu.Lock() + c.conn = conn + c.connMu.Unlock() + + c.connected.Store(true) + if c.cfg.OnConnect != nil { + c.cfg.OnConnect() + } + + defer func() { + c.connMu.Lock() + if c.conn != nil { + c.conn.Close() + c.conn = nil + } + c.connMu.Unlock() + c.connected.Store(false) + }() + + for { + select { + case <-ctx.Done(): + return ctx.Err() + case <-c.stopCh: + return nil + default: + } + + conn.SetReadDeadline(time.Now().Add(60 * time.Second)) + + _, msg, err := conn.ReadMessage() + if err != nil { + return fmt.Errorf("read: %w", err) + } + + c.bytesReceived.Add(int64(len(msg))) + + if err := c.process(ctx, msg); err != nil { + if c.cfg.OnError != nil { + // We don't have the event here since parsing may have failed, + // pass nil to signal a parse/process error + c.cfg.OnError(err, nil) + } + } + } +} + +func (c *Consumer) buildURL(endpoint string) (string, error) { + u, err := url.Parse(endpoint) + if err != nil { + return "", err + } + + q := u.Query() + for _, coll := range c.cfg.WantedCollections { + q.Add("wantedCollections", coll) + } + if c.cfg.Compress { + q.Set("compress", "true") + } + if cursor := c.cursor.Load(); cursor > 0 { + // Rewind 5 seconds to cover any gaps at reconnect + rewind := cursor - (5 * time.Second.Microseconds()) + q.Set("cursor", fmt.Sprintf("%d", rewind)) + } + + u.RawQuery = q.Encode() + return u.String(), nil +} + +func (c *Consumer) process(ctx context.Context, data []byte) error { + // Decompress if enabled and data has zstd magic bytes + if c.cfg.Compress { + if len(data) >= 4 && data[0] == 0x28 && data[1] == 0xB5 && data[2] == 0x2F && data[3] == 0xFD { + decompressed, err := c.zstdDecoder.DecodeAll(data, nil) + if err != nil { + return fmt.Errorf("decompress: %w", err) + } + data = decompressed + } else if len(data) > 0 && data[0] != '{' { + // Try anyway in case magic bytes differ + if decompressed, err := c.zstdDecoder.DecodeAll(data, nil); err == nil { + data = decompressed + } + } + } + + var event Event + if err := json.Unmarshal(data, &event); err != nil { + return fmt.Errorf("unmarshal event: %w", err) + } + + c.eventsReceived.Add(1) + + if event.TimeUS > 0 { + c.cursor.Store(event.TimeUS) + + if c.cfg.CursorStore != nil && c.eventsReceived.Load()%c.cfg.cursorPersistEvery() == 0 { + if err := c.cfg.CursorStore.SetCursor(ctx, event.TimeUS); err != nil { + // Non-fatal: log via OnError if configured + if c.cfg.OnError != nil { + c.cfg.OnError(fmt.Errorf("persist cursor: %w", err), nil) + } + } + } + } + + if err := c.handler(ctx, &event); err != nil { + if c.cfg.OnError != nil { + c.cfg.OnError(err, &event) + } + } + + return nil +} diff --git a/jetstream/jetstream_test.go b/jetstream/jetstream_test.go new file mode 100644 index 0000000..5a974d0 --- /dev/null +++ b/jetstream/jetstream_test.go @@ -0,0 +1,143 @@ +package jetstream + +import ( + "context" + "encoding/json" + "testing" +) + +func TestNew(t *testing.T) { + c := New(&Config{ + WantedCollections: []string{"app.bsky.feed.post"}, + }, func(ctx context.Context, evt *Event) error { + return nil + }) + if c == nil { + t.Fatal("expected non-nil consumer") + } + if c.IsConnected() { + t.Fatal("should not be connected before Start") + } +} + +func TestConfig_Defaults(t *testing.T) { + cfg := &Config{} + if eps := cfg.endpoints(); len(eps) == 0 { + t.Fatal("expected default endpoints") + } + if cfg.cursorPersistEvery() != 1000 { + t.Fatalf("expected 1000, got %d", cfg.cursorPersistEvery()) + } +} + +func TestConfig_CustomEndpoints(t *testing.T) { + cfg := &Config{Endpoints: []string{"wss://custom.example.com/subscribe"}} + eps := cfg.endpoints() + if len(eps) != 1 || eps[0] != "wss://custom.example.com/subscribe" { + t.Fatalf("unexpected endpoints: %v", eps) + } +} + +func TestBuildURL_Collections(t *testing.T) { + c := New(&Config{ + Endpoints: []string{"wss://jetstream1.us-east.bsky.network/subscribe"}, + WantedCollections: []string{"app.bsky.feed.post", "app.bsky.feed.like"}, + }, func(ctx context.Context, evt *Event) error { return nil }) + + u, err := c.buildURL("wss://jetstream1.us-east.bsky.network/subscribe") + if err != nil { + t.Fatal(err) + } + if u == "" { + t.Fatal("expected non-empty URL") + } + // Should contain wantedCollections params + if !contains(u, "wantedCollections=app.bsky.feed.post") { + t.Errorf("URL missing wantedCollections: %s", u) + } +} + +func TestBuildURL_Cursor(t *testing.T) { + c := New(&Config{ + Endpoints: []string{"wss://jetstream1.us-east.bsky.network/subscribe"}, + }, func(ctx context.Context, evt *Event) error { return nil }) + + c.cursor.Store(1000000) + u, err := c.buildURL("wss://jetstream1.us-east.bsky.network/subscribe") + if err != nil { + t.Fatal(err) + } + if !contains(u, "cursor=") { + t.Errorf("URL missing cursor: %s", u) + } +} + +func TestProcess_ValidEvent(t *testing.T) { + var received *Event + c := New(&Config{}, func(ctx context.Context, evt *Event) error { + received = evt + return nil + }) + + evt := Event{ + DID: "did:plc:test", + TimeUS: 1234567890, + Kind: "commit", + Commit: &Commit{ + Operation: "create", + Collection: "app.bsky.feed.post", + RKey: "abc123", + }, + } + data, _ := json.Marshal(evt) + + if err := c.process(context.Background(), data); err != nil { + t.Fatal(err) + } + if received == nil { + t.Fatal("expected handler to be called") + } + if received.DID != "did:plc:test" { + t.Fatalf("got DID %q", received.DID) + } +} + +func TestProcess_UpdatesCursor(t *testing.T) { + c := New(&Config{}, func(ctx context.Context, evt *Event) error { return nil }) + + data, _ := json.Marshal(Event{DID: "did:plc:test", TimeUS: 9999999}) + c.process(context.Background(), data) + + if c.cursor.Load() != 9999999 { + t.Fatalf("cursor not updated: %d", c.cursor.Load()) + } +} + +func TestProcess_InvalidJSON(t *testing.T) { + c := New(&Config{}, func(ctx context.Context, evt *Event) error { return nil }) + err := c.process(context.Background(), []byte("{bad json")) + if err == nil { + t.Fatal("expected error for invalid JSON") + } +} + +func TestStats(t *testing.T) { + c := New(&Config{}, func(ctx context.Context, evt *Event) error { return nil }) + evts, bytes := c.Stats() + if evts != 0 || bytes != 0 { + t.Fatalf("expected zero stats, got events=%d bytes=%d", evts, bytes) + } +} + +func contains(s, substr string) bool { + return len(s) >= len(substr) && (s == substr || len(s) > 0 && containsStr(s, substr)) +} + +func containsStr(s, substr string) bool { + for i := 0; i <= len(s)-len(substr); i++ { + if s[i:i+len(substr)] == substr { + return true + } + } + return false +} diff --git a/oauth.go b/oauth.go index ebc1295..5e36878 100644 --- a/oauth.go +++ b/oauth.go @@ -98,6 +98,7 @@ func (a *OAuthApp) HandleCallback(ctx context.Context, params url.Values) (*Sess // LoginCLI runs a complete loopback OAuth flow for CLI applications. // It opens the user's browser, starts a temporary HTTP server to receive the // callback, and blocks until authentication completes. +// TODO: should this be part of the library? probably not? (removeds `browser` dep) func (a *OAuthApp) LoginCLI(ctx context.Context, handle string) (*SessionInfo, error) { authURL, err := a.app.StartAuthFlow(ctx, handle) if err != nil { @@ -182,8 +183,6 @@ func (a *OAuthApp) ClientMetadata() oauth.ClientMetadata { return meta } -// Store returns the underlying session store, useful for implementing -// features like "list all sessions" or session cleanup. func (a *OAuthApp) Store() oauth.ClientAuthStore { return a.app.Store } diff --git a/oauth_test.go b/oauth_test.go index 0ba1074..d1b280e 100644 --- a/oauth_test.go +++ b/oauth_test.go @@ -6,61 +6,56 @@ import ( "github.com/bluesky-social/indigo/atproto/auth/oauth" ) -func TestNewOAuthApp_Localhost(t *testing.T) { - app, err := NewOAuthApp(OAuthConfig{ - ClientID: "", - RedirectURI: "http://127.0.0.1:12345/callback", - Scopes: []string{"atproto"}, - Store: oauth.NewMemStore(), - }) - if err != nil { - t.Fatal(err) - } - if app == nil { - t.Fatal("expected non-nil app") - } -} - -func TestNewOAuthApp_LocalhostPrefix(t *testing.T) { - app, err := NewOAuthApp(OAuthConfig{ - ClientID: "http://localhost:8080", - RedirectURI: "http://localhost:8080/oauth/callback", - Scopes: []string{"atproto"}, - Store: oauth.NewMemStore(), - }) - if err != nil { - t.Fatal(err) - } - if app == nil { - t.Fatal("expected non-nil app") - } -} - -func TestNewOAuthApp_Public(t *testing.T) { - app, err := NewOAuthApp(OAuthConfig{ - ClientID: "https://example.com/client-metadata.json", - RedirectURI: "https://example.com/oauth/callback", - Scopes: ScopesForCollections("x.y.bean"), - Store: oauth.NewMemStore(), - }) - if err != nil { - t.Fatal(err) - } - if app == nil { - t.Fatal("expected non-nil app") - } -} - -func TestNewOAuthApp_NilStore(t *testing.T) { - app, err := NewOAuthApp(OAuthConfig{ - RedirectURI: "http://127.0.0.1:12345/callback", - Scopes: []string{"atproto"}, - }) - if err != nil { - t.Fatal(err) +func TestNewOAuthApp(t *testing.T) { + tests := []struct { + name string + config OAuthConfig + }{ + { + name: "localhost IP", + config: OAuthConfig{ + ClientID: "", + RedirectURI: "http://127.0.0.1:12345/callback", + Scopes: []string{"atproto"}, + Store: oauth.NewMemStore(), + }, + }, + { + name: "localhost prefix", + config: OAuthConfig{ + ClientID: "http://localhost:8080", + RedirectURI: "http://localhost:8080/oauth/callback", + Scopes: []string{"atproto"}, + Store: oauth.NewMemStore(), + }, + }, + { + name: "public client", + config: OAuthConfig{ + ClientID: "https://example.com/client-metadata.json", + RedirectURI: "https://example.com/oauth/callback", + Scopes: ScopesForCollections("x.y.bean"), + Store: oauth.NewMemStore(), + }, + }, + { + name: "nil store", + config: OAuthConfig{ + RedirectURI: "http://127.0.0.1:12345/callback", + Scopes: []string{"atproto"}, + }, + }, } - if app == nil { - t.Fatal("expected non-nil app") + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + app, err := NewOAuthApp(tc.config) + if err != nil { + t.Fatal(err) + } + if app == nil { + t.Fatal("expected non-nil app") + } + }) } } diff --git a/public.go b/public.go new file mode 100644 index 0000000..760158d --- /dev/null +++ b/public.go @@ -0,0 +1,304 @@ +package atp + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "net" + "net/http" + "net/url" + "slices" + "strings" + "sync" + "time" +) + +const ( + // PublicAPIBase is the Bluesky public API endpoint used for profile and handle lookups. + PublicAPIBase = "https://public.api.bsky.app" + + // PLCDirectory is used to resolve did:plc identifiers to DID documents. + PLCDirectory = "https://plc.directory" +) + +// ErrSSRFBlocked is returned when a request is blocked due to a private/internal destination. +var ErrSSRFBlocked = errors.New("request blocked: potential SSRF detected") + +// 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. +type PublicClient struct { + httpClient *http.Client + pdsCache map[string]string + pdsCacheMu sync.RWMutex +} + +// NewPublicClient creates a PublicClient with a 30-second timeout. +// To add OTel instrumentation, use NewPublicClientWithHTTP and pass an +// otelhttp-wrapped transport. +func NewPublicClient() *PublicClient { + return NewPublicClientWithHTTP(&http.Client{ + Timeout: 30 * time.Second, + }) +} + +// NewPublicClientWithHTTP creates a PublicClient using the provided http.Client. +// 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), + } +} + +// ResolveHandle resolves an AT Protocol handle to a DID string. +func (c *PublicClient) ResolveHandle(ctx context.Context, handle string) (string, error) { + reqURL := fmt.Sprintf("%s/xrpc/com.atproto.identity.resolveHandle?handle=%s", + PublicAPIBase, url.QueryEscape(handle)) + + req, err := http.NewRequestWithContext(ctx, "GET", reqURL, nil) + if err != nil { + return "", fmt.Errorf("build request: %w", err) + } + + resp, err := c.httpClient.Do(req) + if err != nil { + return "", fmt.Errorf("resolve handle: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode == http.StatusNotFound { + return "", fmt.Errorf("handle not found: %s", handle) + } + if resp.StatusCode != http.StatusOK { + return "", fmt.Errorf("resolve handle: HTTP %d", resp.StatusCode) + } + + var result struct { + DID string `json:"did"` + } + if err := json.NewDecoder(resp.Body).Decode(&result); err != nil { + return "", fmt.Errorf("decode response: %w", err) + } + 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. +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.pdsCacheMu.RUnlock() + + var pdsEndpoint string + + switch { + case strings.HasPrefix(did, "did:plc:"): + reqURL := fmt.Sprintf("%s/%s", PLCDirectory, did) + req, err := http.NewRequestWithContext(ctx, "GET", reqURL, nil) + if err != nil { + return "", fmt.Errorf("build request: %w", err) + } + resp, err := c.httpClient.Do(req) + if err != nil { + return "", fmt.Errorf("fetch DID document: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + return "", fmt.Errorf("DID resolution: HTTP %d", resp.StatusCode) + } + + var didDoc struct { + Service []struct { + ID string `json:"id"` + Type string `json:"type"` + ServiceEndpoint string `json:"serviceEndpoint"` + } `json:"service"` + } + if err := json.NewDecoder(resp.Body).Decode(&didDoc); err != nil { + return "", fmt.Errorf("decode DID document: %w", err) + } + for _, svc := range didDoc.Service { + if svc.ID == "#atproto_pds" || svc.Type == "AtprotoPersonalDataServer" { + pdsEndpoint = svc.ServiceEndpoint + break + } + } + + case strings.HasPrefix(did, "did:web:"): + domain := strings.TrimPrefix(did, "did:web:") + domain = strings.ReplaceAll(domain, "%3A", ":") + if idx := strings.Index(domain, "/"); idx != -1 { + domain = domain[:idx] + } + host := domain + if h, _, err := net.SplitHostPort(domain); err == nil { + host = h + } + if err := validateDomain(host); err != nil { + return "", err + } + pdsEndpoint = "https://" + domain + } + + if pdsEndpoint == "" { + return "", fmt.Errorf("could not resolve PDS endpoint for %s", did) + } + + c.pdsCacheMu.Lock() + c.pdsCache[did] = pdsEndpoint + c.pdsCacheMu.Unlock() + + return pdsEndpoint, nil +} + +// PublicProfile is a user's public profile as returned by the Bluesky public API. +type PublicProfile struct { + DID string `json:"did"` + Handle string `json:"handle"` + DisplayName *string `json:"displayName,omitempty"` + Avatar *string `json:"avatar,omitempty"` +} + +// GetProfile fetches a user's public profile by DID or handle. +func (c *PublicClient) GetProfile(ctx context.Context, actor string) (*PublicProfile, error) { + reqURL := fmt.Sprintf("%s/xrpc/app.bsky.actor.getProfile?actor=%s", + PublicAPIBase, url.QueryEscape(actor)) + + req, err := http.NewRequestWithContext(ctx, "GET", reqURL, nil) + if err != nil { + return nil, fmt.Errorf("build request: %w", err) + } + + resp, err := c.httpClient.Do(req) + if err != nil { + return nil, fmt.Errorf("fetch profile: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + return nil, fmt.Errorf("get profile: HTTP %d", resp.StatusCode) + } + + var profile PublicProfile + if err := json.NewDecoder(resp.Body).Decode(&profile); err != nil { + return nil, fmt.Errorf("decode profile: %w", err) + } + return &profile, nil +} + +// ListPublicRecords fetches up to limit records from a public collection. +// Queries the user's PDS directly, so it works with any collection NSID. +func (c *PublicClient) ListPublicRecords(ctx context.Context, did, collection string, limit int) ([]Record, error) { + pdsEndpoint, err := c.GetPDSEndpoint(ctx, did) + if err != nil { + return nil, fmt.Errorf("resolve PDS: %w", err) + } + + reqURL := fmt.Sprintf("%s/xrpc/com.atproto.repo.listRecords?repo=%s&collection=%s&limit=%d", + pdsEndpoint, url.QueryEscape(did), url.QueryEscape(collection), limit) + + req, err := http.NewRequestWithContext(ctx, "GET", reqURL, nil) + if err != nil { + return nil, fmt.Errorf("build request: %w", err) + } + + resp, err := c.httpClient.Do(req) + if err != nil { + return nil, fmt.Errorf("list records: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + return nil, fmt.Errorf("list records: HTTP %d", resp.StatusCode) + } + + var result struct { + Records []struct { + URI string `json:"uri"` + CID string `json:"cid"` + Value map[string]any `json:"value"` + } `json:"records"` + } + if err := json.NewDecoder(resp.Body).Decode(&result); err != nil { + return nil, fmt.Errorf("decode records: %w", err) + } + + records := make([]Record, len(result.Records)) + for i, r := range result.Records { + records[i] = Record{URI: r.URI, CID: r.CID, Value: r.Value} + } + return records, nil +} + +// GetPublicRecord fetches a single public record from a user's PDS. +func (c *PublicClient) GetPublicRecord(ctx context.Context, did, collection, rkey string) (*Record, error) { + pdsEndpoint, err := c.GetPDSEndpoint(ctx, did) + if err != nil { + return nil, fmt.Errorf("resolve PDS: %w", err) + } + + reqURL := fmt.Sprintf("%s/xrpc/com.atproto.repo.getRecord?repo=%s&collection=%s&rkey=%s", + pdsEndpoint, url.QueryEscape(did), url.QueryEscape(collection), url.QueryEscape(rkey)) + + req, err := http.NewRequestWithContext(ctx, "GET", reqURL, nil) + if err != nil { + return nil, fmt.Errorf("build request: %w", err) + } + + resp, err := c.httpClient.Do(req) + if err != nil { + return nil, fmt.Errorf("get record: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + return nil, fmt.Errorf("get record: HTTP %d", resp.StatusCode) + } + + var r struct { + URI string `json:"uri"` + CID string `json:"cid"` + Value map[string]any `json:"value"` + } + if err := json.NewDecoder(resp.Body).Decode(&r); err != nil { + return nil, fmt.Errorf("decode record: %w", err) + } + return &Record{URI: r.URI, CID: r.CID, Value: r.Value}, nil +} + +// isPrivateIP reports whether ip is in a private/reserved range. +func isPrivateIP(ip net.IP) bool { + return ip.IsLoopback() || + ip.IsLinkLocalUnicast() || + ip.IsLinkLocalMulticast() || + ip.IsPrivate() || + ip.IsUnspecified() || + ip.Equal(net.ParseIP("169.254.169.254")) // cloud metadata +} + +// validateDomain blocks requests to private/internal hosts. +func validateDomain(domain string) error { + if domain == "localhost" || strings.HasSuffix(domain, ".local") { + return ErrSSRFBlocked + } + if ip := net.ParseIP(domain); ip != nil { + if isPrivateIP(ip) { + return ErrSSRFBlocked + } + return nil + } + ips, err := net.LookupIP(domain) + if err != nil { + return nil // let the HTTP request fail naturally + } + if slices.ContainsFunc(ips, isPrivateIP) { + return ErrSSRFBlocked + } + return nil +} diff --git a/public_test.go b/public_test.go new file mode 100644 index 0000000..d638126 --- /dev/null +++ b/public_test.go @@ -0,0 +1,62 @@ +package atp + +import ( + "net" + "testing" +) + +func TestIsPrivateIP(t *testing.T) { + cases := []struct { + ip string + private bool + }{ + {"127.0.0.1", true}, + {"::1", true}, + {"10.0.0.1", true}, + {"172.16.0.1", true}, + {"192.168.1.1", true}, + {"169.254.169.254", true}, + {"0.0.0.0", true}, + {"8.8.8.8", false}, + {"1.1.1.1", false}, + } + + for _, tc := range cases { + ip := net.ParseIP(tc.ip) + got := isPrivateIP(ip) + if got != tc.private { + t.Errorf("isPrivateIP(%s) = %v, want %v", tc.ip, got, tc.private) + } + } +} + +func TestValidateDomain_Localhost(t *testing.T) { + if err := validateDomain("localhost"); err != ErrSSRFBlocked { + t.Fatalf("expected ErrSSRFBlocked for localhost, got %v", err) + } +} + +func TestValidateDomain_DotLocal(t *testing.T) { + if err := validateDomain("internal.local"); err != ErrSSRFBlocked { + t.Fatalf("expected ErrSSRFBlocked for .local domain, got %v", err) + } +} + +func TestValidateDomain_PrivateIP(t *testing.T) { + if err := validateDomain("192.168.1.1"); err != ErrSSRFBlocked { + t.Fatalf("expected ErrSSRFBlocked for private IP, got %v", err) + } +} + +func TestValidateDomain_MetadataIP(t *testing.T) { + if err := validateDomain("169.254.169.254"); err != ErrSSRFBlocked { + t.Fatalf("expected ErrSSRFBlocked for metadata IP, got %v", err) + } +} + +func TestNewPublicClient(t *testing.T) { + c := NewPublicClient() + if c == nil { + t.Fatal("expected non-nil client") + } +} diff --git a/record.go b/record.go index de1086b..273879e 100644 --- a/record.go +++ b/record.go @@ -1,6 +1,5 @@ package atp -// Record represents a single record returned from a PDS. type Record struct { URI string CID string diff --git a/scopes.go b/scopes.go index d8a07ac..18013a9 100644 --- a/scopes.go +++ b/scopes.go @@ -6,6 +6,8 @@ package atp // // ScopesForCollections("x.y.bean", "x.y.brew") // // => ["atproto", "repo:x.y.bean", "repo:x.y.brew"] +// +// TODO: add support for collections and more granular control than just full rw func ScopesForCollections(collections ...string) []string { scopes := make([]string, 0, 1+len(collections)) scopes = append(scopes, "atproto") diff --git a/scopes_test.go b/scopes_test.go index a411308..3fc9579 100644 --- a/scopes_test.go +++ b/scopes_test.go @@ -6,18 +6,29 @@ import ( ) func TestScopesForCollections(t *testing.T) { - got := ScopesForCollections("social.arabica.alpha.bean", "social.arabica.alpha.brew") - want := []string{"atproto", "repo:social.arabica.alpha.bean", "repo:social.arabica.alpha.brew"} - if !slices.Equal(got, want) { - t.Fatalf("got %v, want %v", got, want) + tests := []struct { + name string + collections []string + want []string + }{ + { + name: "no collections", + collections: nil, + want: []string{"atproto"}, + }, + { + name: "multiple collections", + collections: []string{"social.arabica.alpha.bean", "social.arabica.alpha.brew"}, + want: []string{"atproto", "repo:social.arabica.alpha.bean", "repo:social.arabica.alpha.brew"}, + }, } -} - -func TestScopesForCollections_Empty(t *testing.T) { - got := ScopesForCollections() - want := []string{"atproto"} - if !slices.Equal(got, want) { - t.Fatalf("got %v, want %v", got, want) + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + got := ScopesForCollections(tc.collections...) + if !slices.Equal(got, tc.want) { + t.Fatalf("got %v, want %v", got, tc.want) + } + }) } } diff --git a/tracing/tracing.go b/tracing/tracing.go index dd5f75b..0f39226 100644 --- a/tracing/tracing.go +++ b/tracing/tracing.go @@ -24,8 +24,9 @@ func tracer() trace.Tracer { // Init creates and registers a tracer provider with an OTLP HTTP exporter. // It reads OTEL_EXPORTER_OTLP_ENDPOINT (default: localhost:4318). -// The serviceName appears in your tracing backend (e.g. "arabica", "solanum"). +// The serviceName appears in your tracing backend (e.g. "arabica"). // Returns the provider so the caller can defer provider.Shutdown(ctx). +// TODO: allow grpc exporting to port 4317 func Init(ctx context.Context, serviceName string) (*sdktrace.TracerProvider, error) { endpoint := os.Getenv("OTEL_EXPORTER_OTLP_ENDPOINT") if endpoint == "" { diff --git a/tracing/tracing_test.go b/tracing/tracing_test.go index c915fab..2652710 100644 --- a/tracing/tracing_test.go +++ b/tracing/tracing_test.go @@ -7,28 +7,29 @@ import ( "go.opentelemetry.io/otel/trace" ) -func TestBoltSpan_NoParent(t *testing.T) { +func TestSpan_NoParent(t *testing.T) { ctx := context.Background() - // Without a parent span, should return a no-op span - _, span := BoltSpan(ctx, "GetSession", "oauth_sessions") - if span.SpanContext().IsValid() { - t.Fatal("expected no-op span without parent") + tests := []struct { + name string + fn func(context.Context) (context.Context, trace.Span) + }{ + {"bolt", func(ctx context.Context) (context.Context, trace.Span) { + return BoltSpan(ctx, "GetSession", "oauth_sessions") + }}, + {"sqlite", func(ctx context.Context) (context.Context, trace.Span) { + return SqliteSpan(ctx, "query", "records") + }}, + {"pds", func(ctx context.Context) (context.Context, trace.Span) { + return PdsSpan(ctx, "createRecord", "x.y.z", "did:plc:test") + }}, } -} - -func TestSqliteSpan_NoParent(t *testing.T) { - ctx := context.Background() - _, span := SqliteSpan(ctx, "query", "records") - if span.SpanContext().IsValid() { - t.Fatal("expected no-op span without parent") - } -} - -func TestPdsSpan_NoParent(t *testing.T) { - ctx := context.Background() - _, span := PdsSpan(ctx, "createRecord", "x.y.z", "did:plc:test") - if span.SpanContext().IsValid() { - t.Fatal("expected no-op span without parent") + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + _, span := tc.fn(ctx) + if span.SpanContext().IsValid() { + t.Fatal("expected no-op span without parent") + } + }) } } diff --git a/uri.go b/uri.go index a4e6623..23caf4c 100644 --- a/uri.go +++ b/uri.go @@ -7,6 +7,7 @@ import ( ) 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) } diff --git a/uri_test.go b/uri_test.go index 63f79df..2495d72 100644 --- a/uri_test.go +++ b/uri_test.go @@ -11,32 +11,61 @@ func TestBuildATURI(t *testing.T) { } func TestParseATURI(t *testing.T) { - did, collection, rkey, err := ParseATURI("at://did:plc:abc/app.bsky.feed.post/3jxy") - if err != nil { - t.Fatal(err) + tests := []struct { + name string + input string + wantDID string + wantColl string + wantRKey string + wantErr bool + }{ + { + name: "valid URI", + input: "at://did:plc:abc/app.bsky.feed.post/3jxy", + wantDID: "did:plc:abc", + wantColl: "app.bsky.feed.post", + wantRKey: "3jxy", + }, + { + name: "invalid URI", + input: "not-a-uri", + wantErr: true, + }, } - if did != "did:plc:abc" || collection != "app.bsky.feed.post" || rkey != "3jxy" { - t.Fatalf("got did=%q collection=%q rkey=%q", did, collection, rkey) - } -} - -func TestParseATURI_Invalid(t *testing.T) { - _, _, _, err := ParseATURI("not-a-uri") - if err == nil { - t.Fatal("expected error for invalid URI") + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + did, collection, rkey, err := ParseATURI(tc.input) + if tc.wantErr { + if err == nil { + t.Fatal("expected error for invalid URI") + } + return + } + 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) + } + }) } } func TestRKeyFromURI(t *testing.T) { - got := RKeyFromURI("at://did:plc:abc/app.bsky.feed.post/3jxy") - if got != "3jxy" { - t.Fatalf("got %q, want %q", got, "3jxy") + tests := []struct { + name string + input string + want string + }{ + {"valid URI", "at://did:plc:abc/app.bsky.feed.post/3jxy", "3jxy"}, + {"invalid URI", "garbage", ""}, } -} - -func TestRKeyFromURI_Invalid(t *testing.T) { - got := RKeyFromURI("garbage") - if got != "" { - t.Fatalf("expected empty string for invalid URI, got %q", got) + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + got := RKeyFromURI(tc.input) + if got != tc.want { + t.Fatalf("got %q, want %q", got, tc.want) + } + }) } }