From 48414be75d53ef54e83728cb8b1283fe37bd938d Mon Sep 17 00:00:00 2001 From: Evan Jarrett Date: Fri, 17 Oct 2025 15:41:20 -0500 Subject: [PATCH] lots of unit testing for xrpc endpoints. start pointing appview to the new endpoints. remove legacy api endpoints --- CLAUDE.md | 8 - Dockerfile.appview | 2 +- Dockerfile.hold | 2 +- cmd/hold/main.go | 52 +-- pkg/appview/middleware/registry.go | 2 +- pkg/appview/storage/proxy_blob_store.go | 257 +++++------ pkg/atproto/client.go | 2 +- pkg/atproto/lexicon.go | 12 +- pkg/auth/hold_remote.go | 8 +- pkg/auth/oauth/client.go | 2 +- pkg/auth/token/handler.go | 8 +- pkg/hold/blobstore_adapter.go | 28 +- pkg/hold/handlers.go | 498 +------------------- pkg/hold/multipart.go | 8 + pkg/hold/pds/auth.go | 127 ++++- pkg/hold/pds/auth_test.go | 587 ++++++++++++++++++++++++ pkg/hold/pds/did_test.go | 274 +++++++++++ pkg/hold/pds/events.go | 20 +- pkg/hold/pds/events_test.go | 384 ++++++++++++++++ pkg/hold/pds/repomgr.go | 16 +- pkg/hold/pds/xrpc.go | 78 +++- pkg/hold/pds/xrpc_multipart_test.go | 65 ++- pkg/hold/pds/xrpc_test.go | 179 ++++++-- pkg/hold/service.go | 21 +- pkg/hold/storage.go | 55 ++- 25 files changed, 1842 insertions(+), 853 deletions(-) create mode 100644 pkg/hold/pds/auth_test.go create mode 100644 pkg/hold/pds/did_test.go create mode 100644 pkg/hold/pds/events_test.go diff --git a/CLAUDE.md b/CLAUDE.md index 815df78..fe20a9f 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -368,14 +368,6 @@ Write access: Key insight: "Private" gates anonymous access, not authenticated access. This reflects ATProto's current limitation (no private PDS records yet). -**Endpoints:** -- `POST /get-presigned-url` - Get download URL for blob -- `POST /put-presigned-url` - Get upload URL for blob -- `GET /blobs/{digest}` - Proxy download (fallback if no presigned URL support) -- `PUT /blobs/{digest}` - Proxy upload (fallback) -- `POST /register` - Manual registration endpoint -- `GET /health` - Health check - **Embedded PDS Endpoints:** Each hold service includes an embedded PDS (Personal Data Server) that stores captain + crew records: diff --git a/Dockerfile.appview b/Dockerfile.appview index 88946dd..26dd855 100644 --- a/Dockerfile.appview +++ b/Dockerfile.appview @@ -1,4 +1,4 @@ -FROM golang:1.25.2-trixie AS builder +FROM docker.io/golang:1.25.2-trixie AS builder RUN apt-get update && \ apt-get install -y --no-install-recommends sqlite3 libsqlite3-dev && \ diff --git a/Dockerfile.hold b/Dockerfile.hold index 26b1969..4b0fcf1 100644 --- a/Dockerfile.hold +++ b/Dockerfile.hold @@ -1,4 +1,4 @@ -FROM golang:1.25.2-trixie AS builder +FROM docker.io/golang:1.25.2-trixie AS builder RUN apt-get update && \ apt-get install -y --no-install-recommends sqlite3 libsqlite3-dev && \ diff --git a/cmd/hold/main.go b/cmd/hold/main.go index 0c4458a..b96f2ed 100644 --- a/cmd/hold/main.go +++ b/cmd/hold/main.go @@ -5,8 +5,6 @@ import ( "fmt" "log" "net/http" - "strconv" - "strings" "atcr.io/pkg/hold" "atcr.io/pkg/hold/pds" @@ -66,7 +64,7 @@ func main() { if holdPDS != nil { holdDID := holdPDS.DID() blobStore := hold.NewHoldServiceBlobStore(service, holdDID) - xrpcHandler = pds.NewXRPCHandler(holdPDS, cfg.Server.PublicURL, blobStore, broadcaster) + xrpcHandler = pds.NewXRPCHandler(holdPDS, cfg.Server.PublicURL, blobStore, broadcaster, nil) } // Setup HTTP routes @@ -82,54 +80,6 @@ func main() { http.NotFound(w, r) }) - mux.HandleFunc("/presigned-url", service.HandlePresignedURL) - mux.HandleFunc("/move", service.HandleMove) - - // Multipart upload endpoints - mux.HandleFunc("/start-multipart", service.HandleStartMultipart) - mux.HandleFunc("/part-presigned-url", service.HandleGetPartURL) - mux.HandleFunc("/complete-multipart", service.HandleCompleteMultipart) - mux.HandleFunc("/abort-multipart", service.HandleAbortMultipart) - - // Buffered multipart part upload endpoint (for when presigned URLs are disabled/unavailable) - mux.HandleFunc("/multipart-parts/", func(w http.ResponseWriter, r *http.Request) { - if r.Method != http.MethodPut { - http.Error(w, "method not allowed", http.StatusMethodNotAllowed) - return - } - - // Parse URL: /multipart-parts/{uploadID}/{partNumber} - path := r.URL.Path[len("/multipart-parts/"):] - parts := strings.Split(path, "/") - if len(parts) != 2 { - http.Error(w, "invalid path format, expected /multipart-parts/{uploadID}/{partNumber}", http.StatusBadRequest) - return - } - - uploadID := parts[0] - partNumber, err := strconv.Atoi(parts[1]) - if err != nil { - http.Error(w, fmt.Sprintf("invalid part number: %v", err), http.StatusBadRequest) - return - } - - // Get DID from query param - did := r.URL.Query().Get("did") - - service.HandleMultipartPartUpload(w, r, uploadID, partNumber, did, service.MultipartMgr) - }) - - mux.HandleFunc("/blobs/", func(w http.ResponseWriter, r *http.Request) { - switch r.Method { - case http.MethodGet, http.MethodHead: - service.HandleProxyGet(w, r) - case http.MethodPut: - service.HandleProxyPut(w, r) - default: - http.Error(w, "method not allowed", http.StatusMethodNotAllowed) - } - }) - // Register XRPC/ATProto PDS endpoints if PDS is initialized if xrpcHandler != nil { log.Printf("Registering ATProto PDS endpoints") diff --git a/pkg/appview/middleware/registry.go b/pkg/appview/middleware/registry.go index 11e69d1..630f385 100644 --- a/pkg/appview/middleware/registry.go +++ b/pkg/appview/middleware/registry.go @@ -60,7 +60,7 @@ type NamespaceResolver struct { distribution.Namespace directory identity.Directory defaultStorageEndpoint string - testMode bool // If true, fallback to default hold when user's hold is unreachable + testMode bool // If true, fallback to default hold when user's hold is unreachable repositories sync.Map // Cache of RoutingRepository instances by key (did:reponame) } diff --git a/pkg/appview/storage/proxy_blob_store.go b/pkg/appview/storage/proxy_blob_store.go index c818adb..1b3a2a1 100644 --- a/pkg/appview/storage/proxy_blob_store.go +++ b/pkg/appview/storage/proxy_blob_store.go @@ -202,9 +202,10 @@ func (p *ProxyBlobStore) Open(ctx context.Context, dgst digest.Digest) (io.ReadS }, nil } -// Put stores a blob +// Put stores a blob using the multipart upload flow +// This ensures all uploads go through the same XRPC path func (p *ProxyBlobStore) Put(ctx context.Context, mediaType string, content []byte) (distribution.Descriptor, error) { - // Check write access + // Check write access (fast-fail before starting multipart upload) if err := p.checkWriteAccess(ctx); err != nil { return distribution.Descriptor{}, err } @@ -212,41 +213,33 @@ func (p *ProxyBlobStore) Put(ctx context.Context, mediaType string, content []by // Calculate digest dgst := digest.FromBytes(content) - // Get upload URL - url, err := p.getUploadURL(ctx, dgst, int64(len(content))) + // Use Create() flow for all uploads (goes through multipart XRPC endpoints) + writer, err := p.Create(ctx) if err != nil { - fmt.Printf("[proxy_blob_store/Put] Failed to get upload URL: digest=%s, error=%v\n", dgst, err) + fmt.Printf("[proxy_blob_store/Put] Failed to create writer: %v\n", err) return distribution.Descriptor{}, err } - // Upload the blob - req, err := http.NewRequestWithContext(ctx, "PUT", url, bytes.NewReader(content)) - if err != nil { - fmt.Printf("[proxy_blob_store/Put] Failed to create request: %v\n", err) + // Write the content + if _, err := writer.Write(content); err != nil { + writer.Cancel(ctx) + fmt.Printf("[proxy_blob_store/Put] Failed to write content: %v\n", err) return distribution.Descriptor{}, err } - req.Header.Set("Content-Type", "application/octet-stream") - resp, err := p.httpClient.Do(req) - if err != nil { - fmt.Printf("[proxy_blob_store/Put] HTTP request failed: %v\n", err) - return distribution.Descriptor{}, err - } - defer resp.Body.Close() - - if resp.StatusCode != http.StatusOK && resp.StatusCode != http.StatusCreated { - bodyBytes, _ := io.ReadAll(resp.Body) - fmt.Printf(" Error Body: %s\n", string(bodyBytes)) - return distribution.Descriptor{}, fmt.Errorf("upload failed with status %d: %s", resp.StatusCode, string(bodyBytes)) - } - - fmt.Printf("[proxy_blob_store/Put] Upload successful: digest=%s, size=%d\n", dgst, len(content)) - - return distribution.Descriptor{ + // Commit with the calculated digest + desc, err := writer.Commit(ctx, distribution.Descriptor{ Digest: dgst, Size: int64(len(content)), MediaType: mediaType, - }, nil + }) + if err != nil { + fmt.Printf("[proxy_blob_store/Put] Failed to commit: %v\n", err) + return distribution.Descriptor{}, err + } + + fmt.Printf("[proxy_blob_store/Put] Upload successful: digest=%s, size=%d\n", dgst, len(content)) + return desc, nil } // Delete removes a blob @@ -348,75 +341,35 @@ func (p *ProxyBlobStore) Resume(ctx context.Context, id string) (distribution.Bl return writer, nil } -// getPresignedURL requests a presigned URL from the storage service for any operation -func (p *ProxyBlobStore) getPresignedURL(ctx context.Context, operation, dgst string, size int64) (string, error) { - reqBody := map[string]any{ - "operation": operation, - "did": p.did, - "digest": dgst, - } - - // Only include size for PUT operations - if size > 0 { - reqBody["size"] = size - } - - body, err := json.Marshal(reqBody) - if err != nil { - return "", err - } - - url := fmt.Sprintf("%s/presigned-url", p.storageEndpoint) - req, err := http.NewRequestWithContext(ctx, "POST", url, bytes.NewReader(body)) - if err != nil { - return "", err - } - req.Header.Set("Content-Type", "application/json") - - resp, err := p.httpClient.Do(req) - if err != nil { - return "", err - } - defer resp.Body.Close() - - if resp.StatusCode != http.StatusOK { - return "", fmt.Errorf("failed to get presigned URL: status %d", resp.StatusCode) - } - - var result struct { - URL string `json:"url"` - } - if err := json.NewDecoder(resp.Body).Decode(&result); err != nil { - return "", err - } - - return result.URL, nil -} - -// getDownloadURL requests a presigned download URL from the storage service +// getDownloadURL returns the XRPC getBlob URL for downloading a blob +// The hold service will redirect to a presigned S3 URL func (p *ProxyBlobStore) getDownloadURL(ctx context.Context, dgst digest.Digest) (string, error) { - return p.getPresignedURL(ctx, "GET", dgst.String(), 0) + // Use XRPC endpoint: GET /xrpc/com.atproto.sync.getBlob?did={holdDID}&cid={digest} + // Per migration doc: hold accepts OCI digest directly as cid parameter (checks for sha256: prefix) + url := fmt.Sprintf("%s/xrpc/com.atproto.sync.getBlob?did=%s&cid=%s", + p.storageEndpoint, p.holdDID, dgst.String()) + return url, nil } -// getHeadURL requests a presigned HEAD URL from the storage service +// getHeadURL returns the XRPC getBlob URL for HEAD requests +// The hold service will redirect to a presigned S3 URL func (p *ProxyBlobStore) getHeadURL(ctx context.Context, dgst digest.Digest) (string, error) { - return p.getPresignedURL(ctx, "HEAD", dgst.String(), 0) + // Same as GET - hold service handles HEAD method on getBlob endpoint + url := fmt.Sprintf("%s/xrpc/com.atproto.sync.getBlob?did=%s&cid=%s", + p.storageEndpoint, p.holdDID, dgst.String()) + return url, nil } -// getUploadURL requests a presigned upload URL from the storage service +// getUploadURL is deprecated - single blob uploads should use Create() instead +// XRPC migration: No direct presigned upload URL endpoint, use multipart flow for all uploads func (p *ProxyBlobStore) getUploadURL(ctx context.Context, dgst digest.Digest, size int64) (string, error) { - fmt.Printf("DEBUG [proxy_blob_store/getUploadURL]: storageEndpoint=%s, digest=%s\n", p.storageEndpoint, dgst) - url, err := p.getPresignedURL(ctx, "PUT", dgst.String(), size) - if err == nil { - fmt.Printf("DEBUG [proxy_blob_store/getUploadURL]: Got presigned URL=%s\n", url) - } - return url, err + return "", fmt.Errorf("single blob upload via Put() not supported with XRPC endpoints - use Create() instead") } -// startMultipartUpload initiates a multipart upload via hold service +// startMultipartUpload initiates a multipart upload via XRPC uploadBlob endpoint func (p *ProxyBlobStore) startMultipartUpload(ctx context.Context, digest string) (string, error) { reqBody := map[string]any{ - "did": p.did, + "action": "start", "digest": digest, } @@ -425,7 +378,7 @@ func (p *ProxyBlobStore) startMultipartUpload(ctx context.Context, digest string return "", err } - url := fmt.Sprintf("%s/start-multipart", p.storageEndpoint) + url := fmt.Sprintf("%s/xrpc/com.atproto.repo.uploadBlob", p.storageEndpoint) req, err := http.NewRequestWithContext(ctx, "POST", url, bytes.NewReader(body)) if err != nil { return "", err @@ -444,7 +397,8 @@ func (p *ProxyBlobStore) startMultipartUpload(ctx context.Context, digest string } var result struct { - UploadID string `json:"upload_id"` + UploadID string `json:"uploadId"` + Mode string `json:"mode"` } if err := json.NewDecoder(resp.Body).Decode(&result); err != nil { return "", err @@ -453,55 +407,70 @@ func (p *ProxyBlobStore) startMultipartUpload(ctx context.Context, digest string return result.UploadID, nil } -// getPartPresignedURL gets a presigned URL for uploading a specific part -func (p *ProxyBlobStore) getPartPresignedURL(ctx context.Context, digest, uploadID string, partNumber int) (string, error) { +// PartUploadInfo contains structured information for uploading a part +type PartUploadInfo struct { + URL string `json:"url"` + Method string `json:"method,omitempty"` + Headers map[string]string `json:"headers,omitempty"` +} + +// getPartUploadInfo gets structured upload info for uploading a specific part via XRPC +func (p *ProxyBlobStore) getPartUploadInfo(ctx context.Context, digest, uploadID string, partNumber int) (*PartUploadInfo, error) { reqBody := map[string]any{ - "did": p.did, - "digest": digest, - "upload_id": uploadID, - "part_number": partNumber, + "action": "part", + "uploadId": uploadID, + "partNumber": partNumber, + "digest": digest, } body, err := json.Marshal(reqBody) if err != nil { - return "", err + return nil, err } - url := fmt.Sprintf("%s/part-presigned-url", p.storageEndpoint) + url := fmt.Sprintf("%s/xrpc/com.atproto.repo.uploadBlob", p.storageEndpoint) req, err := http.NewRequestWithContext(ctx, "POST", url, bytes.NewReader(body)) if err != nil { - return "", err + return nil, err } req.Header.Set("Content-Type", "application/json") resp, err := p.httpClient.Do(req) if err != nil { - return "", err + return nil, err } defer resp.Body.Close() if resp.StatusCode != http.StatusOK { bodyBytes, _ := io.ReadAll(resp.Body) - return "", fmt.Errorf("get part URL failed: status %d, body: %s", resp.StatusCode, string(bodyBytes)) + return nil, fmt.Errorf("get part URL failed: status %d, body: %s", resp.StatusCode, string(bodyBytes)) } - var result struct { - URL string `json:"url"` - } - if err := json.NewDecoder(resp.Body).Decode(&result); err != nil { - return "", err + var uploadInfo PartUploadInfo + if err := json.NewDecoder(resp.Body).Decode(&uploadInfo); err != nil { + return nil, err } - return result.URL, nil + return &uploadInfo, nil } -// completeMultipartUpload completes a multipart upload via hold service +// completeMultipartUpload completes a multipart upload via XRPC uploadBlob endpoint +// The XRPC complete action handles the move from temp to final location internally func (p *ProxyBlobStore) completeMultipartUpload(ctx context.Context, digest, uploadID string, parts []CompletedPart) error { + // Convert parts to XRPC format (partNumber instead of part_number) + xrpcParts := make([]map[string]any, len(parts)) + for i, part := range parts { + xrpcParts[i] = map[string]any{ + "partNumber": part.PartNumber, + "etag": part.ETag, + } + } + reqBody := map[string]any{ - "did": p.did, - "digest": digest, - "upload_id": uploadID, - "parts": parts, + "action": "complete", + "uploadId": uploadID, + "digest": digest, + "parts": xrpcParts, } body, err := json.Marshal(reqBody) @@ -509,7 +478,7 @@ func (p *ProxyBlobStore) completeMultipartUpload(ctx context.Context, digest, up return err } - url := fmt.Sprintf("%s/complete-multipart", p.storageEndpoint) + url := fmt.Sprintf("%s/xrpc/com.atproto.repo.uploadBlob", p.storageEndpoint) req, err := http.NewRequestWithContext(ctx, "POST", url, bytes.NewReader(body)) if err != nil { return err @@ -530,12 +499,12 @@ func (p *ProxyBlobStore) completeMultipartUpload(ctx context.Context, digest, up return nil } -// abortMultipartUpload aborts a multipart upload via hold service +// abortMultipartUpload aborts a multipart upload via XRPC uploadBlob endpoint func (p *ProxyBlobStore) abortMultipartUpload(ctx context.Context, digest, uploadID string) error { reqBody := map[string]any{ - "did": p.did, - "digest": digest, - "upload_id": uploadID, + "action": "abort", + "uploadId": uploadID, + "digest": digest, } body, err := json.Marshal(reqBody) @@ -543,7 +512,7 @@ func (p *ProxyBlobStore) abortMultipartUpload(ctx context.Context, digest, uploa return err } - url := fmt.Sprintf("%s/abort-multipart", p.storageEndpoint) + url := fmt.Sprintf("%s/xrpc/com.atproto.repo.uploadBlob", p.storageEndpoint) req, err := http.NewRequestWithContext(ctx, "POST", url, bytes.NewReader(body)) if err != nil { return err @@ -624,20 +593,31 @@ func (w *ProxyBlobWriter) flushPart() error { ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute) defer cancel() - // Get presigned URL for this part + // Get structured upload info for this part tempDigest := fmt.Sprintf("uploads/temp-%s", w.id) - url, err := w.store.getPartPresignedURL(ctx, tempDigest, w.uploadID, w.partNumber) + uploadInfo, err := w.store.getPartUploadInfo(ctx, tempDigest, w.uploadID, w.partNumber) if err != nil { - return fmt.Errorf("failed to get part presigned URL: %w", err) + return fmt.Errorf("failed to get part upload info: %w", err) } - // Upload part to S3 - req, err := http.NewRequestWithContext(ctx, "PUT", url, bytes.NewReader(w.buffer.Bytes())) + // Determine HTTP method (default to PUT) + method := uploadInfo.Method + if method == "" { + method = "PUT" + } + + // Upload part (either to S3 presigned URL or back to XRPC with headers) + req, err := http.NewRequestWithContext(ctx, method, uploadInfo.URL, bytes.NewReader(w.buffer.Bytes())) if err != nil { return err } req.Header.Set("Content-Type", "application/octet-stream") + // Apply any additional headers from the response (for buffered mode) + for key, value := range uploadInfo.Headers { + req.Header.Set(key, value) + } + resp, err := w.store.httpClient.Do(req) if err != nil { return err @@ -650,9 +630,18 @@ func (w *ProxyBlobWriter) flushPart() error { } // Store ETag for completion + // For buffered mode, ETag might be in JSON response body etag := resp.Header.Get("ETag") if etag == "" { - return fmt.Errorf("no ETag in response") + // Try to parse JSON response for buffered mode + var result struct { + ETag string `json:"etag"` + } + if err := json.NewDecoder(resp.Body).Decode(&result); err == nil && result.ETag != "" { + etag = result.ETag + } else { + return fmt.Errorf("no ETag in response") + } } w.parts = append(w.parts, CompletedPart{ @@ -727,37 +716,13 @@ func (w *ProxyBlobWriter) Commit(ctx context.Context, desc distribution.Descript } } - // Complete multipart upload at temp location + // Complete multipart upload - XRPC complete action handles move internally tempDigest := fmt.Sprintf("uploads/temp-%s", w.id) fmt.Printf("🔒 [Commit] Completing multipart upload: uploadID=%s, parts=%d\n", w.uploadID, len(w.parts)) if err := w.store.completeMultipartUpload(ctx, tempDigest, w.uploadID, w.parts); err != nil { return distribution.Descriptor{}, fmt.Errorf("failed to complete multipart upload: %w", err) } - // Move from temp → final location (server-side S3 copy) - tempPath := fmt.Sprintf("uploads/temp-%s", w.id) - finalPath := desc.Digest.String() - - fmt.Printf("[Commit] Moving blob: %s → %s\n", tempPath, finalPath) - moveURL := fmt.Sprintf("%s/move?from=%s&to=%s&did=%s", - w.store.storageEndpoint, tempPath, finalPath, w.store.did) - - req, err := http.NewRequestWithContext(ctx, "POST", moveURL, nil) - if err != nil { - return distribution.Descriptor{}, fmt.Errorf("failed to create move request: %w", err) - } - - resp, err := w.store.httpClient.Do(req) - if err != nil { - return distribution.Descriptor{}, fmt.Errorf("failed to move blob: %w", err) - } - defer resp.Body.Close() - - if resp.StatusCode != http.StatusOK && resp.StatusCode != http.StatusCreated { - bodyBytes, _ := io.ReadAll(resp.Body) - return distribution.Descriptor{}, fmt.Errorf("move blob failed: status %d, body: %s", resp.StatusCode, string(bodyBytes)) - } - fmt.Printf("[Commit] Upload completed successfully: digest=%s, size=%d, parts=%d\n", desc.Digest, w.size, len(w.parts)) return distribution.Descriptor{ diff --git a/pkg/atproto/client.go b/pkg/atproto/client.go index 4047d20..0d461c7 100644 --- a/pkg/atproto/client.go +++ b/pkg/atproto/client.go @@ -25,7 +25,7 @@ type Client struct { did string accessToken string // For Basic Auth only httpClient *http.Client - useIndigoClient bool // true if using indigo's OAuth client (handles auth automatically) + useIndigoClient bool // true if using indigo's OAuth client (handles auth automatically) indigoClient *atclient.APIClient // indigo's API client for OAuth requests } diff --git a/pkg/atproto/lexicon.go b/pkg/atproto/lexicon.go index 181f667..d2510ec 100644 --- a/pkg/atproto/lexicon.go +++ b/pkg/atproto/lexicon.go @@ -394,12 +394,12 @@ func ResolveHoldDIDFromURL(holdURL string) string { // Uses CBOR encoding for efficient storage in hold's carstore type CaptainRecord struct { Type string `json:"$type" cborgen:"$type"` - Owner string `json:"owner" cborgen:"owner"` // DID of hold owner - Public bool `json:"public" cborgen:"public"` // Public read access - AllowAllCrew bool `json:"allowAllCrew" cborgen:"allowAllCrew"` // Allow any authenticated user to register as crew - DeployedAt string `json:"deployedAt" cborgen:"deployedAt"` // RFC3339 timestamp - Region string `json:"region,omitempty" cborgen:"region,omitempty"` // S3 region (optional) - Provider string `json:"provider,omitempty" cborgen:"provider,omitempty"` // Deployment provider (optional) + Owner string `json:"owner" cborgen:"owner"` // DID of hold owner + Public bool `json:"public" cborgen:"public"` // Public read access + AllowAllCrew bool `json:"allowAllCrew" cborgen:"allowAllCrew"` // Allow any authenticated user to register as crew + DeployedAt string `json:"deployedAt" cborgen:"deployedAt"` // RFC3339 timestamp + Region string `json:"region,omitempty" cborgen:"region,omitempty"` // S3 region (optional) + Provider string `json:"provider,omitempty" cborgen:"provider,omitempty"` // Deployment provider (optional) } // CrewRecord represents a crew member in the hold diff --git a/pkg/auth/hold_remote.go b/pkg/auth/hold_remote.go index a0fc114..bdd2915 100644 --- a/pkg/auth/hold_remote.go +++ b/pkg/auth/hold_remote.go @@ -544,10 +544,10 @@ func (a *RemoteHoldAuthorizer) cacheDenial(holdDID, userDID string) error { // This function handles second+ denials: 1m, 5m, 15m, 1h func getBackoffDuration(denialCount int) time.Duration { backoffs := []time.Duration{ - 1 * time.Minute, // 1st DB denial (2nd overall) - being added soon - 5 * time.Minute, // 2nd DB denial (3rd overall) - probably not happening - 15 * time.Minute, // 3rd DB denial (4th overall) - definitely not soon - 60 * time.Minute, // 4th+ DB denial (5th+ overall) - stop hammering + 1 * time.Minute, // 1st DB denial (2nd overall) - being added soon + 5 * time.Minute, // 2nd DB denial (3rd overall) - probably not happening + 15 * time.Minute, // 3rd DB denial (4th overall) - definitely not soon + 60 * time.Minute, // 4th+ DB denial (5th+ overall) - stop hammering } idx := denialCount - 1 diff --git a/pkg/auth/oauth/client.go b/pkg/auth/oauth/client.go index 4cccb90..2909b69 100644 --- a/pkg/auth/oauth/client.go +++ b/pkg/auth/oauth/client.go @@ -129,7 +129,7 @@ func RedirectURI(baseURL string) string { func GetDefaultScopes() []string { return []string{ "atproto", - "blob:application/vnd.oci.image.manifest.v1+json", + "blob:application/vnd.oci.image.manifest.v1+json", "blob:application/vnd.docker.distribution.manifest.v2+json", fmt.Sprintf("repo:%s", atproto.ManifestCollection), fmt.Sprintf("repo:%s", atproto.TagCollection), diff --git a/pkg/auth/token/handler.go b/pkg/auth/token/handler.go index 4b88cc3..736b63b 100644 --- a/pkg/auth/token/handler.go +++ b/pkg/auth/token/handler.go @@ -18,10 +18,10 @@ import ( // Handler handles /auth/token requests type Handler struct { - issuer *Issuer - validator *atproto.SessionValidator - deviceStore *db.DeviceStore // For validating device secrets - defaultHoldDID string + issuer *Issuer + validator *atproto.SessionValidator + deviceStore *db.DeviceStore // For validating device secrets + defaultHoldDID string } // NewHandler creates a new token handler diff --git a/pkg/hold/blobstore_adapter.go b/pkg/hold/blobstore_adapter.go index b4cc36f..09a894f 100644 --- a/pkg/hold/blobstore_adapter.go +++ b/pkg/hold/blobstore_adapter.go @@ -81,14 +81,34 @@ func (b *HoldServiceBlobStore) StartMultipartUpload(ctx context.Context, digest return uploadID, modeStr, nil } -// GetPartUploadURL returns a presigned URL for uploading a specific part -func (b *HoldServiceBlobStore) GetPartUploadURL(ctx context.Context, uploadID string, partNumber int, did string) (string, error) { +// GetPartUploadURL returns structured upload info for uploading a specific part +func (b *HoldServiceBlobStore) GetPartUploadURL(ctx context.Context, uploadID string, partNumber int, did string) (*pds.PartUploadInfo, error) { session, err := b.service.MultipartMgr.GetSession(uploadID) if err != nil { - return "", err + return nil, err } - return b.service.GetPartUploadURL(ctx, session, partNumber, did) + // For S3Native mode: return presigned URL + if session.Mode == S3Native { + url, err := b.service.GetPartUploadURL(ctx, session, partNumber, did) + if err != nil { + return nil, err + } + return &pds.PartUploadInfo{ + URL: url, + Method: "PUT", + }, nil + } + + // Buffered mode: return XRPC endpoint with headers + return &pds.PartUploadInfo{ + URL: fmt.Sprintf("%s/xrpc/com.atproto.repo.uploadBlob", b.service.config.Server.PublicURL), + Method: "PUT", + Headers: map[string]string{ + "X-Upload-Id": uploadID, + "X-Part-Number": fmt.Sprintf("%d", partNumber), + }, + }, nil } // CompleteMultipartUpload finalizes a multipart upload diff --git a/pkg/hold/handlers.go b/pkg/hold/handlers.go index 1990359..54c009e 100644 --- a/pkg/hold/handlers.go +++ b/pkg/hold/handlers.go @@ -1,496 +1,6 @@ package hold -import ( - "context" - "encoding/json" - "fmt" - "io" - "log" - "net/http" - "time" -) - -// PresignedURLOperation defines the type of presigned URL operation -type PresignedURLOperation string - -const ( - OperationGet PresignedURLOperation = "GET" - OperationHead PresignedURLOperation = "HEAD" - OperationPut PresignedURLOperation = "PUT" -) - -// PresignedURLRequest represents a request for a presigned URL (GET, HEAD, or PUT) -type PresignedURLRequest struct { - Operation PresignedURLOperation `json:"operation"` - DID string `json:"did"` - Digest string `json:"digest"` - Size int64 `json:"size,omitempty"` // Only required for PUT operations -} - -// PresignedURLResponse contains the presigned URL -type PresignedURLResponse struct { - URL string `json:"url"` - ExpiresAt time.Time `json:"expires_at"` -} - -// HandlePresignedURL handles presigned URL requests (GET, HEAD, or PUT) -// Operation type is specified in the request body -func (s *HoldService) HandlePresignedURL(w http.ResponseWriter, r *http.Request) { - if r.Method != http.MethodPost { - http.Error(w, "method not allowed", http.StatusMethodNotAllowed) - return - } - - var req PresignedURLRequest - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - http.Error(w, fmt.Sprintf("invalid request: %v", err), http.StatusBadRequest) - return - } - - // Validate DID authorization based on operation type - var authorized bool - switch req.Operation { - case OperationGet, OperationHead: - authorized = s.isAuthorizedRead(req.DID) - case OperationPut: - authorized = s.isAuthorizedWrite(req.DID) - default: - http.Error(w, "unsupported operation", http.StatusBadRequest) - return - } - - if !authorized { - log.Printf("[HandlePresignedURL:%s] Authorization FAILED", req.Operation) - if req.DID == "" { - http.Error(w, "unauthorized: authentication required", http.StatusUnauthorized) - } else { - http.Error(w, "forbidden: access denied", http.StatusForbidden) - } - return - } - - // Generate presigned URL (15 minute expiry) - ctx := context.Background() - expiry := time.Now().Add(15 * time.Minute) - - url, err := s.getPresignedURL(ctx, req.Operation, req.Digest, req.DID) - if err != nil { - log.Printf("[HandlePresignedURL:%s] getPresignedURL failed: %v", req.Operation, err) - http.Error(w, fmt.Sprintf("failed to generate URL: %v", err), http.StatusInternalServerError) - return - } - - log.Printf("[HandlePresignedURL:%s] Returning URL to client", req.Operation) - - resp := PresignedURLResponse{ - URL: url, - ExpiresAt: expiry, - } - - w.Header().Set("Content-Type", "application/json") - json.NewEncoder(w).Encode(resp) -} - -// HandleProxyGet proxies a blob download through the service -func (s *HoldService) HandleProxyGet(w http.ResponseWriter, r *http.Request) { - if r.Method != http.MethodGet && r.Method != http.MethodHead { - http.Error(w, "method not allowed", http.StatusMethodNotAllowed) - return - } - - // Extract digest from path (e.g., /blobs/sha256:abc123) - digest := r.URL.Path[len("/blobs/"):] - if digest == "" { - http.Error(w, "missing digest", http.StatusBadRequest) - return - } - - // Get DID from query param or header - did := r.URL.Query().Get("did") - if did == "" { - did = r.Header.Get("X-ATCR-DID") - } - log.Printf(" DID: %s", did) - - // Authorize READ access - if !s.isAuthorizedRead(did) { - log.Printf("[HandleProxyGet] Authorization FAILED") - if did == "" { - http.Error(w, "unauthorized: authentication required", http.StatusUnauthorized) - } else { - http.Error(w, "forbidden: access denied", http.StatusForbidden) - } - return - } - - ctx := r.Context() - path := blobPath(digest) - - // For HEAD requests, just check if blob exists - if r.Method == http.MethodHead { - stat, err := s.driver.Stat(ctx, path) - if err != nil { - http.Error(w, "blob not found", http.StatusNotFound) - return - } - w.Header().Set("Content-Type", "application/octet-stream") - w.Header().Set("Content-Length", fmt.Sprintf("%d", stat.Size())) - w.WriteHeader(http.StatusOK) - return - } - - // For GET requests, read and return the blob - content, err := s.driver.GetContent(ctx, path) - if err != nil { - http.Error(w, "blob not found", http.StatusNotFound) - return - } - - w.Header().Set("Content-Type", "application/octet-stream") - w.Write(content) -} - -// HandleMove moves a blob from one path to another -// POST /move?from={path}&to={digest}&did={did} -func (s *HoldService) HandleMove(w http.ResponseWriter, r *http.Request) { - if r.Method != http.MethodPost { - http.Error(w, "method not allowed", http.StatusMethodNotAllowed) - return - } - - fromPath := r.URL.Query().Get("from") - toDigest := r.URL.Query().Get("to") - did := r.URL.Query().Get("did") - - if fromPath == "" || toDigest == "" { - http.Error(w, "missing from or to parameter", http.StatusBadRequest) - return - } - - // Authorize WRITE access - if !s.isAuthorizedWrite(did) { - if did == "" { - http.Error(w, "unauthorized: authentication required", http.StatusUnauthorized) - } else { - http.Error(w, "forbidden: write access denied", http.StatusForbidden) - } - return - } - - ctx := r.Context() - sourcePath := blobPath(fromPath) - destPath := blobPath(toDigest) - - // Try to move using driver's Move operation - if err := s.driver.Move(ctx, sourcePath, destPath); err != nil { - log.Printf("HandleMove: failed to move blob: %v", err) - http.Error(w, fmt.Sprintf("failed to move blob: %v", err), http.StatusInternalServerError) - return - } - - log.Printf("HandleMove: successfully moved blob from=%s to=%s", fromPath, toDigest) - w.WriteHeader(http.StatusOK) -} - -// HandleProxyPut proxies a blob upload through the service -func (s *HoldService) HandleProxyPut(w http.ResponseWriter, r *http.Request) { - if r.Method != http.MethodPut { - http.Error(w, "method not allowed", http.StatusMethodNotAllowed) - return - } - - digest := r.URL.Path[len("/blobs/"):] - if digest == "" { - http.Error(w, "missing digest", http.StatusBadRequest) - return - } - - did := r.URL.Query().Get("did") - if did == "" { - did = r.Header.Get("X-ATCR-DID") - } - - // Authorize WRITE access - if !s.isAuthorizedWrite(did) { - log.Printf("[HandleProxyPut] Authorization FAILED") - if did == "" { - http.Error(w, "unauthorized: authentication required", http.StatusUnauthorized) - } else { - http.Error(w, "forbidden: write access denied", http.StatusForbidden) - } - return - } - - // Stream blob to storage (no buffering) - ctx := r.Context() - path := blobPath(digest) - - // Create writer for streaming - writer, err := s.driver.Writer(ctx, path, false) - if err != nil { - log.Printf("HandleProxyPut: failed to create writer: %v", err) - http.Error(w, "failed to create writer", http.StatusInternalServerError) - return - } - - // Stream directly from request body to storage - written, err := io.Copy(writer, r.Body) - if err != nil { - writer.Cancel(ctx) - log.Printf("HandleProxyPut: failed to write blob: %v", err) - http.Error(w, "failed to write blob", http.StatusInternalServerError) - return - } - - // Commit the write - if err := writer.Commit(ctx); err != nil { - log.Printf("HandleProxyPut: failed to commit blob: %v", err) - http.Error(w, "failed to commit blob", http.StatusInternalServerError) - return - } - - log.Printf("HandleProxyPut: successfully stored blob path=%s, size=%d", digest, written) - w.WriteHeader(http.StatusCreated) -} - -// StartMultipartUploadRequest initiates a multipart upload -type StartMultipartUploadRequest struct { - DID string `json:"did"` - Digest string `json:"digest"` -} - -// StartMultipartUploadResponse contains the multipart upload ID -type StartMultipartUploadResponse struct { - UploadID string `json:"upload_id"` - ExpiresAt time.Time `json:"expires_at"` -} - -// HandleStartMultipart initiates a multipart upload -func (s *HoldService) HandleStartMultipart(w http.ResponseWriter, r *http.Request) { - if r.Method != http.MethodPost { - http.Error(w, "method not allowed", http.StatusMethodNotAllowed) - return - } - - var req StartMultipartUploadRequest - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - http.Error(w, fmt.Sprintf("invalid request: %v", err), http.StatusBadRequest) - return - } - - // Validate DID authorization for WRITE - if !s.isAuthorizedWrite(req.DID) { - if req.DID == "" { - http.Error(w, "unauthorized: authentication required", http.StatusUnauthorized) - } else { - http.Error(w, "forbidden: write access denied", http.StatusForbidden) - } - return - } - - // Start multipart upload with manager (supports both S3Native and Buffered modes) - ctx := r.Context() - uploadID, mode, err := s.StartMultipartUploadWithManager(ctx, req.Digest, s.MultipartMgr) - if err != nil { - http.Error(w, fmt.Sprintf("failed to start multipart upload: %v", err), http.StatusInternalServerError) - return - } - - log.Printf("Started multipart upload: uploadID=%s, mode=%v, digest=%s", uploadID, mode, req.Digest) - - expiry := time.Now().Add(24 * time.Hour) // Multipart uploads can take longer - - resp := StartMultipartUploadResponse{ - UploadID: uploadID, - ExpiresAt: expiry, - } - - w.Header().Set("Content-Type", "application/json") - json.NewEncoder(w).Encode(resp) -} - -// GetPartURLRequest requests a presigned URL for a specific part -type GetPartURLRequest struct { - DID string `json:"did"` - Digest string `json:"digest"` - UploadID string `json:"upload_id"` - PartNumber int `json:"part_number"` -} - -// GetPartURLResponse contains the presigned URL for a part -type GetPartURLResponse struct { - URL string `json:"url"` - ExpiresAt time.Time `json:"expires_at"` -} - -// HandleGetPartURL generates a presigned URL for uploading a specific part -func (s *HoldService) HandleGetPartURL(w http.ResponseWriter, r *http.Request) { - if r.Method != http.MethodPost { - http.Error(w, "method not allowed", http.StatusMethodNotAllowed) - return - } - - var req GetPartURLRequest - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - http.Error(w, fmt.Sprintf("invalid request: %v", err), http.StatusBadRequest) - return - } - - // Validate DID authorization for WRITE - if !s.isAuthorizedWrite(req.DID) { - if req.DID == "" { - http.Error(w, "unauthorized: authentication required", http.StatusUnauthorized) - } else { - http.Error(w, "forbidden: write access denied", http.StatusForbidden) - } - return - } - - // Get multipart session - session, err := s.MultipartMgr.GetSession(req.UploadID) - if err != nil { - http.Error(w, fmt.Sprintf("session not found: %v", err), http.StatusNotFound) - return - } - - // Get part upload URL (presigned for S3Native, proxy for Buffered) - ctx := r.Context() - url, err := s.GetPartUploadURL(ctx, session, req.PartNumber, req.DID) - if err != nil { - http.Error(w, fmt.Sprintf("failed to generate part URL: %v", err), http.StatusInternalServerError) - return - } - - expiry := time.Now().Add(15 * time.Minute) - - resp := GetPartURLResponse{ - URL: url, - ExpiresAt: expiry, - } - - w.Header().Set("Content-Type", "application/json") - json.NewEncoder(w).Encode(resp) -} - -// CompleteMultipartRequest completes a multipart upload -type CompleteMultipartRequest struct { - DID string `json:"did"` - Digest string `json:"digest"` - UploadID string `json:"upload_id"` - Parts []CompletedPart `json:"parts"` -} - -// CompletedPart represents an uploaded part with its ETag -type CompletedPart struct { - PartNumber int `json:"part_number"` - ETag string `json:"etag"` -} - -// HandleCompleteMultipart completes a multipart upload -func (s *HoldService) HandleCompleteMultipart(w http.ResponseWriter, r *http.Request) { - if r.Method != http.MethodPost { - http.Error(w, "method not allowed", http.StatusMethodNotAllowed) - return - } - - var req CompleteMultipartRequest - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - http.Error(w, fmt.Sprintf("invalid request: %v", err), http.StatusBadRequest) - return - } - - // Validate DID authorization for WRITE - if !s.isAuthorizedWrite(req.DID) { - if req.DID == "" { - http.Error(w, "unauthorized: authentication required", http.StatusUnauthorized) - } else { - http.Error(w, "forbidden: write access denied", http.StatusForbidden) - } - return - } - - // Get multipart session - session, err := s.MultipartMgr.GetSession(req.UploadID) - if err != nil { - http.Error(w, fmt.Sprintf("session not found: %v", err), http.StatusNotFound) - return - } - - // For S3Native mode, use parts from request (uploaded directly to S3) - // For Buffered mode, parts are in the session - if session.Mode == S3Native { - // Record parts from AppView's request (they have ETags from S3) - for _, p := range req.Parts { - session.RecordS3Part(p.PartNumber, p.ETag, 0) - } - log.Printf("Recorded %d S3 parts from request for uploadID=%s", len(req.Parts), req.UploadID) - } - - // Complete multipart upload (handles both S3Native and Buffered modes) - ctx := r.Context() - if err := s.CompleteMultipartUploadWithManager(ctx, session, s.MultipartMgr); err != nil { - http.Error(w, fmt.Sprintf("failed to complete multipart upload: %v", err), http.StatusInternalServerError) - return - } - - log.Printf("Completed multipart upload: uploadID=%s, mode=%v", req.UploadID, session.Mode) - - w.WriteHeader(http.StatusOK) - w.Header().Set("Content-Type", "application/json") - json.NewEncoder(w).Encode(map[string]string{ - "status": "completed", - }) -} - -// AbortMultipartRequest aborts an in-progress upload -type AbortMultipartRequest struct { - DID string `json:"did"` - Digest string `json:"digest"` - UploadID string `json:"upload_id"` -} - -// HandleAbortMultipart aborts an in-progress multipart upload -func (s *HoldService) HandleAbortMultipart(w http.ResponseWriter, r *http.Request) { - if r.Method != http.MethodPost { - http.Error(w, "method not allowed", http.StatusMethodNotAllowed) - return - } - - var req AbortMultipartRequest - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - http.Error(w, fmt.Sprintf("invalid request: %v", err), http.StatusBadRequest) - return - } - - // Validate DID authorization for WRITE - if !s.isAuthorizedWrite(req.DID) { - if req.DID == "" { - http.Error(w, "unauthorized: authentication required", http.StatusUnauthorized) - } else { - http.Error(w, "forbidden: write access denied", http.StatusForbidden) - } - return - } - - // Get multipart session - session, err := s.MultipartMgr.GetSession(req.UploadID) - if err != nil { - http.Error(w, fmt.Sprintf("session not found: %v", err), http.StatusNotFound) - return - } - - // Abort multipart upload (handles both S3Native and Buffered modes) - ctx := r.Context() - if err := s.AbortMultipartUploadWithManager(ctx, session, s.MultipartMgr); err != nil { - http.Error(w, fmt.Sprintf("failed to abort multipart upload: %v", err), http.StatusInternalServerError) - return - } - - log.Printf("Aborted multipart upload: uploadID=%s, mode=%v", req.UploadID, session.Mode) - - w.WriteHeader(http.StatusOK) - w.Header().Set("Content-Type", "application/json") - json.NewEncoder(w).Encode(map[string]string{ - "status": "aborted", - }) -} +// This file previously contained legacy HTTP handlers that have been replaced by XRPC endpoints. +// The handlers (HandleProxyGet, HandleProxyPut, HandleMultipartPartUpload) are no longer needed +// as all blob operations now go through the XRPC com.atproto.repo.uploadBlob and +// com.atproto.sync.getBlob endpoints. diff --git a/pkg/hold/multipart.go b/pkg/hold/multipart.go index 56393b5..ef8b663 100644 --- a/pkg/hold/multipart.go +++ b/pkg/hold/multipart.go @@ -24,6 +24,12 @@ const ( Buffered ) +// CompletedPart represents an uploaded part with its ETag +type CompletedPart struct { + PartNumber int `json:"part_number"` + ETag string `json:"etag"` +} + // MultipartSession tracks an in-progress multipart upload type MultipartSession struct { UploadID string // Unique upload ID @@ -270,6 +276,8 @@ func (s *HoldService) GetPartUploadURL(ctx context.Context, session *MultipartSe } // Buffered mode: return proxy endpoint + // url := fmt.Sprintf("%s/xrpc/com.atproto.repo.uploadBlob", s.config.Server.PublicURL) + url := fmt.Sprintf("%s/multipart-parts/%s/%d?did=%s", s.config.Server.PublicURL, session.UploadID, partNumber, did) return url, nil diff --git a/pkg/hold/pds/auth.go b/pkg/hold/pds/auth.go index bac2288..8ae598f 100644 --- a/pkg/hold/pds/auth.go +++ b/pkg/hold/pds/auth.go @@ -7,12 +7,18 @@ import ( "fmt" "io" "net/http" + "slices" "strings" "github.com/bluesky-social/indigo/atproto/identity" "github.com/bluesky-social/indigo/atproto/syntax" ) +// HTTPClient interface allows injecting a custom HTTP client for testing +type HTTPClient interface { + Do(*http.Request) (*http.Response, error) +} + // ValidatedUser represents a successfully validated user from DPoP + OAuth type ValidatedUser struct { DID string @@ -27,7 +33,10 @@ type ValidatedUser struct { // 2. Extract DPoP header (proof JWT) // 3. Call user's PDS to validate token via com.atproto.server.getSession // 4. Return validated user DID -func ValidateDPoPRequest(r *http.Request) (*ValidatedUser, error) { +// +// The httpClient parameter is optional and defaults to http.DefaultClient if nil. +// This allows tests to inject a mock HTTP client. +func ValidateDPoPRequest(r *http.Request, httpClient HTTPClient) (*ValidatedUser, error) { // Extract Authorization header authHeader := r.Header.Get("Authorization") if authHeader == "" { @@ -72,7 +81,7 @@ func ValidateDPoPRequest(r *http.Request) (*ValidatedUser, error) { } // Validate token with the user's PDS - session, err := validateTokenWithPDS(r.Context(), pds, accessToken, dpopProof) + session, err := validateTokenWithPDS(r.Context(), pds, accessToken, dpopProof, httpClient) if err != nil { return nil, fmt.Errorf("token validation failed: %w", err) } @@ -139,7 +148,8 @@ type SessionResponse struct { } // validateTokenWithPDS calls the user's PDS to validate the token -func validateTokenWithPDS(ctx context.Context, pdsURL, accessToken, dpopProof string) (*SessionResponse, error) { +// The httpClient parameter is optional and defaults to http.DefaultClient if nil. +func validateTokenWithPDS(ctx context.Context, pdsURL, accessToken, dpopProof string, httpClient HTTPClient) (*SessionResponse, error) { // Call com.atproto.server.getSession with DPoP headers url := fmt.Sprintf("%s/xrpc/com.atproto.server.getSession", strings.TrimSuffix(pdsURL, "/")) @@ -152,7 +162,13 @@ func validateTokenWithPDS(ctx context.Context, pdsURL, accessToken, dpopProof st req.Header.Set("Authorization", "DPoP "+accessToken) req.Header.Set("DPoP", dpopProof) - resp, err := http.DefaultClient.Do(req) + // Use provided client or default to http.DefaultClient + client := httpClient + if client == nil { + client = http.DefaultClient + } + + resp, err := client.Do(req) if err != nil { return nil, fmt.Errorf("failed to call PDS: %w", err) } @@ -194,10 +210,11 @@ func ResolveDIDToPDS(ctx context.Context, did string) (string, error) { } // ValidateOwnerOrCrewAdmin validates that the request has valid DPoP + OAuth tokens -// and that the authenticated user is either the hold owner or a crew member with crew:admin permission -func ValidateOwnerOrCrewAdmin(r *http.Request, pds *HoldPDS) (*ValidatedUser, error) { +// and that the authenticated user is either the hold owner or a crew member with crew:admin permission. +// The httpClient parameter is optional and defaults to http.DefaultClient if nil. +func ValidateOwnerOrCrewAdmin(r *http.Request, pds *HoldPDS, httpClient HTTPClient) (*ValidatedUser, error) { // Validate DPoP + OAuth token - user, err := ValidateDPoPRequest(r) + user, err := ValidateDPoPRequest(r, httpClient) if err != nil { return nil, fmt.Errorf("authentication failed: %w", err) } @@ -222,10 +239,8 @@ func ValidateOwnerOrCrewAdmin(r *http.Request, pds *HoldPDS) (*ValidatedUser, er for _, member := range crew { if member.Record.Member == user.DID { // Check if this crew member has crew:admin permission - for _, perm := range member.Record.Permissions { - if perm == "crew:admin" { - return user, nil - } + if slices.Contains(member.Record.Permissions, "crew:admin") { + return user, nil } // User is crew but doesn't have admin permission return nil, fmt.Errorf("crew member lacks required 'crew:admin' permission") @@ -235,3 +250,93 @@ func ValidateOwnerOrCrewAdmin(r *http.Request, pds *HoldPDS) (*ValidatedUser, er // User is neither owner nor authorized crew return nil, fmt.Errorf("user is not authorized (must be hold owner or crew admin)") } + +// ValidateBlobWriteAccess validates that the request has valid DPoP + OAuth tokens +// and that the authenticated user is either the hold owner or a crew member with blob:write permission. +// The httpClient parameter is optional and defaults to http.DefaultClient if nil. +func ValidateBlobWriteAccess(r *http.Request, pds *HoldPDS, httpClient HTTPClient) (*ValidatedUser, error) { + // Validate DPoP + OAuth token + user, err := ValidateDPoPRequest(r, httpClient) + if err != nil { + return nil, fmt.Errorf("authentication failed: %w", err) + } + + // Get captain record to check owner and public settings + _, captain, err := pds.GetCaptainRecord(r.Context()) + if err != nil { + return nil, fmt.Errorf("failed to get captain record: %w", err) + } + + // Check if user is the owner (always has write access) + if user.DID == captain.Owner { + return user, nil + } + + // Check if user is crew with blob:write permission + crew, err := pds.ListCrewMembers(r.Context()) + if err != nil { + return nil, fmt.Errorf("failed to check crew membership: %w", err) + } + + for _, member := range crew { + if member.Record.Member == user.DID { + // Check if this crew member has blob:write permission + if slices.Contains(member.Record.Permissions, "blob:write") { + return user, nil + } + // User is crew but doesn't have write permission + return nil, fmt.Errorf("crew member lacks required 'blob:write' permission") + } + } + + // User is neither owner nor authorized crew + return nil, fmt.Errorf("user is not authorized for blob write (must be hold owner or crew with blob:write permission)") +} + +// ValidateBlobReadAccess validates that the request has read access to blobs +// If captain.public = true: No auth required (returns nil user to indicate public access) +// If captain.public = false: Requires valid DPoP + OAuth and (captain OR crew with blob:read permission). +// The httpClient parameter is optional and defaults to http.DefaultClient if nil. +func ValidateBlobReadAccess(r *http.Request, pds *HoldPDS, httpClient HTTPClient) (*ValidatedUser, error) { + // Get captain record to check public setting + _, captain, err := pds.GetCaptainRecord(r.Context()) + if err != nil { + return nil, fmt.Errorf("failed to get captain record: %w", err) + } + + // If hold is public, allow access without authentication + if captain.Public { + return nil, nil // nil user indicates public access + } + + // Private hold - require authentication + user, err := ValidateDPoPRequest(r, httpClient) + if err != nil { + return nil, fmt.Errorf("authentication required for private hold: %w", err) + } + + // Check if user is the owner (always has read access) + if user.DID == captain.Owner { + return user, nil + } + + // Check if user is crew with blob:read permission + crew, err := pds.ListCrewMembers(r.Context()) + if err != nil { + return nil, fmt.Errorf("failed to check crew membership: %w", err) + } + + for _, member := range crew { + if member.Record.Member == user.DID { + // Check if this crew member has blob:read permission + if slices.Contains(member.Record.Permissions, "blob:read") { + return user, nil + } + // User is crew but doesn't have read permission + return nil, fmt.Errorf("crew member lacks required 'blob:read' permission") + } + } + + // User is neither owner nor authorized crew + return nil, fmt.Errorf("user is not authorized for blob read (must be hold owner or crew with blob:read permission)") +} diff --git a/pkg/hold/pds/auth_test.go b/pkg/hold/pds/auth_test.go new file mode 100644 index 0000000..72b434e --- /dev/null +++ b/pkg/hold/pds/auth_test.go @@ -0,0 +1,587 @@ +package pds + +import ( + "encoding/base64" + "encoding/json" + "fmt" + "io" + "net/http" + "net/http/httptest" + "slices" + "strings" + "testing" + "time" + + "github.com/bluesky-social/indigo/atproto/atcrypto" + "github.com/bluesky-social/indigo/atproto/auth/oauth" +) + +// Tests for authorization functions in auth.go + +// mockPDSClient is a mock HTTP client that simulates a PDS server +// It validates DPoP tokens and returns session information +type mockPDSClient struct{} + +func (m *mockPDSClient) Do(req *http.Request) (*http.Response, error) { + // Verify request is for getSession endpoint + if !strings.Contains(req.URL.Path, "/xrpc/com.atproto.server.getSession") { + return &http.Response{ + StatusCode: http.StatusNotFound, + Body: http.NoBody, + }, nil + } + + // Verify DPoP headers are present + authHeader := req.Header.Get("Authorization") + dpopHeader := req.Header.Get("DPoP") + + if authHeader == "" || dpopHeader == "" { + return &http.Response{ + StatusCode: http.StatusUnauthorized, + Body: http.NoBody, + }, nil + } + + // Extract access token from Authorization header + parts := strings.SplitN(authHeader, " ", 2) + if len(parts) != 2 || parts[0] != "DPoP" { + return &http.Response{ + StatusCode: http.StatusUnauthorized, + Body: http.NoBody, + }, nil + } + + accessToken := parts[1] + + // Parse token to extract DID + did, _, err := extractDIDFromToken(accessToken) + if err != nil { + return &http.Response{ + StatusCode: http.StatusBadRequest, + Body: http.NoBody, + }, nil + } + + // Return session response + session := SessionResponse{ + DID: did, + Handle: strings.Replace(did, "did:plc:", "", 1) + ".test", + } + + body, _ := json.Marshal(session) + + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(strings.NewReader(string(body))), + Header: http.Header{"Content-Type": []string{"application/json"}}, + }, nil +} + +// DPoPTestHelper provides utilities for creating valid DPoP requests in tests +type DPoPTestHelper struct { + privKey atcrypto.PrivateKey + did string + pdsURL string +} + +// NewDPoPTestHelper creates a new test helper for the given DID and PDS +func NewDPoPTestHelper(did, pdsURL string) (*DPoPTestHelper, error) { + // Generate a test P-256 key (required for OAuth DPoP) + // Note: ATProto uses K-256 for DID keys, but OAuth DPoP requires P-256 + privKey, err := atcrypto.GeneratePrivateKeyP256() + if err != nil { + return nil, fmt.Errorf("failed to generate key: %w", err) + } + + return &DPoPTestHelper{ + privKey: privKey, + did: did, + pdsURL: pdsURL, + }, nil +} + +// CreateAccessToken creates a mock OAuth access token for testing +// This mimics what a real PDS would issue +func (h *DPoPTestHelper) CreateAccessToken() (string, error) { + // Create access token claims + claims := map[string]any{ + "sub": h.did, // Subject (DID) + "iss": h.pdsURL, // Issuer (PDS URL) + "aud": "atcr", // Audience + "iat": time.Now().Unix(), // Issued at + "exp": time.Now().Add(1 * time.Hour).Unix(), // Expires in 1 hour + } + + // For testing, we create a valid JWT structure without actually validating the signature + // The ValidateDPoPRequest in real use would validate this by calling the PDS + header := base64.RawURLEncoding.EncodeToString([]byte(`{"alg":"ES256K","typ":"JWT"}`)) + payload, err := json.Marshal(claims) + if err != nil { + return "", fmt.Errorf("failed to marshal claims: %w", err) + } + encodedPayload := base64.RawURLEncoding.EncodeToString(payload) + + // Create a mock signature (in real use, the PDS validates this) + signature := base64.RawURLEncoding.EncodeToString([]byte("mock-signature-for-testing")) + + tokenString := fmt.Sprintf("%s.%s.%s", header, encodedPayload, signature) + return tokenString, nil +} + +// CreateDPoPProof creates a DPoP proof JWT for the given HTTP request +func (h *DPoPTestHelper) CreateDPoPProof(method, url string) (string, error) { + return oauth.NewAuthDPoP(method, url, "", h.privKey) +} + +// AddDPoPToRequest adds proper DPoP headers to an HTTP request +func (h *DPoPTestHelper) AddDPoPToRequest(req *http.Request) error { + // Create access token + accessToken, err := h.CreateAccessToken() + if err != nil { + return fmt.Errorf("failed to create access token: %w", err) + } + + // Create DPoP proof for this specific request + dpopProof, err := h.CreateDPoPProof(req.Method, req.URL.String()) + if err != nil { + return fmt.Errorf("failed to create DPoP proof: %w", err) + } + + // Add headers + req.Header.Set("Authorization", "DPoP "+accessToken) + req.Header.Set("DPoP", dpopProof) + + return nil +} + +// AddTestDPoP is a quick helper for common test case: owner with standard PDS +func AddTestDPoP(req *http.Request, did, pdsURL string) error { + helper, err := NewDPoPTestHelper(did, pdsURL) + if err != nil { + return err + } + return helper.AddDPoPToRequest(req) +} + +// TestValidateBlobWriteAccess_Owner tests that the hold owner has write access +func TestValidateBlobWriteAccess_Owner(t *testing.T) { + pds, ctx := setupTestPDS(t) + + ownerDID := "did:plc:owner123" + + // Bootstrap with owner + err := pds.Bootstrap(ctx, ownerDID, true, false) + if err != nil { + t.Fatalf("Failed to bootstrap PDS: %v", err) + } + + // Create DPoP helper for owner + dpopHelper, err := NewDPoPTestHelper(ownerDID, "https://test-pds.example.com") + if err != nil { + t.Fatalf("Failed to create DPoP helper: %v", err) + } + + // Create request with proper DPoP tokens + req := httptest.NewRequest(http.MethodPost, "/test", nil) + if err := dpopHelper.AddDPoPToRequest(req); err != nil { + t.Fatalf("Failed to add DPoP to request: %v", err) + } + + // Use mock PDS client + mockClient := &mockPDSClient{} + + // Test owner has write access + user, err := ValidateBlobWriteAccess(req, pds, mockClient) + if err != nil { + t.Errorf("Expected owner to have write access, got error: %v", err) + } + + if user == nil { + t.Fatal("Expected non-nil user") + } + + if user.DID != ownerDID { + t.Errorf("Expected DID %s, got %s", ownerDID, user.DID) + } + + if !user.Authorized { + t.Error("Expected user to be authorized") + } +} + +// TestValidateBlobWriteAccess_CrewPermissions tests crew permission checking +func TestValidateBlobWriteAccess_CrewPermissions(t *testing.T) { + pds, ctx := setupTestPDS(t) + + ownerDID := "did:plc:owner123" + + // Bootstrap + err := pds.Bootstrap(ctx, ownerDID, true, false) + if err != nil { + t.Fatalf("Failed to bootstrap PDS: %v", err) + } + + // Add crew member with blob:write permission + writerDID := "did:plc:writer123" + _, err = pds.AddCrewMember(ctx, writerDID, "writer", []string{"blob:write"}) + if err != nil { + t.Fatalf("Failed to add crew member: %v", err) + } + + // Add crew member without blob:write permission + readerDID := "did:plc:reader123" + _, err = pds.AddCrewMember(ctx, readerDID, "reader", []string{"blob:read"}) + if err != nil { + t.Fatalf("Failed to add crew member: %v", err) + } + + mockClient := &mockPDSClient{} + + // Test writer (has blob:write permission) can write + t.Run("crew with blob:write can write", func(t *testing.T) { + dpopHelper, err := NewDPoPTestHelper(writerDID, "https://test-pds.example.com") + if err != nil { + t.Fatalf("Failed to create DPoP helper: %v", err) + } + + req := httptest.NewRequest(http.MethodPost, "/test", nil) + if err := dpopHelper.AddDPoPToRequest(req); err != nil { + t.Fatalf("Failed to add DPoP to request: %v", err) + } + + user, err := ValidateBlobWriteAccess(req, pds, mockClient) + if err != nil { + t.Errorf("Expected writer to have write access, got error: %v", err) + } + + if user == nil || user.DID != writerDID { + t.Errorf("Expected user DID %s, got %v", writerDID, user) + } + }) + + // Test reader (no blob:write permission) cannot write + t.Run("crew without blob:write cannot write", func(t *testing.T) { + dpopHelper, err := NewDPoPTestHelper(readerDID, "https://test-pds.example.com") + if err != nil { + t.Fatalf("Failed to create DPoP helper: %v", err) + } + + req := httptest.NewRequest(http.MethodPost, "/test", nil) + if err := dpopHelper.AddDPoPToRequest(req); err != nil { + t.Fatalf("Failed to add DPoP to request: %v", err) + } + + _, err = ValidateBlobWriteAccess(req, pds, mockClient) + if err == nil { + t.Error("Expected reader without blob:write permission to be denied") + } + + if !strings.Contains(err.Error(), "blob:write") { + t.Errorf("Expected error about blob:write permission, got: %v", err) + } + }) +} + +// TestValidateBlobReadAccess_PublicHold tests public hold access +func TestValidateBlobReadAccess_PublicHold(t *testing.T) { + pds, ctx := setupTestPDS(t) + + ownerDID := "did:plc:owner123" + + // Bootstrap with public=true + err := pds.Bootstrap(ctx, ownerDID, true, false) + if err != nil { + t.Fatalf("Failed to bootstrap PDS: %v", err) + } + + // Verify captain record has public=true + _, captain, err := pds.GetCaptainRecord(ctx) + if err != nil { + t.Fatalf("Failed to get captain record: %v", err) + } + + if !captain.Public { + t.Error("Expected public=true for captain record") + } + + // Create request without auth headers (anonymous user) + req := httptest.NewRequest(http.MethodGet, "/test", nil) + + // This should return nil (public access allowed) for public holds + user, err := ValidateBlobReadAccess(req, pds, nil) + if err != nil { + t.Errorf("Expected public access for public hold, got error: %v", err) + } + + // nil user indicates public access + if user != nil { + t.Error("Expected nil user for public access") + } +} + +// TestValidateBlobReadAccess_PrivateHold tests private hold access +func TestValidateBlobReadAccess_PrivateHold(t *testing.T) { + pds, ctx := setupTestPDS(t) + + ownerDID := "did:plc:owner123" + + // Bootstrap with public=false + err := pds.Bootstrap(ctx, ownerDID, false, false) + if err != nil { + t.Fatalf("Failed to bootstrap PDS: %v", err) + } + + // Update captain to be private + _, err = pds.UpdateCaptainRecord(ctx, false, false) + if err != nil { + t.Fatalf("Failed to update captain record: %v", err) + } + + // Verify captain record has public=false + _, captain, err := pds.GetCaptainRecord(ctx) + if err != nil { + t.Fatalf("Failed to get captain record: %v", err) + } + + if captain.Public { + t.Error("Expected public=false for captain record") + } + + // Create request without auth headers (anonymous user) + req := httptest.NewRequest(http.MethodGet, "/test", nil) + + // This should return error (auth required) for private holds + user, err := ValidateBlobReadAccess(req, pds, nil) + if err == nil { + t.Error("Expected error for private hold without auth") + } + + if user != nil { + t.Error("Expected nil user when auth fails") + } +} + +// TestValidateOwnerOrCrewAdmin tests admin permission checking +func TestValidateOwnerOrCrewAdmin(t *testing.T) { + pds, ctx := setupTestPDS(t) + + ownerDID := "did:plc:owner123" + + // Bootstrap + err := pds.Bootstrap(ctx, ownerDID, true, false) + if err != nil { + t.Fatalf("Failed to bootstrap PDS: %v", err) + } + + // Add crew member with crew:admin permission + adminDID := "did:plc:admin123" + _, err = pds.AddCrewMember(ctx, adminDID, "admin", []string{"crew:admin", "blob:write", "blob:read"}) + if err != nil { + t.Fatalf("Failed to add crew admin: %v", err) + } + + // Add crew member without crew:admin permission + writerDID := "did:plc:writer123" + _, err = pds.AddCrewMember(ctx, writerDID, "writer", []string{"blob:write"}) + if err != nil { + t.Fatalf("Failed to add crew writer: %v", err) + } + + // Verify crew records were created + crew, err := pds.ListCrewMembers(ctx) + if err != nil { + t.Fatalf("Failed to list crew members: %v", err) + } + + // Verify admin has crew:admin permission + hasAdminPermission := false + for _, member := range crew { + if member.Record.Member == adminDID { + if slices.Contains(member.Record.Permissions, "crew:admin") { + hasAdminPermission = true + } + } + } + + if !hasAdminPermission { + t.Error("Admin crew member should have crew:admin permission") + } + + // Verify writer does NOT have crew:admin permission + writerHasAdminPermission := false + for _, member := range crew { + if member.Record.Member == writerDID { + if slices.Contains(member.Record.Permissions, "crew:admin") { + writerHasAdminPermission = true + } + } + } + + if writerHasAdminPermission { + t.Error("Writer crew member should NOT have crew:admin permission") + } + + // Test that function requires auth (will fail without DPoP tokens) + req := httptest.NewRequest(http.MethodPost, "/test", nil) + _, err = ValidateOwnerOrCrewAdmin(req, pds, nil) + if err == nil { + t.Error("Expected error for missing auth headers") + } +} + +// TestCrewPermissions tests various permission combinations +func TestCrewPermissions(t *testing.T) { + pds, ctx := setupTestPDS(t) + + ownerDID := "did:plc:owner123" + + // Bootstrap + err := pds.Bootstrap(ctx, ownerDID, true, false) + if err != nil { + t.Fatalf("Failed to bootstrap PDS: %v", err) + } + + tests := []struct { + name string + did string + role string + permissions []string + }{ + { + name: "full admin", + did: "did:plc:fulladmin", + role: "admin", + permissions: []string{"crew:admin", "blob:write", "blob:read"}, + }, + { + name: "writer only", + did: "did:plc:writer", + role: "writer", + permissions: []string{"blob:write"}, + }, + { + name: "reader only", + did: "did:plc:reader", + role: "reader", + permissions: []string{"blob:read"}, + }, + { + name: "read-write", + did: "did:plc:readwrite", + role: "editor", + permissions: []string{"blob:read", "blob:write"}, + }, + } + + // Add all crew members + for _, tt := range tests { + _, err := pds.AddCrewMember(ctx, tt.did, tt.role, tt.permissions) + if err != nil { + t.Fatalf("Failed to add crew member %s: %v", tt.name, err) + } + } + + // Verify all crew members were created + crew, err := pds.ListCrewMembers(ctx) + if err != nil { + t.Fatalf("Failed to list crew members: %v", err) + } + + // Should have: 1 owner (from bootstrap) + 4 test crew members + expectedCount := len(tests) + 1 + if len(crew) != expectedCount { + t.Errorf("Expected %d crew members (owner + %d test members), got %d", + expectedCount, len(tests), len(crew)) + } + + // Verify each crew member has the expected permissions + for _, tt := range tests { + found := false + for _, member := range crew { + if member.Record.Member == tt.did { + found = true + + // Check that all expected permissions are present + for _, expectedPerm := range tt.permissions { + hasPerm := slices.Contains(member.Record.Permissions, expectedPerm) + if !hasPerm { + t.Errorf("Crew member %s missing expected permission %s", + tt.name, expectedPerm) + } + } + + // Verify role + if member.Record.Role != tt.role { + t.Errorf("Crew member %s has role %s, expected %s", + tt.name, member.Record.Role, tt.role) + } + } + } + + if !found { + t.Errorf("Crew member %s not found in list", tt.name) + } + } +} + +// TestCaptainRecordSettings tests captain record public/allowAllCrew settings +func TestCaptainRecordSettings(t *testing.T) { + tests := []struct { + name string + public bool + allowAllCrew bool + }{ + { + name: "public hold, crew approval required", + public: true, + allowAllCrew: false, + }, + { + name: "public hold, open crew", + public: true, + allowAllCrew: true, + }, + { + name: "private hold, crew approval required", + public: false, + allowAllCrew: false, + }, + { + name: "private hold, open crew", + public: false, + allowAllCrew: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + pds, ctx := setupTestPDS(t) + + ownerDID := "did:plc:owner123" + + // Bootstrap with specified settings + err := pds.Bootstrap(ctx, ownerDID, tt.public, tt.allowAllCrew) + if err != nil { + t.Fatalf("Failed to bootstrap PDS: %v", err) + } + + // Verify captain record has expected settings + _, captain, err := pds.GetCaptainRecord(ctx) + if err != nil { + t.Fatalf("Failed to get captain record: %v", err) + } + + if captain.Public != tt.public { + t.Errorf("Expected public=%v, got %v", tt.public, captain.Public) + } + + if captain.AllowAllCrew != tt.allowAllCrew { + t.Errorf("Expected allowAllCrew=%v, got %v", tt.allowAllCrew, captain.AllowAllCrew) + } + + if captain.Owner != ownerDID { + t.Errorf("Expected owner %s, got %s", ownerDID, captain.Owner) + } + }) + } +} diff --git a/pkg/hold/pds/did_test.go b/pkg/hold/pds/did_test.go new file mode 100644 index 0000000..d807456 --- /dev/null +++ b/pkg/hold/pds/did_test.go @@ -0,0 +1,274 @@ +package pds + +import ( + "context" + "encoding/json" + "path/filepath" + "testing" +) + +// TestGenerateDIDFromURL tests DID generation from various URL formats +func TestGenerateDIDFromURL(t *testing.T) { + tests := []struct { + name string + publicURL string + expectedDID string + }{ + { + name: "standard HTTP with standard port", + publicURL: "http://hold.example.com", + expectedDID: "did:web:hold.example.com", + }, + { + name: "standard HTTPS with standard port", + publicURL: "https://hold.example.com", + expectedDID: "did:web:hold.example.com", + }, + { + name: "HTTP with non-standard port", + publicURL: "http://hold.example.com:8080", + expectedDID: "did:web:hold.example.com:8080", + }, + { + name: "HTTPS with non-standard port", + publicURL: "https://hold.example.com:8443", + expectedDID: "did:web:hold.example.com:8443", + }, + { + name: "localhost with port", + publicURL: "http://localhost:8080", + expectedDID: "did:web:localhost:8080", + }, + { + name: "HTTP with explicit port 80", + publicURL: "http://hold.example.com:80", + expectedDID: "did:web:hold.example.com", + }, + { + name: "HTTPS with explicit port 443", + publicURL: "https://hold.example.com:443", + expectedDID: "did:web:hold.example.com", + }, + { + name: "subdomain", + publicURL: "https://hold1.atcr.io", + expectedDID: "did:web:hold1.atcr.io", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + did := GenerateDIDFromURL(tt.publicURL) + if did != tt.expectedDID { + t.Errorf("Expected DID %s, got %s", tt.expectedDID, did) + } + }) + } +} + +// TestGenerateDIDFromURL_InvalidURL tests handling of invalid URLs +func TestGenerateDIDFromURL_InvalidURL(t *testing.T) { + // Invalid URLs get parsed with empty hostname, which defaults to localhost + did := GenerateDIDFromURL("not a url") + if did != "did:web:localhost" { + t.Errorf("Expected did:web:localhost for invalid URL, got %s", did) + } +} + +// TestGenerateDIDDocument tests DID document generation +func TestGenerateDIDDocument(t *testing.T) { + ctx := context.Background() + tmpDir := t.TempDir() + + dbPath := filepath.Join(tmpDir, "pds.db") + keyPath := filepath.Join(tmpDir, "signing-key") + publicURL := "https://hold.example.com" + + pds, err := NewHoldPDS(ctx, "did:web:hold.example.com", publicURL, dbPath, keyPath) + if err != nil { + t.Fatalf("Failed to create PDS: %v", err) + } + + doc, err := pds.GenerateDIDDocument(publicURL) + if err != nil { + t.Fatalf("Failed to generate DID document: %v", err) + } + + // Verify required fields + if doc.ID != "did:web:hold.example.com" { + t.Errorf("Expected DID did:web:hold.example.com, got %s", doc.ID) + } + + // Verify context + if len(doc.Context) != 3 { + t.Errorf("Expected 3 context entries, got %d", len(doc.Context)) + } + + expectedContexts := []string{ + "https://www.w3.org/ns/did/v1", + "https://w3id.org/security/multikey/v1", + "https://w3id.org/security/suites/secp256k1-2019/v1", + } + for i, expected := range expectedContexts { + if doc.Context[i] != expected { + t.Errorf("Expected context[%d] = %s, got %s", i, expected, doc.Context[i]) + } + } + + // Verify alsoKnownAs + if len(doc.AlsoKnownAs) != 1 || doc.AlsoKnownAs[0] != "at://hold.example.com" { + t.Errorf("Expected alsoKnownAs=['at://hold.example.com'], got %v", doc.AlsoKnownAs) + } + + // Verify verification method + if len(doc.VerificationMethod) != 1 { + t.Fatalf("Expected 1 verification method, got %d", len(doc.VerificationMethod)) + } + + vm := doc.VerificationMethod[0] + if vm.ID != "did:web:hold.example.com#atproto" { + t.Errorf("Expected verification method ID did:web:hold.example.com#atproto, got %s", vm.ID) + } + if vm.Type != "Multikey" { + t.Errorf("Expected type Multikey, got %s", vm.Type) + } + if vm.Controller != "did:web:hold.example.com" { + t.Errorf("Expected controller did:web:hold.example.com, got %s", vm.Controller) + } + if vm.PublicKeyMultibase == "" { + t.Error("Expected non-empty publicKeyMultibase") + } + + // Verify authentication + if len(doc.Authentication) != 1 || doc.Authentication[0] != "did:web:hold.example.com#atproto" { + t.Errorf("Expected authentication=['did:web:hold.example.com#atproto'], got %v", doc.Authentication) + } + + // Verify services + if len(doc.Service) != 2 { + t.Fatalf("Expected 2 services, got %d", len(doc.Service)) + } + + // Check PDS service + pdsService := doc.Service[0] + if pdsService.ID != "#atproto_pds" { + t.Errorf("Expected service ID #atproto_pds, got %s", pdsService.ID) + } + if pdsService.Type != "AtprotoPersonalDataServer" { + t.Errorf("Expected service type AtprotoPersonalDataServer, got %s", pdsService.Type) + } + if pdsService.ServiceEndpoint != publicURL { + t.Errorf("Expected service endpoint %s, got %s", publicURL, pdsService.ServiceEndpoint) + } + + // Check hold service + holdService := doc.Service[1] + if holdService.ID != "#atcr_hold" { + t.Errorf("Expected service ID #atcr_hold, got %s", holdService.ID) + } + if holdService.Type != "AtcrHoldService" { + t.Errorf("Expected service type AtcrHoldService, got %s", holdService.Type) + } + if holdService.ServiceEndpoint != publicURL { + t.Errorf("Expected service endpoint %s, got %s", publicURL, holdService.ServiceEndpoint) + } +} + +// TestGenerateDIDDocument_WithPort tests DID document with non-standard port +func TestGenerateDIDDocument_WithPort(t *testing.T) { + ctx := context.Background() + tmpDir := t.TempDir() + + dbPath := filepath.Join(tmpDir, "pds.db") + keyPath := filepath.Join(tmpDir, "signing-key") + publicURL := "https://hold.example.com:8443" + + pds, err := NewHoldPDS(ctx, "did:web:hold.example.com:8443", publicURL, dbPath, keyPath) + if err != nil { + t.Fatalf("Failed to create PDS: %v", err) + } + + doc, err := pds.GenerateDIDDocument(publicURL) + if err != nil { + t.Fatalf("Failed to generate DID document: %v", err) + } + + // Verify DID includes port + if doc.ID != "did:web:hold.example.com:8443" { + t.Errorf("Expected DID did:web:hold.example.com:8443, got %s", doc.ID) + } + + // Verify alsoKnownAs includes port + if doc.AlsoKnownAs[0] != "at://hold.example.com:8443" { + t.Errorf("Expected alsoKnownAs with port, got %s", doc.AlsoKnownAs[0]) + } +} + +// TestMarshalDIDDocument tests DID document JSON marshaling +func TestMarshalDIDDocument(t *testing.T) { + ctx := context.Background() + tmpDir := t.TempDir() + + dbPath := filepath.Join(tmpDir, "pds.db") + keyPath := filepath.Join(tmpDir, "signing-key") + publicURL := "https://hold.example.com" + + pds, err := NewHoldPDS(ctx, "did:web:hold.example.com", publicURL, dbPath, keyPath) + if err != nil { + t.Fatalf("Failed to create PDS: %v", err) + } + + jsonBytes, err := pds.MarshalDIDDocument() + if err != nil { + t.Fatalf("Failed to marshal DID document: %v", err) + } + + // Verify it's valid JSON + var doc map[string]any + if err := json.Unmarshal(jsonBytes, &doc); err != nil { + t.Fatalf("Failed to unmarshal DID document JSON: %v", err) + } + + // Verify required fields + if id, ok := doc["id"].(string); !ok || id != "did:web:hold.example.com" { + t.Errorf("Expected id='did:web:hold.example.com', got %v", doc["id"]) + } + + if _, ok := doc["@context"]; !ok { + t.Error("Expected @context field in JSON") + } + + if _, ok := doc["verificationMethod"]; !ok { + t.Error("Expected verificationMethod field in JSON") + } + + if _, ok := doc["service"]; !ok { + t.Error("Expected service field in JSON") + } + + // Verify pretty-printed (has indentation) + if len(jsonBytes) < 100 { + t.Error("Expected pretty-printed JSON to be reasonably sized") + } +} + +// TestGenerateDIDDocument_InvalidURL tests error handling +func TestGenerateDIDDocument_InvalidURL(t *testing.T) { + ctx := context.Background() + tmpDir := t.TempDir() + + dbPath := filepath.Join(tmpDir, "pds.db") + keyPath := filepath.Join(tmpDir, "signing-key") + publicURL := "https://hold.example.com" + + pds, err := NewHoldPDS(ctx, "did:web:hold.example.com", publicURL, dbPath, keyPath) + if err != nil { + t.Fatalf("Failed to create PDS: %v", err) + } + + // Try to generate DID document with invalid URL + _, err = pds.GenerateDIDDocument("ht!tp://invalid url") + if err == nil { + t.Error("Expected error for invalid URL, got nil") + } +} diff --git a/pkg/hold/pds/events.go b/pkg/hold/pds/events.go index a873c33..a062b2a 100644 --- a/pkg/hold/pds/events.go +++ b/pkg/hold/pds/events.go @@ -19,7 +19,7 @@ type EventBroadcaster struct { eventSeq int64 eventHistory []HistoricalEvent // Ring buffer for cursor backfill maxHistory int - holdDID string // DID of the hold for setting repo field + holdDID string // DID of the hold for setting repo field } // Subscriber represents a WebSocket client subscribed to the firehose @@ -37,15 +37,15 @@ type HistoricalEvent struct { // RepoCommitEvent represents a #commit event in subscribeRepos type RepoCommitEvent struct { - Seq int64 `json:"seq" cborgen:"seq"` - Repo string `json:"repo" cborgen:"repo"` - Commit string `json:"commit" cborgen:"commit"` // CID string - Rev string `json:"rev" cborgen:"rev"` - Since *string `json:"since,omitempty" cborgen:"since,omitempty"` - Blocks []byte `json:"blocks" cborgen:"blocks"` // CAR slice bytes - Ops []*atproto.SyncSubscribeRepos_RepoOp `json:"ops" cborgen:"ops"` - Time string `json:"time" cborgen:"time"` - Type string `json:"$type" cborgen:"$type"` // Always "#commit" + Seq int64 `json:"seq" cborgen:"seq"` + Repo string `json:"repo" cborgen:"repo"` + Commit string `json:"commit" cborgen:"commit"` // CID string + Rev string `json:"rev" cborgen:"rev"` + Since *string `json:"since,omitempty" cborgen:"since,omitempty"` + Blocks []byte `json:"blocks" cborgen:"blocks"` // CAR slice bytes + Ops []*atproto.SyncSubscribeRepos_RepoOp `json:"ops" cborgen:"ops"` + Time string `json:"time" cborgen:"time"` + Type string `json:"$type" cborgen:"$type"` // Always "#commit" } // NewEventBroadcaster creates a new event broadcaster diff --git a/pkg/hold/pds/events_test.go b/pkg/hold/pds/events_test.go new file mode 100644 index 0000000..161e176 --- /dev/null +++ b/pkg/hold/pds/events_test.go @@ -0,0 +1,384 @@ +package pds + +import ( + "context" + "encoding/json" + "testing" + "time" + + atproto "github.com/bluesky-social/indigo/api/atproto" + "github.com/ipfs/go-cid" +) + +// TestNewEventBroadcaster tests event broadcaster creation +func TestNewEventBroadcaster(t *testing.T) { + holdDID := "did:web:hold.example.com" + broadcaster := NewEventBroadcaster(holdDID, 100) + + if broadcaster.holdDID != holdDID { + t.Errorf("Expected holdDID=%s, got %s", holdDID, broadcaster.holdDID) + } + + if broadcaster.eventSeq != 0 { + t.Errorf("Expected initial eventSeq=0, got %d", broadcaster.eventSeq) + } + + if broadcaster.maxHistory != 100 { + t.Errorf("Expected maxHistory=100, got %d", broadcaster.maxHistory) + } + + if len(broadcaster.subscribers) != 0 { + t.Errorf("Expected 0 subscribers initially, got %d", len(broadcaster.subscribers)) + } +} + +// TestNewEventBroadcaster_DefaultHistory tests default history size +func TestNewEventBroadcaster_DefaultHistory(t *testing.T) { + // Zero or negative maxHistory should default to 100 + broadcaster := NewEventBroadcaster("did:web:test", 0) + if broadcaster.maxHistory != 100 { + t.Errorf("Expected default maxHistory=100 for input 0, got %d", broadcaster.maxHistory) + } + + broadcaster2 := NewEventBroadcaster("did:web:test", -5) + if broadcaster2.maxHistory != 100 { + t.Errorf("Expected default maxHistory=100 for negative input, got %d", broadcaster2.maxHistory) + } +} + +// TestGetCurrentSeq tests sequence number tracking +func TestGetCurrentSeq(t *testing.T) { + broadcaster := NewEventBroadcaster("did:web:hold.example.com", 10) + + // Initial seq should be 0 + seq := broadcaster.GetCurrentSeq() + if seq != 0 { + t.Errorf("Expected initial seq=0, got %d", seq) + } + + // After broadcasting, seq should increment + ctx := context.Background() + testCID, _ := cid.Decode("bafyreib2rxk3rkhh5ylyxj3x3gathxt3s32qvwj2lf3qg4kmzr6b7teqke") + + event := &RepoEvent{ + NewRoot: testCID, + Rev: "test-rev-1", + RepoSlice: []byte("test CAR data"), + Ops: []RepoOp{ + { + Kind: EvtKindCreateRecord, + Collection: "io.atcr.hold.crew", + Rkey: "test123", + }, + }, + } + + broadcaster.Broadcast(ctx, event) + + seq = broadcaster.GetCurrentSeq() + if seq != 1 { + t.Errorf("Expected seq=1 after one broadcast, got %d", seq) + } + + // Broadcast again + broadcaster.Broadcast(ctx, event) + + seq = broadcaster.GetCurrentSeq() + if seq != 2 { + t.Errorf("Expected seq=2 after two broadcasts, got %d", seq) + } +} + +// TestBroadcast tests event broadcasting +func TestBroadcast(t *testing.T) { + broadcaster := NewEventBroadcaster("did:web:hold.example.com", 10) + ctx := context.Background() + + testCID, _ := cid.Decode("bafyreib2rxk3rkhh5ylyxj3x3gathxt3s32qvwj2lf3qg4kmzr6b7teqke") + + event := &RepoEvent{ + NewRoot: testCID, + Rev: "test-rev-1", + RepoSlice: []byte("test CAR data"), + Ops: []RepoOp{ + { + Kind: EvtKindCreateRecord, + Collection: "io.atcr.hold.crew", + Rkey: "test123", + RecCid: &testCID, + }, + }, + } + + // Broadcast should not panic without subscribers + broadcaster.Broadcast(ctx, event) + + // Verify sequence incremented + if broadcaster.eventSeq != 1 { + t.Errorf("Expected eventSeq=1, got %d", broadcaster.eventSeq) + } + + // Verify event added to history + if len(broadcaster.eventHistory) != 1 { + t.Errorf("Expected 1 event in history, got %d", len(broadcaster.eventHistory)) + } + + he := broadcaster.eventHistory[0] + if he.Seq != 1 { + t.Errorf("Expected history seq=1, got %d", he.Seq) + } + + if he.Event.Repo != "did:web:hold.example.com" { + t.Errorf("Expected repo=did:web:hold.example.com, got %s", he.Event.Repo) + } + + if he.Event.Type != "#commit" { + t.Errorf("Expected type=#commit, got %s", he.Event.Type) + } + + if len(he.Event.Ops) != 1 { + t.Errorf("Expected 1 op, got %d", len(he.Event.Ops)) + } +} + +// TestAddToHistory_RingBuffer tests ring buffer behavior +func TestAddToHistory_RingBuffer(t *testing.T) { + // Create broadcaster with small history + broadcaster := NewEventBroadcaster("did:web:hold.example.com", 3) + ctx := context.Background() + + testCID, _ := cid.Decode("bafyreib2rxk3rkhh5ylyxj3x3gathxt3s32qvwj2lf3qg4kmzr6b7teqke") + + // Broadcast 5 events (exceeds maxHistory of 3) + for i := 0; i < 5; i++ { + event := &RepoEvent{ + NewRoot: testCID, + Rev: "test-rev", + RepoSlice: []byte("test CAR data"), + Ops: []RepoOp{}, + } + broadcaster.Broadcast(ctx, event) + } + + // Should only keep last 3 events + if len(broadcaster.eventHistory) != 3 { + t.Errorf("Expected 3 events in history (ring buffer), got %d", len(broadcaster.eventHistory)) + } + + // Verify we kept the most recent events (seq 3, 4, 5) + expectedSeqs := []int64{3, 4, 5} + for i, expected := range expectedSeqs { + if broadcaster.eventHistory[i].Seq != expected { + t.Errorf("Expected history[%d].Seq=%d, got %d", i, expected, broadcaster.eventHistory[i].Seq) + } + } + + // Final sequence should be 5 + if broadcaster.eventSeq != 5 { + t.Errorf("Expected eventSeq=5, got %d", broadcaster.eventSeq) + } +} + +// TestConvertToCommitEvent tests event conversion +func TestConvertToCommitEvent(t *testing.T) { + broadcaster := NewEventBroadcaster("did:web:hold.example.com", 10) + + testCID, _ := cid.Decode("bafyreib2rxk3rkhh5ylyxj3x3gathxt3s32qvwj2lf3qg4kmzr6b7teqke") + since := "prev-rev" + + event := &RepoEvent{ + NewRoot: testCID, + Rev: "test-rev-123", + Since: &since, + RepoSlice: []byte("test CAR data"), + Ops: []RepoOp{ + { + Kind: EvtKindCreateRecord, + Collection: "io.atcr.hold.crew", + Rkey: "member1", + RecCid: &testCID, + }, + { + Kind: EvtKindUpdateRecord, + Collection: "io.atcr.hold.captain", + Rkey: "self", + RecCid: &testCID, + }, + { + Kind: EvtKindDeleteRecord, + Collection: "io.atcr.hold.crew", + Rkey: "oldmember", + RecCid: nil, // Deletes don't have CIDs + }, + }, + } + + commitEvent := broadcaster.convertToCommitEvent(event, 42) + + // Verify basic fields + if commitEvent.Seq != 42 { + t.Errorf("Expected seq=42, got %d", commitEvent.Seq) + } + + if commitEvent.Repo != "did:web:hold.example.com" { + t.Errorf("Expected repo=did:web:hold.example.com, got %s", commitEvent.Repo) + } + + if commitEvent.Commit != testCID.String() { + t.Errorf("Expected commit=%s, got %s", testCID.String(), commitEvent.Commit) + } + + if commitEvent.Rev != "test-rev-123" { + t.Errorf("Expected rev=test-rev-123, got %s", commitEvent.Rev) + } + + if commitEvent.Since == nil || *commitEvent.Since != since { + t.Errorf("Expected since=%s, got %v", since, commitEvent.Since) + } + + if string(commitEvent.Blocks) != "test CAR data" { + t.Errorf("Expected blocks='test CAR data', got %s", string(commitEvent.Blocks)) + } + + if commitEvent.Type != "#commit" { + t.Errorf("Expected type=#commit, got %s", commitEvent.Type) + } + + // Verify time is set + if commitEvent.Time == "" { + t.Error("Expected non-empty time") + } + + // Parse time to verify it's valid RFC3339 + _, err := time.Parse(time.RFC3339, commitEvent.Time) + if err != nil { + t.Errorf("Expected valid RFC3339 time, got error: %v", err) + } + + // Verify ops conversion + if len(commitEvent.Ops) != 3 { + t.Fatalf("Expected 3 ops, got %d", len(commitEvent.Ops)) + } + + // Check create op + createOp := commitEvent.Ops[0] + if createOp.Action != "create" { + t.Errorf("Expected action=create, got %s", createOp.Action) + } + if createOp.Path != "io.atcr.hold.crew/member1" { + t.Errorf("Expected path=io.atcr.hold.crew/member1, got %s", createOp.Path) + } + if createOp.Cid == nil { + t.Error("Expected non-nil CID for create op") + } + + // Check update op + updateOp := commitEvent.Ops[1] + if updateOp.Action != "update" { + t.Errorf("Expected action=update, got %s", updateOp.Action) + } + if updateOp.Path != "io.atcr.hold.captain/self" { + t.Errorf("Expected path=io.atcr.hold.captain/self, got %s", updateOp.Path) + } + + // Check delete op + deleteOp := commitEvent.Ops[2] + if deleteOp.Action != "delete" { + t.Errorf("Expected action=delete, got %s", deleteOp.Action) + } + if deleteOp.Path != "io.atcr.hold.crew/oldmember" { + t.Errorf("Expected path=io.atcr.hold.crew/oldmember, got %s", deleteOp.Path) + } + if deleteOp.Cid != nil { + t.Error("Expected nil CID for delete op") + } +} + +// TestConvertToCommitEvent_NoSince tests event without since field +func TestConvertToCommitEvent_NoSince(t *testing.T) { + broadcaster := NewEventBroadcaster("did:web:hold.example.com", 10) + + testCID, _ := cid.Decode("bafyreib2rxk3rkhh5ylyxj3x3gathxt3s32qvwj2lf3qg4kmzr6b7teqke") + + event := &RepoEvent{ + NewRoot: testCID, + Rev: "test-rev-123", + Since: nil, // No since + RepoSlice: []byte("test CAR data"), + Ops: []RepoOp{}, + } + + commitEvent := broadcaster.convertToCommitEvent(event, 1) + + if commitEvent.Since != nil { + t.Errorf("Expected nil since, got %v", commitEvent.Since) + } +} + +// TestSetRepoEventHandler tests handler registration +func TestSetRepoEventHandler(t *testing.T) { + broadcaster := NewEventBroadcaster("did:web:hold.example.com", 10) + + handler := broadcaster.SetRepoEventHandler() + if handler == nil { + t.Fatal("Expected non-nil handler") + } + + // Call handler + ctx := context.Background() + testCID, _ := cid.Decode("bafyreib2rxk3rkhh5ylyxj3x3gathxt3s32qvwj2lf3qg4kmzr6b7teqke") + + event := &RepoEvent{ + NewRoot: testCID, + Rev: "test-rev", + RepoSlice: []byte("test CAR data"), + Ops: []RepoOp{}, + } + + handler(ctx, event) + + // Verify event was broadcast + if broadcaster.eventSeq != 1 { + t.Errorf("Expected eventSeq=1 after handler call, got %d", broadcaster.eventSeq) + } + + if len(broadcaster.eventHistory) != 1 { + t.Errorf("Expected 1 event in history after handler call, got %d", len(broadcaster.eventHistory)) + } +} + +// TestEncodeCBOR tests CBOR encoding (currently JSON) +func TestEncodeCBOR(t *testing.T) { + testCID, _ := cid.Decode("bafyreib2rxk3rkhh5ylyxj3x3gathxt3s32qvwj2lf3qg4kmzr6b7teqke") + + event := &RepoCommitEvent{ + Seq: 1, + Repo: "did:web:hold.example.com", + Commit: testCID.String(), + Rev: "test-rev", + Blocks: []byte("test data"), + Ops: []*atproto.SyncSubscribeRepos_RepoOp{}, + Time: time.Now().Format(time.RFC3339), + Type: "#commit", + } + + encoded, err := encodeCBOR(event) + if err != nil { + t.Fatalf("Failed to encode CBOR: %v", err) + } + + if len(encoded) == 0 { + t.Error("Expected non-empty encoded data") + } + + // Current implementation uses JSON, so verify it's valid JSON + // In future, this would be proper CBOR validation + var decoded RepoCommitEvent + if err := json.Unmarshal(encoded, &decoded); err != nil { + t.Errorf("Failed to decode JSON: %v", err) + } + + if decoded.Seq != 1 { + t.Errorf("Expected decoded seq=1, got %d", decoded.Seq) + } +} diff --git a/pkg/hold/pds/repomgr.go b/pkg/hold/pds/repomgr.go index 8289dd0..162d851 100644 --- a/pkg/hold/pds/repomgr.go +++ b/pkg/hold/pds/repomgr.go @@ -5,14 +5,14 @@ // Reason: The indigo library is unmaintained and contains a critical bug in UpdateRecord // // Modifications from original: -// - Changed package from 'repomgr' to 'pds' for integration with hold service -// - Fixed UpdateRecord bug (line 263): Changed r.PutRecord to r.UpdateRecord -// (UpdateRecord was incorrectly calling PutRecord, causing incorrect MST operations) -// - Removed 5 Prometheus metrics calls (openAndSigCheckDuration, calcDiffDuration, -// writeCarSliceDuration, repoOpsImported) as metrics are not used in this project -// - Added PutRecord method (lines 309-381) for creating records with explicit rkeys -// (like CreateRecord but with specified rkey instead of auto-generated TID) -// Based on streamplace/indigo implementation +// - Changed package from 'repomgr' to 'pds' for integration with hold service +// - Fixed UpdateRecord bug (line 263): Changed r.PutRecord to r.UpdateRecord +// (UpdateRecord was incorrectly calling PutRecord, causing incorrect MST operations) +// - Removed 5 Prometheus metrics calls (openAndSigCheckDuration, calcDiffDuration, +// writeCarSliceDuration, repoOpsImported) as metrics are not used in this project +// - Added PutRecord method (lines 309-381) for creating records with explicit rkeys +// (like CreateRecord but with specified rkey instead of auto-generated TID) +// Based on streamplace/indigo implementation package pds import ( diff --git a/pkg/hold/pds/xrpc.go b/pkg/hold/pds/xrpc.go index 1a989f9..a385596 100644 --- a/pkg/hold/pds/xrpc.go +++ b/pkg/hold/pds/xrpc.go @@ -27,6 +27,7 @@ type XRPCHandler struct { publicURL string blobStore BlobStore broadcaster *EventBroadcaster + httpClient HTTPClient // For testing - allows injecting mock HTTP client } // BlobStore interface wraps the existing hold service storage operations @@ -48,8 +49,8 @@ type BlobStore interface { // Multipart upload operations (used for OCI container layers only) // StartMultipartUpload initiates a multipart upload, returns uploadID and mode StartMultipartUpload(ctx context.Context, digest string) (uploadID string, mode string, err error) - // GetPartUploadURL returns a presigned URL for uploading a specific part - GetPartUploadURL(ctx context.Context, uploadID string, partNumber int, did string) (url string, err error) + // GetPartUploadURL returns structured upload info (URL + optional headers) for a specific part + GetPartUploadURL(ctx context.Context, uploadID string, partNumber int, did string) (*PartUploadInfo, error) // CompleteMultipartUpload finalizes a multipart upload CompleteMultipartUpload(ctx context.Context, uploadID string, parts []PartInfo) error // AbortMultipartUpload cancels a multipart upload @@ -64,13 +65,22 @@ type PartInfo struct { ETag string `json:"etag"` } +// PartUploadInfo contains structured information for uploading a part +// Used for both S3 presigned URLs and buffered mode with headers +type PartUploadInfo struct { + URL string `json:"url"` // URL to PUT the part to + Method string `json:"method,omitempty"` // HTTP method (usually "PUT") + Headers map[string]string `json:"headers,omitempty"` // Additional headers required for the request +} + // NewXRPCHandler creates a new XRPC handler -func NewXRPCHandler(pds *HoldPDS, publicURL string, blobStore BlobStore, broadcaster *EventBroadcaster) *XRPCHandler { +func NewXRPCHandler(pds *HoldPDS, publicURL string, blobStore BlobStore, broadcaster *EventBroadcaster, httpClient HTTPClient) *XRPCHandler { return &XRPCHandler{ pds: pds, publicURL: publicURL, blobStore: blobStore, broadcaster: broadcaster, + httpClient: httpClient, } } @@ -78,8 +88,8 @@ func NewXRPCHandler(pds *HoldPDS, publicURL string, blobStore BlobStore, broadca func corsMiddleware(next http.HandlerFunc) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Access-Control-Allow-Origin", "*") - w.Header().Set("Access-Control-Allow-Methods", "GET, POST, OPTIONS") - w.Header().Set("Access-Control-Allow-Headers", "Content-Type, Authorization") + w.Header().Set("Access-Control-Allow-Methods", "GET, HEAD, POST, PUT, OPTIONS") + w.Header().Set("Access-Control-Allow-Headers", "Content-Type, Authorization, DPoP, X-Upload-Id, X-Part-Number, X-ATCR-DID") // Handle preflight OPTIONS requests if r.Method == http.MethodOptions { @@ -444,7 +454,7 @@ func (h *XRPCHandler) HandleDeleteRecord(w http.ResponseWriter, r *http.Request) } // Validate DPoP + OAuth and check authorization - _, err := ValidateOwnerOrCrewAdmin(r, h.pds) + _, err := ValidateOwnerOrCrewAdmin(r, h.pds, h.httpClient) if err != nil { http.Error(w, fmt.Sprintf("unauthorized: %v", err), http.StatusForbidden) return @@ -733,16 +743,16 @@ func (h *XRPCHandler) HandleUploadBlob(w http.ResponseWriter, r *http.Request) { // Mode 3: Direct blob upload (ATProto-compliant) // Receives raw bytes, computes CID, stores via distribution driver - // TODO: Authentication check - - // Extract DID for ATProto blob storage (per-DID paths) - did := r.URL.Query().Get("did") - if did == "" { - // TODO: Extract from auth context when authentication is implemented - // For now, use hold's DID as fallback - did = h.pds.DID() + // Requires admin-level access (captain or crew admin) + user, err := ValidateOwnerOrCrewAdmin(r, h.pds, h.httpClient) + if err != nil { + http.Error(w, fmt.Sprintf("authorization failed: %v", err), http.StatusForbidden) + return } + // Use authenticated user's DID for ATProto blob storage (per-DID paths) + did := user.DID + // Upload blob directly - blobStore will compute CID and store blobCID, size, err := h.blobStore.UploadBlob(r.Context(), did, r.Body) if err != nil { @@ -770,6 +780,14 @@ func (h *XRPCHandler) HandleUploadBlob(w http.ResponseWriter, r *http.Request) { func (h *XRPCHandler) handleBufferedPartUpload(w http.ResponseWriter, r *http.Request, uploadID, partNumberStr string) { ctx := r.Context() + // Validate blob write access + // This checks DPoP + OAuth tokens and verifies user is captain or crew with blob:write permission + _, err := ValidateBlobWriteAccess(r, h.pds, h.httpClient) + if err != nil { + http.Error(w, fmt.Sprintf("authorization failed: %v", err), http.StatusForbidden) + return + } + // Parse part number partNumber, err := strconv.Atoi(partNumberStr) if err != nil { @@ -816,6 +834,14 @@ func (h *XRPCHandler) handleMultipartOperation(w http.ResponseWriter, r *http.Re return } + // Validate blob write access for all multipart operations + // This checks DPoP + OAuth tokens and verifies user is captain or crew with blob:write permission + user, err := ValidateBlobWriteAccess(r, h.pds, h.httpClient) + if err != nil { + http.Error(w, fmt.Sprintf("authorization failed: %v", err), http.StatusForbidden) + return + } + // Route based on action switch req.Action { case "start": @@ -844,22 +870,14 @@ func (h *XRPCHandler) handleMultipartOperation(w http.ResponseWriter, r *http.Re return } - // Extract DID from query or header (for authorization) - did := r.URL.Query().Get("did") - if did == "" { - did = r.Header.Get("X-ATCR-DID") - } - - url, err := h.blobStore.GetPartUploadURL(ctx, req.UploadID, req.PartNumber, did) + uploadInfo, err := h.blobStore.GetPartUploadURL(ctx, req.UploadID, req.PartNumber, user.DID) if err != nil { http.Error(w, fmt.Sprintf("failed to get part URL: %v", err), http.StatusInternalServerError) return } w.Header().Set("Content-Type", "application/json") - json.NewEncoder(w).Encode(map[string]any{ - "url": url, - }) + json.NewEncoder(w).Encode(uploadInfo) case "complete": // Complete multipart upload @@ -902,6 +920,7 @@ func (h *XRPCHandler) handleMultipartOperation(w http.ResponseWriter, r *http.Re // HandleGetBlob wraps existing presigned download URL logic // Supports both ATProto CIDs and OCI sha256 digests +// Authorization: If captain.public = true, open to all. If false, requires crew with blob:read permission. func (h *XRPCHandler) HandleGetBlob(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodGet && r.Method != http.MethodHead { http.Error(w, "method not allowed", http.StatusMethodNotAllowed) @@ -921,6 +940,15 @@ func (h *XRPCHandler) HandleGetBlob(w http.ResponseWriter, r *http.Request) { return } + // Validate blob read access + // If captain.public = true, returns nil (public access allowed) + // If captain.public = false, validates auth and checks for blob:read permission + _, err := ValidateBlobReadAccess(r, h.pds, h.httpClient) + if err != nil { + http.Error(w, fmt.Sprintf("authorization failed: %v", err), http.StatusForbidden) + return + } + // Flexible digest parsing: accept both CID and sha256 digest formats var digest string if strings.HasPrefix(cidOrDigest, "sha256:") { @@ -1035,7 +1063,7 @@ func (h *XRPCHandler) HandleRequestCrew(w http.ResponseWriter, r *http.Request) } // Validate DPoP + OAuth token from Authorization and DPoP headers - user, err := ValidateDPoPRequest(r) + user, err := ValidateDPoPRequest(r, h.httpClient) if err != nil { http.Error(w, fmt.Sprintf("authentication failed: %v", err), http.StatusUnauthorized) return diff --git a/pkg/hold/pds/xrpc_multipart_test.go b/pkg/hold/pds/xrpc_multipart_test.go index 93396f0..ac3c7d7 100644 --- a/pkg/hold/pds/xrpc_multipart_test.go +++ b/pkg/hold/pds/xrpc_multipart_test.go @@ -7,6 +7,18 @@ import ( "testing" ) +// addTestDPoPAuth adds DPoP authentication headers to a request for testing +func addTestDPoPAuth(t *testing.T, req *http.Request, did string) { + t.Helper() + dpopHelper, err := NewDPoPTestHelper(did, "https://test-pds.example.com") + if err != nil { + t.Fatalf("Failed to create DPoP helper: %v", err) + } + if err := dpopHelper.AddDPoPToRequest(req); err != nil { + t.Fatalf("Failed to add DPoP to request: %v", err) + } +} + // ATCR-Specific Tests: Non-standard multipart upload extensions // // This file contains tests for ATCR's custom multipart upload extensions @@ -28,15 +40,17 @@ func TestHandleUploadBlob_MultipartStart(t *testing.T) { } req := makeXRPCPostRequest("/xrpc/com.atproto.repo.uploadBlob", body) + // Add DPoP authentication - owner has blob:write permission + addTestDPoPAuth(t, req, "did:plc:testowner123") w := httptest.NewRecorder() handler.HandleUploadBlob(w, req) + // Should return 200 OK with upload metadata if w.Code != http.StatusOK { - t.Errorf("Expected status 200, got %d", w.Code) + t.Errorf("Expected status 200 OK, got %d", w.Code) } - // Verify response contains uploadId and mode result := assertJSONResponse(t, w, http.StatusOK) if uploadID, ok := result["uploadId"].(string); !ok || uploadID == "" { @@ -62,6 +76,7 @@ func TestHandleUploadBlob_MultipartStart_MissingDigest(t *testing.T) { } req := makeXRPCPostRequest("/xrpc/com.atproto.repo.uploadBlob", body) + addTestDPoPAuth(t, req, "did:plc:testowner123") w := httptest.NewRecorder() handler.HandleUploadBlob(w, req) @@ -80,7 +95,7 @@ func TestHandleUploadBlob_MultipartPart(t *testing.T) { uploadID := "test-upload-123" partNumber := 1 - did := "did:plc:testuser" + expectedDID := "did:plc:testowner123" // DID from authenticated user body := map[string]any{ "action": "part", @@ -88,30 +103,31 @@ func TestHandleUploadBlob_MultipartPart(t *testing.T) { "partNumber": partNumber, } - req := makeXRPCPostRequest("/xrpc/com.atproto.repo.uploadBlob?did="+did, body) + req := makeXRPCPostRequest("/xrpc/com.atproto.repo.uploadBlob", body) + addTestDPoPAuth(t, req, expectedDID) w := httptest.NewRecorder() handler.HandleUploadBlob(w, req) + // Should return 200 OK with presigned URL if w.Code != http.StatusOK { - t.Errorf("Expected status 200, got %d", w.Code) + t.Errorf("Expected status 200 OK, got %d", w.Code) } - // Verify response contains URL result := assertJSONResponse(t, w, http.StatusOK) if url, ok := result["url"].(string); !ok || url == "" { t.Error("Expected url string in response") } - // Verify blob store was called + // Verify blob store was called with authenticated user's DID if len(blobStore.partURLCalls) != 1 { t.Fatalf("Expected GetPartUploadURL to be called once") } call := blobStore.partURLCalls[0] - if call.uploadID != uploadID || call.partNumber != partNumber || call.did != did { + if call.uploadID != uploadID || call.partNumber != partNumber || call.did != expectedDID { t.Errorf("Expected GetPartUploadURL(%s, %d, %s), got (%s, %d, %s)", - uploadID, partNumber, did, call.uploadID, call.partNumber, call.did) + uploadID, partNumber, expectedDID, call.uploadID, call.partNumber, call.did) } } @@ -150,6 +166,7 @@ func TestHandleUploadBlob_MultipartPart_MissingParams(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { req := makeXRPCPostRequest("/xrpc/com.atproto.repo.uploadBlob", tt.body) + addTestDPoPAuth(t, req, "did:plc:testowner123") w := httptest.NewRecorder() handler.HandleUploadBlob(w, req) @@ -181,15 +198,16 @@ func TestHandleUploadBlob_MultipartComplete(t *testing.T) { } req := makeXRPCPostRequest("/xrpc/com.atproto.repo.uploadBlob", body) + addTestDPoPAuth(t, req, "did:plc:testowner123") w := httptest.NewRecorder() handler.HandleUploadBlob(w, req) + // Should return 200 OK with completion status if w.Code != http.StatusOK { - t.Errorf("Expected status 200, got %d", w.Code) + t.Errorf("Expected status 200 OK, got %d", w.Code) } - // Verify response result := assertJSONResponse(t, w, http.StatusOK) if status, ok := result["status"].(string); !ok || status != "completed" { @@ -237,6 +255,7 @@ func TestHandleUploadBlob_MultipartComplete_MissingParams(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { req := makeXRPCPostRequest("/xrpc/com.atproto.repo.uploadBlob", tt.body) + addTestDPoPAuth(t, req, "did:plc:testowner123") w := httptest.NewRecorder() handler.HandleUploadBlob(w, req) @@ -263,15 +282,16 @@ func TestHandleUploadBlob_MultipartAbort(t *testing.T) { } req := makeXRPCPostRequest("/xrpc/com.atproto.repo.uploadBlob", body) + addTestDPoPAuth(t, req, "did:plc:testowner123") w := httptest.NewRecorder() handler.HandleUploadBlob(w, req) + // Should return 200 OK with abort status if w.Code != http.StatusOK { - t.Errorf("Expected status 200, got %d", w.Code) + t.Errorf("Expected status 200 OK, got %d", w.Code) } - // Verify response result := assertJSONResponse(t, w, http.StatusOK) if status, ok := result["status"].(string); !ok || status != "aborted" { @@ -293,6 +313,7 @@ func TestHandleUploadBlob_MultipartAbort_MissingUploadID(t *testing.T) { } req := makeXRPCPostRequest("/xrpc/com.atproto.repo.uploadBlob", body) + addTestDPoPAuth(t, req, "did:plc:testowner123") w := httptest.NewRecorder() handler.HandleUploadBlob(w, req) @@ -316,15 +337,16 @@ func TestHandleUploadBlob_BufferedPartUpload(t *testing.T) { req := httptest.NewRequest(http.MethodPut, "/xrpc/com.atproto.repo.uploadBlob", bytes.NewReader(data)) req.Header.Set("X-Upload-Id", uploadID) req.Header.Set("X-Part-Number", partNumber) + addTestDPoPAuth(t, req, "did:plc:testowner123") w := httptest.NewRecorder() handler.HandleUploadBlob(w, req) + // Should return 200 OK with ETag if w.Code != http.StatusOK { - t.Errorf("Expected status 200, got %d", w.Code) + t.Errorf("Expected status 200 OK, got %d", w.Code) } - // Verify response contains ETag result := assertJSONResponse(t, w, http.StatusOK) if etag, ok := result["etag"].(string); !ok || etag == "" { @@ -347,11 +369,11 @@ func TestHandleUploadBlob_BufferedPartUpload_MissingHeaders(t *testing.T) { handler, _, _ := setupTestXRPCHandlerWithBlobs(t) tests := []struct { - name string - uploadID string - partNumber string - setUploadID bool - setPartNumber bool + name string + uploadID string + partNumber string + setUploadID bool + setPartNumber bool }{ { name: "missing both headers", @@ -381,6 +403,7 @@ func TestHandleUploadBlob_BufferedPartUpload_MissingHeaders(t *testing.T) { if tt.setPartNumber { req.Header.Set("X-Part-Number", tt.partNumber) } + addTestDPoPAuth(t, req, "did:plc:testowner123") w := httptest.NewRecorder() handler.HandleUploadBlob(w, req) @@ -399,6 +422,7 @@ func TestHandleUploadBlob_BufferedPartUpload_InvalidPartNumber(t *testing.T) { req := httptest.NewRequest(http.MethodPut, "/xrpc/com.atproto.repo.uploadBlob", bytes.NewReader([]byte("data"))) req.Header.Set("X-Upload-Id", "test-123") req.Header.Set("X-Part-Number", "not-a-number") + addTestDPoPAuth(t, req, "did:plc:testowner123") w := httptest.NewRecorder() handler.HandleUploadBlob(w, req) @@ -417,6 +441,7 @@ func TestHandleUploadBlob_UnknownAction(t *testing.T) { } req := makeXRPCPostRequest("/xrpc/com.atproto.repo.uploadBlob", body) + addTestDPoPAuth(t, req, "did:plc:testowner123") w := httptest.NewRecorder() handler.HandleUploadBlob(w, req) diff --git a/pkg/hold/pds/xrpc_test.go b/pkg/hold/pds/xrpc_test.go index f89e199..f930876 100644 --- a/pkg/hold/pds/xrpc_test.go +++ b/pkg/hold/pds/xrpc_test.go @@ -54,8 +54,11 @@ func setupTestXRPCHandler(t *testing.T) (*XRPCHandler, context.Context) { t.Fatalf("Failed to bootstrap PDS: %v", err) } - // Create XRPC handler - handler := NewXRPCHandler(pds, "https://hold.example.com", nil, nil) + // Create mock PDS client for DPoP validation + mockClient := &mockPDSClient{} + + // Create XRPC handler with mock HTTP client + handler := NewXRPCHandler(pds, "https://hold.example.com", nil, nil, mockClient) return handler, ctx } @@ -652,7 +655,8 @@ func TestHandleListRecords_InvalidLimit(t *testing.T) { // TestHandleListRecords_EmptyCollection tests listing empty collection func TestHandleListRecords_EmptyCollection(t *testing.T) { pds, ctx := setupTestPDS(t) // Don't bootstrap - no records created yet - handler := NewXRPCHandler(pds, "https://hold.example.com", nil, nil) + mockClient := &mockPDSClient{} + handler := NewXRPCHandler(pds, "https://hold.example.com", nil, nil, mockClient) // Initialize repo manually (setupTestPDS doesn't call Bootstrap, so no crew members) err := pds.repomgr.InitNewActor(ctx, pds.uid, "", pds.did, "", "", "") @@ -744,31 +748,37 @@ func TestHandleDeleteRecord(t *testing.T) { } req := makeXRPCPostRequest("/xrpc/com.atproto.repo.deleteRecord", body) + + // Add DPoP authentication - owner has admin permission to delete crew + ownerDID := "did:plc:testowner123" + dpopHelper, err := NewDPoPTestHelper(ownerDID, "https://test-pds.example.com") + if err != nil { + t.Fatalf("Failed to create DPoP helper: %v", err) + } + if err := dpopHelper.AddDPoPToRequest(req); err != nil { + t.Fatalf("Failed to add DPoP to request: %v", err) + } + w := httptest.NewRecorder() - // Note: This test will fail auth check since we're not providing DPoP tokens - // For now, we're testing the request parsing and response structure - // A real implementation would need proper auth mocking handler.HandleDeleteRecord(w, req) - // We expect 403 Forbidden due to missing auth - // This tests that the endpoint is parsing JSON body correctly - if w.Code != http.StatusForbidden { - // If somehow auth passes (shouldn't in this test), verify response structure - if w.Code == http.StatusOK { - result := assertJSONResponse(t, w, http.StatusOK) + // Should return 200 OK with commit metadata + if w.Code != http.StatusOK { + t.Errorf("Expected status 200 OK, got %d", w.Code) + } - // Per spec, response should have commit object - if commit, ok := result["commit"].(map[string]any); !ok { - t.Error("Expected commit object in response") - } else { - if cid, ok := commit["cid"].(string); !ok || cid == "" { - t.Error("Expected cid in commit object") - } - if rev, ok := commit["rev"].(string); !ok || rev == "" { - t.Error("Expected rev in commit object") - } - } + result := assertJSONResponse(t, w, http.StatusOK) + + // Per spec, response should have commit object + if commit, ok := result["commit"].(map[string]any); !ok { + t.Error("Expected commit object in response") + } else { + if cid, ok := commit["cid"].(string); !ok || cid == "" { + t.Error("Expected cid in commit object") + } + if rev, ok := commit["rev"].(string); !ok || rev == "" { + t.Error("Expected rev in commit object") } } } @@ -902,7 +912,8 @@ func TestHandleListRepos(t *testing.T) { // Spec: https://docs.bsky.app/docs/api/com-atproto-sync-list-repos func TestHandleListRepos_EmptyRepo(t *testing.T) { pds, ctx := setupTestPDS(t) // Don't bootstrap - handler := NewXRPCHandler(pds, "https://hold.example.com", nil, nil) + mockClient := &mockPDSClient{} + handler := NewXRPCHandler(pds, "https://hold.example.com", nil, nil, mockClient) // setupTestPDS creates the PDS/database but doesn't initialize the repo // Check if implementation returns repos before initialization @@ -1334,14 +1345,14 @@ type mockBlobStore struct { partUploadError error // Track calls - downloadCalls []string // Track digests requested for download - uploadCalls []string // Track digests requested for upload - uploadBlobCalls []uploadBlobCall // Track direct blob uploads - startCalls []string // Track digests for multipart start - partURLCalls []partURLCall - completeCalls []string - abortCalls []string - partUploadCalls []partUploadCall + downloadCalls []string // Track digests requested for download + uploadCalls []string // Track digests requested for upload + uploadBlobCalls []uploadBlobCall // Track direct blob uploads + startCalls []string // Track digests for multipart start + partURLCalls []partURLCall + completeCalls []string + abortCalls []string + partUploadCalls []partUploadCall } type uploadBlobCall struct { @@ -1419,12 +1430,15 @@ func (m *mockBlobStore) StartMultipartUpload(ctx context.Context, digest string) return "test-upload-id", "s3native", nil } -func (m *mockBlobStore) GetPartUploadURL(ctx context.Context, uploadID string, partNumber int, did string) (string, error) { +func (m *mockBlobStore) GetPartUploadURL(ctx context.Context, uploadID string, partNumber int, did string) (*PartUploadInfo, error) { m.partURLCalls = append(m.partURLCalls, partURLCall{uploadID, partNumber, did}) if m.partURLError != nil { - return "", m.partURLError + return nil, m.partURLError } - return "https://s3.example.com/part/" + uploadID, nil + return &PartUploadInfo{ + URL: "https://s3.example.com/part/" + uploadID, + Method: "PUT", + }, nil } func (m *mockBlobStore) CompleteMultipartUpload(ctx context.Context, uploadID string, parts []PartInfo) error { @@ -1451,7 +1465,7 @@ func (m *mockBlobStore) HandleBufferedPartUpload(ctx context.Context, uploadID s return "test-etag-" + uploadID, nil } -// setupTestXRPCHandlerWithBlobs creates handler with mock blob store +// setupTestXRPCHandlerWithBlobs creates handler with mock blob store and mock PDS client func setupTestXRPCHandlerWithBlobs(t *testing.T) (*XRPCHandler, *mockBlobStore, context.Context) { t.Helper() @@ -1488,8 +1502,11 @@ func setupTestXRPCHandlerWithBlobs(t *testing.T) (*XRPCHandler, *mockBlobStore, // Create mock blob store blobStore := newMockBlobStore() - // Create XRPC handler with mock blob store - handler := NewXRPCHandler(pds, "https://hold.example.com", blobStore, nil) + // Create mock PDS client for DPoP validation + mockClient := &mockPDSClient{} + + // Create XRPC handler with mock blob store and mock HTTP client + handler := NewXRPCHandler(pds, "https://hold.example.com", blobStore, nil, mockClient) return handler, blobStore, ctx } @@ -1507,6 +1524,17 @@ func TestHandleUploadBlob(t *testing.T) { // Test standard single blob upload (POST with raw bytes) req := httptest.NewRequest(http.MethodPost, "/xrpc/com.atproto.repo.uploadBlob", bytes.NewReader(blobData)) req.Header.Set("Content-Type", "application/octet-stream") + + // Add DPoP authentication - owner has admin permission for blob upload + ownerDID := "did:plc:testowner123" + dpopHelper, err := NewDPoPTestHelper(ownerDID, "https://test-pds.example.com") + if err != nil { + t.Fatalf("Failed to create DPoP helper: %v", err) + } + if err := dpopHelper.AddDPoPToRequest(req); err != nil { + t.Fatalf("Failed to add DPoP to request: %v", err) + } + w := httptest.NewRecorder() handler.HandleUploadBlob(w, req) @@ -1559,13 +1587,24 @@ func TestHandleUploadBlob_EmptyBody(t *testing.T) { // Empty blob should succeed (edge case) req := httptest.NewRequest(http.MethodPost, "/xrpc/com.atproto.repo.uploadBlob", bytes.NewReader([]byte{})) req.Header.Set("Content-Type", "application/octet-stream") + + // Add DPoP authentication + ownerDID := "did:plc:testowner123" + dpopHelper, err := NewDPoPTestHelper(ownerDID, "https://test-pds.example.com") + if err != nil { + t.Fatalf("Failed to create DPoP helper: %v", err) + } + if err := dpopHelper.AddDPoPToRequest(req); err != nil { + t.Fatalf("Failed to add DPoP to request: %v", err) + } + w := httptest.NewRecorder() handler.HandleUploadBlob(w, req) - // Should succeed with empty blob + // Should return 200 OK for empty blob (edge case) if w.Code != http.StatusOK { - t.Errorf("Expected status 200, got %d", w.Code) + t.Errorf("Expected status 200 OK for empty blob, got %d", w.Code) } // Verify blob store was called with 0 bytes @@ -1600,12 +1639,24 @@ func TestHandleUploadBlob_BlobStoreError(t *testing.T) { req := httptest.NewRequest(http.MethodPost, "/xrpc/com.atproto.repo.uploadBlob", bytes.NewReader([]byte("test data"))) req.Header.Set("Content-Type", "application/octet-stream") + + // Add DPoP authentication + ownerDID := "did:plc:testowner123" + dpopHelper, err := NewDPoPTestHelper(ownerDID, "https://test-pds.example.com") + if err != nil { + t.Fatalf("Failed to create DPoP helper: %v", err) + } + if err := dpopHelper.AddDPoPToRequest(req); err != nil { + t.Fatalf("Failed to add DPoP to request: %v", err) + } + w := httptest.NewRecorder() handler.HandleUploadBlob(w, req) + // Should get 500 Internal Server Error for blob store error if w.Code != http.StatusInternalServerError { - t.Errorf("Expected status 500, got %d", w.Code) + t.Errorf("Expected status 500 for blob store error, got %d", w.Code) } } @@ -1775,3 +1826,49 @@ func TestHandleGetBlob_BlobStoreError(t *testing.T) { t.Errorf("Expected status 500, got %d", w.Code) } } + +// TestHandleGetBlobCORSHeaders tests that CORS headers are set for blob downloads +// // Spec: https://docs.bsky.app/docs/api/com-atproto-sync-get-blob +func TestHandleGetBlob_CORSHeaders(t *testing.T) { + handler, _, ctx := setupTestXRPCHandlerWithBlobs(t) + + // Make hold public + _, err := handler.pds.UpdateCaptainRecord(ctx, true, false) + if err != nil { + t.Fatalf("Failed to update captain: %v", err) + } + + holdDID := "did:web:hold.example.com" + cid := "bafyreib2rxk3rkhh5ylyxj3x3gathxt3s32qvwj2lf3qg4kmzr6b7teqke" + url := fmt.Sprintf("/xrpc/com.atproto.sync.getBlob?did=%s&cid=%s", holdDID, cid) + + // Test GET request + req := httptest.NewRequest(http.MethodGet, url, nil) + w := httptest.NewRecorder() + + // Wrap with CORS middleware + corsHandler := corsMiddleware(handler.HandleGetBlob) + corsHandler(w, req) + + // Verify CORS headers are present + if origin := w.Header().Get("Access-Control-Allow-Origin"); origin != "*" { + t.Errorf("Expected Access-Control-Allow-Origin: *, got %s", origin) + } + + // Test OPTIONS preflight + req2 := httptest.NewRequest(http.MethodOptions, url, nil) + w2 := httptest.NewRecorder() + + corsHandler(w2, req2) + + if w2.Code != http.StatusOK { + t.Errorf("Expected OPTIONS to return 200, got %d", w2.Code) + } + + methods := w2.Header().Get("Access-Control-Allow-Methods") + if !strings.Contains(methods, "GET") || !strings.Contains(methods, "HEAD") { + t.Errorf("Expected Access-Control-Allow-Methods to include GET and HEAD, got %s", methods) + } + + t.Logf("✓ CORS headers correctly set for blob downloads") +} diff --git a/pkg/hold/service.go b/pkg/hold/service.go index dff3665..eff7413 100644 --- a/pkg/hold/service.go +++ b/pkg/hold/service.go @@ -22,14 +22,23 @@ type HoldPDSInterface interface { type HoldService struct { driver storagedriver.StorageDriver config *Config - s3Client *s3.S3 // S3 client for presigned URLs (nil if not S3 storage) - bucket string // S3 bucket name - s3PathPrefix string // S3 path prefix (if any) - MultipartMgr *MultipartManager // Exported for access in route handlers - pds HoldPDSInterface // Embedded PDS for captain/crew records - authorizer auth.HoldAuthorizer // Authorizer for access control + s3Client *s3.S3 // S3 client for presigned URLs (nil if not S3 storage) + bucket string // S3 bucket name + s3PathPrefix string // S3 path prefix (if any) + MultipartMgr *MultipartManager // Exported for access in route handlers + pds HoldPDSInterface // Embedded PDS for captain/crew records + authorizer auth.HoldAuthorizer // Authorizer for access control } +// PresignedURLOperation defines the type of presigned URL operation +type PresignedURLOperation string + +const ( + OperationGet PresignedURLOperation = "GET" + OperationHead PresignedURLOperation = "HEAD" + OperationPut PresignedURLOperation = "PUT" +) + // NewHoldService creates a new hold service // holdPDS must be a *pds.HoldPDS but we use any to avoid import cycle func NewHoldService(cfg *Config, holdPDS any) (*HoldService, error) { diff --git a/pkg/hold/storage.go b/pkg/hold/storage.go index 312d1a6..5e7cdca 100644 --- a/pkg/hold/storage.go +++ b/pkg/hold/storage.go @@ -77,8 +77,12 @@ func (s *HoldService) getPresignedURL(ctx context.Context, operation PresignedUR // Check if presigned URLs are disabled if s.config.Server.DisablePresignedURLs { - log.Printf("Presigned URLs disabled, using proxy URL") - return s.getProxyURL(digest, did), nil + log.Printf("Presigned URLs disabled, using XRPC endpoint") + url := s.getProxyURL(digest, did, operation) + if url == "" { + return "", fmt.Errorf("XRPC proxy not supported for PUT operations - use multipart upload") + } + return url, nil } // Generate presigned URL if S3 client is available @@ -122,19 +126,50 @@ func (s *HoldService) getPresignedURL(ctx context.Context, operation PresignedUR url, err := req.Presign(15 * time.Minute) if err != nil { log.Printf("[getPresignedURL] Presign FAILED for %s: %v", operation, err) - log.Printf(" Falling back to proxy URL") - return s.getProxyURL(digest, did), nil + log.Printf(" Falling back to XRPC endpoint") + proxyURL := s.getProxyURL(digest, did, operation) + if proxyURL == "" { + return "", fmt.Errorf("presign failed and XRPC proxy not supported for PUT operations") + } + return proxyURL, nil } return url, nil } - // Fallback: return proxy URL through this service - return s.getProxyURL(digest, did), nil + // Fallback: return XRPC endpoint through this service + proxyURL := s.getProxyURL(digest, did, operation) + if proxyURL == "" { + return "", fmt.Errorf("S3 client not available and XRPC proxy not supported for PUT operations") + } + return proxyURL, nil } -// getProxyURL returns a proxy URL for blob operations (fallback when presigned URLs unavailable) -func (s *HoldService) getProxyURL(digest, did string) string { - // All operations use the same proxy endpoint - return fmt.Sprintf("%s/blobs/%s?did=%s", s.config.Server.PublicURL, digest, did) +// getProxyURL returns XRPC endpoint for blob operations (fallback when presigned URLs unavailable) +// For GET/HEAD operations, returns the XRPC getBlob endpoint +// For PUT operations, this fallback is no longer supported - use multipart upload instead +func (s *HoldService) getProxyURL(digest, did string, operation PresignedURLOperation) string { + // For read operations, use XRPC getBlob endpoint + if operation == OperationGet || operation == OperationHead { + // Generate hold DID from public URL + holdDID := s.getHoldDID() + return fmt.Sprintf("%s/xrpc/com.atproto.sync.getBlob?did=%s&cid=%s", + s.config.Server.PublicURL, holdDID, digest) + } + + // For PUT operations, proxy fallback is not supported with XRPC + // Clients should use multipart upload flow via com.atproto.repo.uploadBlob + return "" +} + +// getHoldDID generates a did:web from the hold's public URL +func (s *HoldService) getHoldDID() string { + // Convert URL to did:web format + // https://hold01.atcr.io → did:web:hold01.atcr.io + url := s.config.Server.PublicURL + url = strings.TrimPrefix(url, "https://") + url = strings.TrimPrefix(url, "http://") + url = strings.Split(url, "/")[0] // Remove path + url = strings.Split(url, ":")[0] // Remove port + return fmt.Sprintf("did:web:%s", url) }