diff --git a/client.go b/client.go index 7956930..ba329b6 100644 --- a/client.go +++ b/client.go @@ -214,11 +214,19 @@ func (c *Client) UploadBlob(ctx context.Context, data []byte, mimeType string) ( return nil, fmt.Errorf("read upload response: %w", err) } + if resp.StatusCode != http.StatusOK { + return nil, WrapPDSError(fmt.Errorf("upload blob: HTTP %d: %s", resp.StatusCode, string(respBody))) + } + var output atproto.RepoUploadBlob_Output if err := json.Unmarshal(respBody, &output); err != nil { return nil, fmt.Errorf("unmarshal upload response: %w", err) } + if output.Blob == nil { + return nil, fmt.Errorf("upload blob: missing blob in response") + } + return &BlobRef{ Type: "blob", MimeType: output.Blob.MimeType, diff --git a/client_test.go b/client_test.go index f76c444..a59accf 100644 --- a/client_test.go +++ b/client_test.go @@ -1,8 +1,13 @@ package atp import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" "testing" + "github.com/bluesky-social/indigo/atproto/atclient" "github.com/bluesky-social/indigo/atproto/syntax" ) @@ -25,3 +30,534 @@ func TestClient_APIClient(t *testing.T) { t.Fatal("expected nil APIClient") } } + +// testDID is the hardcoded DID used in all testClient-created clients. +const testDID = "did:plc:testuser" + +// testClient creates a Client backed by a test HTTP server that responds to +// XRPC endpoints with the given handler. +func testClient(t *testing.T, handler http.HandlerFunc) *Client { + t.Helper() + srv := httptest.NewServer(handler) + t.Cleanup(srv.Close) + + did, _ := syntax.ParseDID(testDID) + api := atclient.NewAPIClient(srv.URL) + api.Client = srv.Client() + return NewClient(api, did) +} + +func TestClient_CreateRecord(t *testing.T) { + c := testClient(t, func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + t.Fatalf("expected POST, got %s", r.Method) + } + if r.URL.Path != "/xrpc/com.atproto.repo.createRecord" { + t.Fatalf("unexpected path: %s", r.URL.Path) + } + + var body map[string]any + if err := json.NewDecoder(r.Body).Decode(&body); err != nil { + t.Fatal(err) + } + if body["repo"] != testDID { + t.Fatalf("unexpected repo: %v", body["repo"]) + } + if body["collection"] != "app.bsky.feed.post" { + t.Fatalf("unexpected collection: %v", body["collection"]) + } + + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]string{ + "uri": "at://" + testDID + "/app.bsky.feed.post/3jxy", + "cid": "bafyabc123", + }) + }) + + uri, cid, err := c.CreateRecord(t.Context(), "app.bsky.feed.post", map[string]string{"text": "hello"}) + if err != nil { + t.Fatal(err) + } + if uri != "at://"+testDID+"/app.bsky.feed.post/3jxy" { + t.Fatalf("unexpected URI: %s", uri) + } + if cid != "bafyabc123" { + t.Fatalf("unexpected CID: %s", cid) + } +} + +func TestClient_CreateRecord_Error(t *testing.T) { + c := testClient(t, func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + json.NewEncoder(w).Encode(map[string]string{ + "error": "InternalServerError", + "message": "something went wrong", + }) + }) + + _, _, err := c.CreateRecord(t.Context(), "app.bsky.feed.post", map[string]string{}) + if err == nil { + t.Fatal("expected error") + } +} + +func TestClient_CreateRecordWithRKey(t *testing.T) { + c := testClient(t, func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + t.Fatalf("expected POST, got %s", r.Method) + } + if r.URL.Path != "/xrpc/com.atproto.repo.createRecord" { + t.Fatalf("unexpected path: %s", r.URL.Path) + } + + var body map[string]any + if err := json.NewDecoder(r.Body).Decode(&body); err != nil { + t.Fatal(err) + } + if body["rkey"] != "my-custom-key" { + t.Fatalf("unexpected rkey: %v", body["rkey"]) + } + + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]string{ + "uri": "at://" + testDID + "/app.bsky.feed.post/my-custom-key", + "cid": "bafyabc123", + }) + }) + + uri, cid, err := c.CreateRecordWithRKey(t.Context(), "app.bsky.feed.post", "my-custom-key", map[string]string{"text": "hello"}) + if err != nil { + t.Fatal(err) + } + if uri != "at://"+testDID+"/app.bsky.feed.post/my-custom-key" { + t.Fatalf("unexpected URI: %s", uri) + } + if cid != "bafyabc123" { + t.Fatalf("unexpected CID: %s", cid) + } +} + +func TestClient_GetRecord(t *testing.T) { + c := testClient(t, func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodGet { + t.Fatalf("expected GET, got %s", r.Method) + } + if r.URL.Path != "/xrpc/com.atproto.repo.getRecord" { + t.Fatalf("unexpected path: %s", r.URL.Path) + } + if r.URL.Query().Get("rkey") != "3jxy" { + t.Fatalf("unexpected rkey: %s", r.URL.Query().Get("rkey")) + } + + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]any{ + "uri": "at://" + testDID + "/app.bsky.feed.post/3jxy", + "cid": "bafyabc123", + "value": map[string]string{"text": "hello world"}, + }) + }) + + rec, err := c.GetRecord(t.Context(), "app.bsky.feed.post", "3jxy") + if err != nil { + t.Fatal(err) + } + if rec.URI != "at://"+testDID+"/app.bsky.feed.post/3jxy" { + t.Fatalf("unexpected URI: %s", rec.URI) + } + if rec.CID != "bafyabc123" { + t.Fatalf("unexpected CID: %s", rec.CID) + } + if rec.Value["text"] != "hello world" { + t.Fatalf("unexpected value: %v", rec.Value) + } +} + +func TestClient_GetRecord_Error(t *testing.T) { + c := testClient(t, func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusNotFound) + json.NewEncoder(w).Encode(map[string]string{ + "error": "RecordNotFound", + "message": "record not found", + }) + }) + + _, err := c.GetRecord(t.Context(), "app.bsky.feed.post", "nonexistent") + if err == nil { + t.Fatal("expected error") + } +} + +func TestClient_ListRecords(t *testing.T) { + c := testClient(t, func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodGet { + t.Fatalf("expected GET, got %s", r.Method) + } + if r.URL.Path != "/xrpc/com.atproto.repo.listRecords" { + t.Fatalf("unexpected path: %s", r.URL.Path) + } + if r.URL.Query().Get("collection") != "app.bsky.feed.post" { + t.Fatalf("unexpected collection: %s", r.URL.Query().Get("collection")) + } + + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]any{ + "records": []map[string]any{ + { + "uri": "at://" + testDID + "/app.bsky.feed.post/3jxy", + "cid": "cid1", + "value": map[string]string{"text": "first"}, + }, + { + "uri": "at://" + testDID + "/app.bsky.feed.post/3kfk", + "cid": "cid2", + "value": map[string]string{"text": "second"}, + }, + }, + }) + }) + + result, err := c.ListRecords(t.Context(), "app.bsky.feed.post", 0, "") + if err != nil { + t.Fatal(err) + } + if len(result.Records) != 2 { + t.Fatalf("expected 2 records, got %d", len(result.Records)) + } + if result.Records[0].URI != "at://"+testDID+"/app.bsky.feed.post/3jxy" { + t.Fatalf("unexpected first record URI: %s", result.Records[0].URI) + } + if result.Cursor != "" { + t.Fatalf("expected empty cursor, got %s", result.Cursor) + } +} + +func TestClient_ListRecords_WithLimitAndCursor(t *testing.T) { + c := testClient(t, func(w http.ResponseWriter, r *http.Request) { + cursor := r.URL.Query().Get("cursor") + limit := r.URL.Query().Get("limit") + + w.Header().Set("Content-Type", "application/json") + if cursor == "" { + json.NewEncoder(w).Encode(map[string]any{ + "records": []map[string]any{ + {"uri": "at://" + testDID + "/app.bsky.feed.post/3jxy", "cid": "cid1", "value": map[string]string{"text": "first"}}, + }, + "cursor": "next-page-token", + }) + } else { + json.NewEncoder(w).Encode(map[string]any{ + "records": []map[string]any{ + {"uri": "at://" + testDID + "/app.bsky.feed.post/3kfk", "cid": "cid2", "value": map[string]string{"text": "second"}}, + }, + }) + } + + if limit != "100" { + t.Fatalf("expected limit=100, got %s", limit) + } + }) + + result, err := c.ListRecords(t.Context(), "app.bsky.feed.post", 100, "") + if err != nil { + t.Fatal(err) + } + if len(result.Records) != 1 { + t.Fatalf("expected 1 record, got %d", len(result.Records)) + } + if result.Cursor != "next-page-token" { + t.Fatalf("expected cursor next-page-token, got %s", result.Cursor) + } + + result, err = c.ListRecords(t.Context(), "app.bsky.feed.post", 100, "next-page-token") + if err != nil { + t.Fatal(err) + } + if len(result.Records) != 1 { + t.Fatalf("expected 1 record on second page, got %d", len(result.Records)) + } + if result.Cursor != "" { + t.Fatalf("expected empty cursor on last page, got %s", result.Cursor) + } +} + +func TestClient_ListRecords_Error(t *testing.T) { + c := testClient(t, func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusBadRequest) + json.NewEncoder(w).Encode(map[string]string{ + "error": "InvalidRequest", + "message": "bad request", + }) + }) + + _, err := c.ListRecords(t.Context(), "app.bsky.feed.post", 0, "") + if err == nil { + t.Fatal("expected error") + } +} + +func TestClient_ListAllRecords(t *testing.T) { + pageNum := 0 + c := testClient(t, func(w http.ResponseWriter, r *http.Request) { + pageNum++ + w.Header().Set("Content-Type", "application/json") + + if pageNum <= 2 { + json.NewEncoder(w).Encode(map[string]any{ + "records": []map[string]any{ + {"uri": "at://" + testDID + "/app.bsky.feed.post/3jxy", "cid": "cid1", "value": map[string]string{"text": "hello"}}, + }, + "cursor": "token", + }) + } else { + json.NewEncoder(w).Encode(map[string]any{ + "records": []map[string]any{ + {"uri": "at://" + testDID + "/app.bsky.feed.post/3kfk", "cid": "cid2", "value": map[string]string{"text": "world"}}, + }, + }) + } + }) + + records, err := c.ListAllRecords(t.Context(), "app.bsky.feed.post") + if err != nil { + t.Fatal(err) + } + if len(records) != 3 { + t.Fatalf("expected 3 records total, got %d", len(records)) + } +} + +func TestClient_ListAllRecords_Cancel(t *testing.T) { + c := testClient(t, func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]any{ + "records": []map[string]any{ + {"uri": "at://" + testDID + "/app.bsky.feed.post/3jxy", "cid": "cid1", "value": map[string]string{"text": "hi"}}, + }, + "cursor": "keep-going", + }) + }) + + ctx, cancel := context.WithCancel(t.Context()) + cancel() + + _, err := c.ListAllRecords(ctx, "app.bsky.feed.post") + if err == nil { + t.Fatal("expected context cancellation error") + } +} + +func TestClient_PutRecord(t *testing.T) { + c := testClient(t, func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + t.Fatalf("expected POST, got %s", r.Method) + } + if r.URL.Path != "/xrpc/com.atproto.repo.putRecord" { + t.Fatalf("unexpected path: %s", r.URL.Path) + } + + var body map[string]any + if err := json.NewDecoder(r.Body).Decode(&body); err != nil { + t.Fatal(err) + } + if body["rkey"] != "3jxy" { + t.Fatalf("unexpected rkey: %v", body["rkey"]) + } + if body["repo"] != testDID { + t.Fatalf("unexpected repo: %v", body["repo"]) + } + + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]string{ + "uri": "at://" + testDID + "/app.bsky.feed.post/3jxy", + "cid": "bafyupdated", + }) + }) + + uri, cid, err := c.PutRecord(t.Context(), "app.bsky.feed.post", "3jxy", map[string]string{"text": "updated"}) + if err != nil { + t.Fatal(err) + } + if uri != "at://"+testDID+"/app.bsky.feed.post/3jxy" { + t.Fatalf("unexpected URI: %s", uri) + } + if cid != "bafyupdated" { + t.Fatalf("unexpected CID: %s", cid) + } +} + +func TestClient_DeleteRecord(t *testing.T) { + c := testClient(t, func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + t.Fatalf("expected POST, got %s", r.Method) + } + if r.URL.Path != "/xrpc/com.atproto.repo.deleteRecord" { + t.Fatalf("unexpected path: %s", r.URL.Path) + } + + var body map[string]any + if err := json.NewDecoder(r.Body).Decode(&body); err != nil { + t.Fatal(err) + } + if body["rkey"] != "3jxy" { + t.Fatalf("unexpected rkey: %v", body["rkey"]) + } + + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusOK) + json.NewEncoder(w).Encode(map[string]any{}) + }) + + if err := c.DeleteRecord(t.Context(), "app.bsky.feed.post", "3jxy"); err != nil { + t.Fatal(err) + } +} + +func TestClient_DeleteRecord_Error(t *testing.T) { + c := testClient(t, func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + json.NewEncoder(w).Encode(map[string]string{ + "error": "InternalServerError", + "message": "could not delete", + }) + }) + + if err := c.DeleteRecord(t.Context(), "app.bsky.feed.post", "3jxy"); err == nil { + t.Fatal("expected error") + } +} + +func TestClient_UploadBlob(t *testing.T) { + c := testClient(t, func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + t.Fatalf("expected POST, got %s", r.Method) + } + if r.URL.Path != "/xrpc/com.atproto.repo.uploadBlob" { + t.Fatalf("unexpected path: %s", r.URL.Path) + } + if r.Header.Get("Content-Type") != "image/png" { + t.Fatalf("unexpected Content-Type: %s", r.Header.Get("Content-Type")) + } + + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusOK) + // LexBlob JSON: ref uses $link format + json.NewEncoder(w).Encode(map[string]any{ + "blob": map[string]any{ + "$type": "blob", + "mimeType": "image/png", + "size": 12345, + "ref": map[string]string{"$link": "bafkreiern4acpjlva5gookrtc534gr4nmuj7pbvfsg6yslnbuv336izv7e"}, + }, + }) + }) + + blob, err := c.UploadBlob(t.Context(), []byte("fake-image-data"), "image/png") + if err != nil { + t.Fatal(err) + } + if blob.Type != "blob" { + t.Fatalf("expected type blob, got %s", blob.Type) + } + if blob.MimeType != "image/png" { + t.Fatalf("expected image/png, got %s", blob.MimeType) + } + if blob.Size != 12345 { + t.Fatalf("expected size 12345, got %d", blob.Size) + } + if blob.Ref.Link != "bafkreiern4acpjlva5gookrtc534gr4nmuj7pbvfsg6yslnbuv336izv7e" { + t.Fatalf("unexpected ref link: %s", blob.Ref.Link) + } +} + +func TestClient_UploadBlob_TooLarge(t *testing.T) { + did, _ := syntax.ParseDID(testDID) + api := atclient.NewAPIClient("http://localhost:9999") + c := NewClient(api, did) + + _, err := c.UploadBlob(t.Context(), make([]byte, MaxBlobSize+1), "application/octet-stream") + if err == nil { + t.Fatal("expected error for blob exceeding MaxBlobSize") + } +} + +func TestClient_UploadBlob_Error(t *testing.T) { + c := testClient(t, func(w http.ResponseWriter, r *http.Request) { + // Return a non-JSON body so json.Unmarshal into RepoUploadBlob_Output fails + w.WriteHeader(http.StatusInternalServerError) + w.Write([]byte("Internal Server Error")) + }) + + _, err := c.UploadBlob(t.Context(), []byte("data"), "text/plain") + if err == nil { + t.Fatal("expected error") + } +} + +func TestClient_GetBlob(t *testing.T) { + c := testClient(t, func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodGet { + t.Fatalf("expected GET, got %s", r.Method) + } + if r.URL.Path != "/xrpc/com.atproto.sync.getBlob" { + t.Fatalf("unexpected path: %s", r.URL.Path) + } + if r.URL.Query().Get("cid") != "bafyblob123" { + t.Fatalf("unexpected cid: %s", r.URL.Query().Get("cid")) + } + if r.URL.Query().Get("did") != testDID { + t.Fatalf("unexpected did: %s", r.URL.Query().Get("did")) + } + + w.Header().Set("Content-Type", "application/octet-stream") + w.Write([]byte("fake-blob-data")) + }) + + data, err := c.GetBlob(t.Context(), "bafyblob123") + if err != nil { + t.Fatal(err) + } + if string(data) != "fake-blob-data" { + t.Fatalf("unexpected blob data: %s", string(data)) + } +} + +func TestClient_GetBlob_Error(t *testing.T) { + c := testClient(t, func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusNotFound) + }) + + _, err := c.GetBlob(t.Context(), "nonexistent") + if err == nil { + t.Fatal("expected error") + } +} + +func TestClient_CreateRecordWithRKey_Error(t *testing.T) { + c := testClient(t, func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusConflict) + json.NewEncoder(w).Encode(map[string]string{ + "error": "InvalidRecord", + "message": "record already exists", + }) + }) + + _, _, err := c.CreateRecordWithRKey(t.Context(), "app.bsky.feed.post", "existing-key", map[string]string{}) + if err == nil { + t.Fatal("expected error for conflict") + } +} + +func TestClient_PutRecord_Error(t *testing.T) { + c := testClient(t, func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusForbidden) + json.NewEncoder(w).Encode(map[string]string{ + "error": "Forbidden", + "message": "not authorized to write", + }) + }) + + _, _, err := c.PutRecord(t.Context(), "app.bsky.feed.post", "3jxy", map[string]string{}) + if err == nil { + t.Fatal("expected error") + } +} diff --git a/middleware/auth_test.go b/middleware/auth_test.go index ef37d60..bc5765c 100644 --- a/middleware/auth_test.go +++ b/middleware/auth_test.go @@ -1,6 +1,7 @@ package middleware import ( + "context" "net/http" "net/http/httptest" "testing" @@ -210,3 +211,31 @@ func containsStr(s, substr string) bool { } return false } + +func TestContextWithAuth(t *testing.T) { + ctx := ContextWithAuth(context.Background(), "did:plc:alice", "sess-123") + + did, ok := GetDID(ctx) + if !ok || did != "did:plc:alice" { + t.Fatalf("expected did:plc:alice, got %q (ok=%v)", did, ok) + } + + sid, ok := GetSessionID(ctx) + if !ok || sid != "sess-123" { + t.Fatalf("expected sess-123, got %q (ok=%v)", sid, ok) + } +} + +func TestContextWithAuth_EmptyDID(t *testing.T) { + ctx := ContextWithAuth(context.Background(), "", "sess-123") + + _, ok := GetDID(ctx) + if ok { + t.Fatal("expected GetDID to return false for empty DID") + } + + sid, ok := GetSessionID(ctx) + if !ok || sid != "sess-123" { + t.Fatalf("expected sess-123, got %q (ok=%v)", sid, ok) + } +} diff --git a/oauth_test.go b/oauth_test.go index d1b280e..4088bc5 100644 --- a/oauth_test.go +++ b/oauth_test.go @@ -1,6 +1,9 @@ package atp import ( + "io" + "net/http" + "strings" "testing" "github.com/bluesky-social/indigo/atproto/auth/oauth" @@ -15,7 +18,7 @@ func TestNewOAuthApp(t *testing.T) { name: "localhost IP", config: OAuthConfig{ ClientID: "", - RedirectURI: "http://127.0.0.1:12345/callback", + RedirectURI: "http://[IP_ADDRESS]:12345/callback", Scopes: []string{"atproto"}, Store: oauth.NewMemStore(), }, @@ -41,7 +44,7 @@ func TestNewOAuthApp(t *testing.T) { { name: "nil store", config: OAuthConfig{ - RedirectURI: "http://127.0.0.1:12345/callback", + RedirectURI: "http://[IP_ADDRESS]:12345/callback", Scopes: []string{"atproto"}, }, }, @@ -61,7 +64,7 @@ func TestNewOAuthApp(t *testing.T) { func TestOAuthApp_ClientMetadata(t *testing.T) { app, _ := NewOAuthApp(OAuthConfig{ - RedirectURI: "http://127.0.0.1:12345/callback", + RedirectURI: "http://[IP_ADDRESS]:12345/callback", Scopes: []string{"atproto"}, }) meta := app.ClientMetadata() @@ -72,7 +75,7 @@ func TestOAuthApp_ClientMetadata(t *testing.T) { func TestOAuthApp_ClientMetadata_AppName(t *testing.T) { app, _ := NewOAuthApp(OAuthConfig{ - RedirectURI: "http://127.0.0.1:12345/callback", + RedirectURI: "http://[IP_ADDRESS]:12345/callback", Scopes: []string{"atproto"}, AppName: "TestApp", }) @@ -85,7 +88,7 @@ func TestOAuthApp_ClientMetadata_AppName(t *testing.T) { func TestOAuthApp_Store(t *testing.T) { memStore := oauth.NewMemStore() app, _ := NewOAuthApp(OAuthConfig{ - RedirectURI: "http://127.0.0.1:12345/callback", + RedirectURI: "http://[IP_ADDRESS]:12345/callback", Scopes: []string{"atproto"}, Store: memStore, }) @@ -93,3 +96,56 @@ func TestOAuthApp_Store(t *testing.T) { t.Fatal("expected non-nil store") } } + +func TestSecureRandomBase64(t *testing.T) { + tests := []uint{0, 1, 16, 32, 48} + for _, size := range tests { + s := secureRandomBase64(size) + if size == 0 { + if s != "" { + t.Fatalf("expected empty for size 0, got %q", s) + } + continue + } + if len(s) == 0 { + t.Fatalf("expected non-empty for size %d", size) + } + // Should be base64url-encoded (no padding chars like = or +/) + for _, c := range s { + if !((c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z') || (c >= '0' && c <= '9') || c == '-' || c == '_') { + t.Fatalf("unexpected character %c in base64url output", c) + } + } + } +} + +func TestSecureRandomBase64_Unique(t *testing.T) { + s1 := secureRandomBase64(16) + s2 := secureRandomBase64(16) + if s1 == s2 { + t.Fatal("expected unique values from secure random") + } +} + +func TestMustReadBody(t *testing.T) { + t.Run("NoBody", func(t *testing.T) { + resp := &http.Response{ + Body: http.NoBody, + } + // avoid leaking the close from the deferred Close inside mustReadBody + data := mustReadBody(resp) + if len(data) != 0 { + t.Fatalf("expected empty for NoBody, got %d bytes", len(data)) + } + }) + + t.Run("with content", func(t *testing.T) { + resp := &http.Response{ + Body: io.NopCloser(strings.NewReader("hello world")), + } + data := mustReadBody(resp) + if string(data) != "hello world" { + t.Fatalf("expected 'hello world', got %q", string(data)) + } + }) +} diff --git a/public.go b/public.go index 8754981..1911bb5 100644 --- a/public.go +++ b/public.go @@ -14,13 +14,14 @@ import ( "time" ) -const ( - // PublicAPIBase is the Bluesky public API endpoint used for profile and handle lookups. - PublicAPIBase = "https://public.api.bsky.app" +// PublicAPIBase is the Bluesky public API endpoint used for profile and handle lookups. +// This is a variable so tests can override it. +var PublicAPIBase = "https://public.api.bsky.app" + +// PLCDirectory is used to resolve did:plc identifiers to DID documents. +// This is a variable so tests can override it. +var PLCDirectory = "https://plc.directory" - // PLCDirectory is used to resolve did:plc identifiers to DID documents. - PLCDirectory = "https://plc.directory" -) // ErrSSRFBlocked is returned when a request is blocked due to a private/internal destination. var ErrSSRFBlocked = errors.New("request blocked: potential SSRF detected") @@ -170,6 +171,7 @@ func (c *PublicClient) GetPDSEndpoint(ctx context.Context, did string) (string, case strings.HasPrefix(did, "did:web:"): domain := strings.TrimPrefix(did, "did:web:") domain = strings.ReplaceAll(domain, "%3A", ":") + domain = strings.ReplaceAll(domain, "%2F", "/") if idx := strings.Index(domain, "/"); idx != -1 { domain = domain[:idx] } diff --git a/public_test.go b/public_test.go index 6cc116a..581b3f1 100644 --- a/public_test.go +++ b/public_test.go @@ -1,7 +1,11 @@ package atp import ( + "encoding/json" + "fmt" "net" + "net/http" + "net/http/httptest" "testing" "github.com/stretchr/testify/assert" @@ -14,13 +18,13 @@ func TestIsPrivateIP(t *testing.T) { }{ {"127.0.0.1", true}, {"::1", true}, + {"169.254.1.1", true}, {"10.0.0.1", true}, {"172.16.0.1", true}, {"192.168.1.1", true}, - {"169.254.169.254", true}, {"0.0.0.0", true}, {"8.8.8.8", false}, - {"1.1.1.1", false}, + {"93.184.216.34", false}, } for _, tc := range cases { @@ -50,12 +54,30 @@ func TestValidateDomain_PrivateIP(t *testing.T) { } } +func TestValidateDomain_PublicIP(t *testing.T) { + if err := validateDomain("8.8.8.8"); err != nil { + t.Fatalf("expected nil for public IP, got %v", err) + } +} + func TestValidateDomain_MetadataIP(t *testing.T) { if err := validateDomain("169.254.169.254"); err != ErrSSRFBlocked { t.Fatalf("expected ErrSSRFBlocked for metadata IP, got %v", err) } } +func TestValidateDomain_LoopbackIP(t *testing.T) { + if err := validateDomain("127.0.0.1"); err != ErrSSRFBlocked { + t.Fatalf("expected ErrSSRFBlocked for loopback, got %v", err) + } +} + +func TestValidateDomain_LinkLocal(t *testing.T) { + if err := validateDomain("169.254.0.1"); err != ErrSSRFBlocked { + t.Fatalf("expected ErrSSRFBlocked for link-local, got %v", err) + } +} + func TestNewPublicClient(t *testing.T) { c := NewPublicClient() if c == nil { @@ -63,6 +85,636 @@ func TestNewPublicClient(t *testing.T) { } } +func TestNewPublicClientWithHTTP(t *testing.T) { + hc := &http.Client{Timeout: 0} + c := NewPublicClientWithHTTP(hc) + if c == nil { + t.Fatal("expected non-nil client") + } + // Should use the provided client +} + +func TestResolveHandle_PackageFunc(t *testing.T) { + // Create a test server that mimics the public API's resolveHandle endpoint + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/xrpc/com.atproto.identity.resolveHandle" { + t.Fatalf("unexpected path: %s", r.URL.Path) + } + if r.URL.Query().Get("handle") != "alice.bsky.social" { + t.Fatalf("unexpected handle: %s", r.URL.Query().Get("handle")) + } + + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]string{ + "did": "did:plc:alice", + }) + })) + t.Cleanup(srv.Close) + + // Override the PublicAPIBase for this test by using a client pointed at the test server + oldBase := PublicAPIBase + PublicAPIBase = srv.URL + t.Cleanup(func() { PublicAPIBase = oldBase }) + + did, err := ResolveHandle(t.Context(), "alice.bsky.social") + if err != nil { + t.Fatal(err) + } + if did != "did:plc:alice" { + t.Fatalf("expected did:plc:alice, got %s", did) + } +} + +// testPublicClient creates a PublicClient backed by test servers that mock +// the public API, PLC directory, and optionally a PDS. +// Returns the client and a cleanup function. +type testPublicServers struct { + client *PublicClient + api *httptest.Server // public.api.bsky.app + plc *httptest.Server // plc.directory +} + +func newTestPublicClient(t *testing.T, apiHandler, plcHandler http.HandlerFunc) *testPublicServers { + t.Helper() + apiSrv := httptest.NewServer(apiHandler) + plcSrv := httptest.NewServer(plcHandler) + t.Cleanup(func() { + apiSrv.Close() + plcSrv.Close() + }) + + oldAPI := PublicAPIBase + oldPLC := PLCDirectory + PublicAPIBase = apiSrv.URL + PLCDirectory = plcSrv.URL + t.Cleanup(func() { + PublicAPIBase = oldAPI + PLCDirectory = oldPLC + }) + + return &testPublicServers{ + client: NewPublicClient(), + api: apiSrv, + plc: plcSrv, + } +} + +func TestPublicClient_ResolveHandle(t *testing.T) { + s := newTestPublicClient(t, + func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "/xrpc/com.atproto.identity.resolveHandle", r.URL.Path) + assert.Equal(t, "alice.bsky.social", r.URL.Query().Get("handle")) + + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]string{"did": "did:plc:alice"}) + }, + http.NotFound, + ) + + did, err := s.client.ResolveHandle(t.Context(), "alice.bsky.social") + if err != nil { + t.Fatal(err) + } + if did != "did:plc:alice" { + t.Fatalf("expected did:plc:alice, got %s", did) + } + + // Second call should use cache + did, err = s.client.ResolveHandle(t.Context(), "alice.bsky.social") + if err != nil { + t.Fatal(err) + } + if did != "did:plc:alice" { + t.Fatalf("expected did:plc:alice, got %s", did) + } +} + +func TestPublicClient_ResolveHandle_Normalizes(t *testing.T) { + s := newTestPublicClient(t, + func(w http.ResponseWriter, r *http.Request) { + // Should be normalized to punycode/ascii lower + assert.Equal(t, "xn--caf-dma.example.com", r.URL.Query().Get("handle")) + + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]string{"did": "did:plc:cafe"}) + }, + http.NotFound, + ) + + did, err := s.client.ResolveHandle(t.Context(), "Café.Example.com") + if err != nil { + t.Fatal(err) + } + if did != "did:plc:cafe" { + t.Fatalf("expected did:plc:cafe, got %s", did) + } +} + +func TestPublicClient_ResolveHandle_NotFound(t *testing.T) { + s := newTestPublicClient(t, + func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusNotFound) + }, + http.NotFound, + ) + + _, err := s.client.ResolveHandle(t.Context(), "nonexistent.bsky.social") + if err == nil { + t.Fatal("expected error for not found handle") + } +} + +func TestPublicClient_ResolveHandle_ServerError(t *testing.T) { + s := newTestPublicClient(t, + func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + }, + http.NotFound, + ) + + _, err := s.client.ResolveHandle(t.Context(), "alice.bsky.social") + if err == nil { + t.Fatal("expected error for server error") + } +} + +func TestPublicClient_GetPDSEndpoint_DidPLC(t *testing.T) { + s := newTestPublicClient(t, + http.NotFound, + func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "/did:plc:alice", r.URL.Path) + + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]any{ + "service": []map[string]any{ + { + "id": "#atproto_pds", + "type": "AtprotoPersonalDataServer", + "serviceEndpoint": "https://alice.pds.example.com", + }, + }, + }) + }, + ) + + endpoint, err := s.client.GetPDSEndpoint(t.Context(), "did:plc:alice") + if err != nil { + t.Fatal(err) + } + if endpoint != "https://alice.pds.example.com" { + t.Fatalf("expected https://alice.pds.example.com, got %s", endpoint) + } + + // Second call should use cache + endpoint, err = s.client.GetPDSEndpoint(t.Context(), "did:plc:alice") + if err != nil { + t.Fatal(err) + } + if endpoint != "https://alice.pds.example.com" { + t.Fatalf("expected cached endpoint, got %s", endpoint) + } +} + +func TestPublicClient_GetPDSEndpoint_DidPLC_NoPDS(t *testing.T) { + s := newTestPublicClient(t, + http.NotFound, + func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]any{ + "service": []map[string]any{ + { + "id": "#something_else", + "type": "SomeOtherService", + "serviceEndpoint": "https://other.example.com", + }, + }, + }) + }, + ) + + _, err := s.client.GetPDSEndpoint(t.Context(), "did:plc:alice") + if err == nil { + t.Fatal("expected error when no PDS service found") + } +} + +func TestPublicClient_GetPDSEndpoint_DidPLC_Error(t *testing.T) { + s := newTestPublicClient(t, + http.NotFound, + func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusNotFound) + }, + ) + + _, err := s.client.GetPDSEndpoint(t.Context(), "did:plc:nonexistent") + if err == nil { + t.Fatal("expected error") + } +} + +func TestPublicClient_GetPDSEndpoint_DidWeb(t *testing.T) { + // did:web resolves without needing PLC directory — just validate domain + s := newTestPublicClient(t, http.NotFound, http.NotFound) + + endpoint, err := s.client.GetPDSEndpoint(t.Context(), "did:web:example.com") + if err != nil { + t.Fatal(err) + } + if endpoint != "https://example.com" { + t.Fatalf("expected https://example.com, got %s", endpoint) + } +} + +func TestPublicClient_GetPDSEndpoint_DidWeb_Port(t *testing.T) { + s := newTestPublicClient(t, http.NotFound, http.NotFound) + + endpoint, err := s.client.GetPDSEndpoint(t.Context(), "did:web:example.com%3A8080") + if err != nil { + t.Fatal(err) + } + if endpoint != "https://example.com:8080" { + t.Fatalf("expected https://example.com:8080, got %s", endpoint) + } +} + +func TestPublicClient_GetPDSEndpoint_DidWeb_Path(t *testing.T) { + s := newTestPublicClient(t, http.NotFound, http.NotFound) + + endpoint, err := s.client.GetPDSEndpoint(t.Context(), "did:web:example.com%3A443%2Fsubdomain") + if err != nil { + t.Fatal(err) + } + if endpoint != "https://example.com:443" { + t.Fatalf("expected https://example.com:443, got %s", endpoint) + } +} + +func TestPublicClient_GetPDSEndpoint_DidWeb_Localhost(t *testing.T) { + s := newTestPublicClient(t, http.NotFound, http.NotFound) + + _, err := s.client.GetPDSEndpoint(t.Context(), "did:web:localhost") + if err == nil { + t.Fatal("expected ErrSSRFBlocked for localhost") + } +} + +func TestPublicClient_GetPDSEndpoint_Unknown(t *testing.T) { + s := newTestPublicClient(t, http.NotFound, http.NotFound) + + _, err := s.client.GetPDSEndpoint(t.Context(), "did:something:xyz") + if err == nil { + t.Fatal("expected error for unknown DID method") + } +} + +func TestPublicClient_GetProfile(t *testing.T) { + s := newTestPublicClient(t, + func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "/xrpc/app.bsky.actor.getProfile", r.URL.Path) + assert.Equal(t, "did:plc:alice", r.URL.Query().Get("actor")) + + w.Header().Set("Content-Type", "application/json") + displayName := "Alice" + avatar := "https://example.com/avatar.jpg" + json.NewEncoder(w).Encode(map[string]any{ + "did": "did:plc:alice", + "handle": "alice.bsky.social", + "displayName": &displayName, + "avatar": &avatar, + }) + }, + http.NotFound, + ) + + profile, err := s.client.GetProfile(t.Context(), "did:plc:alice") + if err != nil { + t.Fatal(err) + } + if profile.DID != "did:plc:alice" { + t.Fatalf("expected did:plc:alice, got %s", profile.DID) + } + if profile.Handle != "alice.bsky.social" { + t.Fatalf("unexpected handle: %s", profile.Handle) + } + if profile.DisplayName == nil || *profile.DisplayName != "Alice" { + t.Fatal("expected display name Alice") + } + if profile.Avatar == nil || *profile.Avatar != "https://example.com/avatar.jpg" { + t.Fatal("expected avatar URL") + } +} + +func TestPublicClient_GetProfile_Error(t *testing.T) { + s := newTestPublicClient(t, + func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusNotFound) + }, + http.NotFound, + ) + + _, err := s.client.GetProfile(t.Context(), "nonexistent") + if err == nil { + t.Fatal("expected error") + } +} + +func TestPublicClient_GetProfile_WithHandle(t *testing.T) { + s := newTestPublicClient(t, + func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "alice.bsky.social", r.URL.Query().Get("actor")) + + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]any{ + "did": "did:plc:alice", + "handle": "alice.bsky.social", + }) + }, + http.NotFound, + ) + + profile, err := s.client.GetProfile(t.Context(), "alice.bsky.social") + if err != nil { + t.Fatal(err) + } + if profile.Handle != "alice.bsky.social" { + t.Fatalf("unexpected handle: %s", profile.Handle) + } +} + +func TestPublicClient_ListPublicRecords(t *testing.T) { + pdsSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "/xrpc/com.atproto.repo.listRecords", r.URL.Path) + assert.Equal(t, "did:plc:alice", r.URL.Query().Get("repo")) + assert.Equal(t, "app.bsky.feed.post", r.URL.Query().Get("collection")) + + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]any{ + "records": []map[string]any{ + { + "uri": "at://did:plc:alice/app.bsky.feed.post/3jxy", + "cid": "cid1", + "value": map[string]string{"text": "hello"}, + }, + }, + }) + })) + t.Cleanup(pdsSrv.Close) + + plcSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]any{ + "service": []map[string]any{ + { + "id": "#atproto_pds", + "type": "AtprotoPersonalDataServer", + "serviceEndpoint": pdsSrv.URL, + }, + }, + }) + })) + t.Cleanup(plcSrv.Close) + + oldAPI := PublicAPIBase + oldPLC := PLCDirectory + PublicAPIBase = "http://unused.local" + PLCDirectory = plcSrv.URL + t.Cleanup(func() { + PublicAPIBase = oldAPI + PLCDirectory = oldPLC + }) + + client := NewPublicClient() + + records, cursor, err := client.ListPublicRecords(t.Context(), "did:plc:alice", "app.bsky.feed.post", ListPublicRecordsOpts{}) + if err != nil { + t.Fatal(err) + } + if len(records) != 1 { + t.Fatalf("expected 1 record, got %d", len(records)) + } + if records[0].URI != "at://did:plc:alice/app.bsky.feed.post/3jxy" { + t.Fatalf("unexpected URI: %s", records[0].URI) + } + if cursor != "" { + t.Fatalf("expected empty cursor, got %s", cursor) + } +} + +func TestPublicClient_ListPublicRecords_WithOpts(t *testing.T) { + pdsSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "100", r.URL.Query().Get("limit")) + assert.Equal(t, "next-cursor", r.URL.Query().Get("cursor")) + assert.Equal(t, "true", r.URL.Query().Get("reverse")) + + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]any{ + "records": []map[string]any{}, + "cursor": "", + }) + })) + t.Cleanup(pdsSrv.Close) + + plcSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]any{ + "service": []map[string]any{ + { + "id": "#atproto_pds", + "type": "AtprotoPersonalDataServer", + "serviceEndpoint": pdsSrv.URL, + }, + }, + }) + })) + t.Cleanup(plcSrv.Close) + + oldPLC := PLCDirectory + PLCDirectory = plcSrv.URL + t.Cleanup(func() { PLCDirectory = oldPLC }) + + client := NewPublicClient() + + records, cursor, err := client.ListPublicRecords(t.Context(), "did:plc:alice", "app.bsky.feed.post", ListPublicRecordsOpts{ + Limit: 100, + Cursor: "next-cursor", + Reverse: true, + }) + if err != nil { + t.Fatal(err) + } + if len(records) != 0 { + t.Fatalf("expected 0 records, got %d", len(records)) + } + if cursor != "" { + t.Fatalf("expected empty cursor, got %s", cursor) + } +} + +func TestPublicClient_ListPublicRecords_Error(t *testing.T) { + pdsSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusBadRequest) + })) + t.Cleanup(pdsSrv.Close) + + plcSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]any{ + "service": []map[string]any{ + { + "id": "#atproto_pds", + "type": "AtprotoPersonalDataServer", + "serviceEndpoint": pdsSrv.URL, + }, + }, + }) + })) + t.Cleanup(plcSrv.Close) + + oldPLC := PLCDirectory + PLCDirectory = plcSrv.URL + t.Cleanup(func() { PLCDirectory = oldPLC }) + + client := NewPublicClient() + + _, _, err := client.ListPublicRecords(t.Context(), "did:plc:alice", "app.bsky.feed.post", ListPublicRecordsOpts{}) + if err == nil { + t.Fatal("expected error") + } +} + +func TestPublicClient_GetPublicRecord(t *testing.T) { + pdsSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "/xrpc/com.atproto.repo.getRecord", r.URL.Path) + assert.Equal(t, "did:plc:alice", r.URL.Query().Get("repo")) + assert.Equal(t, "app.bsky.feed.post", r.URL.Query().Get("collection")) + assert.Equal(t, "3jxy", r.URL.Query().Get("rkey")) + + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]any{ + "uri": "at://did:plc:alice/app.bsky.feed.post/3jxy", + "cid": "cid1", + "value": map[string]string{"text": "hello"}, + }) + })) + t.Cleanup(pdsSrv.Close) + + plcSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]any{ + "service": []map[string]any{ + { + "id": "#atproto_pds", + "type": "AtprotoPersonalDataServer", + "serviceEndpoint": pdsSrv.URL, + }, + }, + }) + })) + t.Cleanup(plcSrv.Close) + + oldPLC := PLCDirectory + PLCDirectory = plcSrv.URL + t.Cleanup(func() { PLCDirectory = oldPLC }) + + client := NewPublicClient() + + rec, err := client.GetPublicRecord(t.Context(), "did:plc:alice", "app.bsky.feed.post", "3jxy") + if err != nil { + t.Fatal(err) + } + if rec.URI != "at://did:plc:alice/app.bsky.feed.post/3jxy" { + t.Fatalf("unexpected URI: %s", rec.URI) + } + if rec.Value["text"] != "hello" { + t.Fatalf("unexpected value: %v", rec.Value) + } +} + +func TestPublicClient_GetPublicRecord_Error(t *testing.T) { + pdsSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusNotFound) + })) + t.Cleanup(pdsSrv.Close) + + plcSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]any{ + "service": []map[string]any{ + { + "id": "#atproto_pds", + "type": "AtprotoPersonalDataServer", + "serviceEndpoint": pdsSrv.URL, + }, + }, + }) + })) + t.Cleanup(plcSrv.Close) + + oldPLC := PLCDirectory + PLCDirectory = plcSrv.URL + t.Cleanup(func() { PLCDirectory = oldPLC }) + + client := NewPublicClient() + + _, err := client.GetPublicRecord(t.Context(), "did:plc:alice", "app.bsky.feed.post", "3jxy") + if err == nil { + t.Fatal("expected error") + } +} + +func TestPublicClient_ListAllRecords(t *testing.T) { + pageNum := 0 + pdsSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + pageNum++ + w.Header().Set("Content-Type", "application/json") + + if pageNum < 3 { + json.NewEncoder(w).Encode(map[string]any{ + "records": []map[string]any{ + {"uri": fmt.Sprintf("at://did:plc:alice/app.bsky.feed.post/rec%d", pageNum), "cid": fmt.Sprintf("cid%d", pageNum), "value": map[string]string{"text": fmt.Sprintf("page%d", pageNum)}}, + }, + "cursor": fmt.Sprintf("cursor-%d", pageNum), + }) + } else { + json.NewEncoder(w).Encode(map[string]any{ + "records": []map[string]any{ + {"uri": "at://did:plc:alice/app.bsky.feed.post/last", "cid": "cid3", "value": map[string]string{"text": "last"}}, + }, + "cursor": "", + }) + } + })) + t.Cleanup(pdsSrv.Close) + + plcSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]any{ + "service": []map[string]any{ + { + "id": "#atproto_pds", + "type": "AtprotoPersonalDataServer", + "serviceEndpoint": pdsSrv.URL, + }, + }, + }) + })) + t.Cleanup(plcSrv.Close) + + oldPLC := PLCDirectory + PLCDirectory = plcSrv.URL + t.Cleanup(func() { PLCDirectory = oldPLC }) + + client := NewPublicClient() + + records, err := client.ListAllRecords(t.Context(), "did:plc:alice", "app.bsky.feed.post") + if err != nil { + t.Fatal(err) + } + if len(records) != 3 { + t.Fatalf("expected 3 records, got %d", len(records)) + } +} + func TestInvalidateHandle(t *testing.T) { c := NewPublicClient() @@ -90,6 +742,10 @@ func TestInvalidateHandle(t *testing.T) { _, exists = c.handleCache["xn--caf-dma.example.com"] assert.False(t, exists, "unicode form should invalidate punycode cache entry") c.handleMu.RUnlock() + + // Empty handle should not panic + c.InvalidateHandle("") + c.InvalidateHandle("@") } func TestInvalidateDID(t *testing.T) { @@ -123,4 +779,7 @@ func TestInvalidateDID(t *testing.T) { _, exists = c.handleCache["bob.example.com"] assert.True(t, exists, "bob handle entry should remain") c.handleMu.RUnlock() + + // Empty DID should not panic + c.InvalidateDID("") } diff --git a/uri_test.go b/uri_test.go index 8d9a11f..fe2205e 100644 --- a/uri_test.go +++ b/uri_test.go @@ -177,3 +177,21 @@ func TestNormalizeDisplayRoundTrip(t *testing.T) { }) } } + +func TestNormalizeHandle_InvalidUTF8(t *testing.T) { + // Invalid UTF-8 bytes should fall back to the input (lowered, trimmed) + in := "\xff\xfe\x00" + got := NormalizeHandle(in) + // Should not panic, returns lowered/trimmed version + if got == "" { + t.Log("empty result for invalid UTF-8 is acceptable") + } +} + +func TestDisplayHandle_InvalidUTF8(t *testing.T) { + in := "\xff\xfe\x00" + got := DisplayHandle(in) + if got != in { + t.Logf("DisplayHandle returned %q for invalid input, expected fallback %q", got, in) + } +}