From 1cab448fc22421697fe3a42ac78bfac3c74d6aba Mon Sep 17 00:00:00 2001 From: Tim Miller Date: Sat, 3 Oct 2026 14:38:56 +0900 Subject: [PATCH] Bug fixes --- docs/docs/authentication.md | 22 + docs/docs/common-operations.md | 20 + docs/docs/creating-clients.md | 59 ++ docs/docs/identity-resolution.md | 46 ++ docs/docs/oauth.md | 50 ++ docs/docs/project-setup.md | 6 +- skills/carpanet/SKILL.md | 19 +- src/CarpaNet.OAuth/ATProtoOAuthClient.cs | 299 +++++++---- src/CarpaNet.OAuth/AtprotoSyntax.cs | 253 +++++++++ .../AuthorizationServerDiscovery.cs | 42 +- src/CarpaNet.OAuth/CarpaNet.OAuth.csproj | 4 + src/CarpaNet.OAuth/DPoPTokenProvider.cs | 168 +++++- src/CarpaNet.OAuth/IOAuthStateStore.cs | 7 + src/CarpaNet.OAuth/OAuthClientConfig.cs | 26 +- src/CarpaNet.OAuth/OAuthException.cs | 12 +- src/CarpaNet.OAuth/OAuthJsonContext.cs | 1 + .../OAuthProtectedResourceMetadata.cs | 41 ++ src/CarpaNet.OAuth/OAuthSession.cs | 260 ++++++++- src/CarpaNet.OAuth/README.md | 16 +- .../Scopes/AccountPermission.cs | 202 +++++++ src/CarpaNet.OAuth/Scopes/AtprotoScope.cs | 164 ++++++ src/CarpaNet.OAuth/Scopes/BlobPermission.cs | 240 +++++++++ .../Scopes/IAtprotoOAuthScope.cs | 13 + .../Scopes/IdentityPermission.cs | 117 ++++ src/CarpaNet.OAuth/Scopes/IncludeScope.cs | 130 +++++ src/CarpaNet.OAuth/Scopes/RepoPermission.cs | 208 +++++++ src/CarpaNet.OAuth/Scopes/RpcPermission.cs | 164 ++++++ src/CarpaNet.OAuth/Scopes/ScopeHelpers.cs | 97 ++++ src/CarpaNet.OAuth/Scopes/ScopeSet.cs | 296 ++++++++++ .../Scopes/ScopeStringSyntax.cs | 416 ++++++++++++++ .../Generation/ApiGenerator.cs | 176 +++++- .../Generation/CborContextGenerator.cs | 271 +++++++++- .../Generation/JsonContextGenerator.cs | 89 ++- .../Generation/ObjectGenerator.cs | 15 +- .../Generation/UnionGenerator.cs | 122 ++++- src/CarpaNet.SourceGen/LexiconGenerator.cs | 97 +++- src/CarpaNet.SourceGen/TypeRegistry.cs | 4 +- .../build/CarpaNet.SourceGen.targets | 6 + src/CarpaNet/ATProtoClient.cs | 296 ++++++---- src/CarpaNet/ATProtoClientXrpcExtensions.cs | 311 +++++++++++ .../Auth/INotifySessionInvalidated.cs | 55 ++ src/CarpaNet/Auth/SessionTokenProvider.cs | 73 ++- .../Blob/ATProtoClientBlobExtensions.cs | 58 +- src/CarpaNet/Blob/BlobRef.cs | 16 + src/CarpaNet/Http/HttpClientFactory.cs | 66 ++- src/CarpaNet/Http/ProgressReportingStream.cs | 135 +++++ src/CarpaNet/Http/RateLimitHandler.cs | 138 +++-- src/CarpaNet/Http/XrpcHttpHandler.cs | 37 ++ src/CarpaNet/IXrpcRequestClient.cs | 30 ++ src/CarpaNet/Identity/DnsOverHttpsResolver.cs | 308 +++++++++++ src/CarpaNet/Identity/DnsResolverDefaults.cs | 47 ++ src/CarpaNet/Identity/IdentityJsonContext.cs | 12 + src/CarpaNet/Identity/IdentityResolver.cs | 160 ++++-- .../Identity/IdentityResolverOptions.cs | 95 ++++ src/CarpaNet/Identity/XrpcHandleResolver.cs | 122 +++++ src/CarpaNet/README.md | 23 +- src/CarpaNet/ScopedATProtoClient.cs | 120 +++++ src/CarpaNet/XrpcBody.cs | 175 ++++++ src/CarpaNet/XrpcRequest.cs | 52 ++ src/CarpaNet/XrpcRequestOptions.cs | 151 ++++++ src/CarpaNet/build/CarpaNet.targets | 6 + .../Auth/SessionInvalidationTests.cs | 163 ++++++ .../Blob/BlobPipelineTests.cs | 102 ++++ .../CarpaNet.UnitTests.csproj | 1 + .../Generation/BinaryXrpcGenerationTests.cs | 454 ++++++++++++++++ .../CrossNamespaceAndArrayParameterTests.cs | 99 ++++ .../Generation/GeneratorBuildPropertyTests.cs | 152 ++++++ .../Generation/GeneratorTestHarness.cs | 198 +++++++ .../Generation/OpenUnionGenerationTests.cs | 279 ++++++++++ .../Http/XrpcRequestPipelineTests.cs | 507 ++++++++++++++++++ .../Identity/DnsOverHttpsResolverTests.cs | 301 +++++++++++ .../Identity/HandleResolutionOrderTests.cs | 266 +++++++++ .../OAuth/OAuthCallbackValidationTests.cs | 373 +++++++++++++ .../OAuth/OAuthClientPipelineTests.cs | 343 ++++++++++++ .../OAuth/Scopes/AccountPermissionTests.cs | 115 ++++ .../OAuth/Scopes/BlobPermissionTests.cs | 98 ++++ .../OAuth/Scopes/IdentityPermissionTests.cs | 60 +++ .../OAuth/Scopes/IncludeScopeTests.cs | 93 ++++ .../OAuth/Scopes/MimeTests.cs | 75 +++ .../OAuth/Scopes/RepoPermissionTests.cs | 142 +++++ .../OAuth/Scopes/RpcPermissionTests.cs | 164 ++++++ .../OAuth/Scopes/ScopeSetTests.cs | 118 ++++ .../OAuth/Scopes/ScopeSyntaxTests.cs | 111 ++++ tests/CarpaNet.UnitTests/TypeRegistryTests.cs | 4 +- 84 files changed, 10458 insertions(+), 424 deletions(-) create mode 100644 src/CarpaNet.OAuth/AtprotoSyntax.cs create mode 100644 src/CarpaNet.OAuth/OAuthProtectedResourceMetadata.cs create mode 100644 src/CarpaNet.OAuth/Scopes/AccountPermission.cs create mode 100644 src/CarpaNet.OAuth/Scopes/AtprotoScope.cs create mode 100644 src/CarpaNet.OAuth/Scopes/BlobPermission.cs create mode 100644 src/CarpaNet.OAuth/Scopes/IAtprotoOAuthScope.cs create mode 100644 src/CarpaNet.OAuth/Scopes/IdentityPermission.cs create mode 100644 src/CarpaNet.OAuth/Scopes/IncludeScope.cs create mode 100644 src/CarpaNet.OAuth/Scopes/RepoPermission.cs create mode 100644 src/CarpaNet.OAuth/Scopes/RpcPermission.cs create mode 100644 src/CarpaNet.OAuth/Scopes/ScopeHelpers.cs create mode 100644 src/CarpaNet.OAuth/Scopes/ScopeSet.cs create mode 100644 src/CarpaNet.OAuth/Scopes/ScopeStringSyntax.cs create mode 100644 src/CarpaNet/ATProtoClientXrpcExtensions.cs create mode 100644 src/CarpaNet/Auth/INotifySessionInvalidated.cs create mode 100644 src/CarpaNet/Http/ProgressReportingStream.cs create mode 100644 src/CarpaNet/IXrpcRequestClient.cs create mode 100644 src/CarpaNet/Identity/DnsOverHttpsResolver.cs create mode 100644 src/CarpaNet/Identity/DnsResolverDefaults.cs create mode 100644 src/CarpaNet/Identity/IdentityJsonContext.cs create mode 100644 src/CarpaNet/Identity/IdentityResolverOptions.cs create mode 100644 src/CarpaNet/Identity/XrpcHandleResolver.cs create mode 100644 src/CarpaNet/ScopedATProtoClient.cs create mode 100644 src/CarpaNet/XrpcBody.cs create mode 100644 src/CarpaNet/XrpcRequest.cs create mode 100644 src/CarpaNet/XrpcRequestOptions.cs create mode 100644 tests/CarpaNet.UnitTests/Auth/SessionInvalidationTests.cs create mode 100644 tests/CarpaNet.UnitTests/Blob/BlobPipelineTests.cs create mode 100644 tests/CarpaNet.UnitTests/Generation/BinaryXrpcGenerationTests.cs create mode 100644 tests/CarpaNet.UnitTests/Generation/CrossNamespaceAndArrayParameterTests.cs create mode 100644 tests/CarpaNet.UnitTests/Generation/GeneratorBuildPropertyTests.cs create mode 100644 tests/CarpaNet.UnitTests/Generation/GeneratorTestHarness.cs create mode 100644 tests/CarpaNet.UnitTests/Generation/OpenUnionGenerationTests.cs create mode 100644 tests/CarpaNet.UnitTests/Http/XrpcRequestPipelineTests.cs create mode 100644 tests/CarpaNet.UnitTests/Identity/DnsOverHttpsResolverTests.cs create mode 100644 tests/CarpaNet.UnitTests/Identity/HandleResolutionOrderTests.cs create mode 100644 tests/CarpaNet.UnitTests/OAuth/OAuthCallbackValidationTests.cs create mode 100644 tests/CarpaNet.UnitTests/OAuth/OAuthClientPipelineTests.cs create mode 100644 tests/CarpaNet.UnitTests/OAuth/Scopes/AccountPermissionTests.cs create mode 100644 tests/CarpaNet.UnitTests/OAuth/Scopes/BlobPermissionTests.cs create mode 100644 tests/CarpaNet.UnitTests/OAuth/Scopes/IdentityPermissionTests.cs create mode 100644 tests/CarpaNet.UnitTests/OAuth/Scopes/IncludeScopeTests.cs create mode 100644 tests/CarpaNet.UnitTests/OAuth/Scopes/MimeTests.cs create mode 100644 tests/CarpaNet.UnitTests/OAuth/Scopes/RepoPermissionTests.cs create mode 100644 tests/CarpaNet.UnitTests/OAuth/Scopes/RpcPermissionTests.cs create mode 100644 tests/CarpaNet.UnitTests/OAuth/Scopes/ScopeSetTests.cs create mode 100644 tests/CarpaNet.UnitTests/OAuth/Scopes/ScopeSyntaxTests.cs diff --git a/docs/docs/authentication.md b/docs/docs/authentication.md index 0a148e2..da2c50d 100644 --- a/docs/docs/authentication.md +++ b/docs/docs/authentication.md @@ -42,3 +42,25 @@ if (client.TokenProvider is { } provider) }; } ``` + +## When a Session Ends + +Both `SessionTokenProvider` and the OAuth `DPoPTokenProvider` implement `INotifySessionInvalidated`. +`SessionInvalidated` is raised once when the server rejects the refresh token (for example +`ExpiredToken`, `InvalidToken` or `invalid_grant`). The provider drops its tokens first; delete any +stored session and ask the user to sign in again. Network errors, 5xx responses and rate limits do +not raise it. + +```csharp +if (client.TokenProvider is INotifySessionInvalidated notifier) +{ + notifier.SessionInvalidated += (sender, args) => + { + // args.Did, args.Reason + RemoveStoredSession(args.Did); + }; +} +``` + +`RefreshAsync` always refreshes, even when the current access token has not expired yet; concurrent +callers share one refresh. diff --git a/docs/docs/common-operations.md b/docs/docs/common-operations.md index ac107d5..c9a4c15 100644 --- a/docs/docs/common-operations.md +++ b/docs/docs/common-operations.md @@ -71,6 +71,26 @@ foreach (var item in timeline.Feed) } ``` +## Upload and Download Blobs + +Blob calls use the client's own authentication, so they work with app-password and OAuth (DPoP) +clients alike. The upload streams the content without buffering it. + +```csharp +using CarpaNet.Blob; +using CarpaNet.Http; + +await using var file = File.OpenRead("photo.jpg"); +var progress = new Progress(sent => Console.WriteLine($"{sent} bytes")); +var blobRef = await client.UploadBlobAsync(new ProgressReportingStream(file, progress), "image/jpeg"); + +// Generated records use ATBlob +ATBlob image = blobRef.ToATBlob(); + +// Another account's blob is fetched from that account's PDS, without your credentials. +byte[] data = await client.DownloadBlobAsync(new ATDid(ownerDid), image.Ref); +``` + ## AT Protocol Types CarpaNet provides strongly-typed wrappers for AT Protocol identifiers: diff --git a/docs/docs/creating-clients.md b/docs/docs/creating-clients.md index 1e3c55e..d939c94 100644 --- a/docs/docs/creating-clients.md +++ b/docs/docs/creating-clients.md @@ -47,6 +47,11 @@ var client = ATProtoClientFactory.Create(new ATProtoClientOptions }); ``` +When CarpaNet creates the `HttpClient` itself, the rate-limit, timeout and `UserAgent` options are +applied to it. When you pass your own `HttpClient`, compose its handlers yourself (for example with +`HttpClientFactory.Create(new HttpClientFactoryOptions { ... })`); `UserAgent` is then added to +each request unless the `HttpClient` already sets one. + ## Restoring a Session ```csharp @@ -78,3 +83,57 @@ if (client.TokenProvider is { } provider) }; } ``` + +## Per-Request Options: Proxies, Labelers, Headers and Other Services + +`WithRequestOptions` (and the shortcuts below) return a `ScopedATProtoClient` that shares the +session and HTTP pipeline of the client it wraps. Every call made through it, including the +generated API methods, uses the options. + +```csharp +// Send app.bsky calls to the Bluesky AppView through the PDS. +var appview = client.WithProxy("did:web:api.bsky.app#bsky_appview"); + +// Let the PDS answer itself (no atproto-proxy header), e.g. for preferences. +var pds = client.WithoutProxy(); + +// Extra headers for one kind of call. +var feeds = appview.WithHeader("Accept-Language", "en,de"); + +// Accepted labelers for this scope only. +var scoped = appview.WithAcceptLabelers(new[] { AcceptLabelersHeader.Redact(modDid), myLabelerDid }); + +// A different service. Session credentials are never sent there; add any token yourself. +var video = client + .WithServiceUrl(new Uri("https://video.bsky.app")) + .WithHeader("Authorization", $"Bearer {serviceAuthToken}"); +``` + +Options set on an outer scope win over inner scopes and over the proxy a generated method would use +(the `chat.bsky.*` methods proxy to the chat service by default). + +To change the accepted labelers for every request on a client, call `SetLabelerDids`: + +```csharp +client.SetLabelerDids(new[] { AcceptLabelersHeader.Redact(modDid), subscribedLabelerDid }); +``` + +### Where credentials are sent + +The access token (or DPoP proof) is attached only to requests that go to the session's own PDS. +A query with a `repo` parameter for another account is sent to that account's PDS without +credentials, and so is anything sent with `WithServiceUrl`. + +### Binary bodies and responses + +Procedures with a non-JSON body (`com.atproto.repo.uploadBlob`, `app.bsky.video.uploadPart`) and +queries with a non-JSON response (`com.atproto.sync.getBlob`) are generated to take a `Stream` and +return `byte[]`. They can also be called directly: + +```csharp +var output = await client.PostBinaryAsync(nsid, proxyServiceDid: null, parameters, stream, "video/mp4"); +byte[] car = await client.GetBytesAsync("com.atproto.sync.getRepo", null, parameters); +``` + +A body from a seekable stream is replayed after a token refresh; a non-seekable body is sent once. +Wrap a stream in `ProgressReportingStream` to report upload progress. diff --git a/docs/docs/identity-resolution.md b/docs/docs/identity-resolution.md index 026a265..e5ff898 100644 --- a/docs/docs/identity-resolution.md +++ b/docs/docs/identity-resolution.md @@ -19,3 +19,49 @@ var didDoc2 = await resolver.ResolveAsync("did:plc:z72i7hdynmk6r22z27h6tvur"); ``` The `ATProtoClient` creates an `IdentityResolver` automatically (configurable via `ATProtoClientOptions.CreateIdentityResolver`). + +## Handle resolution methods + +A handle is resolved to a DID with these methods. They are tried in order, and the first method that returns a DID wins: + +1. **DNS** – the TXT record at `_atproto.`, through an `IDnsResolver`. +2. **Well-known** – `https:///.well-known/atproto-did`. +3. **XRPC** – `com.atproto.identity.resolveHandle` on a service that you configure, such as your PDS or `https://public.api.bsky.app`. This method is skipped if no service URL is set. + +Use `IdentityResolverOptions` to change the order or to add the XRPC service: + +```csharp +var resolver = new IdentityResolver(httpClient, new IdentityResolverOptions +{ + Cache = new MemoryIdentityCache(), + HandleResolutionServiceUrl = IdentityResolverOptions.PublicBlueskyAppViewUrl, + // Optional. The default order is Dns, WellKnown, Xrpc. + HandleResolutionOrder = new[] { HandleResolutionMethod.Xrpc }, +}); +``` + +> **Trust:** A DID from the XRPC method is only as trustworthy as the service. The client does not check the handle's DNS record or well-known file. `ResolveAsync` still checks that the DID document claims the handle (`alsoKnownAs`), for all methods. + +## DNS resolvers + +| Resolver | Transport | Use | +| --- | --- | --- | +| `DefaultDnsResolver` | Raw UDP to `1.1.1.1` and `8.8.8.8` | Desktop, server and mobile | +| `DnsOverHttpsResolver` | HTTPS JSON API (`application/dns-json`) to Cloudflare, then Google | Browsers (WebAssembly), and networks that block UDP | + +If you do not give a DNS resolver, `IdentityResolver` calls `DnsResolverDefaults.CreateDefault(httpClient)`. This returns `DnsOverHttpsResolver` in a browser or under WASI, and `DefaultDnsResolver` on other platforms. + +```csharp +// Force DNS-over-HTTPS, with custom endpoints and a 3-second timeout for each endpoint +var dns = new DnsOverHttpsResolver( + httpClient, + new[] { DnsOverHttpsResolver.GoogleEndpoint, DnsOverHttpsResolver.CloudflareEndpoint }, + TimeSpan.FromSeconds(3)); +var resolver = new IdentityResolver(httpClient, new IdentityResolverOptions { DnsResolver = dns }); +``` + +`DnsOverHttpsResolver` tries the next endpoint if a request fails, times out, returns malformed JSON or returns a DNS error such as SERVFAIL. An NXDOMAIN answer returns an empty list. + +### Browsers + +In a browser, the well-known request is usually blocked by CORS. For reliable handle resolution, set `HandleResolutionServiceUrl` to your PDS or to an AppView. diff --git a/docs/docs/oauth.md b/docs/docs/oauth.md index 81cbb83..44e1466 100644 --- a/docs/docs/oauth.md +++ b/docs/docs/oauth.md @@ -58,6 +58,56 @@ var authUrl = await oauthSession.AuthorizeAsync(userHandle); var atClient = await oauthSession.CallbackAsync(Request.Url.ToString()); ``` +## Building the Scope + +`OAuthClientConfig.Scope` is a space-separated string. You can build it with `ScopeSet` (namespace `CarpaNet.OAuth.Scopes`). `ScopeSet` formats each permission in the normalized atproto scope syntax. + +```csharp +using CarpaNet.OAuth.Scopes; + +var scopes = new ScopeSet() + .AddAtproto() // atproto (required) + .AddRepo("app.bsky.feed.post", RepoActions.Create) // repo:app.bsky.feed.post?action=create + .AddBlob("image/*") // blob:image/* + .AddRpc("app.bsky.actor.getProfile", + "did:web:api.bsky.app#bsky_appview") // rpc:app.bsky.actor.getProfile?aud=did:web:api.bsky.app%23bsky_appview + .AddAccount(AccountAttribute.Email) // account:email + .AddIdentity(IdentityAttribute.Handle) // identity:handle + .AddInclude("com.example.authBasic"); // include:com.example.authBasic + +config.SetScope(scopes); // throws if "atproto" is missing +``` + +To read or check scopes: + +- `AtprotoScope.IsValid(value)` and `AtprotoScope.Normalize(scope)` validate and normalize scope strings. +- `RepoPermission.TryParse`, `RpcPermission.TryParse`, `BlobPermission.TryParse`, `AccountPermission.TryParse`, `IdentityPermission.TryParse` and `IncludeScope.TryParse` parse one scope value. +- `ScopeSet.Parse(grantedScope).MatchesRepo("app.bsky.feed.post", RepoActions.Create)` checks a granted scope. `MatchesRpc`, `MatchesBlob`, `MatchesAccount` and `MatchesIdentity` do the same for the other resources. + +The library does not expand `include:` scopes into the permissions of their lexicon permission set. + +## Callback Validation + +`CallbackAsync` validates the authorization response before it creates the session. When a check fails, it throws `OAuthCallbackException` (with `AppState` set) and does not store a session. + +- **`iss` parameter (RFC 9207).** If the callback has an `iss` parameter, it must equal the issuer that the flow started with (`issuer_mismatch`). If the server metadata has `authorization_response_iss_parameter_supported: true`, the parameter is required (`missing_iss`). These checks occur before the code is exchanged. +- **Token subject.** The `sub` of the token response must be an atproto DID (`invalid_sub`). If you started the flow with a handle or DID, `sub` must be that account's DID (`sub_mismatch`). The library resolves the DID document of `sub` and reads the protected resource metadata of its PDS. The `authorization_servers` list must contain the issuer that issued the tokens (`sub_issuer_mismatch`). When one of these checks fails, the library revokes the tokens. +- **PDS URL.** The session uses the PDS from the DID document of `sub`. This is also true when you start the flow with a PDS or entryway URL (for example `https://bsky.social`). + +If you implement a persistent `IOAuthStateStore`, also store `OAuthStateData.ExpectedSub`. If it is not stored, the `sub_mismatch` check does not occur. + +## Token Refresh + +Before each refresh, the session resolves the account's DID document again (bypassing the cache) and +checks that its PDS still names the same authorization server. If the account has moved to a PDS +behind another authorization server, the refresh is not attempted and `SessionInvalidated` is raised +with the reason `issuer_mismatch`; a failed lookup only fails that refresh. The refreshed tokens use +the PDS from the DID document as their audience. + +DPoP proofs are single use. When the session's `HttpClient` includes `RateLimitHandler`, the OAuth +client registers a callback (`RateLimitHandler.SetRetryPreparer`) so each 429 retry is signed with a +new proof. A DPoP-signed request without such a callback is not retried by the handler. + ## Restoring an OAuth Session ```csharp diff --git a/docs/docs/project-setup.md b/docs/docs/project-setup.md index a8c8bd8..3ae30cd 100644 --- a/docs/docs/project-setup.md +++ b/docs/docs/project-setup.md @@ -77,8 +77,8 @@ This scans your lexicons for `ref` fields pointing to external NSIDs, resolves t |----------|---------|-------------| | `CarpaNet_JsonContextName` | `ATProtoJsonContext` | Name of the generated JSON serializer context | | `CarpaNet_CborContextName` | `ATProtoCborContext` | Name of the generated CBOR serializer context | -| `CarpaNet_SourceGen_RootNamespace` | Project namespace | Root namespace for generated code | -| `CarpaNet_SourceGen_EmitValidationAttributes` | `false` | Emit `[ATStringLength]`, `[Range]` attributes | +| `CarpaNet_RootNamespace` | None (from NSID) | Root namespace prefix for generated code | +| `CarpaNet_EmitValidationAttributes` | `true` | Emit `[ATStringLength]`, `[ATRange]` validation attributes | | `CarpaNet_LexiconAutoResolve` | `false` | Auto-resolve transitive lexicon dependencies | | `CarpaNet_LexiconAutoResolveMaxDepth` | `10` | Max iterations for transitive resolution | | `CarpaNet_LexiconCacheDir` | `obj/lexicon-cache/` | Cache directory for resolved lexicons | @@ -87,6 +87,8 @@ This scans your lexicons for `ref` fields pointing to external NSIDs, resolves t | `CarpaNet_PlcDirectoryUrl` | `https://plc.directory` | PLC directory URL | | `CarpaNet_DnsServers` | (empty) | Semicolon-separated DNS server IPs | +The older `CarpaNet_SourceGen_RootNamespace`, `CarpaNet_SourceGen_JsonContextName`, `CarpaNet_SourceGen_CborContextName` and `CarpaNet_SourceGen_EmitValidationAttributes` names still work. If both names are set, the `CarpaNet_*` name wins. + ## Inspecting Generated Code Roslyn allows for emiting the compiler generated files. This makes it easy to debug (and for LLMs to inspect, as it were.) diff --git a/skills/carpanet/SKILL.md b/skills/carpanet/SKILL.md index c696566..d9aa6ce 100644 --- a/skills/carpanet/SKILL.md +++ b/skills/carpanet/SKILL.md @@ -92,8 +92,8 @@ This scans your lexicons for `ref` fields pointing to external NSIDs, resolves t |----------|---------|-------------| | `CarpaNet_JsonContextName` | `ATProtoJsonContext` | Name of the generated JSON serializer context | | `CarpaNet_CborContextName` | `ATProtoCborContext` | Name of the generated CBOR serializer context | -| `CarpaNet_SourceGen_RootNamespace` | Project namespace | Root namespace for generated code | -| `CarpaNet_SourceGen_EmitValidationAttributes` | `false` | Emit `[ATStringLength]`, `[Range]` attributes | +| `CarpaNet_RootNamespace` | None (from NSID) | Root namespace prefix for generated code | +| `CarpaNet_EmitValidationAttributes` | `true` | Emit `[ATStringLength]`, `[ATRange]` validation attributes | | `CarpaNet_LexiconAutoResolve` | `false` | Auto-resolve transitive lexicon dependencies | | `CarpaNet_LexiconAutoResolveMaxDepth` | `10` | Max iterations for transitive resolution | | `CarpaNet_LexiconCacheDir` | `obj/lexicon-cache/` | Cache directory for resolved lexicons | @@ -102,6 +102,8 @@ This scans your lexicons for `ref` fields pointing to external NSIDs, resolves t | `CarpaNet_PlcDirectoryUrl` | `https://plc.directory` | PLC directory URL | | `CarpaNet_DnsServers` | (empty) | Semicolon-separated DNS server IPs | +The older `CarpaNet_SourceGen_RootNamespace`, `CarpaNet_SourceGen_JsonContextName`, `CarpaNet_SourceGen_CborContextName` and `CarpaNet_SourceGen_EmitValidationAttributes` names still work. If both names are set, the `CarpaNet_*` name wins. + ### Inspecting Generated Code ```xml @@ -547,6 +549,19 @@ var didDoc2 = await resolver.ResolveAsync("did:plc:z72i7hdynmk6r22z27h6tvur"); The `ATProtoClient` creates an `IdentityResolver` automatically (configurable via `ATProtoClientOptions.CreateIdentityResolver`). +Handle resolution tries DNS TXT, then HTTPS well-known, then the `com.atproto.identity.resolveHandle` XRPC method (only if a service URL is set). In browsers (WASM), use DNS-over-HTTPS and an XRPC service, because UDP is unavailable and well-known requests are usually blocked by CORS: + +```csharp +var resolver = new IdentityResolver(httpClient, new IdentityResolverOptions +{ + Cache = new MemoryIdentityCache(), + DnsResolver = new DnsOverHttpsResolver(httpClient), // default in browsers via DnsResolverDefaults.CreateDefault + HandleResolutionServiceUrl = IdentityResolverOptions.PublicBlueskyAppViewUrl, // or the user's PDS +}); +``` + +A DID from the XRPC method is only as trustworthy as the service. `ResolveAsync` still checks that the DID document claims the handle. + --- ## Repository & CAR File Reading diff --git a/src/CarpaNet.OAuth/ATProtoOAuthClient.cs b/src/CarpaNet.OAuth/ATProtoOAuthClient.cs index 03b6850..047ef65 100644 --- a/src/CarpaNet.OAuth/ATProtoOAuthClient.cs +++ b/src/CarpaNet.OAuth/ATProtoOAuthClient.cs @@ -1,5 +1,6 @@ using System; using System.Collections.Generic; +using System.Linq; using System.Net.Http; using System.Text.Json; using System.Text.Json.Serialization.Metadata; @@ -17,15 +18,17 @@ namespace CarpaNet.OAuth; /// /// Represents an authenticated OAuth session. /// -public sealed class ATProtoOAuthClient : IATProtoClient, IDisposable +public sealed class ATProtoOAuthClient : IATProtoClient, IXrpcRequestClient, IDisposable { private readonly DPoPTokenProvider _tokenProvider; private readonly HttpClient _httpClient; + private readonly bool _ownsHttpClient; private readonly JsonSerializerOptions _jsonOptions; private readonly IdentityResolver? _identityResolver; private readonly ILogger _logger; private bool _disposed; private readonly OAuthSession _session; + private IReadOnlyList? _labelerDids; /// /// Gets the user's DID. @@ -66,9 +69,23 @@ public sealed class ATProtoOAuthClient : IATProtoClient, IDisposable public IdentityResolver? IdentityResolver => _identityResolver; /// - /// Gets the list of labeler DIDs whose labels should be included in responses. + /// Gets the labeler DIDs sent in the atproto-accept-labelers header. + /// Change it with . /// - public IReadOnlyList? LabelerDids { get; } + public IReadOnlyList? LabelerDids => Volatile.Read(ref _labelerDids); + + /// + public JsonSerializerOptions JsonOptions => _jsonOptions; + + /// + /// Replaces the labeler DIDs sent in the atproto-accept-labelers header on later requests. + /// Entries may carry parameters such as ;redact. + /// + /// The labeler DIDs, or null to send no header. + public void SetLabelerDids(IEnumerable? labelerDids) + { + Volatile.Write(ref _labelerDids, labelerDids?.ToArray()); + } internal ATProtoOAuthClient( string did, @@ -79,15 +96,17 @@ internal ATProtoOAuthClient( IdentityResolver identityResolver, JsonSerializerOptions? jsonOptions = null, IReadOnlyList? labelerDids = null, - ILoggerFactory? loggerFactory = null) + ILoggerFactory? loggerFactory = null, + HttpClient? httpClient = null) { Did = did ?? throw new ArgumentNullException(nameof(did)); BaseUrl = new Uri(pdsUrl); _tokenProvider = tokenProvider ?? throw new ArgumentNullException(nameof(tokenProvider)); _session = session ?? throw new ArgumentNullException(nameof(session)); AppState = appState; - LabelerDids = labelerDids; - _httpClient = new HttpClient(); + _labelerDids = labelerDids?.ToArray(); + _httpClient = httpClient ?? new HttpClient(); + _ownsHttpClient = httpClient == null; var factory = loggerFactory ?? NullLoggerFactory.Instance; _logger = factory.CreateLogger(); _identityResolver = identityResolver; @@ -99,85 +118,211 @@ internal ATProtoOAuthClient( } /// - public async Task GetAsync( + public Task GetAsync( string nsid, IEnumerable>? parameters = null, CancellationToken cancellationToken = default) { ThrowIfDisposed(); - _logger.LogDebug("OAuth GET {Nsid}", nsid); - - var url = (await XrpcHttpHandler.BuildUrlAsync(BaseUrl, nsid, parameters, _identityResolver, _logger, cancellationToken).ConfigureAwait(false)).ToString(); - using var request = await _tokenProvider.CreateDPoPRequestAsync(HttpMethod.Get, url).ConfigureAwait(false); - XrpcHttpHandler.AddCommonHeaders(request, null, LabelerDids); - - var response = await SendWithRetryAsync(request, url, cancellationToken).ConfigureAwait(false); - return await XrpcHttpHandler.ProcessResponseAsync(response, _jsonOptions, _logger, cancellationToken).ConfigureAwait(false); + return this.QueryAsync(nsid, parameters, null, cancellationToken); } /// - public async Task GetAsync( + public Task GetAsync( string nsid, string proxyServiceDid, IEnumerable>? parameters = null, CancellationToken cancellationToken = default) { ThrowIfDisposed(); - - var url = (await XrpcHttpHandler.BuildUrlAsync(BaseUrl, nsid, parameters, _identityResolver, _logger, cancellationToken).ConfigureAwait(false)).ToString(); - using var request = await _tokenProvider.CreateDPoPRequestAsync(HttpMethod.Get, url).ConfigureAwait(false); - XrpcHttpHandler.AddCommonHeaders(request, proxyServiceDid, LabelerDids); - - var response = await SendWithRetryAsync(request, url, cancellationToken).ConfigureAwait(false); - return await XrpcHttpHandler.ProcessResponseAsync(response, _jsonOptions, _logger, cancellationToken).ConfigureAwait(false); + return this.QueryAsync(nsid, parameters, new XrpcRequestOptions { ProxyServiceDid = proxyServiceDid }, cancellationToken); } /// - public async Task PostAsync( + public Task PostAsync( string nsid, TInput? input, CancellationToken cancellationToken = default) { ThrowIfDisposed(); - _logger.LogDebug("OAuth POST {Nsid}", nsid); - - var url = XrpcHttpHandler.BuildUrl(BaseUrl, nsid).ToString(); - using var request = await _tokenProvider.CreateDPoPRequestAsync(HttpMethod.Post, url).ConfigureAwait(false); - XrpcHttpHandler.AddCommonHeaders(request, null, LabelerDids); - - if (input != null) - { - var typeInfo = (JsonTypeInfo)_jsonOptions.GetTypeInfo(typeof(TInput)); - var json = JsonSerializer.Serialize(input, typeInfo); - request.Content = new StringContent(json, System.Text.Encoding.UTF8, "application/json"); - } - - var response = await SendWithRetryAsync(request, url, cancellationToken).ConfigureAwait(false); - return await XrpcHttpHandler.ProcessResponseAsync(response, _jsonOptions, _logger, cancellationToken).ConfigureAwait(false); + return this.ProcedureAsync(nsid, null, input, null, cancellationToken); } /// - public async Task PostAsync( + public Task PostAsync( string nsid, string proxyServiceDid, TInput? input, CancellationToken cancellationToken = default) { ThrowIfDisposed(); + return this.ProcedureAsync(nsid, null, input, new XrpcRequestOptions { ProxyServiceDid = proxyServiceDid }, cancellationToken); + } + + /// + /// + /// + /// Requests go to the session's PDS unless is set. + /// The DPoP-bound access token and proof are attached only when the request goes to the + /// session's own PDS, so credentials are never sent to another server. + /// + /// + /// A 401 from the PDS is retried once with a fresh nonce when the server asks for one + /// (use_dpop_nonce), and otherwise once after a token refresh, when the body can be replayed. + /// + /// + public async Task SendXrpcAsync(XrpcRequest request, CancellationToken cancellationToken = default) + { + ThrowIfDisposed(); + if (request == null) + { + throw new ArgumentNullException(nameof(request)); + } - var url = XrpcHttpHandler.BuildUrl(BaseUrl, nsid).ToString(); - using var request = await _tokenProvider.CreateDPoPRequestAsync(HttpMethod.Post, url).ConfigureAwait(false); - XrpcHttpHandler.AddCommonHeaders(request, proxyServiceDid, LabelerDids); + var options = request.Options; + var proxy = options?.EffectiveProxyServiceDid; + var url = await ResolveRequestUrlAsync(request, proxy, cancellationToken).ConfigureAwait(false); + var urlString = url.ToString(); + var attachCredentials = options?.ServiceUrl == null + && !HasAuthorizationHeader(options) + && XrpcHttpHandler.IsSameOrigin(url, BaseUrl); - if (input != null) + _logger.LogDebug("OAuth {Method} {Nsid}", request.Method, request.Nsid); + + if (attachCredentials) { - var typeInfo = (JsonTypeInfo)_jsonOptions.GetTypeInfo(typeof(TInput)); - var json = JsonSerializer.Serialize(input, typeInfo); - request.Content = new StringContent(json, System.Text.Encoding.UTF8, "application/json"); + // Refreshes the token first when it is about to expire. + await _tokenProvider.GetAccessTokenAsync(cancellationToken).ConfigureAwait(false); } - var response = await SendWithRetryAsync(request, url, cancellationToken).ConfigureAwait(false); - return await XrpcHttpHandler.ProcessResponseAsync(response, _jsonOptions, _logger, cancellationToken).ConfigureAwait(false); + var retriedNonce = false; + var refreshed = false; + while (true) + { + var response = await SendOnceAsync(request, url, urlString, proxy, attachCredentials, cancellationToken).ConfigureAwait(false); + + if (!attachCredentials + || response.StatusCode != System.Net.HttpStatusCode.Unauthorized + || (request.Body != null && !request.Body.IsReplayable)) + { + return response; + } + + if (!retriedNonce && IsUseDPoPNonceChallenge(response)) + { + // The nonce from the response is already cached; send again with a new proof. + _logger.LogDebug("DPoP nonce required, retrying with the new nonce"); + retriedNonce = true; + response.Dispose(); + continue; + } + + if (refreshed) + { + return response; + } + + _logger.LogWarning("Received 401, refreshing DPoP token and retrying"); + try + { + await _tokenProvider.RefreshAsync(cancellationToken).ConfigureAwait(false); + } + catch (Exception ex) when (ex is TokenRefreshException || ex is InvalidOperationException) + { + _logger.LogWarning("DPoP token refresh failed, returning 401"); + return response; + } + + refreshed = true; + response.Dispose(); + } + } + + private async Task ResolveRequestUrlAsync(XrpcRequest request, string? proxy, CancellationToken cancellationToken) + { + var serviceUrl = request.Options?.ServiceUrl; + if (serviceUrl != null) + { + return XrpcHttpHandler.BuildUrl(serviceUrl, request.Nsid, request.Parameters); + } + + if (request.Method == HttpMethod.Get && proxy == null) + { + return await XrpcHttpHandler.BuildUrlAsync( + BaseUrl, request.Nsid, request.Parameters, + _identityResolver, _logger, cancellationToken).ConfigureAwait(false); + } + + return XrpcHttpHandler.BuildUrl(BaseUrl, request.Nsid, request.Parameters); + } + + private async Task SendOnceAsync( + XrpcRequest request, + Uri url, + string urlString, + string? proxy, + bool attachCredentials, + CancellationToken cancellationToken) + { + using var message = attachCredentials + ? await _tokenProvider.CreateDPoPRequestAsync(request.Method, urlString).ConfigureAwait(false) + : new HttpRequestMessage(request.Method, url); + + XrpcHttpHandler.AddCommonHeaders(message, proxy, request.Options?.AcceptLabelers ?? LabelerDids); + XrpcHttpHandler.AddCustomHeaders(message, request.Options?.Headers); + + if (attachCredentials) + { + // A DPoP proof is single use; if RateLimitHandler retries this request it must sign a new one. + RateLimitHandler.SetRetryPreparer(message, retry => _tokenProvider.AddDPoPHeadersAsync(retry)); + } + + if (request.Body != null) + { + message.Content = request.Body.CreateContent() + ?? throw new InvalidOperationException("The request body has already been sent and cannot be replayed."); + } + + var response = await _httpClient.SendAsync(message, cancellationToken).ConfigureAwait(false); + if (attachCredentials) + { + _tokenProvider.UpdateNonceFromResponse(response, urlString); + } + + return response; + } + + private static bool IsUseDPoPNonceChallenge(HttpResponseMessage response) + { + foreach (var challenge in response.Headers.WwwAuthenticate) + { + if (string.Equals(challenge.Scheme, "DPoP", StringComparison.OrdinalIgnoreCase) + && challenge.Parameter != null + && challenge.Parameter.IndexOf("use_dpop_nonce", StringComparison.Ordinal) >= 0) + { + return true; + } + } + + return false; + } + + private static bool HasAuthorizationHeader(XrpcRequestOptions? options) + { + if (options?.Headers == null) + { + return false; + } + + foreach (var header in options.Headers) + { + if (header.Key.Equals("Authorization", StringComparison.OrdinalIgnoreCase)) + { + return true; + } + } + + return false; } /// @@ -204,53 +349,6 @@ public async Task SignOutAsync(CancellationToken cancellationToken = default) Dispose(); } - private async Task SendWithRetryAsync( - HttpRequestMessage request, - string url, - CancellationToken cancellationToken) - { - var response = await _httpClient.SendAsync(request, cancellationToken).ConfigureAwait(false); - _tokenProvider.UpdateNonceFromResponse(response, url); - - // Handle 401 with token refresh - if (response.StatusCode == System.Net.HttpStatusCode.Unauthorized) - { - _logger.LogWarning("Received 401, refreshing DPoP token and retrying"); - // Try to refresh the token - await _tokenProvider.RefreshAsync(cancellationToken).ConfigureAwait(false); - - // Create a new request (can't reuse the old one) - using var retryRequest = await _tokenProvider.CreateDPoPRequestAsync(request.Method, url).ConfigureAwait(false); - - if (request.Content != null) - { - // Clone content - var contentBytes = await request.Content.ReadAsByteArrayAsync().ConfigureAwait(false); - retryRequest.Content = new ByteArrayContent(contentBytes); - - foreach (var header in request.Content.Headers) - { - retryRequest.Content.Headers.TryAddWithoutValidation(header.Key, header.Value); - } - } - - // Copy custom headers - foreach (var header in request.Headers) - { - if (header.Key != "Authorization" && header.Key != "DPoP") - { - retryRequest.Headers.TryAddWithoutValidation(header.Key, header.Value); - } - } - - response.Dispose(); - response = await _httpClient.SendAsync(retryRequest, cancellationToken).ConfigureAwait(false); - _tokenProvider.UpdateNonceFromResponse(response, url); - } - - return response; - } - private void ThrowIfDisposed() { if (_disposed) @@ -269,7 +367,12 @@ public void Dispose() _disposed = true; _tokenProvider.Dispose(); - _identityResolver?.Dispose(); - _httpClient.Dispose(); + + // The identity resolver belongs to the OAuthSession that created this client and may be + // shared with other clients, so it is not disposed here. + if (_ownsHttpClient) + { + _httpClient.Dispose(); + } } } diff --git a/src/CarpaNet.OAuth/AtprotoSyntax.cs b/src/CarpaNet.OAuth/AtprotoSyntax.cs new file mode 100644 index 0000000..2994c4b --- /dev/null +++ b/src/CarpaNet.OAuth/AtprotoSyntax.cs @@ -0,0 +1,253 @@ +using System; + +namespace CarpaNet.OAuth; + +/// +/// Allocation-free syntax checks for atproto identifiers, mirroring the validation rules of the +/// reference TypeScript packages (@atproto/did and @atproto/syntax). +/// +internal static class AtprotoSyntax +{ + private const string DidPlcPrefix = "did:plc:"; + private const string DidWebPrefix = "did:web:"; + private const int DidPlcLength = 32; + + /// + /// Whether the value is a DID using one of the atproto blessed methods (did:plc or did:web). + /// + public static bool IsAtprotoDid(string? value) + { + if (value == null) + { + return false; + } + + if (value.StartsWith(DidPlcPrefix, StringComparison.Ordinal)) + { + return IsDidPlc(value); + } + + if (value.StartsWith(DidWebPrefix, StringComparison.Ordinal)) + { + return IsAtprotoDidWeb(value); + } + + return false; + } + + /// + /// Whether the value is an absolute DID reference (did:...#fragment) whose DID uses an atproto method. + /// + public static bool IsAtprotoDidRefAbsolute(string? value) + { + if (value == null) + { + return false; + } + + var hashIndex = value.IndexOf('#'); + if (hashIndex == -1 || hashIndex == value.Length - 1) + { + return false; // No fragment, or empty fragment + } + + if (value.IndexOf('#', hashIndex + 1) != -1) + { + return false; // More than one '#' + } + + return IsFragment(value, hashIndex + 1, value.Length) && + IsAtprotoDid(value.Substring(0, hashIndex)); + } + + /// + /// Whether the value is a syntactically valid NSID. + /// + public static bool IsNsid(string? value) + { + if (value == null || value.Length < 5 || value.Length > 253 + 1 + 63) + { + return false; + } + + var segmentCount = 0; + var segmentStart = 0; + for (var i = 0; i <= value.Length; i++) + { + if (i < value.Length && value[i] != '.') + { + var c = value[i]; + if (!IsAsciiLetterOrDigit(c) && c != '-') + { + return false; + } + + continue; + } + + var length = i - segmentStart; + if (length < 1 || length > 63) + { + return false; + } + + var first = value[segmentStart]; + var last = value[i - 1]; + if (first == '-' || last == '-') + { + return false; + } + + if (segmentCount == 0 && IsAsciiDigit(first)) + { + return false; // First segment may not start with a digit + } + + if (i == value.Length) + { + // Name segment: letters and digits only, no leading digit + if (IsAsciiDigit(first) || value.IndexOf('-', segmentStart) != -1) + { + return false; + } + } + + segmentCount++; + segmentStart = i + 1; + } + + return segmentCount >= 3; + } + + private static bool IsDidPlc(string value) + { + if (value.Length != DidPlcLength) + { + return false; + } + + for (var i = DidPlcPrefix.Length; i < DidPlcLength; i++) + { + var c = value[i]; + if (!((c >= 'a' && c <= 'z') || (c >= '2' && c <= '7'))) + { + return false; + } + } + + return true; + } + + private static bool IsAtprotoDidWeb(string value) + { + var start = DidWebPrefix.Length; + if (value.Length > 2048 || value.Length == start || value[start] == ':') + { + return false; + } + + // Method-specific identifier characters (DID spec) + for (var i = start; i < value.Length; i++) + { + var c = value[i]; + if (IsAsciiLetterOrDigit(c) || c == '.' || c == '-' || c == '_') + { + continue; + } + + if (c == ':') + { + // Atproto does not allow path components in Web DIDs + return false; + } + + if (c == '%') + { + if (i + 2 >= value.Length || + !IsUpperHexDigit(value[i + 1]) || + !IsUpperHexDigit(value[i + 2])) + { + return false; + } + + i += 2; + continue; + } + + return false; + } + + // Atproto does not allow port numbers in Web DIDs, except for localhost + var isLocalhost = string.Equals(value, "did:web:localhost", StringComparison.Ordinal) || + value.StartsWith("did:web:localhost%3A", StringComparison.Ordinal); + if (!isLocalhost && value.IndexOf("%3A", start, StringComparison.Ordinal) != -1) + { + return false; + } + + var host = Uri.UnescapeDataString(value.Substring(start)); + return Uri.TryCreate("https://" + host, UriKind.Absolute, out _); + } + + private static bool IsFragment(string value, int start, int end) + { + for (var i = start; i < end; i++) + { + var c = value[i]; + if (IsAsciiLetterOrDigit(c)) + { + continue; + } + + switch (c) + { + // unreserved + case '-': + case '.': + case '_': + case '~': + // sub-delims + case '!': + case '$': + case '&': + case '\'': + case '(': + case ')': + case '*': + case '+': + case ',': + case ';': + case '=': + // pchar extra / fragment extra + case ':': + case '@': + case '/': + case '?': + continue; + case '%': + if (i + 2 >= end || !IsHexDigit(value[i + 1]) || !IsHexDigit(value[i + 2])) + { + return false; + } + + i += 2; + continue; + default: + return false; + } + } + + return true; + } + + internal static bool IsAsciiDigit(char c) => c >= '0' && c <= '9'; + + internal static bool IsAsciiLetterOrDigit(char c) => + (c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z') || (c >= '0' && c <= '9'); + + internal static bool IsHexDigit(char c) => + (c >= '0' && c <= '9') || (c >= 'a' && c <= 'f') || (c >= 'A' && c <= 'F'); + + private static bool IsUpperHexDigit(char c) => + (c >= '0' && c <= '9') || (c >= 'A' && c <= 'F'); +} diff --git a/src/CarpaNet.OAuth/AuthorizationServerDiscovery.cs b/src/CarpaNet.OAuth/AuthorizationServerDiscovery.cs index f0e5110..cc93980 100644 --- a/src/CarpaNet.OAuth/AuthorizationServerDiscovery.cs +++ b/src/CarpaNet.OAuth/AuthorizationServerDiscovery.cs @@ -152,7 +152,28 @@ public async Task DiscoverAuthorizationServerAsync( ThrowIfDisposed(); _logger.LogDebug("Discovering authorization server for {ResourceUrl}", resourceUrl); + var metadata = await GetProtectedResourceMetadataAsync(resourceUrl, cancellationToken).ConfigureAwait(false); + + // Use the first authorization server + return metadata.AuthorizationServers![0] + ?? throw new OAuthException("invalid_authorization_server", "Authorization server URL is null."); + } + + /// + /// Fetches OAuth protected resource metadata (RFC 9728) from a PDS. This is not cached, + /// since it is used to verify which authorization server currently protects a resource. + /// + /// The PDS URL. + /// Cancellation token. + /// The protected resource metadata. is guaranteed to be non-empty. + public async Task GetProtectedResourceMetadataAsync( + string resourceUrl, + CancellationToken cancellationToken = default) + { + ThrowIfDisposed(); + var wellKnownUrl = GetProtectedResourceWellKnownUrl(resourceUrl); + _logger.LogDebug("Fetching protected resource metadata from {Url}", wellKnownUrl); var response = await _httpClient.GetAsync(wellKnownUrl, cancellationToken).ConfigureAwait(false); if (!response.IsSuccessStatusCode) @@ -163,19 +184,28 @@ public async Task DiscoverAuthorizationServerAsync( } var content = await response.Content.ReadAsStringAsync().ConfigureAwait(false); - using var doc = JsonDocument.Parse(content); - if (!doc.RootElement.TryGetProperty("authorization_servers", out var servers) || - servers.GetArrayLength() == 0) + OAuthProtectedResourceMetadata? metadata; + try + { + metadata = JsonSerializer.Deserialize(content, OAuthJsonContext.Default.OAuthProtectedResourceMetadata); + } + catch (JsonException ex) + { + throw new OAuthException( + "resource_metadata_parse_failed", + $"Failed to parse protected resource metadata from {wellKnownUrl}.", + ex); + } + + if (metadata?.AuthorizationServers == null || metadata.AuthorizationServers.Length == 0) { throw new OAuthException( "no_authorization_server", "Protected resource does not specify any authorization servers."); } - // Use the first authorization server - return servers[0].GetString() - ?? throw new OAuthException("invalid_authorization_server", "Authorization server URL is null."); + return metadata; } /// diff --git a/src/CarpaNet.OAuth/CarpaNet.OAuth.csproj b/src/CarpaNet.OAuth/CarpaNet.OAuth.csproj index 9e7d7c4..84b7d5f 100644 --- a/src/CarpaNet.OAuth/CarpaNet.OAuth.csproj +++ b/src/CarpaNet.OAuth/CarpaNet.OAuth.csproj @@ -20,6 +20,10 @@ + + + + diff --git a/src/CarpaNet.OAuth/DPoPTokenProvider.cs b/src/CarpaNet.OAuth/DPoPTokenProvider.cs index 8cf89d2..9a3a424 100644 --- a/src/CarpaNet.OAuth/DPoPTokenProvider.cs +++ b/src/CarpaNet.OAuth/DPoPTokenProvider.cs @@ -8,6 +8,7 @@ using CarpaNet.OAuth.Crypto; using CarpaNet.OAuth.Storage; using CarpaNet.Auth; +using CarpaNet.Identity; using Microsoft.Extensions.Logging; using Microsoft.Extensions.Logging.Abstractions; @@ -16,11 +17,12 @@ namespace CarpaNet.OAuth; /// /// Token provider that uses OAuth 2.0 with DPoP for ATProtocol. /// -public sealed class DPoPTokenProvider : ITokenProvider, IDisposable +public sealed class DPoPTokenProvider : ITokenProvider, INotifySessionInvalidated, IDisposable { private readonly HttpClient _httpClient; private readonly IOAuthSessionStore _sessionStore; private readonly AuthorizationServerDiscovery _discovery; + private readonly IdentityResolver? _identityResolver; private readonly DPoPNonceCache _nonceCache; private readonly TimeSpan _refreshBuffer; private readonly string? _clientId; @@ -33,6 +35,7 @@ public sealed class DPoPTokenProvider : ITokenProvider, IDisposable private DPoPKeyPair? _dpopKey; private TokenSet? _tokenSet; private OAuthAuthorizationServerMetadata? _serverMetadata; + private bool _invalidated; private bool _disposed; /// @@ -62,6 +65,14 @@ public sealed class DPoPTokenProvider : ITokenProvider, IDisposable /// public event EventHandler? TokenRefreshed; + /// + /// + /// Raised when the token endpoint rejects the refresh token (invalid_grant, + /// invalid_token, unauthorized_client or invalid_client). The token set is + /// dropped first; the stored session is left for the caller to delete. + /// + public event EventHandler? SessionInvalidated; + /// /// Creates a new DPoP token provider. /// @@ -81,9 +92,11 @@ public DPoPTokenProvider( string? clientId = null, string? redirectUri = null, string? scope = null, - ILoggerFactory? loggerFactory = null) + ILoggerFactory? loggerFactory = null, + IdentityResolver? identityResolver = null) { _httpClient = httpClient ?? throw new ArgumentNullException(nameof(httpClient)); + _identityResolver = identityResolver; _sessionStore = sessionStore ?? throw new ArgumentNullException(nameof(sessionStore)); _discovery = discovery ?? new AuthorizationServerDiscovery(httpClient, loggerFactory: loggerFactory); _nonceCache = new DPoPNonceCache(); @@ -113,6 +126,7 @@ public async Task RestoreSessionAsync(string sub, CancellationToken cancel _sub = sub; _tokenSet = sessionData.TokenSet; + _invalidated = false; _dpopKey = DPoPKeyPair.Import(sessionData.DPoPKey); // Fetch server metadata @@ -138,6 +152,7 @@ internal async Task SetupAsync( _sub = sub; _tokenSet = tokenSet; + _invalidated = false; _dpopKey = dpopKey; _serverMetadata = serverMetadata; @@ -187,16 +202,34 @@ public async Task RefreshAsync(CancellationToken cancellationToken = default) } _logger.LogDebug("Refreshing DPoP token for {Sub}", _sub); + + // Remember the token this caller saw, so a refresh that another caller completed + // while this one waited for the lock is not repeated. + var staleAccessToken = _tokenSet.AccessToken; + SessionInvalidatedEventArgs? invalidated = null; + await _refreshLock.WaitAsync(cancellationToken).ConfigureAwait(false); try { // Double-check after acquiring lock - if (HasValidToken) + if (_tokenSet == null) + { + throw new InvalidOperationException("No refresh token available."); + } + + if (!string.Equals(_tokenSet.AccessToken, staleAccessToken, StringComparison.Ordinal) && HasValidToken) { - _logger.LogDebug("Token still valid after lock, skipping refresh"); + _logger.LogDebug("Token was refreshed by another caller, skipping refresh"); return; } + // Before refreshing, check that this authorization server is still the authority for + // the account (as the reference TypeScript client does). The refresh request stays the + // last async step, so its result can be stored right away. + var audience = _identityResolver != null + ? await VerifyIssuerAsync(_tokenSet.Sub.Length > 0 ? _tokenSet.Sub : _sub!, cancellationToken).ConfigureAwait(false) + : _tokenSet.Audience; + var tokenEndpoint = _serverMetadata.TokenEndpoint; // Build refresh request @@ -220,7 +253,7 @@ public async Task RefreshAsync(CancellationToken cancellationToken = default) var oldRefreshToken = _tokenSet.RefreshToken; _tokenSet = newTokenSet; _tokenSet.Issuer = _serverMetadata.Issuer; - _tokenSet.Audience = PdsUrl?.ToString() ?? string.Empty; + _tokenSet.Audience = audience; // Keep old refresh token if new one not provided if (string.IsNullOrEmpty(_tokenSet.RefreshToken)) @@ -247,9 +280,22 @@ public async Task RefreshAsync(CancellationToken cancellationToken = default) _sub ?? string.Empty, null)); // OAuth doesn't return handle } + catch (InvalidOperationException) + { + throw; + } catch (Exception ex) { _logger.LogError("DPoP token refresh failed for {Sub}", _sub); + if (ex is OAuthException oauthError && IsRejectedRefresh(oauthError.ErrorCode)) + { + invalidated = Invalidate(oauthError.ErrorCode, ex); + } + else if (ex is IssuerVerificationException { IsPermanent: true } issuerError) + { + invalidated = Invalidate(issuerError.Code, ex); + } + throw new TokenRefreshException( "refresh_failed", ex.Message, @@ -259,7 +305,119 @@ public async Task RefreshAsync(CancellationToken cancellationToken = default) finally { _refreshLock.Release(); + + if (invalidated != null) + { + SessionInvalidated?.Invoke(this, invalidated); + } + } + } + + /// + /// Resolves the account's DID document (bypassing the cache) and checks that its PDS still names + /// this session's issuer as an authorization server. + /// + /// The account's PDS URL, the audience for the refreshed tokens. + private async Task VerifyIssuerAsync(string sub, CancellationToken cancellationToken) + { + var issuer = _serverMetadata!.Issuer; + + DidDocument didDoc; + OAuthProtectedResourceMetadata resourceMetadata; + string pdsUrl; + try + { + didDoc = await _identityResolver!.ResolveDidAsync(sub, skipCache: true, cancellationToken).ConfigureAwait(false); + } + catch (OperationCanceledException) + { + throw; + } + catch (Exception ex) + { + throw new IssuerVerificationException("identity_resolution_failed", $"Could not resolve '{sub}': {ex.Message}", isPermanent: false, ex); + } + + if (!string.Equals(didDoc.Id, sub, StringComparison.Ordinal)) + { + throw new IssuerVerificationException("invalid_sub", $"DID document id '{didDoc.Id}' does not match '{sub}'.", isPermanent: true); + } + + var endpoint = didDoc.PdsEndpoint; + if (string.IsNullOrEmpty(endpoint) || !Uri.TryCreate(endpoint, UriKind.Absolute, out _)) + { + throw new IssuerVerificationException("pds_not_found", $"No PDS endpoint in the DID document of '{sub}'.", isPermanent: true); } + + pdsUrl = endpoint!.TrimEnd('/'); + try + { + resourceMetadata = await _discovery.GetProtectedResourceMetadataAsync(pdsUrl, cancellationToken).ConfigureAwait(false); + } + catch (OperationCanceledException) + { + throw; + } + catch (Exception ex) + { + throw new IssuerVerificationException("resource_metadata_failed", $"Could not read the protected resource metadata of {pdsUrl}: {ex.Message}", isPermanent: false, ex); + } + + if (resourceMetadata.AuthorizationServers != null) + { + foreach (var server in resourceMetadata.AuthorizationServers) + { + if (string.Equals(server, issuer, StringComparison.Ordinal)) + { + return pdsUrl; + } + } + } + + // The account moved to a PDS with another authorization server, or the DID now points + // somewhere hostile. Either way these tokens must not be refreshed here. + _logger.LogWarning("PDS {PdsUrl} of {Sub} is no longer protected by issuer {Issuer}", pdsUrl, sub, issuer); + throw new IssuerVerificationException("issuer_mismatch", $"The PDS of '{sub}' ({pdsUrl}) is not protected by issuer '{issuer}'.", isPermanent: true); + } + + private sealed class IssuerVerificationException : Exception + { + public IssuerVerificationException(string code, string message, bool isPermanent, Exception? inner = null) + : base(message, inner) + { + Code = code; + IsPermanent = isPermanent; + } + + public string Code { get; } + + public bool IsPermanent { get; } + } + + private static bool IsRejectedRefresh(string errorCode) + { + return errorCode == "invalid_grant" + || errorCode == "invalid_token" + || errorCode == "unauthorized_client" + || errorCode == "invalid_client"; + } + + /// + /// Drops the token set after the server rejected the refresh token. Returns the event to raise, + /// or null when the event was already raised for this session. + /// + private SessionInvalidatedEventArgs? Invalidate(string reason, Exception exception) + { + _tokenSet = null; + + if (_invalidated) + { + return null; + } + + _invalidated = true; + _logger.LogWarning("OAuth session for {Sub} was rejected by the server ({Reason})", _sub, reason); + return new SessionInvalidatedEventArgs(_sub, reason, exception); } /// diff --git a/src/CarpaNet.OAuth/IOAuthStateStore.cs b/src/CarpaNet.OAuth/IOAuthStateStore.cs index fc07555..d5ddcf1 100644 --- a/src/CarpaNet.OAuth/IOAuthStateStore.cs +++ b/src/CarpaNet.OAuth/IOAuthStateStore.cs @@ -40,6 +40,13 @@ public sealed class OAuthStateData /// public string? PdsUrl { get; set; } + /// + /// The DID of the account the flow was started for, when authorization was started from a + /// handle or DID. The token response's sub must match it. Null when authorization + /// was started from a PDS or entryway URL, in which case any account may sign in. + /// + public string? ExpectedSub { get; set; } + /// /// When this state expires. /// diff --git a/src/CarpaNet.OAuth/OAuthClientConfig.cs b/src/CarpaNet.OAuth/OAuthClientConfig.cs index 9bb4e19..e335d56 100644 --- a/src/CarpaNet.OAuth/OAuthClientConfig.cs +++ b/src/CarpaNet.OAuth/OAuthClientConfig.cs @@ -3,6 +3,7 @@ using System.Text.Json; using CarpaNet.Identity; using CarpaNet.OAuth.Crypto; +using CarpaNet.OAuth.Scopes; using CarpaNet.OAuth.Storage; using Microsoft.Extensions.Logging; @@ -25,10 +26,33 @@ public sealed class OAuthClientConfig public string RedirectUri { get; set; } = string.Empty; /// - /// The scope to request (default: "atproto"). + /// The scope to request (default: "atproto"), as a space-separated string. + /// Use to build it from a . /// public string Scope { get; set; } = "atproto"; + /// + /// Sets from a . + /// + /// The scopes to request. Must contain atproto. + /// This configuration, for chaining. + /// The set does not contain the atproto scope. + public OAuthClientConfig SetScope(ScopeSet scopes) + { + if (scopes == null) + { + throw new ArgumentNullException(nameof(scopes)); + } + + if (!scopes.Contains(AtprotoScope.Atproto)) + { + throw new ArgumentException("atproto OAuth requires the 'atproto' scope.", nameof(scopes)); + } + + Scope = scopes.ToString(); + return this; + } + /// /// The HttpClient to use for requests. If not provided, a new one will be created. /// diff --git a/src/CarpaNet.OAuth/OAuthException.cs b/src/CarpaNet.OAuth/OAuthException.cs index 8b35fa4..011411b 100644 --- a/src/CarpaNet.OAuth/OAuthException.cs +++ b/src/CarpaNet.OAuth/OAuthException.cs @@ -46,7 +46,8 @@ private static string FormatMessage(string errorCode, string? errorDescription) } /// -/// Exception thrown when an OAuth callback contains an error. +/// Exception thrown when an OAuth callback contains an error or fails validation +/// (for example an iss mismatch or an unverifiable token subject). /// public class OAuthCallbackException : OAuthException { @@ -63,6 +64,15 @@ public OAuthCallbackException(string errorCode, string? errorDescription, string { AppState = appState; } + + /// + /// Creates a new OAuth callback exception with an inner exception. + /// + public OAuthCallbackException(string errorCode, string? errorDescription, string? appState, Exception innerException) + : base(errorCode, errorDescription, innerException) + { + AppState = appState; + } } /// diff --git a/src/CarpaNet.OAuth/OAuthJsonContext.cs b/src/CarpaNet.OAuth/OAuthJsonContext.cs index 4a6a435..cc220de 100644 --- a/src/CarpaNet.OAuth/OAuthJsonContext.cs +++ b/src/CarpaNet.OAuth/OAuthJsonContext.cs @@ -7,6 +7,7 @@ namespace CarpaNet.OAuth; /// JSON serialization context for OAuth types. /// [JsonSerializable(typeof(OAuthAuthorizationServerMetadata))] +[JsonSerializable(typeof(OAuthProtectedResourceMetadata))] [JsonSerializable(typeof(OAuthTokenResponse))] [JsonSerializable(typeof(OAuthClientMetadata))] [JsonSerializable(typeof(JsonWebKeySet))] diff --git a/src/CarpaNet.OAuth/OAuthProtectedResourceMetadata.cs b/src/CarpaNet.OAuth/OAuthProtectedResourceMetadata.cs new file mode 100644 index 0000000..fcc4b1b --- /dev/null +++ b/src/CarpaNet.OAuth/OAuthProtectedResourceMetadata.cs @@ -0,0 +1,41 @@ +using System; +using System.Text.Json.Serialization; + +namespace CarpaNet.OAuth; + +/// +/// OAuth 2.0 Protected Resource Metadata (RFC 9728), as served by a PDS at +/// /.well-known/oauth-protected-resource. +/// +public sealed class OAuthProtectedResourceMetadata +{ + /// + /// The protected resource's resource identifier (URL). + /// + [JsonPropertyName("resource")] + public string? Resource { get; set; } + + /// + /// Issuer identifiers of the authorization servers that protect this resource. + /// + [JsonPropertyName("authorization_servers")] + public string[]? AuthorizationServers { get; set; } + + /// + /// Scopes supported by the protected resource. + /// + [JsonPropertyName("scopes_supported")] + public string[]? ScopesSupported { get; set; } + + /// + /// Methods supported for sending bearer tokens to the resource. + /// + [JsonPropertyName("bearer_methods_supported")] + public string[]? BearerMethodsSupported { get; set; } + + /// + /// URL of human-readable documentation for the resource. + /// + [JsonPropertyName("resource_documentation")] + public string? ResourceDocumentation { get; set; } +} diff --git a/src/CarpaNet.OAuth/OAuthSession.cs b/src/CarpaNet.OAuth/OAuthSession.cs index 78f0cbd..2b69dad 100644 --- a/src/CarpaNet.OAuth/OAuthSession.cs +++ b/src/CarpaNet.OAuth/OAuthSession.cs @@ -26,6 +26,7 @@ public sealed class OAuthSession : IDisposable private readonly IOAuthSessionStore _sessionStore; private readonly AuthorizationServerDiscovery _discovery; private readonly IdentityResolver _identityResolver; + private readonly bool _ownsIdentityResolver; private readonly ILogger _logger; private readonly ILoggerFactory _loggerFactory; private bool _disposed; @@ -47,7 +48,8 @@ public OAuthSession(OAuthClientConfig config) _stateStore = config.StateStore ?? new MemoryOAuthStateStore(); _sessionStore = config.SessionStore ?? new MemoryOAuthSessionStore(); _discovery = new AuthorizationServerDiscovery(_httpClient, loggerFactory: _loggerFactory); - _identityResolver = config.IdentityResolver ?? new IdentityResolver(_httpClient, dnsResolver: new CarpaNet.Identity.DefaultDnsResolver(), cache: new MemoryIdentityCache(), loggerFactory: _loggerFactory); + _ownsIdentityResolver = config.IdentityResolver == null; + _identityResolver = config.IdentityResolver ?? new IdentityResolver(_httpClient, cache: new MemoryIdentityCache(), loggerFactory: _loggerFactory); } /// @@ -66,7 +68,7 @@ public async Task AuthorizeAsync( _logger.LogInformation("Starting OAuth authorization for {Input}", input); // Resolve identity to find PDS and authorization server - var (pdsUrl, issuer, serverMetadata) = await ResolveIdentityAsync(input, cancellationToken).ConfigureAwait(false); + var (pdsUrl, issuer, serverMetadata, expectedSub) = await ResolveIdentityAsync(input, cancellationToken).ConfigureAwait(false); // Generate PKCE var (verifier, challenge) = Pkce.Generate(); @@ -85,6 +87,7 @@ public async Task AuthorizeAsync( Verifier = verifier, AppState = appState, PdsUrl = pdsUrl, + ExpectedSub = expectedSub, ExpiresAt = DateTimeOffset.UtcNow + _config.StateExpiration }; @@ -158,9 +161,21 @@ public async Task AuthorizeAsync( /// /// Handles the OAuth callback and exchanges the code for tokens. /// + /// + /// The callback is validated before the session is created: + /// + /// The iss parameter (RFC 9207) must match the issuer the flow was started with. It is + /// required when the server advertises authorization_response_iss_parameter_supported. + /// The token response's sub must be an atproto DID, must match the account the flow was + /// started for (when started from a handle or DID), and the PDS in its DID document must be protected by the + /// issuer that issued the tokens. The session's PDS URL is taken from that DID document. + /// + /// When sub validation fails, the tokens are revoked and the session is not stored. + /// /// The full callback URL with query parameters. /// Cancellation token. /// The OAuth session. + /// The callback contains an error, or fails iss or sub validation. public async Task CallbackAsync( string callbackUrl, CancellationToken cancellationToken = default) @@ -189,9 +204,10 @@ public async Task CallbackAsync( throw new OAuthCallbackException(error, errorDescription, appState); } - // Get code and state + // Get code, state and issuer var code = query["code"]; var stateParam = query["state"]; + var issParam = query["iss"]; if (string.IsNullOrEmpty(code)) { @@ -220,6 +236,9 @@ public async Task CallbackAsync( storedState.Issuer, cancellationToken).ConfigureAwait(false); + // Validate the iss parameter (RFC 9207) before redeeming the code + ValidateIssuerParameter(issParam, storedState, serverMetadata); + _logger.LogDebug("Exchanging authorization code"); // Exchange code for tokens var tokenSet = await ExchangeCodeAsync( @@ -230,7 +249,23 @@ public async Task CallbackAsync( cancellationToken).ConfigureAwait(false); tokenSet.Issuer = storedState.Issuer; - tokenSet.Audience = storedState.PdsUrl ?? string.Empty; + + // The token response MUST be verified before its "sub" can be trusted. The session's + // PDS (DPoP audience) is the one from the sub's DID document, not the URL the flow + // was started from (which may be an entryway). + try + { + tokenSet.Audience = await VerifySubjectAsync( + tokenSet.Sub, + storedState, + serverMetadata, + cancellationToken).ConfigureAwait(false); + } + catch + { + await TryRevokeTokenAsync(serverMetadata, tokenSet, dpopKey).ConfigureAwait(false); + throw; + } // Create token provider var tokenProvider = new DPoPTokenProvider( @@ -241,7 +276,8 @@ public async Task CallbackAsync( _config.ClientId, _config.RedirectUri, _config.Scope, - loggerFactory: _loggerFactory); + loggerFactory: _loggerFactory, + identityResolver: _identityResolver); await tokenProvider.SetupAsync( tokenSet.Sub, @@ -262,7 +298,8 @@ await tokenProvider.SetupAsync( identityResolver: _identityResolver, _config.JsonOptions, _config.LabelerDids, - loggerFactory: _loggerFactory); + loggerFactory: _loggerFactory, + httpClient: _httpClient); } catch { @@ -292,7 +329,8 @@ await tokenProvider.SetupAsync( _config.ClientId, _config.RedirectUri, _config.Scope, - loggerFactory: _loggerFactory); + loggerFactory: _loggerFactory, + identityResolver: _identityResolver); var restored = await tokenProvider.RestoreSessionAsync(sub, cancellationToken).ConfigureAwait(false); if (!restored) @@ -310,7 +348,8 @@ await tokenProvider.SetupAsync( identityResolver: _identityResolver, _config.JsonOptions, _config.LabelerDids, - loggerFactory: _loggerFactory); + loggerFactory: _loggerFactory, + httpClient: _httpClient); } /// @@ -337,22 +376,12 @@ public async Task RevokeAsync(string sub, CancellationToken cancellationToken = sessionData.TokenSet.Issuer, cancellationToken).ConfigureAwait(false); - if (!string.IsNullOrEmpty(serverMetadata.RevocationEndpoint) && - !string.IsNullOrEmpty(sessionData.TokenSet.RefreshToken)) + if (!string.IsNullOrEmpty(sessionData.TokenSet.RefreshToken)) { var dpopKey = DPoPKeyPair.Import(sessionData.DPoPKey); try { - using var request = new HttpRequestMessage(HttpMethod.Post, serverMetadata.RevocationEndpoint); - - var nonce = new DPoPNonceCache().Get(serverMetadata.RevocationEndpoint!); - var proof = await dpopKey.CreateProofAsync("POST", serverMetadata.RevocationEndpoint!, nonce).ConfigureAwait(false); - request.Headers.Add("DPoP", proof); - - var content = $"token={Uri.EscapeDataString(sessionData.TokenSet.RefreshToken)}&token_type_hint=refresh_token"; - request.Content = new StringContent(content, Encoding.UTF8, "application/x-www-form-urlencoded"); - - await _httpClient.SendAsync(request, cancellationToken).ConfigureAwait(false); + await TryRevokeTokenAsync(serverMetadata, sessionData.TokenSet, dpopKey).ConfigureAwait(false); } finally { @@ -370,12 +399,13 @@ public async Task RevokeAsync(string sub, CancellationToken cancellationToken = await _sessionStore.DeleteAsync(sub, cancellationToken).ConfigureAwait(false); } - private async Task<(string pdsUrl, string issuer, OAuthAuthorizationServerMetadata metadata)> ResolveIdentityAsync( + private async Task<(string pdsUrl, string issuer, OAuthAuthorizationServerMetadata metadata, string? did)> ResolveIdentityAsync( string input, CancellationToken cancellationToken) { string pdsUrl; string issuer; + string? did = null; // Check if input is a URL if (Uri.TryCreate(input, UriKind.Absolute, out var inputUri) && @@ -388,8 +418,11 @@ public async Task RevokeAsync(string sub, CancellationToken cancellationToken = // Resolve handle or DID to PDS var didDoc = await _identityResolver.ResolveAsync(input, cancellationToken).ConfigureAwait(false); - pdsUrl = didDoc.PdsEndpoint + pdsUrl = didDoc.PdsEndpoint?.TrimEnd('/') ?? throw new OAuthException("pds_not_found", $"No PDS URL found for: {input}"); + + // The account the user must sign in as + did = !string.IsNullOrEmpty(didDoc.Id) ? didDoc.Id : null; } else { @@ -404,7 +437,182 @@ public async Task RevokeAsync(string sub, CancellationToken cancellationToken = // Get server metadata var metadata = await _discovery.GetMetadataAsync(issuer, cancellationToken).ConfigureAwait(false); - return (pdsUrl, issuer, metadata); + return (pdsUrl, issuer, metadata, did); + } + + /// + /// Validates the RFC 9207 iss authorization response parameter. + /// + private void ValidateIssuerParameter( + string? issParam, + OAuthStateData storedState, + OAuthAuthorizationServerMetadata serverMetadata) + { + if (issParam != null) + { + if (!IsIssuer(issParam, storedState, serverMetadata)) + { + _logger.LogWarning("Callback issuer mismatch: expected {Expected}, got {Actual}", storedState.Issuer, issParam); + throw new OAuthCallbackException( + "issuer_mismatch", + $"Callback issuer '{issParam}' does not match expected issuer '{storedState.Issuer}'.", + storedState.AppState); + } + } + else if (serverMetadata.AuthorizationResponseIssParameterSupported) + { + _logger.LogWarning("Callback is missing the iss parameter required by {Issuer}", storedState.Issuer); + throw new OAuthCallbackException( + "missing_iss", + "The iss parameter is missing from the authorization response.", + storedState.AppState); + } + } + + /// + /// Verifies that the token response's subject is an atproto DID whose PDS is protected by the + /// issuer that issued the tokens. + /// + /// The user's PDS URL (the resource server and DPoP audience). + private async Task VerifySubjectAsync( + string sub, + OAuthStateData storedState, + OAuthAuthorizationServerMetadata serverMetadata, + CancellationToken cancellationToken) + { + if (!AtprotoSyntax.IsAtprotoDid(sub)) + { + throw new OAuthCallbackException( + "invalid_sub", + $"Token response subject '{sub}' is not a valid atproto DID.", + storedState.AppState); + } + + if (storedState.ExpectedSub != null && !string.Equals(sub, storedState.ExpectedSub, StringComparison.Ordinal)) + { + _logger.LogWarning("Token subject {Sub} does not match expected {Expected}", sub, storedState.ExpectedSub); + throw new OAuthCallbackException( + "sub_mismatch", + $"Token response subject '{sub}' does not match the account authorization was started for ('{storedState.ExpectedSub}').", + storedState.AppState); + } + + string pdsUrl; + OAuthProtectedResourceMetadata resourceMetadata; + try + { + // Always resolve fresh: a stale DID document could point to a previous PDS + var didDoc = await _identityResolver.ResolveDidAsync(sub, skipCache: true, cancellationToken).ConfigureAwait(false); + + if (!string.Equals(didDoc.Id, sub, StringComparison.Ordinal)) + { + throw new OAuthCallbackException( + "invalid_sub", + $"DID document id '{didDoc.Id}' does not match token subject '{sub}'.", + storedState.AppState); + } + + var endpoint = didDoc.PdsEndpoint; + if (string.IsNullOrEmpty(endpoint) || + !Uri.TryCreate(endpoint, UriKind.Absolute, out var endpointUri) || + (endpointUri.Scheme != "https" && endpointUri.Scheme != "http")) + { + throw new OAuthCallbackException( + "pds_not_found", + $"No valid PDS endpoint found in the DID document of '{sub}'.", + storedState.AppState); + } + + pdsUrl = endpoint!.TrimEnd('/'); + resourceMetadata = await _discovery.GetProtectedResourceMetadataAsync(pdsUrl, cancellationToken).ConfigureAwait(false); + } + catch (OAuthCallbackException) + { + throw; + } + catch (OperationCanceledException) + { + throw; + } + catch (Exception ex) + { + throw new OAuthCallbackException( + "sub_verification_failed", + $"Failed to verify token subject '{sub}': {ex.Message}", + storedState.AppState, + ex); + } + + foreach (var server in resourceMetadata.AuthorizationServers!) + { + if (server != null && IsIssuer(server, storedState, serverMetadata)) + { + _logger.LogDebug("Token subject {Sub} verified: PDS={PdsUrl}", sub, pdsUrl); + return pdsUrl; + } + } + + // Best case: the user switched PDS. Worst case: a malicious server is trying to + // impersonate a user. Either way, these tokens must not be used. + _logger.LogWarning("PDS {PdsUrl} of {Sub} is not protected by issuer {Issuer}", pdsUrl, sub, storedState.Issuer); + throw new OAuthCallbackException( + "sub_issuer_mismatch", + $"The PDS of '{sub}' ({pdsUrl}) is not protected by issuer '{storedState.Issuer}'.", + storedState.AppState); + } + + private static bool IsIssuer(string value, OAuthStateData storedState, OAuthAuthorizationServerMetadata serverMetadata) + { + // The stored issuer and the metadata issuer were verified to be the same issuer + // (case-insensitively) when the metadata was fetched; accept either exact spelling. + return string.Equals(value, serverMetadata.Issuer, StringComparison.Ordinal) || + string.Equals(value, storedState.Issuer, StringComparison.Ordinal); + } + + /// + /// Best-effort revocation of a token set at the authorization server. Revoking the refresh + /// token revokes the whole grant; the access token is used when no refresh token exists. + /// + private async Task TryRevokeTokenAsync( + OAuthAuthorizationServerMetadata serverMetadata, + TokenSet tokenSet, + DPoPKeyPair dpopKey) + { + var endpoint = serverMetadata.RevocationEndpoint; + if (string.IsNullOrEmpty(endpoint)) + { + return; + } + + var (token, hint) = !string.IsNullOrEmpty(tokenSet.RefreshToken) + ? (tokenSet.RefreshToken!, "refresh_token") + : (tokenSet.AccessToken, "access_token"); + + if (string.IsNullOrEmpty(token)) + { + return; + } + + try + { + using var request = new HttpRequestMessage(HttpMethod.Post, endpoint); + var proof = await dpopKey.CreateProofAsync("POST", endpoint!, null).ConfigureAwait(false); + request.Headers.Add("DPoP", proof); + + var content = BuildFormContent(new Dictionary + { + ["token"] = token, + ["token_type_hint"] = hint, + ["client_id"] = _config.ClientId + }); + request.Content = new StringContent(content, Encoding.UTF8, "application/x-www-form-urlencoded"); + + using var response = await _httpClient.SendAsync(request, CancellationToken.None).ConfigureAwait(false); + } + catch (Exception ex) + { + _logger.LogWarning("Token revocation failed: {Message}", ex.Message); + } } private async Task PushAuthorizationRequestAsync( @@ -639,7 +847,11 @@ public void Dispose() _disposed = true; _discovery.Dispose(); - _identityResolver?.Dispose(); + + if (_ownsIdentityResolver) + { + _identityResolver.Dispose(); + } if (_ownsHttpClient) { diff --git a/src/CarpaNet.OAuth/README.md b/src/CarpaNet.OAuth/README.md index 2084c82..a943000 100644 --- a/src/CarpaNet.OAuth/README.md +++ b/src/CarpaNet.OAuth/README.md @@ -29,4 +29,18 @@ var authUrl = await oauthClient.AuthorizeAsync(handle); // ... redirect user to authUrl, receive callback ... var session = await oauthClient.CallbackAsync(callbackUrl); // session implements IATProtoClient -``` \ No newline at end of file +``` + +`CallbackAsync` validates the `iss` callback parameter (RFC 9207) and the token `sub`. The session's PDS URL is taken from the DID document of `sub`, also when authorization started from an entryway URL. + +### Scopes + +Use `ScopeSet` (namespace `CarpaNet.OAuth.Scopes`) to build the scope string with the atproto permission syntax: + +```csharp +config.SetScope(new ScopeSet() + .AddAtproto() + .AddRepo("app.bsky.feed.post", RepoActions.Create) + .AddBlob("image/*")); +// "atproto repo:app.bsky.feed.post?action=create blob:image/*" +``` diff --git a/src/CarpaNet.OAuth/Scopes/AccountPermission.cs b/src/CarpaNet.OAuth/Scopes/AccountPermission.cs new file mode 100644 index 0000000..335165d --- /dev/null +++ b/src/CarpaNet.OAuth/Scopes/AccountPermission.cs @@ -0,0 +1,202 @@ +using System; +using System.Collections.Generic; + +namespace CarpaNet.OAuth.Scopes; + +/// +/// Account attributes that can be granted with an account: scope. +/// +public enum AccountAttribute +{ + /// The account email (email). + Email, + + /// The account repository (repo), e.g. for import. + Repo, + + /// The account status (status), e.g. activation. + Status, +} + +/// +/// Actions that can be granted with an account: scope. +/// +[Flags] +public enum AccountActions +{ + /// No actions. Not valid in a scope. + None = 0, + + /// Read the attribute (read, the default). + Read = 1, + + /// Manage (and read) the attribute (manage). + Manage = 2, +} + +/// +/// The account:<attr>[?action=read|manage] permission scope. +/// +public sealed class AccountPermission : IAtprotoOAuthScope +{ + /// + /// The scope prefix. + /// + public const string Prefix = "account"; + + /// + /// Creates an account permission. + /// + /// The account attribute. + /// The granted actions (default: ). + /// The attribute or actions are not valid. + public AccountPermission(AccountAttribute attribute, AccountActions actions = AccountActions.Read) + { + if (ToValue(attribute) == null) + { + throw new ArgumentException($"Unknown account attribute: {attribute}", nameof(attribute)); + } + + if (actions == AccountActions.None || (actions & ~(AccountActions.Read | AccountActions.Manage)) != 0) + { + throw new ArgumentException($"Invalid account actions: {actions}", nameof(actions)); + } + + Attribute = attribute; + Actions = actions; + } + + /// + /// The account attribute. + /// + public AccountAttribute Attribute { get; } + + /// + /// The granted actions. + /// + public AccountActions Actions { get; } + + /// + /// Whether this permission allows the given action on the given attribute. + /// implies . + /// + public bool Matches(AccountAttribute attribute, AccountActions action) + { + return Attribute == attribute && + ((Actions & AccountActions.Manage) != 0 || (action != AccountActions.None && (Actions & action) == action)); + } + + /// + /// Parses an account: scope string. + /// + /// The scope string. + /// The parsed permission, or null when invalid. + /// True when the scope is a valid account: scope. + public static bool TryParse(string? scope, out AccountPermission? permission) + { + permission = null; + if (!ScopeStringSyntax.IsScopeStringFor(scope, Prefix)) + { + return false; + } + + var syntax = ScopeStringSyntax.Parse(scope!); + if (syntax == null || !syntax.HasOnlyKeys("attr", "action")) + { + return false; + } + + if (!syntax.TryGetPositionalSingle("attr", out var attrValue) || attrValue == null) + { + return false; + } + + var attribute = ParseAttribute(attrValue); + if (attribute == null) + { + return false; + } + + var actions = AccountActions.Read; + var actionValues = syntax.GetMulti("action"); + if (actionValues != null) + { + actions = AccountActions.None; + foreach (var value in actionValues) + { + switch (value) + { + case "read": + actions |= AccountActions.Read; + break; + case "manage": + actions |= AccountActions.Manage; + break; + default: + return false; + } + } + } + + permission = new AccountPermission(attribute.Value, actions); + return true; + } + + /// + /// Parses an account: scope string. + /// + /// The scope is not a valid account: scope. + public static AccountPermission Parse(string scope) + { + return TryParse(scope, out var permission) + ? permission! + : throw new FormatException($"Invalid account scope: '{scope}'"); + } + + /// + /// Gets the scope string needed to perform the given action on the given attribute. + /// + public static string ScopeNeededFor(AccountAttribute attribute, AccountActions action) + { + return new AccountPermission(attribute, action).ToString(); + } + + /// + public override string ToString() + { + List>? parameters = null; + + // "read" alone is the default and is omitted + if (Actions != AccountActions.Read) + { + parameters = new List>(2); + if ((Actions & AccountActions.Read) != 0) + { + parameters.Add(new KeyValuePair("action", "read")); + } + + if ((Actions & AccountActions.Manage) != 0) + { + parameters.Add(new KeyValuePair("action", "manage")); + } + } + + return ScopeStringSyntax.Format(Prefix, ToValue(Attribute), parameters); + } + + private static AccountAttribute? ParseAttribute(string value) => value switch + { + "email" => AccountAttribute.Email, + "repo" => AccountAttribute.Repo, + "status" => AccountAttribute.Status, + _ => null, + }; + + private static string? ToValue(AccountAttribute attribute) => attribute switch + { + AccountAttribute.Email => "email", + AccountAttribute.Repo => "repo", + AccountAttribute.Status => "status", + _ => null, + }; +} diff --git a/src/CarpaNet.OAuth/Scopes/AtprotoScope.cs b/src/CarpaNet.OAuth/Scopes/AtprotoScope.cs new file mode 100644 index 0000000..91f255d --- /dev/null +++ b/src/CarpaNet.OAuth/Scopes/AtprotoScope.cs @@ -0,0 +1,164 @@ +using System; +using System.Collections.Generic; +using System.Text; + +namespace CarpaNet.OAuth.Scopes; + +/// +/// Parsing, validation and normalization of atproto OAuth scope values. Port of +/// @atproto/oauth-scopes (without permission-set resolution). +/// +public static class AtprotoScope +{ + /// + /// The atproto scope, required in every atproto OAuth request. + /// + public const string Atproto = "atproto"; + + /// + /// The transition:generic scope: broad access similar to an app password. + /// + public const string TransitionGeneric = "transition:generic"; + + /// + /// The transition:chat.bsky scope: access to chat.bsky.* methods (with ). + /// + public const string TransitionChatBsky = "transition:chat.bsky"; + + /// + /// The transition:email scope: read access to the account email. + /// + public const string TransitionEmail = "transition:email"; + + /// + /// Whether the value is one of the static scopes (atproto and the transition: scopes). + /// + public static bool IsStaticScope(string? value) + { + return value == Atproto || + value == TransitionGeneric || + value == TransitionChatBsky || + value == TransitionEmail; + } + + /// + /// Whether the value is for the given scope prefix (resource): equal to it, or followed by : or ?. + /// + /// The scope value. + /// The prefix, e.g. repo. + public static bool IsScopeStringFor(string? value, string prefix) + { + if (prefix == null) + { + throw new ArgumentNullException(nameof(prefix)); + } + + return ScopeStringSyntax.IsScopeStringFor(value, prefix); + } + + /// + /// Whether a single scope value is a valid atproto scope: a static scope, or a permission scope whose + /// parameters are all valid and understood. + /// + public static bool IsValid(string? value) + { + return IsStaticScope(value) || TryParse(value, out _); + } + + /// + /// Parses a single permission scope value (account, blob, identity, include, + /// repo or rpc). Static scopes such as atproto are not permission scopes and return false. + /// + /// The scope value. + /// The parsed scope (for example a ), or null. + /// True when the value is a valid permission scope. + public static bool TryParse(string? value, out IAtprotoOAuthScope? scope) + { + scope = null; + if (string.IsNullOrEmpty(value)) + { + return false; + } + + // Dispatch on the first char to avoid trying every parser + switch (value![0]) + { + case 'a' when AccountPermission.TryParse(value, out var account): + scope = account; + return true; + case 'b' when BlobPermission.TryParse(value, out var blob): + scope = blob; + return true; + case 'i' when IdentityPermission.TryParse(value, out var identity): + scope = identity; + return true; + case 'i' when IncludeScope.TryParse(value, out var include): + scope = include; + return true; + case 'r' when RepoPermission.TryParse(value, out var repo): + scope = repo; + return true; + case 'r' when RpcPermission.TryParse(value, out var rpc): + scope = rpc; + return true; + default: + return false; + } + } + + /// + /// Normalizes a single scope value. + /// + /// The normalized value, or null when the value is not a valid atproto scope. + public static string? NormalizeValue(string? value) + { + if (IsStaticScope(value)) + { + return value; + } + + return TryParse(value, out var scope) ? scope!.ToString() : null; + } + + /// + /// Normalizes a space-separated scope string: each value is normalized, invalid values are dropped, + /// duplicates are removed and the result is sorted (ordinal). + /// + /// The space-separated scope string. + /// The normalized space-separated scope string. + public static string Normalize(string? scope) + { + if (string.IsNullOrEmpty(scope)) + { + return string.Empty; + } + + var values = new SortedSet(StringComparer.Ordinal); + foreach (var value in scope!.Split(' ')) + { + var normalized = NormalizeValue(value); + if (normalized != null) + { + values.Add(normalized); + } + } + + return Join(values); + } + + internal static string Join(IEnumerable values) + { + var sb = new StringBuilder(); + foreach (var value in values) + { + if (sb.Length > 0) + { + sb.Append(' '); + } + + sb.Append(value); + } + + return sb.ToString(); + } +} diff --git a/src/CarpaNet.OAuth/Scopes/BlobPermission.cs b/src/CarpaNet.OAuth/Scopes/BlobPermission.cs new file mode 100644 index 0000000..cd4faaa --- /dev/null +++ b/src/CarpaNet.OAuth/Scopes/BlobPermission.cs @@ -0,0 +1,240 @@ +using System; +using System.Collections.Generic; + +namespace CarpaNet.OAuth.Scopes; + +/// +/// The blob:<mime pattern> (or blob?accept=...&accept=...) permission scope. +/// +public sealed class BlobPermission : IAtprotoOAuthScope +{ + /// + /// The scope prefix. + /// + public const string Prefix = "blob"; + + private const string AnyMime = "*/*"; + + private readonly string[] _accept; + + /// + /// Creates a blob permission. + /// + /// Accepted MIME patterns: */*, type/* or type/subtype. + /// No pattern, or an invalid pattern, was given. + public BlobPermission(params string[] accept) + : this((IEnumerable)accept) + { + } + + /// + /// Creates a blob permission. + /// + /// Accepted MIME patterns: */*, type/* or type/subtype. + /// No pattern, or an invalid pattern, was given. + public BlobPermission(IEnumerable accept) + { + _accept = ScopeHelpers.ToValidatedArray(accept, IsAccept, nameof(accept)); + } + + private BlobPermission(string[] accept, bool _) + { + _accept = accept; + } + + /// + /// The accepted MIME patterns, as given. + /// + public IReadOnlyList Accept => _accept; + + /// + /// Whether this permission allows uploading a blob of the given MIME type. + /// + public bool Matches(string mime) + { + return MatchesAnyAccept(_accept, mime); + } + + /// + /// Whether the value is a concrete MIME type (type/subtype, no wildcard). + /// + public static bool IsMime(string? value) + { + return IsStringSlashString(value) && value!.IndexOf('*') == -1; + } + + /// + /// Whether the value is a valid accept pattern: */*, type/* or type/subtype. + /// + public static bool IsAccept(string? value) + { + if (value == AnyMime) + { + return true; + } + + if (!IsStringSlashString(value)) + { + return false; + } + + return value!.IndexOf('*') == -1 || value.EndsWith("/*", StringComparison.Ordinal); + } + + /// + /// Whether the MIME type matches the accept pattern. Returns false for an invalid MIME type. + /// + public static bool MatchesAccept(string accept, string mime) + { + return IsMime(mime) && MatchesAcceptUnsafe(accept, mime); + } + + /// + /// Whether the MIME type matches any of the accept patterns. Returns false for an invalid MIME type. + /// + public static bool MatchesAnyAccept(IEnumerable accept, string mime) + { + if (!IsMime(mime)) + { + return false; + } + + foreach (var pattern in accept) + { + if (MatchesAcceptUnsafe(pattern, mime)) + { + return true; + } + } + + return false; + } + + /// + /// Parses a blob: scope string. + /// + /// The scope string. + /// The parsed permission, or null when invalid. + /// True when the scope is a valid blob: scope. + public static bool TryParse(string? scope, out BlobPermission? permission) + { + permission = null; + if (!ScopeStringSyntax.IsScopeStringFor(scope, Prefix)) + { + return false; + } + + var syntax = ScopeStringSyntax.Parse(scope!); + if (syntax == null || !syntax.HasOnlyKeys("accept")) + { + return false; + } + + if (!syntax.TryGetPositionalMulti("accept", out var accept) || accept == null) + { + return false; + } + + foreach (var value in accept) + { + if (!IsAccept(value)) + { + return false; + } + } + + permission = new BlobPermission(accept.ToArray(), false); + return true; + } + + /// + /// Parses a blob: scope string. + /// + /// The scope is not a valid blob: scope. + public static BlobPermission Parse(string scope) + { + return TryParse(scope, out var permission) + ? permission! + : throw new FormatException($"Invalid blob scope: '{scope}'"); + } + + /// + /// Gets the scope string needed to upload a blob of the given MIME type. The input is not validated. + /// + public static string ScopeNeededFor(string mime) + { + return new BlobPermission(new[] { mime }, false).ToString(); + } + + /// + public override string ToString() + { + var normalized = Normalize(_accept); + var parameters = new List>(normalized.Length > 1 ? normalized.Length : 0); + var positional = ScopeHelpers.AddPositionalMulti("accept", normalized, parameters); + return ScopeStringSyntax.Format(Prefix, positional, parameters); + } + + private static string[] Normalize(string[] accept) + { + // A more concise representation of the accept values + if (ScopeHelpers.Contains(accept, AnyMime)) + { + return new[] { AnyMime }; + } + + var lower = new string[accept.Length]; + for (var i = 0; i < accept.Length; i++) + { + lower[i] = accept[i].ToLowerInvariant(); + } + + // Drop "type/subtype" values made redundant by a "type/*" value + var kept = new List(lower.Length); + foreach (var value in lower) + { + if (!value.EndsWith("/*", StringComparison.Ordinal)) + { + var slash = value.IndexOf('/'); + var wildcard = value.Substring(0, slash) + "/*"; + if (ScopeHelpers.Contains(lower, wildcard)) + { + continue; + } + } + + kept.Add(value); + } + + return ScopeHelpers.SortedUnique(kept); + } + + private static bool MatchesAcceptUnsafe(string accept, string mime) + { + if (accept == AnyMime) + { + return true; + } + + if (accept.EndsWith("/*", StringComparison.Ordinal)) + { + return mime.StartsWith(accept.Substring(0, accept.Length - 1), StringComparison.Ordinal); + } + + return string.Equals(accept, mime, StringComparison.Ordinal); + } + + private static bool IsStringSlashString(string? value) + { + if (value == null) + { + return false; + } + + var slash = value.IndexOf('/'); + return slash > 0 && + slash < value.Length - 1 && + value.IndexOf('/', slash + 1) == -1 && + value.IndexOf(' ') == -1; + } +} diff --git a/src/CarpaNet.OAuth/Scopes/IAtprotoOAuthScope.cs b/src/CarpaNet.OAuth/Scopes/IAtprotoOAuthScope.cs new file mode 100644 index 0000000..4e4ef95 --- /dev/null +++ b/src/CarpaNet.OAuth/Scopes/IAtprotoOAuthScope.cs @@ -0,0 +1,13 @@ +namespace CarpaNet.OAuth.Scopes; + +/// +/// A parsed atproto OAuth scope value (for example a or an ). +/// +public interface IAtprotoOAuthScope +{ + /// + /// Formats the scope as its normalized scope string. + /// + /// The scope string, e.g. repo:app.bsky.feed.post?action=create. + string ToString(); +} diff --git a/src/CarpaNet.OAuth/Scopes/IdentityPermission.cs b/src/CarpaNet.OAuth/Scopes/IdentityPermission.cs new file mode 100644 index 0000000..ec7aa10 --- /dev/null +++ b/src/CarpaNet.OAuth/Scopes/IdentityPermission.cs @@ -0,0 +1,117 @@ +using System; + +namespace CarpaNet.OAuth.Scopes; + +/// +/// Identity attributes that can be granted with an identity: scope. +/// +public enum IdentityAttribute +{ + /// The account handle (handle). + Handle, + + /// All identity attributes, including the DID document (*). + All, +} + +/// +/// The identity:<handle|*> permission scope. +/// +public sealed class IdentityPermission : IAtprotoOAuthScope +{ + /// + /// The scope prefix. + /// + public const string Prefix = "identity"; + + /// + /// Creates an identity permission. + /// + /// The identity attribute. + /// The attribute is not valid. + public IdentityPermission(IdentityAttribute attribute) + { + if (attribute != IdentityAttribute.Handle && attribute != IdentityAttribute.All) + { + throw new ArgumentException($"Unknown identity attribute: {attribute}", nameof(attribute)); + } + + Attribute = attribute; + } + + /// + /// The identity attribute. + /// + public IdentityAttribute Attribute { get; } + + /// + /// Whether this permission allows access to the given attribute. + /// + public bool Matches(IdentityAttribute attribute) + { + return Attribute == IdentityAttribute.All || Attribute == attribute; + } + + /// + /// Parses an identity: scope string. + /// + /// The scope string. + /// The parsed permission, or null when invalid. + /// True when the scope is a valid identity: scope. + public static bool TryParse(string? scope, out IdentityPermission? permission) + { + permission = null; + if (!ScopeStringSyntax.IsScopeStringFor(scope, Prefix)) + { + return false; + } + + var syntax = ScopeStringSyntax.Parse(scope!); + if (syntax == null || !syntax.HasOnlyKeys("attr")) + { + return false; + } + + if (!syntax.TryGetPositionalSingle("attr", out var value)) + { + return false; + } + + switch (value) + { + case "handle": + permission = new IdentityPermission(IdentityAttribute.Handle); + return true; + case "*": + permission = new IdentityPermission(IdentityAttribute.All); + return true; + default: + return false; + } + } + + /// + /// Parses an identity: scope string. + /// + /// The scope is not a valid identity: scope. + public static IdentityPermission Parse(string scope) + { + return TryParse(scope, out var permission) + ? permission! + : throw new FormatException($"Invalid identity scope: '{scope}'"); + } + + /// + /// Gets the scope string needed to access the given attribute. + /// + public static string ScopeNeededFor(IdentityAttribute attribute) + { + return new IdentityPermission(attribute).ToString(); + } + + /// + public override string ToString() + { + return ScopeStringSyntax.Format(Prefix, Attribute == IdentityAttribute.All ? "*" : "handle", null); + } +} diff --git a/src/CarpaNet.OAuth/Scopes/IncludeScope.cs b/src/CarpaNet.OAuth/Scopes/IncludeScope.cs new file mode 100644 index 0000000..d8918f6 --- /dev/null +++ b/src/CarpaNet.OAuth/Scopes/IncludeScope.cs @@ -0,0 +1,130 @@ +using System; +using System.Collections.Generic; + +namespace CarpaNet.OAuth.Scopes; + +/// +/// The include:<nsid>[?aud=<did#service>] scope, which requests the permissions of a +/// lexicon-defined permission set. Resolving the permission set is not done by this type. +/// +public sealed class IncludeScope : IAtprotoOAuthScope +{ + /// + /// The scope prefix. + /// + public const string Prefix = "include"; + + /// + /// Creates an include scope. + /// + /// The NSID of the permission set. + /// Optional service audience (an atproto DID reference such as did:web:example.com#service) + /// inherited by rpc permissions of the set. + /// The NSID or audience is not valid. + public IncludeScope(string nsid, string? aud = null) + { + if (!AtprotoSyntax.IsNsid(nsid)) + { + throw new ArgumentException($"Invalid NSID: '{nsid}'", nameof(nsid)); + } + + if (aud != null && !AtprotoSyntax.IsAtprotoDidRefAbsolute(aud)) + { + throw new ArgumentException($"Invalid include audience: '{aud}'", nameof(aud)); + } + + Nsid = nsid; + Aud = aud; + } + + /// + /// The NSID of the permission set. + /// + public string Nsid { get; } + + /// + /// The optional service audience. + /// + public string? Aud { get; } + + /// + /// Whether the given NSID is under the namespace authority of this permission set (same NSID group, + /// i.e. everything up to the last . of ). A permission set may only grant + /// permissions for NSIDs under its own authority. + /// + public bool IsParentAuthorityOf(string otherNsid) + { + if (otherNsid == null || otherNsid == ScopeHelpers.Wildcard) + { + return false; + } + + var groupPrefixEnd = Nsid.LastIndexOf('.'); + if (groupPrefixEnd == -1) + { + return false; + } + + // otherNsid must be longer than the group prefix (including the dot) + if (groupPrefixEnd >= otherNsid.Length - 1) + { + return false; + } + + return string.CompareOrdinal(Nsid, 0, otherNsid, 0, groupPrefixEnd + 1) == 0; + } + + /// + /// Parses an include: scope string. + /// + /// The scope string. + /// The parsed scope, or null when invalid. + /// True when the scope is a valid include: scope. + public static bool TryParse(string? scope, out IncludeScope? include) + { + include = null; + if (!ScopeStringSyntax.IsScopeStringFor(scope, Prefix)) + { + return false; + } + + var syntax = ScopeStringSyntax.Parse(scope!); + if (syntax == null || !syntax.HasOnlyKeys("nsid", "aud")) + { + return false; + } + + if (!syntax.TryGetPositionalSingle("nsid", out var nsid) || nsid == null || !AtprotoSyntax.IsNsid(nsid)) + { + return false; + } + + if (!syntax.TryGetSingle("aud", out var aud) || (aud != null && !AtprotoSyntax.IsAtprotoDidRefAbsolute(aud))) + { + return false; + } + + include = new IncludeScope(nsid, aud); + return true; + } + + /// + /// Parses an include: scope string. + /// + /// The scope is not a valid include: scope. + public static IncludeScope Parse(string scope) + { + return TryParse(scope, out var include) + ? include! + : throw new FormatException($"Invalid include scope: '{scope}'"); + } + + /// + public override string ToString() + { + var parameters = Aud == null + ? null + : new[] { new KeyValuePair("aud", Aud) }; + return ScopeStringSyntax.Format(Prefix, Nsid, parameters); + } +} diff --git a/src/CarpaNet.OAuth/Scopes/RepoPermission.cs b/src/CarpaNet.OAuth/Scopes/RepoPermission.cs new file mode 100644 index 0000000..6880d68 --- /dev/null +++ b/src/CarpaNet.OAuth/Scopes/RepoPermission.cs @@ -0,0 +1,208 @@ +using System; +using System.Collections.Generic; + +namespace CarpaNet.OAuth.Scopes; + +/// +/// Record actions that can be granted with a repo: scope. +/// +[Flags] +public enum RepoActions +{ + /// No actions. Not valid in a scope. + None = 0, + + /// Create records (create). + Create = 1, + + /// Update records (update). + Update = 2, + + /// Delete records (delete). + Delete = 4, + + /// All actions (the default when no action is given). + All = Create | Update | Delete, +} + +/// +/// The repo:<collection|*>[?action=create&action=update&action=delete] permission scope. +/// +public sealed class RepoPermission : IAtprotoOAuthScope +{ + /// + /// The scope prefix. + /// + public const string Prefix = "repo"; + + private readonly string[] _collections; + + /// + /// Creates a repo permission for a single collection. + /// + /// The collection NSID, or * for any collection. + /// The granted actions (default: all). + /// The collection or actions are not valid. + public RepoPermission(string collection, RepoActions actions = RepoActions.All) + : this(new[] { collection }, actions) + { + } + + /// + /// Creates a repo permission. + /// + /// Collection NSIDs, or * for any collection. + /// The granted actions (default: all). + /// A collection or the actions are not valid. + public RepoPermission(IEnumerable collections, RepoActions actions = RepoActions.All) + { + if (actions == RepoActions.None || (actions & ~RepoActions.All) != 0) + { + throw new ArgumentException($"Invalid repo actions: {actions}", nameof(actions)); + } + + _collections = ScopeHelpers.ToValidatedArray(collections, IsCollection, nameof(collections)); + Actions = actions; + } + + private RepoPermission(string[] collections, RepoActions actions, bool _) + { + _collections = collections; + Actions = actions; + } + + /// + /// The collections, as given (* means any collection). + /// + public IReadOnlyList Collections => _collections; + + /// + /// The granted actions. + /// + public RepoActions Actions { get; } + + /// + /// Whether this permission allows the given action on the given collection. + /// + public bool Matches(string collection, RepoActions action) + { + return action != RepoActions.None && + (Actions & action) == action && + (ScopeHelpers.Contains(_collections, ScopeHelpers.Wildcard) || ScopeHelpers.Contains(_collections, collection)); + } + + /// + /// Parses a repo: scope string. + /// + /// The scope string. + /// The parsed permission, or null when invalid. + /// True when the scope is a valid repo: scope. + public static bool TryParse(string? scope, out RepoPermission? permission) + { + permission = null; + if (!ScopeStringSyntax.IsScopeStringFor(scope, Prefix)) + { + return false; + } + + var syntax = ScopeStringSyntax.Parse(scope!); + if (syntax == null || !syntax.HasOnlyKeys("collection", "action")) + { + return false; + } + + if (!syntax.TryGetPositionalMulti("collection", out var collections) || collections == null) + { + return false; + } + + foreach (var collection in collections) + { + if (!IsCollection(collection)) + { + return false; + } + } + + var actions = RepoActions.All; + var actionValues = syntax.GetMulti("action"); + if (actionValues != null) + { + actions = RepoActions.None; + foreach (var value in actionValues) + { + var action = ParseAction(value); + if (action == RepoActions.None) + { + return false; + } + + actions |= action; + } + } + + permission = new RepoPermission(collections.ToArray(), actions, false); + return true; + } + + /// + /// Parses a repo: scope string. + /// + /// The scope is not a valid repo: scope. + public static RepoPermission Parse(string scope) + { + return TryParse(scope, out var permission) + ? permission! + : throw new FormatException($"Invalid repo scope: '{scope}'"); + } + + /// + /// Gets the scope string needed to perform the given action on the given collection. The input is not validated. + /// + public static string ScopeNeededFor(string collection, RepoActions action) + { + return new RepoPermission(new[] { collection }, action, false).ToString(); + } + + /// + public override string ToString() + { + var collections = ScopeHelpers.NormalizeWildcardList(_collections); + var parameters = new List>(); + var positional = ScopeHelpers.AddPositionalMulti("collection", collections, parameters); + + // All actions is the default and is omitted + if (Actions != RepoActions.All) + { + if ((Actions & RepoActions.Create) != 0) + { + parameters.Add(new KeyValuePair("action", "create")); + } + + if ((Actions & RepoActions.Update) != 0) + { + parameters.Add(new KeyValuePair("action", "update")); + } + + if ((Actions & RepoActions.Delete) != 0) + { + parameters.Add(new KeyValuePair("action", "delete")); + } + } + + return ScopeStringSyntax.Format(Prefix, positional, parameters); + } + + private static bool IsCollection(string value) + { + return value == ScopeHelpers.Wildcard || AtprotoSyntax.IsNsid(value); + } + + private static RepoActions ParseAction(string value) => value switch + { + "create" => RepoActions.Create, + "update" => RepoActions.Update, + "delete" => RepoActions.Delete, + _ => RepoActions.None, + }; +} diff --git a/src/CarpaNet.OAuth/Scopes/RpcPermission.cs b/src/CarpaNet.OAuth/Scopes/RpcPermission.cs new file mode 100644 index 0000000..e784acc --- /dev/null +++ b/src/CarpaNet.OAuth/Scopes/RpcPermission.cs @@ -0,0 +1,164 @@ +using System; +using System.Collections.Generic; + +namespace CarpaNet.OAuth.Scopes; + +/// +/// The rpc:<lxm|*>?aud=<did#service|*> (or rpc?lxm=...&lxm=...&aud=...) permission scope, +/// which allows calling XRPC methods on a service through the PDS (service proxying). +/// +public sealed class RpcPermission : IAtprotoOAuthScope +{ + /// + /// The scope prefix. + /// + public const string Prefix = "rpc"; + + private readonly string[] _lxm; + + /// + /// Creates an RPC permission for a single method. + /// + /// The service: an atproto DID reference such as did:web:api.bsky.app#bsky_appview, or *. + /// The method NSID, or * for any method. + /// The audience or method is not valid, or both are wildcards. + public RpcPermission(string aud, string lxm) + : this(aud, new[] { lxm }) + { + } + + /// + /// Creates an RPC permission. + /// + /// The service: an atproto DID reference such as did:web:api.bsky.app#bsky_appview, or *. + /// Method NSIDs, or * for any method. + /// The audience or a method is not valid, or both are wildcards. + public RpcPermission(string aud, IEnumerable lxm) + { + if (!IsAud(aud)) + { + throw new ArgumentException($"Invalid rpc audience: '{aud}'", nameof(aud)); + } + + _lxm = ScopeHelpers.ToValidatedArray(lxm, IsLxm, nameof(lxm)); + + if (aud == ScopeHelpers.Wildcard && ScopeHelpers.Contains(_lxm, ScopeHelpers.Wildcard)) + { + throw new ArgumentException("rpc:*?aud=* is not allowed.", nameof(lxm)); + } + + Aud = aud; + } + + private RpcPermission(string aud, string[] lxm, bool _) + { + Aud = aud; + _lxm = lxm; + } + + /// + /// The service audience (* means any service). + /// + public string Aud { get; } + + /// + /// The lexicon methods, as given (* means any method). + /// + public IReadOnlyList Lxm => _lxm; + + /// + /// Whether this permission allows calling the given method on the given service. + /// + public bool Matches(string lxm, string aud) + { + return (Aud == ScopeHelpers.Wildcard || string.Equals(Aud, aud, StringComparison.Ordinal)) && + (ScopeHelpers.Contains(_lxm, ScopeHelpers.Wildcard) || ScopeHelpers.Contains(_lxm, lxm)); + } + + /// + /// Parses an rpc: scope string. + /// + /// The scope string. + /// The parsed permission, or null when invalid. + /// True when the scope is a valid rpc: scope. + public static bool TryParse(string? scope, out RpcPermission? permission) + { + permission = null; + if (!ScopeStringSyntax.IsScopeStringFor(scope, Prefix)) + { + return false; + } + + var syntax = ScopeStringSyntax.Parse(scope!); + if (syntax == null || !syntax.HasOnlyKeys("lxm", "aud")) + { + return false; + } + + if (!syntax.TryGetPositionalMulti("lxm", out var lxm) || lxm == null) + { + return false; + } + + foreach (var value in lxm) + { + if (!IsLxm(value)) + { + return false; + } + } + + if (!syntax.TryGetSingle("aud", out var aud) || aud == null || !IsAud(aud)) + { + return false; + } + + // rpc:*?aud=* is forbidden + if (aud == ScopeHelpers.Wildcard && ScopeHelpers.Contains(lxm, ScopeHelpers.Wildcard)) + { + return false; + } + + permission = new RpcPermission(aud, lxm.ToArray(), false); + return true; + } + + /// + /// Parses an rpc: scope string. + /// + /// The scope is not a valid rpc: scope. + public static RpcPermission Parse(string scope) + { + return TryParse(scope, out var permission) + ? permission! + : throw new FormatException($"Invalid rpc scope: '{scope}'"); + } + + /// + /// Gets the scope string needed to call the given method on the given service. The input is not validated. + /// + public static string ScopeNeededFor(string lxm, string aud) + { + return new RpcPermission(aud, new[] { lxm }, false).ToString(); + } + + /// + public override string ToString() + { + var lxm = ScopeHelpers.NormalizeWildcardList(_lxm); + var parameters = new List>(lxm.Length + 1); + var positional = ScopeHelpers.AddPositionalMulti("lxm", lxm, parameters); + parameters.Add(new KeyValuePair("aud", Aud)); + return ScopeStringSyntax.Format(Prefix, positional, parameters); + } + + private static bool IsLxm(string value) + { + return value == ScopeHelpers.Wildcard || AtprotoSyntax.IsNsid(value); + } + + private static bool IsAud(string? value) + { + return value == ScopeHelpers.Wildcard || AtprotoSyntax.IsAtprotoDidRefAbsolute(value); + } +} diff --git a/src/CarpaNet.OAuth/Scopes/ScopeHelpers.cs b/src/CarpaNet.OAuth/Scopes/ScopeHelpers.cs new file mode 100644 index 0000000..8377ec3 --- /dev/null +++ b/src/CarpaNet.OAuth/Scopes/ScopeHelpers.cs @@ -0,0 +1,97 @@ +using System; +using System.Collections.Generic; + +namespace CarpaNet.OAuth.Scopes; + +/// +/// Shared helpers for the scope permission types. +/// +internal static class ScopeHelpers +{ + public const string Wildcard = "*"; + + public static bool Contains(IReadOnlyList values, string value) + { + for (var i = 0; i < values.Count; i++) + { + if (string.Equals(values[i], value, StringComparison.Ordinal)) + { + return true; + } + } + + return false; + } + + /// + /// Copies the values into an array, validating each one. + /// + public static string[] ToValidatedArray(IEnumerable values, Func validate, string paramName) + { + if (values == null) + { + throw new ArgumentNullException(paramName); + } + + var list = new List(values); + if (list.Count == 0) + { + throw new ArgumentException("At least one value is required.", paramName); + } + + foreach (var value in list) + { + if (!validate(value)) + { + throw new ArgumentException($"Invalid value: '{value}'", paramName); + } + } + + return list.ToArray(); + } + + /// + /// Normalizes a list of NSID-or-wildcard values: a wildcard absorbs every other value, + /// otherwise duplicates are removed and values are sorted (ordinal). + /// + public static string[] NormalizeWildcardList(IReadOnlyList values) + { + if (values.Count > 1 && Contains(values, Wildcard)) + { + return new[] { Wildcard }; + } + + return SortedUnique(values); + } + + public static string[] SortedUnique(IReadOnlyList values) + { + if (values.Count == 1) + { + return new[] { values[0] }; + } + + var set = new SortedSet(values, StringComparer.Ordinal); + var result = new string[set.Count]; + set.CopyTo(result); + return result; + } + + /// + /// Adds a multi-valued parameter, either as the positional value (single value) or as repeated named parameters. + /// + public static string? AddPositionalMulti(string key, string[] values, List> parameters) + { + if (values.Length == 1) + { + return values[0]; + } + + foreach (var value in values) + { + parameters.Add(new KeyValuePair(key, value)); + } + + return null; + } +} diff --git a/src/CarpaNet.OAuth/Scopes/ScopeSet.cs b/src/CarpaNet.OAuth/Scopes/ScopeSet.cs new file mode 100644 index 0000000..745c553 --- /dev/null +++ b/src/CarpaNet.OAuth/Scopes/ScopeSet.cs @@ -0,0 +1,296 @@ +using System; +using System.Collections; +using System.Collections.Generic; + +namespace CarpaNet.OAuth.Scopes; + +/// +/// An ordered set of OAuth scope values, used to build the space-separated scope parameter +/// (see ) and to check granted scopes. +/// +/// +/// +/// var scopes = new ScopeSet() +/// .AddAtproto() +/// .AddRepo("app.bsky.feed.post", RepoActions.Create) +/// .AddBlob("image/*") +/// .AddRpc("app.bsky.actor.getProfile", "did:web:api.bsky.app#bsky_appview"); +/// config.SetScope(scopes); +/// // "atproto repo:app.bsky.feed.post?action=create blob:image/* rpc:app.bsky.actor.getProfile?aud=did:web:api.bsky.app%23bsky_appview" +/// +/// +public sealed class ScopeSet : IReadOnlyCollection +{ + private readonly List _values = new(); + private readonly HashSet _set = new(StringComparer.Ordinal); + + /// + /// Creates an empty scope set. + /// + public ScopeSet() + { + } + + /// + /// Creates a scope set from scope values. + /// + /// The scope values (each without spaces). + public ScopeSet(IEnumerable scopes) + { + if (scopes == null) + { + throw new ArgumentNullException(nameof(scopes)); + } + + foreach (var scope in scopes) + { + Add(scope); + } + } + + /// + /// Parses a space-separated scope string (e.g. a token response's scope). Values are kept as-is, + /// including values this library does not understand. + /// + public static ScopeSet Parse(string? scope) + { + var set = new ScopeSet(); + if (!string.IsNullOrEmpty(scope)) + { + foreach (var value in scope!.Split(' ')) + { + if (value.Length > 0) + { + set.AddCore(value); + } + } + } + + return set; + } + + /// + public int Count => _values.Count; + + /// + /// Adds a raw scope value. Unknown values are allowed, for forward compatibility. + /// + /// The value is empty or contains whitespace. + public ScopeSet Add(string scope) + { + if (string.IsNullOrEmpty(scope)) + { + throw new ArgumentException("Scope value cannot be empty.", nameof(scope)); + } + + foreach (var c in scope) + { + if (char.IsWhiteSpace(c)) + { + throw new ArgumentException($"Scope value cannot contain whitespace: '{scope}'", nameof(scope)); + } + } + + AddCore(scope); + return this; + } + + /// + /// Adds a parsed scope, formatted in its normalized form. + /// + public ScopeSet Add(IAtprotoOAuthScope scope) + { + if (scope == null) + { + throw new ArgumentNullException(nameof(scope)); + } + + AddCore(scope.ToString()); + return this; + } + + /// + /// Adds the atproto scope (required in every request). + /// + public ScopeSet AddAtproto() => AddCore(AtprotoScope.Atproto); + + /// + /// Adds the transition:generic scope. + /// + public ScopeSet AddTransitionGeneric() => AddCore(AtprotoScope.TransitionGeneric); + + /// + /// Adds the transition:chat.bsky scope. + /// + public ScopeSet AddTransitionChatBsky() => AddCore(AtprotoScope.TransitionChatBsky); + + /// + /// Adds the transition:email scope. + /// + public ScopeSet AddTransitionEmail() => AddCore(AtprotoScope.TransitionEmail); + + /// + /// Adds a repo: permission. + /// + /// The collection NSID, or *. + /// The granted actions (default: all). + public ScopeSet AddRepo(string collection, RepoActions actions = RepoActions.All) => + Add(new RepoPermission(collection, actions)); + + /// + /// Adds an rpc: permission. + /// + /// The method NSID, or *. + /// The service DID reference (e.g. did:web:api.bsky.app#bsky_appview), or *. + public ScopeSet AddRpc(string lxm, string aud) => Add(new RpcPermission(aud, lxm)); + + /// + /// Adds a blob: permission. + /// + /// Accepted MIME patterns, e.g. image/*. + public ScopeSet AddBlob(params string[] accept) => Add(new BlobPermission(accept)); + + /// + /// Adds an account: permission. + /// + public ScopeSet AddAccount(AccountAttribute attribute, AccountActions actions = AccountActions.Read) => + Add(new AccountPermission(attribute, actions)); + + /// + /// Adds an identity: permission. + /// + public ScopeSet AddIdentity(IdentityAttribute attribute) => Add(new IdentityPermission(attribute)); + + /// + /// Adds an include: scope. + /// + /// The permission set NSID. + /// Optional service DID reference. + public ScopeSet AddInclude(string nsid, string? aud = null) => Add(new IncludeScope(nsid, aud)); + + /// + /// Removes a scope value. + /// + /// True when the value was present. + public bool Remove(string scope) + { + if (scope == null || !_set.Remove(scope)) + { + return false; + } + + _values.Remove(scope); + return true; + } + + /// + /// Whether the set contains exactly this scope value. + /// + public bool Contains(string scope) => scope != null && _set.Contains(scope); + + /// + /// Whether any repo: scope in the set allows the action on the collection. + /// + public bool MatchesRepo(string collection, RepoActions action) + { + foreach (var value in _values) + { + if (RepoPermission.TryParse(value, out var p) && p!.Matches(collection, action)) + { + return true; + } + } + + return false; + } + + /// + /// Whether any rpc: scope in the set allows calling the method on the service. + /// + public bool MatchesRpc(string lxm, string aud) + { + foreach (var value in _values) + { + if (RpcPermission.TryParse(value, out var p) && p!.Matches(lxm, aud)) + { + return true; + } + } + + return false; + } + + /// + /// Whether any blob: scope in the set allows uploading the MIME type. + /// + public bool MatchesBlob(string mime) + { + foreach (var value in _values) + { + if (BlobPermission.TryParse(value, out var p) && p!.Matches(mime)) + { + return true; + } + } + + return false; + } + + /// + /// Whether any account: scope in the set allows the action on the attribute. + /// + public bool MatchesAccount(AccountAttribute attribute, AccountActions action) + { + foreach (var value in _values) + { + if (AccountPermission.TryParse(value, out var p) && p!.Matches(attribute, action)) + { + return true; + } + } + + return false; + } + + /// + /// Whether any identity: scope in the set allows access to the attribute. + /// + public bool MatchesIdentity(IdentityAttribute attribute) + { + foreach (var value in _values) + { + if (IdentityPermission.TryParse(value, out var p) && p!.Matches(attribute)) + { + return true; + } + } + + return false; + } + + /// + /// Returns the normalized scope string: values normalized, invalid values dropped, sorted. + /// + public string ToNormalizedString() => AtprotoScope.Normalize(ToString()); + + /// + /// Returns the space-separated scope string, in insertion order. + /// + public override string ToString() => AtprotoScope.Join(_values); + + /// + public IEnumerator GetEnumerator() => _values.GetEnumerator(); + + /// + IEnumerator IEnumerable.GetEnumerator() => GetEnumerator(); + + private ScopeSet AddCore(string scope) + { + if (_set.Add(scope)) + { + _values.Add(scope); + } + + return this; + } +} diff --git a/src/CarpaNet.OAuth/Scopes/ScopeStringSyntax.cs b/src/CarpaNet.OAuth/Scopes/ScopeStringSyntax.cs new file mode 100644 index 0000000..537071f --- /dev/null +++ b/src/CarpaNet.OAuth/Scopes/ScopeStringSyntax.cs @@ -0,0 +1,416 @@ +using System; +using System.Collections.Generic; +using System.Text; + +namespace CarpaNet.OAuth.Scopes; + +/// +/// Parsed form of an atproto scope string: prefix[:positional][?key=value&...]. +/// Port of ScopeStringSyntax from @atproto/oauth-scopes. +/// +internal sealed class ScopeStringSyntax +{ + private readonly List>? _params; + + private ScopeStringSyntax(string prefix, string? positional, List>? parameters) + { + Prefix = prefix; + Positional = positional; + _params = parameters; + } + + /// + /// The scope prefix (resource name), e.g. repo. + /// + public string Prefix { get; } + + /// + /// The decoded positional parameter (after :), or null when absent. + /// + public string? Positional { get; } + + /// + /// Whether the scope string is for the given prefix, i.e. equals it or is followed by : or ?. + /// + public static bool IsScopeStringFor(string? value, string prefix) + { + if (value == null) + { + return false; + } + + if (value.Length > prefix.Length) + { + var next = value[prefix.Length]; + return (next == ':' || next == '?') && value.StartsWith(prefix, StringComparison.Ordinal); + } + + return string.Equals(value, prefix, StringComparison.Ordinal); + } + + /// + /// Parses a scope string. Returns null when the positional parameter contains malformed percent-encoding. + /// + public static ScopeStringSyntax? Parse(string scope) + { + var paramIdx = scope.IndexOf('?'); + var colonIdx = scope.IndexOf(':'); + var prefixEnd = paramIdx == -1 ? colonIdx : colonIdx == -1 ? paramIdx : Math.Min(paramIdx, colonIdx); + + if (prefixEnd == -1) + { + return new ScopeStringSyntax(scope, null, null); + } + + var prefix = scope.Substring(0, prefixEnd); + + string? positional = null; + if (colonIdx != -1 && (paramIdx == -1 || colonIdx < paramIdx)) + { + var end = paramIdx == -1 ? scope.Length : paramIdx; + positional = DecodeComponent(scope, colonIdx + 1, end); + if (positional == null) + { + return null; + } + } + + List>? parameters = null; + if (paramIdx != -1 && paramIdx < scope.Length - 1) + { + parameters = ParseQuery(scope, paramIdx + 1); + } + + return new ScopeStringSyntax(prefix, positional, parameters); + } + + /// + /// Whether every named parameter key is one of the allowed keys. + /// + public bool HasOnlyKeys(string key1, string? key2 = null) + { + if (_params == null) + { + return true; + } + + foreach (var kvp in _params) + { + if (!string.Equals(kvp.Key, key1, StringComparison.Ordinal) && + !string.Equals(kvp.Key, key2, StringComparison.Ordinal)) + { + return false; + } + } + + return true; + } + + /// + /// Gets all values of a named parameter, or null when absent. + /// + public List? GetMulti(string key) + { + if (_params == null) + { + return null; + } + + List? values = null; + foreach (var kvp in _params) + { + if (string.Equals(kvp.Key, key, StringComparison.Ordinal)) + { + (values ??= new List(1)).Add(kvp.Value); + } + } + + return values; + } + + /// + /// Gets a single-valued named parameter. + /// + /// False when the parameter is present more than once. + public bool TryGetSingle(string key, out string? value) + { + value = null; + if (_params == null) + { + return true; + } + + foreach (var kvp in _params) + { + if (string.Equals(kvp.Key, key, StringComparison.Ordinal)) + { + if (value != null) + { + return false; + } + + value = kvp.Value; + } + } + + return true; + } + + /// + /// Gets a single-valued parameter that may be given positionally or by name (but not both). + /// + /// False when the syntax is invalid for this parameter. + public bool TryGetPositionalSingle(string key, out string? value) + { + if (!TryGetSingle(key, out value)) + { + return false; + } + + if (value != null) + { + return Positional == null; + } + + value = Positional; + return true; + } + + /// + /// Gets a multi-valued parameter that may be given positionally (as a single value) or by name (but not both). + /// + /// False when the syntax is invalid for this parameter. + public bool TryGetPositionalMulti(string key, out List? values) + { + values = GetMulti(key); + if (values != null) + { + return Positional == null; + } + + if (Positional != null) + { + values = new List(1) { Positional }; + } + + return true; + } + + /// + /// Formats a scope string, encoding components the same way as the reference implementation + /// (encodeURIComponent for the positional parameter, URLSearchParams for named + /// parameters, then un-escaping : / + , @ %). + /// + public static string Format(string prefix, string? positional, IReadOnlyList>? parameters) + { + var sb = new StringBuilder(prefix.Length + 32); + sb.Append(prefix); + + if (positional != null) + { + sb.Append(':'); + AppendEncoded(sb, positional, form: false); + } + + if (parameters != null && parameters.Count > 0) + { + sb.Append('?'); + for (var i = 0; i < parameters.Count; i++) + { + if (i > 0) + { + sb.Append('&'); + } + + AppendEncoded(sb, parameters[i].Key, form: true); + sb.Append('='); + AppendEncoded(sb, parameters[i].Value, form: true); + } + } + + return sb.ToString(); + } + + private static List> ParseQuery(string scope, int start) + { + // application/x-www-form-urlencoded parsing, as done by URLSearchParams + var result = new List>(2); + var pos = start; + while (pos <= scope.Length) + { + var amp = scope.IndexOf('&', pos); + var end = amp == -1 ? scope.Length : amp; + + if (end > pos) + { + var eq = scope.IndexOf('=', pos, end - pos); + string key; + string value; + if (eq == -1) + { + key = DecodeForm(scope, pos, end); + value = string.Empty; + } + else + { + key = DecodeForm(scope, pos, eq); + value = DecodeForm(scope, eq + 1, end); + } + + result.Add(new KeyValuePair(key, value)); + } + + pos = end + 1; + } + + return result; + } + + private static string DecodeForm(string value, int start, int end) + { + var needsDecode = false; + for (var i = start; i < end; i++) + { + if (value[i] == '%' || value[i] == '+') + { + needsDecode = true; + break; + } + } + + var part = value.Substring(start, end - start); + if (!needsDecode) + { + return part; + } + + // Lenient: malformed escapes are kept as-is (like URLSearchParams) + return Uri.UnescapeDataString(part.Replace('+', ' ')); + } + + /// + /// Strict percent-decoding (like decodeURIComponent). Returns null on malformed input. + /// + private static string? DecodeComponent(string value, int start, int end) + { + var hasEscape = false; + for (var i = start; i < end; i++) + { + if (value[i] == '%') + { + if (i + 2 >= end || + !AtprotoSyntax.IsHexDigit(value[i + 1]) || + !AtprotoSyntax.IsHexDigit(value[i + 2])) + { + return null; + } + + hasEscape = true; + i += 2; + } + } + + var part = value.Substring(start, end - start); + return hasEscape ? Uri.UnescapeDataString(part) : part; + } + + private static void AppendEncoded(StringBuilder sb, string value, bool form) + { + for (var i = 0; i < value.Length; i++) + { + var c = value[i]; + if (IsUnencoded(c, form)) + { + sb.Append(c); + } + else if (form && c == ' ') + { + sb.Append('+'); + } + else + { + AppendPercentEncoded(sb, value, ref i); + } + } + } + + private static bool IsUnencoded(char c, bool form) + { + if (AtprotoSyntax.IsAsciiLetterOrDigit(c)) + { + return true; + } + + switch (c) + { + // Unreserved in both encodeURIComponent and URLSearchParams + case '-': + case '_': + case '.': + case '*': + // Chars the reference implementation normalizes back after encoding + case ':': + case '/': + case '+': + case ',': + case '@': + case '%': + return true; + // Unreserved in encodeURIComponent only + case '!': + case '~': + case '\'': + case '(': + case ')': + return !form; + default: + return false; + } + } + + private static void AppendPercentEncoded(StringBuilder sb, string value, ref int index) + { + int codePoint = value[index]; + + if (char.IsHighSurrogate(value[index]) && index + 1 < value.Length && char.IsLowSurrogate(value[index + 1])) + { + codePoint = char.ConvertToUtf32(value[index], value[index + 1]); + index++; + } + else if (char.IsSurrogate(value[index])) + { + codePoint = 0xFFFD; // Lone surrogate + } + + if (codePoint < 0x80) + { + AppendByte(sb, codePoint); + } + else if (codePoint < 0x800) + { + AppendByte(sb, 0xC0 | (codePoint >> 6)); + AppendByte(sb, 0x80 | (codePoint & 0x3F)); + } + else if (codePoint < 0x10000) + { + AppendByte(sb, 0xE0 | (codePoint >> 12)); + AppendByte(sb, 0x80 | ((codePoint >> 6) & 0x3F)); + AppendByte(sb, 0x80 | (codePoint & 0x3F)); + } + else + { + AppendByte(sb, 0xF0 | (codePoint >> 18)); + AppendByte(sb, 0x80 | ((codePoint >> 12) & 0x3F)); + AppendByte(sb, 0x80 | ((codePoint >> 6) & 0x3F)); + AppendByte(sb, 0x80 | (codePoint & 0x3F)); + } + } + + private static void AppendByte(StringBuilder sb, int b) + { + const string hex = "0123456789ABCDEF"; + sb.Append('%'); + sb.Append(hex[(b >> 4) & 0xF]); + sb.Append(hex[b & 0xF]); + } +} diff --git a/src/CarpaNet.SourceGen/Generation/ApiGenerator.cs b/src/CarpaNet.SourceGen/Generation/ApiGenerator.cs index c6a74f9..1963310 100644 --- a/src/CarpaNet.SourceGen/Generation/ApiGenerator.cs +++ b/src/CarpaNet.SourceGen/Generation/ApiGenerator.cs @@ -21,7 +21,8 @@ public static void GenerateParametersClass( string className, LexiconDefinition def, string currentNsid, - TypeRegistry registry) + TypeRegistry registry, + bool emitValidationAttributes = true) { if (def.Parameters == null || def.Parameters.Properties == null) { @@ -48,7 +49,8 @@ public static void GenerateParametersClass( currentNsid, registry, requiredProps, - new List()); + new List(), + emitValidationAttributes: emitValidationAttributes); } // Generate ToQueryParameters method for converting to query parameters @@ -92,6 +94,14 @@ private static void GenerateToQueryParametersMethod( { sb.AppendLine($"list.Add(new System.Collections.Generic.KeyValuePair(\"{jsonName}\", item));"); } + else if (prop.Value.Items?.Type == "integer") + { + sb.AppendLine($"list.Add(new System.Collections.Generic.KeyValuePair(\"{jsonName}\", item.ToString(System.Globalization.CultureInfo.InvariantCulture)));"); + } + else if (prop.Value.Items?.Type == "boolean") + { + sb.AppendLine($"list.Add(new System.Collections.Generic.KeyValuePair(\"{jsonName}\", item ? \"true\" : \"false\"));"); + } else { sb.AppendLine($"list.Add(new System.Collections.Generic.KeyValuePair(\"{jsonName}\", item?.ToString() ?? \"\"));"); @@ -213,7 +223,8 @@ public static void GenerateInputClass( string className, LexiconIO input, string currentNsid, - TypeRegistry registry) + TypeRegistry registry, + bool emitValidationAttributes = true) { if (input.Schema == null) { @@ -253,7 +264,8 @@ public static void GenerateInputClass( currentNsid, registry, requiredProps, - nullableProps); + nullableProps, + emitValidationAttributes: emitValidationAttributes); } } @@ -268,7 +280,8 @@ public static void GenerateOutputClass( string className, LexiconIO output, string currentNsid, - TypeRegistry registry) + TypeRegistry registry, + bool emitValidationAttributes = true) { if (output.Schema == null) { @@ -308,7 +321,8 @@ public static void GenerateOutputClass( currentNsid, registry, requiredProps, - nullableProps); + nullableProps, + emitValidationAttributes: emitValidationAttributes); } } @@ -439,8 +453,85 @@ public static bool IsInputRef(LexiconDefinition def) return null; } + /// + /// The JSON media type used by XRPC for structured bodies. + /// + private const string JsonEncoding = "application/json"; + + /// + /// Returns true if an XRPC encoding is application/json (ignoring case and media type parameters). + /// + public static bool IsJsonEncoding(string? encoding) + { + if (string.IsNullOrWhiteSpace(encoding)) + { + return false; + } + + var mediaType = encoding!; + var semicolon = mediaType.IndexOf(';'); + if (semicolon >= 0) + { + mediaType = mediaType.Substring(0, semicolon); + } + + return string.Equals(mediaType.Trim(), JsonEncoding, StringComparison.OrdinalIgnoreCase); + } + + /// + /// Returns true if a procedure takes a raw (non-JSON) request body: its input declares an encoding + /// that is not application/json, or declares an encoding without a schema + /// (e.g. com.atproto.repo.uploadBlob with */*). + /// + public static bool HasBinaryInput(LexiconDefinition def) + { + var input = def.Input; + if (input == null || string.IsNullOrWhiteSpace(input.Encoding)) + { + return false; + } + + return input.Schema == null || !IsJsonEncoding(input.Encoding); + } + + /// + /// Returns true if a query declares an output encoding other than application/json + /// (e.g. com.atproto.sync.getBlob with */*, or a CAR file). + /// + public static bool HasBinaryOutput(LexiconDefinition def) + { + var output = def.Output; + return output != null + && !string.IsNullOrWhiteSpace(output.Encoding) + && !IsJsonEncoding(output.Encoding); + } + + /// + /// Gets the default Content-Type for a binary input: the declared encoding when it is a concrete + /// media type, otherwise (wildcards such as */* or video/*) application/octet-stream. + /// + private static string GetDefaultContentType(LexiconIO input) + { + var encoding = input.Encoding?.Trim(); + if (string.IsNullOrEmpty(encoding) || encoding!.IndexOf('*') >= 0 || encoding.IndexOf('"') >= 0 || encoding.IndexOf('\\') >= 0) + { + return "application/octet-stream"; + } + + return encoding; + } + + /// + /// Returns true if the definition declares at least one query parameter. + /// + private static bool HasParameters(LexiconDefinition def) + { + return def.Parameters?.Properties != null && def.Parameters.Properties.Count > 0; + } + /// /// Generates extension method for a query. + /// Queries with a non-JSON output encoding return the raw response body as byte[]. /// public static void GenerateQueryExtension( SourceBuilder sb, @@ -450,13 +541,19 @@ public static void GenerateQueryExtension( string currentNamespace, TypeRegistry registry) { - var hasParameters = def.Parameters?.Properties != null && def.Parameters.Properties.Count > 0; + var hasParameters = HasParameters(def); var hasOutput = def.Output?.Schema != null; + var isBinaryOutput = HasBinaryOutput(def); var proxyServiceDid = GetProxyServiceDid(currentNsid); var outputType = hasOutput ? ResolveOutputType(def, currentNsid, currentNamespace, className, registry) : "object"; // Use full namespace path in method name to avoid collisions var methodName = $"{currentNamespace.Replace(".", "")}{className}Async"; + if (isBinaryOutput) + { + outputType = "byte[]"; + } + sb.WriteSummary(def.Description); sb.AppendLine($"public static async System.Threading.Tasks.Task<{outputType}> {methodName}("); sb.Indent(); @@ -472,7 +569,19 @@ public static void GenerateQueryExtension( sb.OpenBrace(); - if (proxyServiceDid != null) + if (isBinaryOutput) + { + // Non-JSON output (blob, CAR file, ...): return the raw response body + sb.AppendLine("return await global::CarpaNet.ATProtoClientXrpcExtensions.GetBytesAsync("); + sb.Indent(); + sb.AppendLine("client,"); + sb.AppendLine($"\"{currentNsid}\","); + sb.AppendLine($"{proxyServiceDid ?? "null"},"); + sb.AppendLine(hasParameters ? "parameters?.ToQueryParameters()," : "null,"); + sb.AppendLine("cancellationToken);"); + sb.Unindent(); + } + else if (proxyServiceDid != null) { // Use the proxy overload for services that require proxying sb.AppendLine($"return await client.GetAsync<{outputType}>("); @@ -498,6 +607,11 @@ public static void GenerateQueryExtension( /// /// Generates extension method for a procedure. + /// + /// Binary input (see ): takes a Stream body and content type and calls PostBinaryAsync. + /// JSON (or no) input with query parameters: adds a parameters argument and calls PostWithParametersAsync. + /// Otherwise: JSON input via IATProtoClient.PostAsync. + /// /// public static void GenerateProcedureExtension( SourceBuilder sb, @@ -507,30 +621,70 @@ public static void GenerateProcedureExtension( string currentNamespace, TypeRegistry registry) { - var hasInput = def.Input?.Schema != null; + var isBinaryInput = HasBinaryInput(def); + var hasInput = !isBinaryInput && def.Input?.Schema != null; var hasOutput = def.Output?.Schema != null; + var hasParameters = HasParameters(def); var proxyServiceDid = GetProxyServiceDid(currentNsid); // Use full namespace path in method name to avoid collisions var methodName = $"{currentNamespace.Replace(".", "")}{className}Async"; var returnType = hasOutput ? ResolveOutputType(def, currentNsid, currentNamespace, className, registry) : "object"; var inputType = hasInput ? ResolveInputType(def, currentNsid, currentNamespace, className, registry) : "object"; + var parametersType = $"{currentNamespace}.{className}Parameters"; sb.WriteSummary(def.Description); sb.AppendLine($"public static async System.Threading.Tasks.Task<{returnType}> {methodName}("); sb.Indent(); sb.AppendLine("this CarpaNet.IATProtoClient client,"); - if (hasInput) + if (isBinaryInput) + { + sb.AppendLine("System.IO.Stream body,"); + sb.AppendLine($"string contentType = \"{GetDefaultContentType(def.Input!)}\","); + } + else if (hasInput) { sb.AppendLine($"{inputType} input,"); } + if (hasParameters) + { + sb.AppendLine($"{parametersType}? parameters = null,"); + } + sb.AppendLine("System.Threading.CancellationToken cancellationToken = default)"); sb.Unindent(); sb.OpenBrace(); - if (proxyServiceDid != null) + if (isBinaryInput) + { + // Raw request body (blob, video, ...) sent with the given Content-Type + sb.AppendLine($"return await global::CarpaNet.ATProtoClientXrpcExtensions.PostBinaryAsync<{returnType}>("); + sb.Indent(); + sb.AppendLine("client,"); + sb.AppendLine($"\"{currentNsid}\","); + sb.AppendLine($"{proxyServiceDid ?? "null"},"); + sb.AppendLine(hasParameters ? "parameters?.ToQueryParameters()," : "null,"); + sb.AppendLine("body,"); + sb.AppendLine("contentType,"); + sb.AppendLine("cancellationToken);"); + sb.Unindent(); + } + else if (hasParameters) + { + // JSON body plus query string parameters + sb.AppendLine($"return await global::CarpaNet.ATProtoClientXrpcExtensions.PostWithParametersAsync<{inputType}, {returnType}>("); + sb.Indent(); + sb.AppendLine("client,"); + sb.AppendLine($"\"{currentNsid}\","); + sb.AppendLine($"{proxyServiceDid ?? "null"},"); + sb.AppendLine("parameters?.ToQueryParameters(),"); + sb.AppendLine(hasInput ? "input," : "null,"); + sb.AppendLine("cancellationToken);"); + sb.Unindent(); + } + else if (proxyServiceDid != null) { // Use the proxy overload for services that require proxying sb.AppendLine($"return await client.PostAsync<{inputType}, {returnType}>("); diff --git a/src/CarpaNet.SourceGen/Generation/CborContextGenerator.cs b/src/CarpaNet.SourceGen/Generation/CborContextGenerator.cs index b90aa09..f82aa1d 100644 --- a/src/CarpaNet.SourceGen/Generation/CborContextGenerator.cs +++ b/src/CarpaNet.SourceGen/Generation/CborContextGenerator.cs @@ -184,9 +184,31 @@ public static void GenerateCborTypeInfo( sb.AppendLine(); } + /// + /// Marker added to the generated-types set when a union type info that reads in place is emitted, + /// signalling that the shared helper must be generated. + /// + public const string UnionHelperMarker = ""; + + /// + /// Name of the internal helper class emitted into the CBOR context namespace for data unions. + /// + public const string UnionHelperClassName = "UnionCborHelper"; + /// /// Generates a CborUnionTypeInfo subclass for a union type. /// + /// + /// When true (unions inside data), members with a known $type are read directly from the caller's + /// reader, so the reader stays positioned inside its enclosing map or array and following values can + /// still be read. Subscription message unions pass false: their bodies carry no $type (the frame + /// header does), so they keep the base class behavior. + /// + /// + /// When true (open unions in data, requires ), any non-null value + /// that is not a member with a known $type (unknown or missing $type) is read into the + /// union's Unknown_* class with its raw DAG-CBOR bytes, and written back unchanged. + /// public static void GenerateCborUnionTypeInfo( SourceBuilder sb, string qualifiedTypeName, @@ -195,7 +217,9 @@ public static void GenerateCborUnionTypeInfo( string currentNsid, TypeRegistry registry, GeneratorOptions options, - HashSet? generatedTypes = null) + HashSet? generatedTypes = null, + bool dispatchInPlace = false, + bool preserveUnknown = false) { generatedTypes ??= new HashSet(); @@ -231,6 +255,247 @@ public static void GenerateCborUnionTypeInfo( sb.AppendLine("protected override System.Collections.Generic.IReadOnlyDictionary DerivedTypes => _derivedTypes;"); + if (dispatchInPlace && refs.Count > 0) + { + generatedTypes.Add(UnionHelperMarker); + var unknownType = preserveUnknown + ? ResolveToGlobalType(UnionGenerator.GetUnknownTypeName(qualifiedTypeName)) + : null; + + sb.AppendLine(); + sb.WriteSummary(unknownType != null + ? "Reads a union member in place; members with an unknown $type are kept as raw DAG-CBOR." + : "Reads a union member in place, so the reader stays positioned within its enclosing container."); + sb.AppendLine($"public override {globalType}? Read(ref CarpaNet.Cbor.DagCborReader reader)"); + sb.OpenBrace(); + sb.AppendLine("var state = reader.PeekState();"); + sb.AppendLine("string? discriminator = null;"); + sb.AppendLine("if (state == System.Formats.Cbor.CborReaderState.StartMap)"); + sb.OpenBrace(); + sb.AppendLine($"discriminator = {UnionHelperClassName}.PeekTypeDiscriminator(reader.GetRemainingData());"); + sb.AppendLine("if (discriminator != null && _derivedTypes.TryGetValue(discriminator, out var typeInfo))"); + sb.OpenBrace(); + sb.AppendLine($"return ({globalType}?)typeInfo.ReadObject(ref reader);"); + sb.CloseBrace(); + sb.CloseBrace(); + sb.AppendLine(); + + if (unknownType != null) + { + // Open union: anything other than null or a known member is kept verbatim + sb.AppendLine("if (state != System.Formats.Cbor.CborReaderState.Null)"); + sb.OpenBrace(); + sb.AppendLine($"var rawCbor = {UnionHelperClassName}.ReadRawValue(ref reader);"); + sb.AppendLine($"return new {unknownType}(discriminator ?? string.Empty, {UnionHelperClassName}.ToJson(rawCbor), rawCbor);"); + sb.CloseBrace(); + sb.AppendLine(); + } + + sb.AppendLine("return base.Read(ref reader);"); + sb.CloseBrace(); + + sb.AppendLine(); + sb.WriteSummary(unknownType != null + ? "Writes a union member with its $type; unknown members are written back from their raw data." + : "Writes a union member with its $type discriminator."); + sb.AppendLine($"public override void Write(ref CarpaNet.Cbor.DagCborWriter writer, {globalType}? value)"); + sb.OpenBrace(); + + if (unknownType != null) + { + sb.AppendLine($"if (value is {unknownType} unknown)"); + sb.OpenBrace(); + sb.AppendLine($"{UnionHelperClassName}.WriteRawValue(ref writer, unknown.Raw, unknown.RawCbor);"); + sb.AppendLine("return;"); + sb.CloseBrace(); + sb.AppendLine(); + } + + // Non-record members do not write $type themselves, but union members must carry it + sb.AppendLine("if (value != null)"); + sb.OpenBrace(); + sb.AppendLine("var runtimeType = value.GetType();"); + sb.AppendLine("foreach (var kvp in _derivedTypes)"); + sb.OpenBrace(); + sb.AppendLine("if (kvp.Value.TargetType == runtimeType && kvp.Value.TypeDiscriminator == null)"); + sb.OpenBrace(); + sb.AppendLine($"{UnionHelperClassName}.WriteWithTypeDiscriminator(ref writer, kvp.Key, kvp.Value, value);"); + sb.AppendLine("return;"); + sb.CloseBrace(); + sb.CloseBrace(); + sb.CloseBrace(); + sb.AppendLine(); + sb.AppendLine("base.Write(ref writer, value);"); + sb.CloseBrace(); + } + + sb.CloseBrace(); + sb.AppendLine(); + } + + /// + /// Generates the internal helper used by data union type infos to peek $type, + /// add $type to members that do not write it, capture raw DAG-CBOR values, + /// and copy them back out unchanged. + /// + public static void GenerateUnionHelper(SourceBuilder sb) + { + sb.WriteSummary("Helpers for reading union members in place and preserving unknown open-union members in DAG-CBOR."); + sb.AppendLine($"internal static class {UnionHelperClassName}"); + sb.OpenBrace(); + + // PeekTypeDiscriminator + sb.WriteSummary("Returns the $type text value of the CBOR map at the start of the data, or null if it has none."); + sb.AppendLine("public static string? PeekTypeDiscriminator(System.ReadOnlyMemory data)"); + sb.OpenBrace(); + sb.AppendLine("var reader = new CarpaNet.Cbor.DagCborReader(data);"); + sb.AppendLine("var count = reader.ReadStartMap();"); + sb.AppendLine("var remaining = count ?? int.MaxValue;"); + sb.AppendLine("while (remaining > 0 && reader.PeekState() != System.Formats.Cbor.CborReaderState.EndMap)"); + sb.OpenBrace(); + sb.AppendLine("if (reader.PeekState() != System.Formats.Cbor.CborReaderState.TextString)"); + sb.OpenBrace(); + sb.AppendLine("return null;"); + sb.CloseBrace(); + sb.AppendLine(); + sb.AppendLine("var key = reader.ReadTextString();"); + sb.AppendLine("if (key == \"$type\")"); + sb.OpenBrace(); + sb.AppendLine("return reader.PeekState() == System.Formats.Cbor.CborReaderState.TextString ? reader.ReadTextString() : null;"); + sb.CloseBrace(); + sb.AppendLine(); + sb.AppendLine("reader.SkipValue();"); + sb.AppendLine("remaining--;"); + sb.CloseBrace(); + sb.AppendLine(); + sb.AppendLine("return null;"); + sb.CloseBrace(); + sb.AppendLine(); + + // ReadRawValue + sb.WriteSummary("Reads the next complete value and returns its encoded bytes."); + sb.AppendLine("public static byte[] ReadRawValue(ref CarpaNet.Cbor.DagCborReader reader)"); + sb.OpenBrace(); + sb.AppendLine("var data = reader.GetRemainingData();"); + sb.AppendLine("var start = reader.BytesRead;"); + sb.AppendLine("reader.SkipValue();"); + sb.AppendLine("return data.Slice(0, reader.BytesRead - start).ToArray();"); + sb.CloseBrace(); + sb.AppendLine(); + + // ToJson + sb.WriteSummary("Converts an encoded CBOR value to JSON (CID links and byte strings become strings)."); + sb.AppendLine("public static System.Text.Json.JsonElement ToJson(byte[] rawCbor)"); + sb.OpenBrace(); + sb.AppendLine("var reader = new CarpaNet.Cbor.DagCborReader(rawCbor);"); + sb.AppendLine("return new CarpaNet.Cbor.Converters.JsonElementCborConverter().ReadTyped(ref reader);"); + sb.CloseBrace(); + sb.AppendLine(); + + // WriteRawValue + sb.WriteSummary("Writes an unknown member: its original CBOR bytes when available, otherwise its JSON value converted to CBOR."); + sb.AppendLine("public static void WriteRawValue(ref CarpaNet.Cbor.DagCborWriter writer, System.Text.Json.JsonElement raw, byte[]? rawCbor)"); + sb.OpenBrace(); + sb.AppendLine("if (rawCbor != null)"); + sb.OpenBrace(); + sb.AppendLine("var reader = new CarpaNet.Cbor.DagCborReader(rawCbor);"); + sb.AppendLine("CopyValue(ref reader, ref writer);"); + sb.AppendLine("return;"); + sb.CloseBrace(); + sb.AppendLine(); + sb.AppendLine("new CarpaNet.Cbor.Converters.JsonElementCborConverter().WriteTyped(ref writer, raw);"); + sb.CloseBrace(); + sb.AppendLine(); + + // WriteWithTypeDiscriminator + sb.WriteSummary("Writes an object through its type info, adding a leading $type entry to the map."); + sb.AppendLine("public static void WriteWithTypeDiscriminator(ref CarpaNet.Cbor.DagCborWriter writer, string discriminator, CarpaNet.Cbor.ICborTypeInfo typeInfo, object value)"); + sb.OpenBrace(); + sb.AppendLine("var temp = new CarpaNet.Cbor.DagCborWriter();"); + sb.AppendLine("typeInfo.WriteObject(ref temp, value);"); + sb.AppendLine("var reader = new CarpaNet.Cbor.DagCborReader(temp.Encode());"); + sb.AppendLine("if (reader.PeekState() != System.Formats.Cbor.CborReaderState.StartMap)"); + sb.OpenBrace(); + sb.AppendLine("CopyValue(ref reader, ref writer);"); + sb.AppendLine("return;"); + sb.CloseBrace(); + sb.AppendLine(); + sb.AppendLine("var count = reader.ReadStartMap();"); + sb.AppendLine("writer.WriteStartMap(count + 1);"); + sb.AppendLine("writer.WriteTextString(\"$type\");"); + sb.AppendLine("writer.WriteTextString(discriminator);"); + sb.AppendLine("while (reader.PeekState() != System.Formats.Cbor.CborReaderState.EndMap)"); + sb.OpenBrace(); + sb.AppendLine("CopyValue(ref reader, ref writer); // key"); + sb.AppendLine("CopyValue(ref reader, ref writer); // value"); + sb.CloseBrace(); + sb.AppendLine("reader.ReadEndMap();"); + sb.AppendLine("writer.WriteEndMap();"); + sb.CloseBrace(); + sb.AppendLine(); + + // CopyValue + sb.WriteSummary("Copies one CBOR value (recursively) from the reader to the writer."); + sb.AppendLine("private static void CopyValue(ref CarpaNet.Cbor.DagCborReader reader, ref CarpaNet.Cbor.DagCborWriter writer)"); + sb.OpenBrace(); + sb.AppendLine("switch (reader.PeekState())"); + sb.OpenBrace(); + sb.AppendLine("case System.Formats.Cbor.CborReaderState.UnsignedInteger:"); + sb.AppendLine(" writer.WriteUInt64(reader.ReadUInt64());"); + sb.AppendLine(" break;"); + sb.AppendLine("case System.Formats.Cbor.CborReaderState.NegativeInteger:"); + sb.AppendLine(" writer.WriteInt64(reader.ReadInt64());"); + sb.AppendLine(" break;"); + sb.AppendLine("case System.Formats.Cbor.CborReaderState.HalfPrecisionFloat:"); + sb.AppendLine("case System.Formats.Cbor.CborReaderState.SinglePrecisionFloat:"); + sb.AppendLine("case System.Formats.Cbor.CborReaderState.DoublePrecisionFloat:"); + sb.AppendLine(" writer.WriteDouble(reader.ReadDouble());"); + sb.AppendLine(" break;"); + sb.AppendLine("case System.Formats.Cbor.CborReaderState.TextString:"); + sb.AppendLine(" writer.WriteTextString(reader.ReadTextString());"); + sb.AppendLine(" break;"); + sb.AppendLine("case System.Formats.Cbor.CborReaderState.ByteString:"); + sb.AppendLine(" writer.WriteByteString(reader.ReadByteString());"); + sb.AppendLine(" break;"); + sb.AppendLine("case System.Formats.Cbor.CborReaderState.Boolean:"); + sb.AppendLine(" writer.WriteBoolean(reader.ReadBoolean());"); + sb.AppendLine(" break;"); + sb.AppendLine("case System.Formats.Cbor.CborReaderState.Null:"); + sb.AppendLine(" reader.ReadNull();"); + sb.AppendLine(" writer.WriteNull();"); + sb.AppendLine(" break;"); + sb.AppendLine("case System.Formats.Cbor.CborReaderState.Tag:"); + sb.AppendLine(" writer.WriteTag(reader.ReadTag());"); + sb.AppendLine(" CopyValue(ref reader, ref writer);"); + sb.AppendLine(" break;"); + sb.AppendLine("case System.Formats.Cbor.CborReaderState.StartArray:"); + sb.OpenBrace(); + sb.AppendLine("writer.WriteStartArray(reader.ReadStartArray());"); + sb.AppendLine("while (reader.PeekState() != System.Formats.Cbor.CborReaderState.EndArray)"); + sb.OpenBrace(); + sb.AppendLine("CopyValue(ref reader, ref writer);"); + sb.CloseBrace(); + sb.AppendLine("reader.ReadEndArray();"); + sb.AppendLine("writer.WriteEndArray();"); + sb.AppendLine("break;"); + sb.CloseBrace(); + sb.AppendLine("case System.Formats.Cbor.CborReaderState.StartMap:"); + sb.OpenBrace(); + sb.AppendLine("writer.WriteStartMap(reader.ReadStartMap());"); + sb.AppendLine("while (reader.PeekState() != System.Formats.Cbor.CborReaderState.EndMap)"); + sb.OpenBrace(); + sb.AppendLine("CopyValue(ref reader, ref writer); // key"); + sb.AppendLine("CopyValue(ref reader, ref writer); // value"); + sb.CloseBrace(); + sb.AppendLine("reader.ReadEndMap();"); + sb.AppendLine("writer.WriteEndMap();"); + sb.AppendLine("break;"); + sb.CloseBrace(); + sb.AppendLine("default:"); + sb.AppendLine(" throw new System.InvalidOperationException($\"Unsupported CBOR state when copying an unknown union member: {reader.PeekState()}\");"); + sb.CloseBrace(); + sb.CloseBrace(); + sb.CloseBrace(); sb.AppendLine(); } @@ -270,7 +535,7 @@ private static void GenerateNestedTypeInfosForProperty( var interfaceShort = $"I{cleanParent}{cleanProp}"; var interfaceQualified = string.IsNullOrEmpty(ns) ? interfaceShort : $"{ns}.{interfaceShort}"; var interfaceSuffix = ToClassSuffix(interfaceQualified); - GenerateCborUnionTypeInfo(sb, interfaceQualified, interfaceSuffix, prop.Refs, currentNsid, registry, options, generatedTypes); + GenerateCborUnionTypeInfo(sb, interfaceQualified, interfaceSuffix, prop.Refs, currentNsid, registry, options, generatedTypes, dispatchInPlace: true, preserveUnknown: prop.Closed != true); } // Handle arrays @@ -297,7 +562,7 @@ private static void GenerateNestedTypeInfosForProperty( var interfaceShort = $"I{cleanParent}{cleanProp}"; var interfaceQualified = string.IsNullOrEmpty(ns) ? interfaceShort : $"{ns}.{interfaceShort}"; var interfaceSuffix = ToClassSuffix(interfaceQualified); - GenerateCborUnionTypeInfo(sb, interfaceQualified, interfaceSuffix, prop.Items.Refs, currentNsid, registry, options, generatedTypes); + GenerateCborUnionTypeInfo(sb, interfaceQualified, interfaceSuffix, prop.Items.Refs, currentNsid, registry, options, generatedTypes, dispatchInPlace: true, preserveUnknown: prop.Items.Closed != true); } } } diff --git a/src/CarpaNet.SourceGen/Generation/JsonContextGenerator.cs b/src/CarpaNet.SourceGen/Generation/JsonContextGenerator.cs index d1509b3..41228b7 100644 --- a/src/CarpaNet.SourceGen/Generation/JsonContextGenerator.cs +++ b/src/CarpaNet.SourceGen/Generation/JsonContextGenerator.cs @@ -202,8 +202,8 @@ public static void GenerateJsonTypeInfo( /// /// Generates a CreateTypeInfo factory method for a union interface with polymorphism options. - /// For open unions (isClosed=false), generates a custom JsonConverter that gracefully - /// returns null for unknown $type discriminator values instead of throwing. + /// For open unions (isClosed=false), generates a custom JsonConverter that reads unknown + /// $type discriminator values into the union's Unknown_* class instead of throwing. /// public static void GenerateJsonUnionTypeInfo( SourceBuilder sb, @@ -253,7 +253,7 @@ public static void GenerateJsonUnionTypeInfo( } else { - // Open unions: generate a custom converter that returns null for unknown $type values + // Open unions: generate a custom converter that preserves unknown $type values GenerateOpenUnionConverter(sb, qualifiedTypeName, methodSuffix, refs, currentNsid, registry); sb.AppendLine($"private static global::System.Text.Json.Serialization.Metadata.JsonTypeInfo Create_{methodSuffix}_TypeInfo(global::System.Text.Json.JsonSerializerOptions options)"); @@ -265,7 +265,11 @@ public static void GenerateJsonUnionTypeInfo( } /// - /// Generates a JsonConverter class for an open union interface that gracefully handles unknown $type values. + /// Generates a JsonConverter class for an open union interface. + /// Members with an unknown or missing $type are read into the generated Unknown_* class + /// (see ) and written back verbatim, so that + /// round-tripping data written by newer clients is lossless. Known members are written with + /// their $type discriminator. /// private static void GenerateOpenUnionConverter( SourceBuilder sb, @@ -275,6 +279,30 @@ private static void GenerateOpenUnionConverter( string currentNsid, TypeRegistry registry) { + var unknownTypeName = UnionGenerator.GetUnknownTypeName(qualifiedTypeName); + + // Resolve each member once: (C# type, $type discriminator). Duplicates would produce + // unreachable (compile-error) switch arms, so keep only the first occurrence. + var readMembers = new List<(string TypeName, string Discriminator)>(); + var writeMembers = new List<(string TypeName, string Discriminator)>(); + var seenTypes = new HashSet(); + var seenDiscriminators = new HashSet(); + foreach (var refString in refs) + { + var typeName = registry.ResolveToCSharpType(refString, currentNsid); + var discriminator = GetTypeDiscriminator(refString, currentNsid, registry); + if (seenDiscriminators.Add(discriminator)) + { + readMembers.Add((typeName, discriminator)); + } + + // Type patterns are only valid for generated classes (which implement the interface) + if (registry.RefGeneratesClass(refString, currentNsid) && seenTypes.Add(typeName)) + { + writeMembers.Add((typeName, discriminator)); + } + } + sb.AppendLine($"private sealed class Converter_{methodSuffix} : global::System.Text.Json.Serialization.JsonConverter"); sb.OpenBrace(); @@ -282,20 +310,24 @@ private static void GenerateOpenUnionConverter( sb.AppendLine($"public override global::{qualifiedTypeName}? Read(ref global::System.Text.Json.Utf8JsonReader reader, global::System.Type typeToConvert, global::System.Text.Json.JsonSerializerOptions options)"); sb.OpenBrace(); sb.AppendLine("var element = global::System.Text.Json.JsonElement.ParseValue(ref reader);"); - sb.AppendLine("if (!element.TryGetProperty(\"$type\", out var typeProp))"); - sb.AppendLine(" return null;"); - sb.AppendLine("var typeStr = typeProp.GetString();"); + sb.AppendLine("string? typeStr = null;"); + sb.AppendLine("if (element.ValueKind == global::System.Text.Json.JsonValueKind.Object"); + sb.AppendLine(" && element.TryGetProperty(\"$type\", out var typeProp)"); + sb.AppendLine(" && typeProp.ValueKind == global::System.Text.Json.JsonValueKind.String)"); + sb.OpenBrace(); + sb.AppendLine("typeStr = typeProp.GetString();"); + sb.CloseBrace(); + sb.AppendLine(); sb.AppendLine("return typeStr switch"); sb.OpenBrace(); - foreach (var refString in refs) + foreach (var (typeName, discriminator) in readMembers) { - var typeName = registry.ResolveToCSharpType(refString, currentNsid); - var discriminator = GetTypeDiscriminator(refString, currentNsid, registry); sb.AppendLine($"\"{discriminator}\" => (global::{qualifiedTypeName}?)global::System.Text.Json.JsonSerializer.Deserialize(element, options.GetTypeInfo(typeof(global::{typeName}))),"); } - sb.AppendLine("_ => null,"); + // Unknown or missing $type: keep the whole value so it can be written back unchanged + sb.AppendLine($"_ => new global::{unknownTypeName}(typeStr ?? string.Empty, element),"); sb.CloseBrace(withSemicolon: true); sb.CloseBrace(); sb.AppendLine(); @@ -303,7 +335,40 @@ private static void GenerateOpenUnionConverter( // Write method sb.AppendLine($"public override void Write(global::System.Text.Json.Utf8JsonWriter writer, global::{qualifiedTypeName} value, global::System.Text.Json.JsonSerializerOptions options)"); sb.OpenBrace(); - sb.AppendLine("global::System.Text.Json.JsonSerializer.Serialize(writer, value, options.GetTypeInfo(value.GetType()));"); + sb.AppendLine($"if (value is global::{unknownTypeName} unknown)"); + sb.OpenBrace(); + sb.AppendLine("unknown.Raw.WriteTo(writer);"); + sb.AppendLine("return;"); + sb.CloseBrace(); + sb.AppendLine(); + + // Known members: emit $type first, then the member's own properties + sb.AppendLine("string? typeId = value switch"); + sb.OpenBrace(); + foreach (var (typeName, discriminator) in writeMembers) + { + sb.AppendLine($"global::{typeName} => \"{discriminator}\","); + } + sb.AppendLine("_ => null,"); + sb.CloseBrace(withSemicolon: true); + sb.AppendLine(); + sb.AppendLine("var element = global::System.Text.Json.JsonSerializer.SerializeToElement(value, options.GetTypeInfo(value.GetType()));"); + sb.AppendLine("if (typeId == null || element.ValueKind != global::System.Text.Json.JsonValueKind.Object)"); + sb.OpenBrace(); + sb.AppendLine("element.WriteTo(writer);"); + sb.AppendLine("return;"); + sb.CloseBrace(); + sb.AppendLine(); + sb.AppendLine("writer.WriteStartObject();"); + sb.AppendLine("writer.WriteString(\"$type\", typeId);"); + sb.AppendLine("foreach (var property in element.EnumerateObject())"); + sb.OpenBrace(); + sb.AppendLine("if (!property.NameEquals(\"$type\"))"); + sb.OpenBrace(); + sb.AppendLine("property.WriteTo(writer);"); + sb.CloseBrace(); + sb.CloseBrace(); + sb.AppendLine("writer.WriteEndObject();"); sb.CloseBrace(); sb.CloseBrace(); diff --git a/src/CarpaNet.SourceGen/Generation/ObjectGenerator.cs b/src/CarpaNet.SourceGen/Generation/ObjectGenerator.cs index 6520314..a7a80bf 100644 --- a/src/CarpaNet.SourceGen/Generation/ObjectGenerator.cs +++ b/src/CarpaNet.SourceGen/Generation/ObjectGenerator.cs @@ -25,7 +25,8 @@ public static void GenerateClass( TypeRegistry registry, bool isRecord = false, string? recordType = null, - string? typeId = null) + string? typeId = null, + bool emitValidationAttributes = true) { var requiredProps = def.Required ?? new List(); var nullableProps = def.Nullable ?? new List(); @@ -63,7 +64,7 @@ public static void GenerateClass( foreach (var prop in properties) { - GenerateProperty(sb, className, prop.Key, prop.Value, currentNsid, registry, requiredProps, nullableProps, isRecord); + GenerateProperty(sb, className, prop.Key, prop.Value, currentNsid, registry, requiredProps, nullableProps, isRecord, emitValidationAttributes); } sb.CloseBrace(); @@ -112,7 +113,8 @@ public static void GenerateProperty( TypeRegistry registry, List requiredProps, List nullableProps, - bool isRecord = false) + bool isRecord = false, + bool emitValidationAttributes = true) { var isRequired = requiredProps.Contains(propertyName) || def.IsRequired; var isNullable = nullableProps.Contains(propertyName) || !isRequired; @@ -140,8 +142,11 @@ public static void GenerateProperty( // JSON property name attribute sb.WriteAttribute($"System.Text.Json.Serialization.JsonPropertyName(\"{propertyName}\")"); - // Add validation attributes - WriteValidationAttributes(sb, def); + // Add validation attributes (opt-out via CarpaNet_EmitValidationAttributes=false) + if (emitValidationAttributes) + { + WriteValidationAttributes(sb, def); + } // Get the C# type var typeName = GetPropertyType(def, currentNsid, registry, className, propertyName); diff --git a/src/CarpaNet.SourceGen/Generation/UnionGenerator.cs b/src/CarpaNet.SourceGen/Generation/UnionGenerator.cs index 74f41c4..2a53a3c 100644 --- a/src/CarpaNet.SourceGen/Generation/UnionGenerator.cs +++ b/src/CarpaNet.SourceGen/Generation/UnionGenerator.cs @@ -16,7 +16,8 @@ public static class UnionGenerator /// /// Generates a union interface with JsonPolymorphic attributes. /// For open unions, no polymorphic attributes are emitted since a custom JsonConverter - /// in the generated JSON context handles unknown $type values gracefully. + /// in the generated JSON context handles unknown $type values; an Unknown_* class + /// (see ) is emitted next to the interface to hold them. /// public static void GenerateUnionInterface( SourceBuilder sb, @@ -49,6 +50,125 @@ public static void GenerateUnionInterface( sb.AppendLine($"public interface {interfaceName}"); sb.OpenBrace(); sb.CloseBrace(); + + // Open unions with at least one known member get a catch-all class for members whose + // $type this code does not know, so they survive a read-modify-write round trip. + if (HasUnknownMemberType(def)) + { + sb.AppendLine(); + GenerateUnknownMemberClass(sb, interfaceName); + } + } + + /// + /// Returns true when a union definition gets a generated Unknown_* member class: + /// open unions (no closed: true) that list at least one ref. Open unions without refs + /// are mapped to JsonElement directly and have no interface-based converter. + /// + public static bool HasUnknownMemberType(LexiconDefinition def) + { + return def.Closed != true && def.Refs != null && def.Refs.Count > 0; + } + + /// + /// Gets the name of the class that holds unknown members of an open union. + /// + /// + /// Naming rule: strip the leading I from the interface name and prefix Unknown_ + /// (IDefsPreferences becomes Unknown_DefsPreferences). Lexicon-derived type names + /// are PascalCase with _ and - removed, so they only ever contain an underscore as + /// their first character; a name with an underscore after Unknown cannot collide with them. + /// Accepts either a short or a namespace-qualified interface name and keeps the namespace. + /// + public static string GetUnknownTypeName(string interfaceName) + { + var lastDot = interfaceName.LastIndexOf('.'); + var ns = lastDot >= 0 ? interfaceName.Substring(0, lastDot + 1) : string.Empty; + var shortName = NsidHelper.StripEscapePrefix(lastDot >= 0 ? interfaceName.Substring(lastDot + 1) : interfaceName); + + if (shortName.Length > 1 && shortName[0] == 'I') + { + shortName = shortName.Substring(1); + } + + return $"{ns}Unknown_{shortName}"; + } + + /// + /// Generates the sealed class that represents an open-union member with an unrecognized + /// (or missing) $type. It keeps the raw JSON so that it can be written back verbatim. + /// + private static void GenerateUnknownMemberClass(SourceBuilder sb, string interfaceName) + { + var className = GetUnknownTypeName(interfaceName); + + sb.AppendLine("/// "); + sb.AppendLine($"/// A member of the open union whose $type is not known to this generated code."); + sb.AppendLine("/// The original data is kept in and is written back unchanged on serialization,"); + sb.AppendLine("/// so data from newer lexicon versions is not lost in a read-modify-write cycle."); + sb.AppendLine("/// "); + sb.AppendLine($"public sealed class {className} : {interfaceName}"); + sb.OpenBrace(); + + sb.AppendLine("/// "); + sb.AppendLine($"/// Initializes a new instance of the class."); + sb.AppendLine("/// "); + sb.AppendLine("/// The $type value of the member, or an empty string if it had none."); + sb.AppendLine("/// The complete JSON value of the member, including $type. It is cloned, so it can outlive its source document."); + sb.AppendLine("/// The original DAG-CBOR encoding of the member, when it was read from CBOR; otherwise ."); + sb.AppendLine("/// is an undefined (default) element."); + sb.AppendLine($"public {className}(string type, System.Text.Json.JsonElement raw, byte[]? rawCbor = null)"); + sb.OpenBrace(); + sb.AppendLine("if (raw.ValueKind == System.Text.Json.JsonValueKind.Undefined)"); + sb.OpenBrace(); + sb.AppendLine("throw new System.ArgumentException(\"The raw JSON value must not be undefined.\", nameof(raw));"); + sb.CloseBrace(); + sb.AppendLine(); + sb.AppendLine("Type = type ?? string.Empty;"); + sb.AppendLine("Raw = raw.Clone();"); + sb.AppendLine("RawCbor = rawCbor;"); + sb.CloseBrace(); + sb.AppendLine(); + + sb.AppendLine("/// "); + sb.AppendLine("/// Gets the $type value of the member, or an empty string if the member had no $type."); + sb.AppendLine("/// "); + sb.AppendLine("public string Type { get; }"); + sb.AppendLine(); + + sb.AppendLine("/// "); + sb.AppendLine("/// Gets the complete JSON value of the member, including $type."); + sb.AppendLine("/// "); + sb.AppendLine("public System.Text.Json.JsonElement Raw { get; }"); + sb.AppendLine(); + + sb.AppendLine("/// "); + sb.AppendLine("/// Gets the original DAG-CBOR encoding of the member when it was read from CBOR, or ."); + sb.AppendLine("/// When set, CBOR serialization writes these bytes back instead of converting ."); + sb.AppendLine("/// "); + sb.AppendLine("public byte[]? RawCbor { get; }"); + sb.AppendLine(); + + sb.AppendLine("/// "); + sb.AppendLine("/// Returns the raw JSON value of the member."); + sb.AppendLine("/// "); + sb.AppendLine("public System.Text.Json.JsonElement ToJson() => Raw;"); + sb.AppendLine(); + + sb.AppendLine("/// "); + sb.AppendLine($"/// Creates an instance from a JSON value, reading $type from it when present."); + sb.AppendLine("/// "); + sb.AppendLine($"public static {className} FromJson(System.Text.Json.JsonElement element)"); + sb.OpenBrace(); + sb.AppendLine("var type = element.ValueKind == System.Text.Json.JsonValueKind.Object"); + sb.AppendLine(" && element.TryGetProperty(\"$type\", out var typeProp)"); + sb.AppendLine(" && typeProp.ValueKind == System.Text.Json.JsonValueKind.String"); + sb.AppendLine(" ? typeProp.GetString() ?? string.Empty"); + sb.AppendLine(" : string.Empty;"); + sb.AppendLine($"return new {className}(type, element);"); + sb.CloseBrace(); + + sb.CloseBrace(); } /// diff --git a/src/CarpaNet.SourceGen/LexiconGenerator.cs b/src/CarpaNet.SourceGen/LexiconGenerator.cs index 715d4e8..65bdfbd 100644 --- a/src/CarpaNet.SourceGen/LexiconGenerator.cs +++ b/src/CarpaNet.SourceGen/LexiconGenerator.cs @@ -83,25 +83,27 @@ public void Initialize(IncrementalGeneratorInitializationContext context) { var options = new GeneratorOptions(); - if (provider.GlobalOptions.TryGetValue("build_property.CarpaNet_RootNamespace", out var rootNs) + var globalOptions = provider.GlobalOptions; + + if (TryGetBuildProperty(globalOptions, "RootNamespace", out var rootNs) && !string.IsNullOrWhiteSpace(rootNs)) { options.RootNamespace = rootNs; } - if (provider.GlobalOptions.TryGetValue("build_property.CarpaNet_EmitValidationAttributes", out var emitValidation) + if (TryGetBuildProperty(globalOptions, "EmitValidationAttributes", out var emitValidation) && bool.TryParse(emitValidation, out var emitValidationValue)) { options.EmitValidationAttributes = emitValidationValue; } - if (provider.GlobalOptions.TryGetValue("build_property.CarpaNet_CborContextName", out var cborContextName) + if (TryGetBuildProperty(globalOptions, "CborContextName", out var cborContextName) && !string.IsNullOrWhiteSpace(cborContextName)) { options.CborContextName = cborContextName; } - if (provider.GlobalOptions.TryGetValue("build_property.CarpaNet_JsonContextName", out var jsonContextName) + if (TryGetBuildProperty(globalOptions, "JsonContextName", out var jsonContextName) && !string.IsNullOrWhiteSpace(jsonContextName)) { options.JsonContextName = jsonContextName; @@ -123,6 +125,34 @@ public void Initialize(IncrementalGeneratorInitializationContext context) context.RegisterSourceOutput(combined, static (ctx, data) => GenerateSource(ctx, data.Left, data.Right)); } + /// + /// Reads a generator MSBuild property, preferring the current CarpaNet_{name} spelling and + /// falling back to the legacy CarpaNet_SourceGen_{name} spelling for backward compatibility. + /// Empty values are treated as unset so that a blank current property does not hide a legacy one. + /// + private static bool TryGetBuildProperty( + Microsoft.CodeAnalysis.Diagnostics.AnalyzerConfigOptions globalOptions, + string name, + out string value) + { + if (globalOptions.TryGetValue($"build_property.CarpaNet_{name}", out var current) + && !string.IsNullOrWhiteSpace(current)) + { + value = current; + return true; + } + + if (globalOptions.TryGetValue($"build_property.CarpaNet_SourceGen_{name}", out var legacy) + && !string.IsNullOrWhiteSpace(legacy)) + { + value = legacy; + return true; + } + + value = string.Empty; + return false; + } + private static void GenerateSource( SourceProductionContext context, ImmutableArray<(string Path, LexiconDocument? Document)> lexicons, @@ -283,7 +313,7 @@ private static string GenerateNamespaceSource( { try { - GenerateDefinitions(sb, nsid, doc, ns, registry, context); + GenerateDefinitions(sb, nsid, doc, ns, registry, context, options); } catch (Exception ex) { @@ -310,8 +340,11 @@ private static void GenerateDefinitions( LexiconDocument doc, string currentNamespace, TypeRegistry registry, - SourceProductionContext context) + SourceProductionContext context, + GeneratorOptions options) { + var emitValidation = options.EmitValidationAttributes; + foreach (var def in doc.Defs) { var defName = def.Key; @@ -326,24 +359,24 @@ private static void GenerateDefinitions( switch (defValue.Type) { case "record": - GenerateRecordType(sb, className, defValue, nsid, registry); + GenerateRecordType(sb, className, defValue, nsid, registry, emitValidation); break; case "query": - GenerateQueryType(sb, className, defValue, nsid, currentNamespace, registry); + GenerateQueryType(sb, className, defValue, nsid, currentNamespace, registry, emitValidation); break; case "procedure": - GenerateProcedureType(sb, className, defValue, nsid, currentNamespace, registry); + GenerateProcedureType(sb, className, defValue, nsid, currentNamespace, registry, emitValidation); break; case "subscription": - GenerateSubscriptionType(sb, className, defValue, nsid, currentNamespace, registry); + GenerateSubscriptionType(sb, className, defValue, nsid, currentNamespace, registry, emitValidation); break; case "object": var objectTypeId = defName == "main" ? nsid : $"{nsid}#{defName}"; - ObjectGenerator.GenerateClass(sb, className, defValue, nsid, registry, typeId: objectTypeId); + ObjectGenerator.GenerateClass(sb, className, defValue, nsid, registry, typeId: objectTypeId, emitValidationAttributes: emitValidation); sb.AppendLine(); break; @@ -377,7 +410,8 @@ private static void GenerateRecordType( string className, LexiconDefinition def, string nsid, - TypeRegistry registry) + TypeRegistry registry, + bool emitValidation) { if (def.Record == null) { @@ -392,7 +426,8 @@ private static void GenerateRecordType( registry, isRecord: true, recordType: nsid, - typeId: nsid); + typeId: nsid, + emitValidationAttributes: emitValidation); sb.AppendLine(); } @@ -403,19 +438,20 @@ private static void GenerateQueryType( LexiconDefinition def, string nsid, string currentNamespace, - TypeRegistry registry) + TypeRegistry registry, + bool emitValidation) { // Generate parameters class if needed if (def.Parameters?.Properties != null && def.Parameters.Properties.Count > 0) { - ApiGenerator.GenerateParametersClass(sb, className, def, nsid, registry); + ApiGenerator.GenerateParametersClass(sb, className, def, nsid, registry, emitValidation); sb.AppendLine(); } // Generate output class (skip if schema is a direct ref — the referenced type already exists) if (def.Output?.Schema != null && !ApiGenerator.IsOutputRef(def)) { - ApiGenerator.GenerateOutputClass(sb, className, def.Output, nsid, registry); + ApiGenerator.GenerateOutputClass(sb, className, def.Output, nsid, registry, emitValidation); sb.AppendLine(); } @@ -433,19 +469,27 @@ private static void GenerateProcedureType( LexiconDefinition def, string nsid, string currentNamespace, - TypeRegistry registry) + TypeRegistry registry, + bool emitValidation) { + // Generate parameters class if needed (procedures may take query parameters, e.g. app.bsky.video.uploadPart) + if (def.Parameters?.Properties != null && def.Parameters.Properties.Count > 0) + { + ApiGenerator.GenerateParametersClass(sb, className, def, nsid, registry, emitValidation); + sb.AppendLine(); + } + // Generate input class (skip if schema is a direct ref — the referenced type already exists) if (def.Input?.Schema != null && !ApiGenerator.IsInputRef(def)) { - ApiGenerator.GenerateInputClass(sb, className, def.Input, nsid, registry); + ApiGenerator.GenerateInputClass(sb, className, def.Input, nsid, registry, emitValidation); sb.AppendLine(); } // Generate output class (skip if schema is a direct ref — the referenced type already exists) if (def.Output?.Schema != null && !ApiGenerator.IsOutputRef(def)) { - ApiGenerator.GenerateOutputClass(sb, className, def.Output, nsid, registry); + ApiGenerator.GenerateOutputClass(sb, className, def.Output, nsid, registry, emitValidation); sb.AppendLine(); } @@ -463,12 +507,13 @@ private static void GenerateSubscriptionType( LexiconDefinition def, string nsid, string currentNamespace, - TypeRegistry registry) + TypeRegistry registry, + bool emitValidation) { // Generate parameters class if needed if (def.Parameters?.Properties != null && def.Parameters.Properties.Count > 0) { - ApiGenerator.GenerateParametersClass(sb, className, def, nsid, registry); + ApiGenerator.GenerateParametersClass(sb, className, def, nsid, registry, emitValidation); sb.AppendLine(); } @@ -629,7 +674,7 @@ private static string GenerateCborContextSource( var interfaceName = $"I{className}"; var fullName = $"{ns}.{interfaceName}"; var suffix = CborContextGenerator.ToClassSuffix(fullName); - CborContextGenerator.GenerateCborUnionTypeInfo(sb, fullName, suffix, defValue.Refs, nsid, registry, options, generatedTypes); + CborContextGenerator.GenerateCborUnionTypeInfo(sb, fullName, suffix, defValue.Refs, nsid, registry, options, generatedTypes, dispatchInPlace: true, preserveUnknown: defValue.Closed != true); typesToRegister.Add((suffix, fullName, null)); } break; @@ -690,7 +735,7 @@ private static string GenerateCborContextSource( var interfaceName = $"I{className}"; var fullName = $"{ns}.{interfaceName}"; var suffix = CborContextGenerator.ToClassSuffix(fullName); - CborContextGenerator.GenerateCborUnionTypeInfo(sb, fullName, suffix, defValue.Items.Refs, nsid, registry, options, generatedTypes); + CborContextGenerator.GenerateCborUnionTypeInfo(sb, fullName, suffix, defValue.Items.Refs, nsid, registry, options, generatedTypes, dispatchInPlace: true, preserveUnknown: defValue.Items.Closed != true); typesToRegister.Add((suffix, fullName, null)); } break; @@ -699,6 +744,12 @@ private static string GenerateCborContextSource( } } + // Shared helper for union type infos that read in place / preserve unknown members + if (generatedTypes.Contains(CborContextGenerator.UnionHelperMarker)) + { + CborContextGenerator.GenerateUnionHelper(sb); + } + // Generate the context class sb.WriteSummary("CBOR serialization context for all AT Protocol types."); sb.AppendLine($"public partial class {options.CborContextName} : CarpaNet.Cbor.CborSerializerContext"); diff --git a/src/CarpaNet.SourceGen/TypeRegistry.cs b/src/CarpaNet.SourceGen/TypeRegistry.cs index 84dcce0..55e43a7 100644 --- a/src/CarpaNet.SourceGen/TypeRegistry.cs +++ b/src/CarpaNet.SourceGen/TypeRegistry.cs @@ -265,7 +265,7 @@ public string ResolveToCSharpType(string? refString, string currentNsid) LexiconTypeKind.Any => "System.Text.Json.JsonElement", LexiconTypeKind.Object => typeInfo.FullCSharpTypeName, LexiconTypeKind.Record => typeInfo.FullCSharpTypeName, - LexiconTypeKind.Union => $"I{typeInfo.CSharpTypeName}", // Unions become interfaces + LexiconTypeKind.Union => $"{typeInfo.CSharpNamespace}.I{typeInfo.CSharpTypeName}", // Unions become interfaces (qualified: the ref may cross namespaces) // For array types, resolve the item type LexiconTypeKind.Array => ResolveArrayType(typeInfo, currentNsid), // For ref types, recursively resolve @@ -401,7 +401,7 @@ private string ResolveArrayType(TypeInfo typeInfo, string currentNsid) var itemType = items.Type switch { "ref" when items.Ref != null => ResolveToCSharpType(items.Ref, typeInfo.Nsid), - "union" when items.Refs != null => $"I{typeInfo.CSharpTypeName}", // Generate interface name for union items + "union" when items.Refs != null => $"{typeInfo.CSharpNamespace}.I{typeInfo.CSharpTypeName}", // Interface for union items, qualified because the array def may be referenced from another namespace "string" => MapStringType(items), "integer" => "long", "boolean" => "bool", diff --git a/src/CarpaNet.SourceGen/build/CarpaNet.SourceGen.targets b/src/CarpaNet.SourceGen/build/CarpaNet.SourceGen.targets index 5225378..0662f9a 100644 --- a/src/CarpaNet.SourceGen/build/CarpaNet.SourceGen.targets +++ b/src/CarpaNet.SourceGen/build/CarpaNet.SourceGen.targets @@ -29,6 +29,12 @@ + + + + + + diff --git a/src/CarpaNet/ATProtoClient.cs b/src/CarpaNet/ATProtoClient.cs index 42911ec..c3a8045 100644 --- a/src/CarpaNet/ATProtoClient.cs +++ b/src/CarpaNet/ATProtoClient.cs @@ -21,7 +21,7 @@ namespace CarpaNet; /// /// Default Implementation of . /// -public sealed class ATProtoClient : IATProtoClient, IDisposable +public sealed class ATProtoClient : IATProtoClient, IXrpcRequestClient, IDisposable { private readonly HttpClient _httpClient; private readonly JsonSerializerOptions _jsonOptions; @@ -32,6 +32,8 @@ public sealed class ATProtoClient : IATProtoClient, IDisposable private readonly Uri? _configuredBaseUrl; private readonly ILogger _logger; private readonly ILoggerFactory _loggerFactory; + private readonly string? _userAgent; + private IReadOnlyList? _labelerDids; private bool _disposed; /// @@ -47,9 +49,23 @@ public sealed class ATProtoClient : IATProtoClient, IDisposable public IdentityResolver? IdentityResolver { get; } /// - /// Gets the optional list of labeler DIDs to accept labels from. + /// Gets the labeler DIDs sent in the atproto-accept-labelers header. + /// Change it with . /// - public IReadOnlyList? LabelerDids { get; } + public IReadOnlyList? LabelerDids => Volatile.Read(ref _labelerDids); + + /// + public JsonSerializerOptions JsonOptions => _jsonOptions; + + /// + /// Replaces the labeler DIDs sent in the atproto-accept-labelers header on later requests. + /// Entries may carry parameters such as ;redact (see ). + /// + /// The labeler DIDs, or null to send no header. + public void SetLabelerDids(IEnumerable? labelerDids) + { + Volatile.Write(ref _labelerDids, labelerDids?.ToArray()); + } /// /// Gets the token provider, if any. @@ -193,14 +209,9 @@ public static async Task CreateWithSessionAsync( options = options?.Clone() ?? new ATProtoClientOptions(); // Create HttpClient if not provided - var httpClient = options.HttpClient ?? new HttpClient(); + var httpClient = options.HttpClient ?? CreateOwnedHttpClient(options); var ownsHttpClient = options.HttpClient == null; - if (options.Timeout.HasValue) - { - httpClient.Timeout = options.Timeout.Value; - } - // Create session token provider and login var tokenProvider = new SessionTokenProvider(httpClient, sessionStore: options.SessionStore, loggerFactory: options.LoggerFactory); try @@ -247,14 +258,9 @@ public static ATProtoClient CreateWithRestoredSession( { options = options?.Clone() ?? new ATProtoClientOptions(); - var httpClient = options.HttpClient ?? new HttpClient(); + var httpClient = options.HttpClient ?? CreateOwnedHttpClient(options); var ownsHttpClient = options.HttpClient == null; - if (options.Timeout.HasValue) - { - httpClient.Timeout = options.Timeout.Value; - } - var tokenProvider = new SessionTokenProvider(httpClient, sessionStore: options.SessionStore, loggerFactory: options.LoggerFactory); tokenProvider.RestoreSession(accessJwt, refreshJwt, did, handle, pdsUrl); @@ -277,14 +283,9 @@ public static ATProtoClient Create(ATProtoClientOptions? options = null) options.BaseUrl ??= new Uri(BlueskyServices.PublicAppView); - var httpClient = options.HttpClient ?? new HttpClient(); + var httpClient = options.HttpClient ?? CreateOwnedHttpClient(options); var ownsHttpClient = options.HttpClient == null; - if (options.Timeout.HasValue) - { - httpClient.Timeout = options.Timeout.Value; - } - var tokenProvider = new SessionTokenProvider(httpClient, sessionStore: options.SessionStore, loggerFactory: options.LoggerFactory); options.TokenProvider = tokenProvider; options.HttpClient = httpClient; @@ -313,13 +314,9 @@ private ATProtoClient(ATProtoClientOptions options, bool ownsHttpClient, bool ow _loggerFactory = options.LoggerFactory ?? NullLoggerFactory.Instance; _logger = _loggerFactory.CreateLogger(); - var httpClient = options.HttpClient ?? new HttpClient(); + var httpClient = options.HttpClient ?? CreateOwnedHttpClient(options); _ownsHttpClient = ownsHttpClient; - - if (options.Timeout.HasValue && ownsHttpClient) - { - httpClient.Timeout = options.Timeout.Value; - } + _userAgent = options.UserAgent; var jsonOptions = options.JsonOptions ?? throw new ArgumentException("JsonOptions must be provided.", nameof(options)); var cborContext = options.CborContext ?? throw new ArgumentException("CborContext must be provided.", nameof(options)); @@ -329,7 +326,7 @@ private ATProtoClient(ATProtoClientOptions options, bool ownsHttpClient, bool ow _cborContext = cborContext; _tokenProvider = options.TokenProvider; _autoRetryOnAuthFailure = options.AutoRetryOnAuthFailure; - LabelerDids = options.LabelerDids; + _labelerDids = options.LabelerDids?.ToArray(); // Store configured base URL (null means dynamic - will use token provider's PDS URL) _configuredBaseUrl = options.BaseUrl ?? _tokenProvider?.PdsUrl; @@ -341,7 +338,7 @@ private ATProtoClient(ATProtoClientOptions options, bool ownsHttpClient, bool ow } else if (options.CreateIdentityResolver) { - IdentityResolver = new IdentityResolver(httpClient, dnsResolver: new DefaultDnsResolver(), cache: new MemoryIdentityCache(), loggerFactory: _loggerFactory); + IdentityResolver = new IdentityResolver(httpClient, cache: new MemoryIdentityCache(), loggerFactory: _loggerFactory); } _ownsTokenProvider = ownsTokenProvider; @@ -354,90 +351,47 @@ private ATProtoClient(ATProtoClientOptions options, bool ownsHttpClient, bool ow #region IATProtoClient Implementation /// - public async Task GetAsync( + public Task GetAsync( string nsid, IEnumerable>? parameters = null, CancellationToken cancellationToken = default) { ThrowIfDisposed(); - _logger.LogDebug("Sending GET {Nsid}", nsid); - - var url = await XrpcHttpHandler.BuildUrlAsync( - this.TokenProvider?.PdsUrl ?? BaseUrl, nsid, parameters, - this.IdentityResolver, _logger, cancellationToken).ConfigureAwait(false); - - using var request = XrpcHttpHandler.CreateGetRequest(url, proxyServiceDid: null, LabelerDids); - await AddAuthHeaderAsync(request, cancellationToken).ConfigureAwait(false); - - var response = await SendWithRetryAsync(request, cancellationToken).ConfigureAwait(false); - return await XrpcHttpHandler.ProcessResponseAsync(response, _jsonOptions, _logger, cancellationToken).ConfigureAwait(false); + return this.QueryAsync(nsid, parameters, null, cancellationToken); } /// - public async Task GetAsync( + public Task GetAsync( string nsid, string proxyServiceDid, IEnumerable>? parameters = null, CancellationToken cancellationToken = default) { ThrowIfDisposed(); - - var url = await XrpcHttpHandler.BuildUrlAsync( - this.TokenProvider?.PdsUrl ?? BaseUrl, nsid, parameters, - this.IdentityResolver, _logger, cancellationToken).ConfigureAwait(false); - using var request = XrpcHttpHandler.CreateGetRequest(url, proxyServiceDid, LabelerDids); - await AddAuthHeaderAsync(request, cancellationToken).ConfigureAwait(false); - - var response = await SendWithRetryAsync(request, cancellationToken).ConfigureAwait(false); - return await XrpcHttpHandler.ProcessResponseAsync(response, _jsonOptions, _logger, cancellationToken).ConfigureAwait(false); + return this.QueryAsync(nsid, parameters, new XrpcRequestOptions { ProxyServiceDid = proxyServiceDid }, cancellationToken); } /// - public async Task PostAsync( + public Task PostAsync( string nsid, TInput? input, CancellationToken cancellationToken = default) { ThrowIfDisposed(); - _logger.LogDebug("Sending POST {Nsid}", nsid); - - if (_tokenProvider == null) - { - throw new InvalidOperationException( - "Cannot make POST requests without authentication. " + - "Use CreateWithSessionAsync or provide a TokenProvider."); - } - - var url = XrpcHttpHandler.BuildUrl(this.TokenProvider?.PdsUrl ?? BaseUrl, nsid); - using var request = XrpcHttpHandler.CreatePostRequest(url, input, _jsonOptions, proxyServiceDid: null, LabelerDids); - await AddAuthHeaderAsync(request, cancellationToken).ConfigureAwait(false); - - var response = await SendWithRetryAsync(request, cancellationToken).ConfigureAwait(false); - return await XrpcHttpHandler.ProcessResponseAsync(response, _jsonOptions, _logger, cancellationToken).ConfigureAwait(false); + ThrowIfNoTokenProvider(); + return this.ProcedureAsync(nsid, null, input, null, cancellationToken); } /// - public async Task PostAsync( + public Task PostAsync( string nsid, string proxyServiceDid, TInput? input, CancellationToken cancellationToken = default) { ThrowIfDisposed(); - - if (_tokenProvider == null) - { - throw new InvalidOperationException( - "Cannot make POST requests without authentication. " + - "Use CreateWithSessionAsync or provide a TokenProvider."); - } - - var url = XrpcHttpHandler.BuildUrl(this.TokenProvider?.PdsUrl ?? BaseUrl, nsid); - using var request = XrpcHttpHandler.CreatePostRequest(url, input, _jsonOptions, proxyServiceDid, LabelerDids); - await AddAuthHeaderAsync(request, cancellationToken).ConfigureAwait(false); - - var response = await SendWithRetryAsync(request, cancellationToken).ConfigureAwait(false); - return await XrpcHttpHandler.ProcessResponseAsync(response, _jsonOptions, _logger, cancellationToken).ConfigureAwait(false); + ThrowIfNoTokenProvider(); + return this.ProcedureAsync(nsid, null, input, new XrpcRequestOptions { ProxyServiceDid = proxyServiceDid }, cancellationToken); } /// @@ -469,58 +423,180 @@ public async IAsyncEnumerable SubscribeAsync( #endregion - #region Private Methods + #region IXrpcRequestClient Implementation - private async Task AddAuthHeaderAsync(HttpRequestMessage request, CancellationToken cancellationToken) + /// + /// + /// + /// Requests go to the session's PDS unless is set. + /// A query with a repo parameter and no proxy is sent to that repo's PDS when an + /// is configured. + /// + /// + /// The access token is attached only when the request goes to the session's own PDS, so + /// credentials are never sent to another server. A 401 from the PDS triggers one token refresh + /// and retry when the body can be replayed. + /// + /// + public async Task SendXrpcAsync(XrpcRequest request, CancellationToken cancellationToken = default) { - if (_tokenProvider == null) - return; + ThrowIfDisposed(); + if (request == null) + throw new ArgumentNullException(nameof(request)); - var token = await _tokenProvider.GetAccessTokenAsync(cancellationToken).ConfigureAwait(false); - if (!string.IsNullOrEmpty(token)) - { - request.Headers.Authorization = new AuthenticationHeaderValue("Bearer", token); - } - } + var options = request.Options; + var proxy = options?.EffectiveProxyServiceDid; + var url = await ResolveRequestUrlAsync(request, proxy, cancellationToken).ConfigureAwait(false); + var attachCredentials = ShouldAttachCredentials(url, options); - private async Task SendWithRetryAsync( - HttpRequestMessage request, - CancellationToken cancellationToken) - { - var response = await _httpClient.SendAsync(request, cancellationToken).ConfigureAwait(false); + _logger.LogDebug("Sending {Method} {Nsid}", request.Method, request.Nsid); + + var response = await SendOnceAsync(request, url, proxy, attachCredentials, isRetry: false, cancellationToken).ConfigureAwait(false); - // Retry on 401 if auto-retry is enabled and we have a token provider if (_autoRetryOnAuthFailure && + attachCredentials && response.StatusCode == System.Net.HttpStatusCode.Unauthorized && - _tokenProvider != null) + (request.Body == null || request.Body.IsReplayable)) { _logger.LogWarning("Received 401, refreshing token and retrying"); try { - await _tokenProvider.RefreshAsync(cancellationToken).ConfigureAwait(false); - - // Create a new request and retry - using var retryRequest = XrpcHttpHandler.CloneRequest(request, "Authorization"); - await AddAuthHeaderAsync(retryRequest, cancellationToken).ConfigureAwait(false); - - response.Dispose(); - return await _httpClient.SendAsync(retryRequest, cancellationToken).ConfigureAwait(false); + await _tokenProvider!.RefreshAsync(cancellationToken).ConfigureAwait(false); } - catch (AuthenticationException) + catch (ATProtoException) { + // Includes AuthenticationException and the 400 a rejected refresh token gets. _logger.LogWarning("Token refresh failed, returning 401"); - // Refresh failed, return original 401 response + return response; } catch (InvalidOperationException) { _logger.LogWarning("Token refresh failed, returning 401"); - // No refresh token available + return response; } + + response.Dispose(); + response = await SendOnceAsync(request, url, proxy, attachCredentials, isRetry: true, cancellationToken).ConfigureAwait(false); } return response; } + #endregion + + #region Private Methods + + private static HttpClient CreateOwnedHttpClient(ATProtoClientOptions options) + { + return HttpClientFactory.Create(new HttpClientFactoryOptions + { + Timeout = options.Timeout, + UserAgent = options.UserAgent, + EnableRateLimitHandler = options.EnableRateLimitHandler, + AutoRetryOnRateLimit = options.AutoRetryOnRateLimit, + RateLimitMaxRetries = options.RateLimitMaxRetries, + LoggerFactory = options.LoggerFactory, + }); + } + + private void ThrowIfNoTokenProvider() + { + if (_tokenProvider == null) + { + throw new InvalidOperationException( + "Cannot make POST requests without authentication. " + + "Use CreateWithSessionAsync or provide a TokenProvider."); + } + } + + private async Task ResolveRequestUrlAsync(XrpcRequest request, string? proxy, CancellationToken cancellationToken) + { + var serviceUrl = request.Options?.ServiceUrl; + if (serviceUrl != null) + { + return XrpcHttpHandler.BuildUrl(serviceUrl, request.Nsid, request.Parameters); + } + + var baseUrl = _tokenProvider?.PdsUrl ?? BaseUrl; + if (request.Method == HttpMethod.Get && proxy == null) + { + return await XrpcHttpHandler.BuildUrlAsync( + baseUrl, request.Nsid, request.Parameters, + IdentityResolver, _logger, cancellationToken).ConfigureAwait(false); + } + + return XrpcHttpHandler.BuildUrl(baseUrl, request.Nsid, request.Parameters); + } + + private bool ShouldAttachCredentials(Uri url, XrpcRequestOptions? options) + { + if (_tokenProvider == null || options?.ServiceUrl != null || HasAuthorizationHeader(options)) + { + return false; + } + + var credentialOrigin = _tokenProvider.PdsUrl ?? _configuredBaseUrl; + return credentialOrigin != null && XrpcHttpHandler.IsSameOrigin(url, credentialOrigin); + } + + private static bool HasAuthorizationHeader(XrpcRequestOptions? options) + { + if (options?.Headers == null) + { + return false; + } + + foreach (var header in options.Headers) + { + if (header.Key.Equals("Authorization", StringComparison.OrdinalIgnoreCase)) + { + return true; + } + } + + return false; + } + + private async Task SendOnceAsync( + XrpcRequest request, + Uri url, + string? proxy, + bool attachCredentials, + bool isRetry, + CancellationToken cancellationToken) + { + using var message = new HttpRequestMessage(request.Method, url); + XrpcHttpHandler.AddCommonHeaders(message, proxy, request.Options?.AcceptLabelers ?? LabelerDids); + XrpcHttpHandler.AddCustomHeaders(message, request.Options?.Headers); + + if (_userAgent != null && message.Headers.UserAgent.Count == 0 && _httpClient.DefaultRequestHeaders.UserAgent.Count == 0) + { + message.Headers.TryAddWithoutValidation("User-Agent", _userAgent); + } + + if (request.Body != null) + { + message.Content = request.Body.CreateContent() + ?? throw new InvalidOperationException("The request body has already been sent and cannot be replayed."); + } + + if (attachCredentials) + { + var token = await _tokenProvider!.GetAccessTokenAsync(cancellationToken).ConfigureAwait(false); + if (!string.IsNullOrEmpty(token)) + { + message.Headers.Authorization = new AuthenticationHeaderValue("Bearer", token); + } + } + + if (isRetry) + { + _logger.LogDebug("Retrying {Method} {Nsid}", request.Method, request.Nsid); + } + + return await _httpClient.SendAsync(message, cancellationToken).ConfigureAwait(false); + } + private void ThrowIfDisposed() { #if NET8_0_OR_GREATER diff --git a/src/CarpaNet/ATProtoClientXrpcExtensions.cs b/src/CarpaNet/ATProtoClientXrpcExtensions.cs new file mode 100644 index 0000000..1c1c7dd --- /dev/null +++ b/src/CarpaNet/ATProtoClientXrpcExtensions.cs @@ -0,0 +1,311 @@ +using System; +using System.Collections.Generic; +using System.IO; +using System.Linq; +using System.Net.Http; +using System.Text.Json.Serialization.Metadata; +using System.Threading; +using System.Threading.Tasks; +using CarpaNet.Http; + +namespace CarpaNet; + +/// +/// XRPC calls beyond the JSON query/procedure methods on : +/// binary bodies, binary responses, procedures with query parameters, and per-request options. +/// +/// +/// These require a client that implements +/// (, the OAuth client and do). +/// Generated API methods call the first three methods. +/// +public static class ATProtoClientXrpcExtensions +{ + #region Used by generated code + + /// + /// Calls a procedure that has a JSON body and query parameters. + /// + /// The input type. + /// The output type. + /// The client. + /// The NSID of the procedure. + /// The service to proxy to, or null. + /// The query parameters. + /// The request body. + /// Cancellation token. + public static Task PostWithParametersAsync( + this IATProtoClient client, + string nsid, + string? proxyServiceDid, + IEnumerable>? parameters, + TInput? input, + CancellationToken cancellationToken = default) + { + if (client is IXrpcRequestClient xrpc) + { + return xrpc.ProcedureAsync(nsid, parameters, input, ProxyOptions(proxyServiceDid), cancellationToken); + } + + if (parameters == null || !parameters.Any()) + { + return proxyServiceDid == null + ? client.PostAsync(nsid, input, cancellationToken) + : client.PostAsync(nsid, proxyServiceDid, input, cancellationToken); + } + + throw NotSupported(client, "procedures with query parameters"); + } + + /// + /// Calls a procedure whose body is not JSON (for example com.atproto.repo.uploadBlob). + /// + /// The output type. + /// The client. + /// The NSID of the procedure. + /// The service to proxy to, or null. + /// The query parameters, or null. + /// The body. It is read from its current position and not disposed. + /// The MIME type of the body. + /// Cancellation token. + public static Task PostBinaryAsync( + this IATProtoClient client, + string nsid, + string? proxyServiceDid, + IEnumerable>? parameters, + Stream body, + string contentType, + CancellationToken cancellationToken = default) + { + var xrpc = AsXrpcClient(client, "binary request bodies"); + return xrpc.ProcedureBinaryAsync(nsid, parameters, body, contentType, ProxyOptions(proxyServiceDid), cancellationToken); + } + + /// + /// Calls a query whose response is not JSON (for example com.atproto.sync.getBlob). + /// + /// The client. + /// The NSID of the query. + /// The service to proxy to, or null. + /// The query parameters, or null. + /// Cancellation token. + /// The response body. + public static Task GetBytesAsync( + this IATProtoClient client, + string nsid, + string? proxyServiceDid, + IEnumerable>? parameters, + CancellationToken cancellationToken = default) + { + var xrpc = AsXrpcClient(client, "binary responses"); + return xrpc.QueryBytesAsync(nsid, parameters, ProxyOptions(proxyServiceDid), cancellationToken); + } + + #endregion + + #region Request helpers + + /// + /// Calls a query and deserializes its JSON response. + /// + public static async Task QueryAsync( + this IXrpcRequestClient client, + string nsid, + IEnumerable>? parameters, + XrpcRequestOptions? options, + CancellationToken cancellationToken = default) + { + var request = new XrpcRequest(HttpMethod.Get, nsid) { Parameters = parameters, Options = options }; + using var response = await client.SendXrpcAsync(request, cancellationToken).ConfigureAwait(false); + return await XrpcHttpHandler.ProcessResponseAsync(response, client.JsonOptions, cancellationToken: cancellationToken).ConfigureAwait(false); + } + + /// + /// Calls a procedure with an optional JSON body and deserializes its JSON response. + /// + public static async Task ProcedureAsync( + this IXrpcRequestClient client, + string nsid, + IEnumerable>? parameters, + TInput? input, + XrpcRequestOptions? options, + CancellationToken cancellationToken = default) + { + var request = new XrpcRequest(HttpMethod.Post, nsid) { Parameters = parameters, Options = options }; + if (input != null) + { + var typeInfo = (JsonTypeInfo)client.JsonOptions.GetTypeInfo(typeof(TInput)); + request.Body = XrpcBody.FromJson(input, typeInfo); + } + + using var response = await client.SendXrpcAsync(request, cancellationToken).ConfigureAwait(false); + return await XrpcHttpHandler.ProcessResponseAsync(response, client.JsonOptions, cancellationToken: cancellationToken).ConfigureAwait(false); + } + + /// + /// Calls a procedure with a binary body and deserializes its JSON response. + /// + public static async Task ProcedureBinaryAsync( + this IXrpcRequestClient client, + string nsid, + IEnumerable>? parameters, + Stream body, + string contentType, + XrpcRequestOptions? options, + CancellationToken cancellationToken = default) + { + var request = new XrpcRequest(HttpMethod.Post, nsid) + { + Parameters = parameters, + Body = XrpcBody.FromStream(body, contentType), + Options = options, + }; + + using var response = await client.SendXrpcAsync(request, cancellationToken).ConfigureAwait(false); + return await XrpcHttpHandler.ProcessResponseAsync(response, client.JsonOptions, cancellationToken: cancellationToken).ConfigureAwait(false); + } + + /// + /// Calls a query and returns its response body as bytes. + /// + public static async Task QueryBytesAsync( + this IXrpcRequestClient client, + string nsid, + IEnumerable>? parameters, + XrpcRequestOptions? options, + CancellationToken cancellationToken = default) + { + var acceptAny = new XrpcRequestOptions + { + Headers = new Dictionary(StringComparer.OrdinalIgnoreCase) { ["Accept"] = "*/*" }, + }; + + var request = new XrpcRequest(HttpMethod.Get, nsid) + { + Parameters = parameters, + Options = XrpcRequestOptions.Combine(options, acceptAny), + }; + + using var response = await client.SendXrpcAsync(request, cancellationToken).ConfigureAwait(false); + if (!response.IsSuccessStatusCode) + { + await XrpcHttpHandler.ThrowForErrorResponseAsync(response, cancellationToken: cancellationToken).ConfigureAwait(false); + } + +#if NET8_0_OR_GREATER + return await response.Content.ReadAsByteArrayAsync(cancellationToken).ConfigureAwait(false); +#else + return await response.Content.ReadAsByteArrayAsync().ConfigureAwait(false); +#endif + } + + #endregion + + #region Scoping + + /// + /// Returns a client that applies to every request, including the + /// generated API methods. Options set here override a generated method's default proxy. + /// + /// The client to wrap. + /// The options to apply. + /// A client that shares 's session and HTTP pipeline. + public static ScopedATProtoClient WithRequestOptions(this IATProtoClient client, XrpcRequestOptions options) + { + if (options == null) + { + throw new ArgumentNullException(nameof(options)); + } + + if (client is ScopedATProtoClient scoped) + { + return new ScopedATProtoClient(scoped.Inner, XrpcRequestOptions.Combine(options, scoped.Options)!); + } + + return new ScopedATProtoClient(client, options); + } + + /// + /// Returns a client that proxies every request to + /// (for example did:web:api.bsky.app#bsky_appview). + /// + public static ScopedATProtoClient WithProxy(this IATProtoClient client, string serviceDid) + { + if (string.IsNullOrEmpty(serviceDid)) + { + throw new ArgumentException("Service DID cannot be null or empty.", nameof(serviceDid)); + } + + return client.WithRequestOptions(new XrpcRequestOptions { ProxyServiceDid = serviceDid }); + } + + /// + /// Returns a client that sends every request without an atproto-proxy header, + /// so the PDS handles it itself. + /// + public static ScopedATProtoClient WithoutProxy(this IATProtoClient client) + => client.WithRequestOptions(new XrpcRequestOptions { DisableProxy = true }); + + /// + /// Returns a client that sends in the + /// atproto-accept-labelers header instead of the client's own list. + /// + public static ScopedATProtoClient WithAcceptLabelers(this IATProtoClient client, IEnumerable labelerDids) + { + if (labelerDids == null) + { + throw new ArgumentNullException(nameof(labelerDids)); + } + + return client.WithRequestOptions(new XrpcRequestOptions { AcceptLabelers = labelerDids.ToArray() }); + } + + /// + /// Returns a client that adds a header to every request. + /// + public static ScopedATProtoClient WithHeader(this IATProtoClient client, string name, string value) + { + if (string.IsNullOrEmpty(name)) + { + throw new ArgumentException("Header name cannot be null or empty.", nameof(name)); + } + + return client.WithRequestOptions(new XrpcRequestOptions + { + Headers = new Dictionary(StringComparer.OrdinalIgnoreCase) { [name] = value }, + }); + } + + /// + /// Returns a client that sends every request to another service. Session credentials are not + /// sent to that service. + /// + public static ScopedATProtoClient WithServiceUrl(this IATProtoClient client, Uri serviceUrl) + { + if (serviceUrl == null) + { + throw new ArgumentNullException(nameof(serviceUrl)); + } + + return client.WithRequestOptions(new XrpcRequestOptions { ServiceUrl = serviceUrl }); + } + + #endregion + + private static XrpcRequestOptions? ProxyOptions(string? proxyServiceDid) + => proxyServiceDid == null ? null : new XrpcRequestOptions { ProxyServiceDid = proxyServiceDid }; + + private static IXrpcRequestClient AsXrpcClient(IATProtoClient client, string feature) + { + if (client is IXrpcRequestClient xrpc) + { + return xrpc; + } + + throw NotSupported(client, feature); + } + + private static NotSupportedException NotSupported(IATProtoClient client, string feature) + => new NotSupportedException( + $"{client.GetType().Name} does not implement {nameof(IXrpcRequestClient)}, which is required for {feature}."); +} diff --git a/src/CarpaNet/Auth/INotifySessionInvalidated.cs b/src/CarpaNet/Auth/INotifySessionInvalidated.cs new file mode 100644 index 0000000..375a822 --- /dev/null +++ b/src/CarpaNet/Auth/INotifySessionInvalidated.cs @@ -0,0 +1,55 @@ +using System; + +namespace CarpaNet.Auth; + +/// +/// Implemented by token providers that can tell when a session has ended for good +/// (for example, a revoked or expired refresh token). +/// +/// +/// and the OAuth DPoP token provider implement this. +/// A temporary failure (network error, 5xx, rate limit) does not raise the event. +/// +public interface INotifySessionInvalidated +{ + /// + /// Raised once when a token refresh is rejected by the server and the session cannot continue. + /// The provider drops its tokens before raising it; the caller decides whether to delete + /// stored session data and ask the user to sign in again. + /// + event EventHandler? SessionInvalidated; +} + +/// +/// Event arguments for . +/// +public sealed class SessionInvalidatedEventArgs : EventArgs +{ + /// + /// Creates new event arguments. + /// + /// The DID of the session that ended, if known. + /// The server's error code (such as ExpiredToken or invalid_grant), or a status description. + /// The exception from the failed refresh, if any. + public SessionInvalidatedEventArgs(string? did, string reason, Exception? exception) + { + Did = did; + Reason = reason; + Exception = exception; + } + + /// + /// Gets the DID of the session that ended, if known. + /// + public string? Did { get; } + + /// + /// Gets the server's error code, or a status description when there was none. + /// + public string Reason { get; } + + /// + /// Gets the exception from the failed refresh, if any. + /// + public Exception? Exception { get; } +} diff --git a/src/CarpaNet/Auth/SessionTokenProvider.cs b/src/CarpaNet/Auth/SessionTokenProvider.cs index e924a99..09a9a27 100644 --- a/src/CarpaNet/Auth/SessionTokenProvider.cs +++ b/src/CarpaNet/Auth/SessionTokenProvider.cs @@ -17,7 +17,7 @@ namespace CarpaNet.Auth; /// Token provider that uses ATProtocol session tokens (createSession/refreshSession). /// Suitable for App Passwords and direct username/password authentication. /// -public sealed class SessionTokenProvider : ITokenProvider, IDisposable +public sealed class SessionTokenProvider : ITokenProvider, INotifySessionInvalidated, IDisposable { private readonly HttpClient _httpClient; private readonly bool _ownsHttpClient; @@ -32,6 +32,7 @@ public sealed class SessionTokenProvider : ITokenProvider, IDisposable private string? _did; private string? _handle; private Uri? _pdsUrl; + private bool _invalidated; private bool _disposed; /// @@ -66,6 +67,14 @@ public sealed class SessionTokenProvider : ITokenProvider, IDisposable /// public event EventHandler? TokenRefreshed; + /// + /// + /// Raised when com.atproto.server.refreshSession answers 400 or 401 (for example + /// ExpiredToken or InvalidToken). The access and refresh tokens are cleared first; + /// and are kept. + /// + public event EventHandler? SessionInvalidated; + /// /// Creates a new SessionTokenProvider with default settings. /// Creates a new HttpClient that will be disposed with this provider. @@ -182,6 +191,7 @@ public void RestoreSession(string accessJwt, string refreshJwt, string did, stri _handle = handle; _pdsUrl = pdsUrl ?? throw new ArgumentNullException(nameof(pdsUrl)); _accessExpiry = ParseJwtExpiry(accessJwt); + _invalidated = false; } /// @@ -214,27 +224,51 @@ public async Task RefreshAsync(CancellationToken cancellationToken = default) throw new InvalidOperationException("No refresh token available. Call LoginAsync first."); } + // Remember the token this caller saw, so a refresh that another caller completed + // while this one waited for the lock is not repeated. + var staleAccessJwt = _accessJwt; + SessionInvalidatedEventArgs? invalidated = null; + // Use lock to prevent concurrent refresh attempts await _refreshLock.WaitAsync(cancellationToken).ConfigureAwait(false); try { // Double-check after acquiring lock - if (HasValidToken) + if (!string.Equals(_accessJwt, staleAccessJwt, StringComparison.Ordinal) && HasValidToken) { - _logger.LogDebug("Token still valid after lock, skipping refresh"); + _logger.LogDebug("Token was refreshed by another caller, skipping refresh"); return; } - var url = XrpcHttpHandler.BuildUrl(_pdsUrl, "com.atproto.server.refreshSession"); + var refreshJwt = _refreshJwt; + var pdsUrl = _pdsUrl; + if (string.IsNullOrEmpty(refreshJwt) || pdsUrl == null) + { + throw new InvalidOperationException("No refresh token available. Call LoginAsync first."); + } + + var url = XrpcHttpHandler.BuildUrl(pdsUrl, "com.atproto.server.refreshSession"); using var httpRequest = new HttpRequestMessage(HttpMethod.Post, url); - httpRequest.Headers.Authorization = new System.Net.Http.Headers.AuthenticationHeaderValue("Bearer", _refreshJwt); + httpRequest.Headers.Authorization = new System.Net.Http.Headers.AuthenticationHeaderValue("Bearer", refreshJwt); var response = await _httpClient.SendAsync(httpRequest, cancellationToken).ConfigureAwait(false); if (!response.IsSuccessStatusCode) { _logger.LogWarning("Session token refresh failed with HTTP {StatusCode}", (int)response.StatusCode); - await XrpcHttpHandler.ThrowForErrorResponseAsync(response, _logger, cancellationToken).ConfigureAwait(false); + + // 400 and 401 mean the server rejected the refresh token; anything else is temporary. + var rejected = response.StatusCode == System.Net.HttpStatusCode.BadRequest + || response.StatusCode == System.Net.HttpStatusCode.Unauthorized; + try + { + await XrpcHttpHandler.ThrowForErrorResponseAsync(response, _logger, cancellationToken).ConfigureAwait(false); + } + catch (ATProtoException ex) when (rejected) + { + invalidated = Invalidate(ex.ErrorCode ?? $"HTTP {(int)response.StatusCode}", ex); + throw; + } } #if NET8_0_OR_GREATER @@ -250,18 +284,40 @@ public async Task RefreshAsync(CancellationToken cancellationToken = default) throw new ATProtoException("Failed to parse refresh session response."); } - await UpdateSessionAsync(session, _pdsUrl).ConfigureAwait(false); + await UpdateSessionAsync(session, pdsUrl).ConfigureAwait(false); _logger.LogInformation("Session token refreshed for {Did}", session.Did); } finally { _refreshLock.Release(); + + if (invalidated != null) + { + SessionInvalidated?.Invoke(this, invalidated); + } } } /// - /// Clears the current session. + /// Drops the tokens after the server rejected the refresh token. Returns the event to raise, + /// or null when the event was already raised for this session. /// + private SessionInvalidatedEventArgs? Invalidate(string reason, Exception exception) + { + _accessJwt = null; + _refreshJwt = null; + _accessExpiry = default; + + if (_invalidated) + { + return null; + } + + _invalidated = true; + _logger.LogWarning("Session for {Did} was rejected by the server ({Reason})", _did, reason); + return new SessionInvalidatedEventArgs(_did, reason, exception); + } + public void ClearSession() { _accessJwt = null; @@ -314,6 +370,7 @@ public async Task RestoreSessionAsync(string sub, CancellationToken cancel private async Task UpdateSessionAsync(SessionResponse session, Uri pdsUrl) { + _invalidated = false; _accessJwt = session.AccessJwt; _refreshJwt = session.RefreshJwt; _did = session.Did; diff --git a/src/CarpaNet/Blob/ATProtoClientBlobExtensions.cs b/src/CarpaNet/Blob/ATProtoClientBlobExtensions.cs index 25df9bb..af3737c 100644 --- a/src/CarpaNet/Blob/ATProtoClientBlobExtensions.cs +++ b/src/CarpaNet/Blob/ATProtoClientBlobExtensions.cs @@ -1,4 +1,5 @@ using System; +using System.Collections.Generic; using System.IO; using System.Net.Http; using System.Net.Http.Headers; @@ -15,11 +16,21 @@ namespace CarpaNet.Blob; /// public static class ATProtoClientBlobExtensions { + private const string UploadBlobNsid = "com.atproto.repo.uploadBlob"; + private const string GetBlobNsid = "com.atproto.sync.getBlob"; + /// /// Uploads a blob to the user's PDS. /// + /// + /// With a client that implements (including OAuth clients), the + /// upload goes through the client's own authentication (Bearer or DPoP) and streams the content + /// without buffering it. A seekable stream can be replayed after a token refresh. To report + /// progress, wrap the stream in a . + /// Use to put the result in a generated record. + /// /// The ATProto client. - /// The blob content stream. + /// The blob content stream. It is read from its current position and not disposed. /// The MIME type of the blob. /// Cancellation token. /// A reference to the uploaded blob. @@ -35,6 +46,17 @@ public static async Task UploadBlobAsync( throw new InvalidOperationException("Blob upload requires authentication."); } + if (client is IXrpcRequestClient xrpc) + { + var request = new XrpcRequest(HttpMethod.Post, UploadBlobNsid) + { + Body = XrpcBody.FromStream(content, mimeType), + }; + + using var response = await xrpc.SendXrpcAsync(request, cancellationToken).ConfigureAwait(false); + return await ReadUploadResponseAsync(response, cancellationToken).ConfigureAwait(false); + } + // Get the token provider to access the access token var tokenProvider = client.TokenProvider ?? throw new InvalidOperationException("No token provider available."); @@ -91,6 +113,12 @@ public static async Task UploadBlobFromFileAsync( /// /// Downloads a blob from a PDS. /// + /// + /// With a client that implements , a blob owned by another + /// account is fetched from that account's PDS (found with the client's + /// ) without session credentials; the user's own + /// blobs are fetched from their PDS with the client's authentication. + /// /// The ATProto client. /// The DID of the repo that owns the blob. /// The CID of the blob to download. @@ -102,6 +130,28 @@ public static async Task DownloadBlobAsync( ATCid cid, CancellationToken cancellationToken = default) { + if (client is IXrpcRequestClient xrpc) + { + XrpcRequestOptions? options = null; + var owner = did.ToString(); + if (!string.Equals(owner, client.AuthenticatedDid, StringComparison.Ordinal) && client.IdentityResolver != null) + { + var didDoc = await client.IdentityResolver.ResolveAsync(owner, cancellationToken).ConfigureAwait(false); + if (didDoc.PdsEndpoint != null) + { + options = new XrpcRequestOptions { ServiceUrl = new Uri(didDoc.PdsEndpoint) }; + } + } + + var parameters = new[] + { + new KeyValuePair("did", owner), + new KeyValuePair("cid", cid.ToString()), + }; + + return await xrpc.QueryBytesAsync(GetBlobNsid, parameters, options, cancellationToken).ConfigureAwait(false); + } + var url = new Uri(client.BaseUrl, $"/xrpc/com.atproto.sync.getBlob?did={did}&cid={cid}"); using var request = new HttpRequestMessage(HttpMethod.Get, url); @@ -164,8 +214,12 @@ private static async Task UploadBlobInternalAsync( request.Headers.Authorization = new AuthenticationHeaderValue("Bearer", accessToken); } - var response = await client.HttpClient.SendAsync(request, cancellationToken).ConfigureAwait(false); + using var response = await client.HttpClient.SendAsync(request, cancellationToken).ConfigureAwait(false); + return await ReadUploadResponseAsync(response, cancellationToken).ConfigureAwait(false); + } + private static async Task ReadUploadResponseAsync(HttpResponseMessage response, CancellationToken cancellationToken) + { if (!response.IsSuccessStatusCode) { var errorContent = await response.Content.ReadAsStringAsync().ConfigureAwait(false); diff --git a/src/CarpaNet/Blob/BlobRef.cs b/src/CarpaNet/Blob/BlobRef.cs index 30845ab..9badd3a 100644 --- a/src/CarpaNet/Blob/BlobRef.cs +++ b/src/CarpaNet/Blob/BlobRef.cs @@ -31,6 +31,22 @@ public sealed class BlobRef /// [JsonPropertyName("size")] public long Size { get; set; } + + /// + /// Converts this reference to the type used by generated records. + /// + /// The blob. + /// If the reference has no CID. + public ATBlob ToATBlob() + { + var link = Ref?.Link; + if (string.IsNullOrEmpty(link)) + { + throw new InvalidOperationException("The blob reference has no CID."); + } + + return new ATBlob(new ATCid(link!), MimeType ?? string.Empty, Size); + } } /// diff --git a/src/CarpaNet/Http/HttpClientFactory.cs b/src/CarpaNet/Http/HttpClientFactory.cs index 9573a8d..eaad3af 100644 --- a/src/CarpaNet/Http/HttpClientFactory.cs +++ b/src/CarpaNet/Http/HttpClientFactory.cs @@ -60,34 +60,16 @@ public static HttpMessageHandler CreateHandler(HttpClientFactoryOptions? options HttpMessageHandler handler; #if NET5_0_OR_GREATER - // Use SocketsHttpHandler for best performance - var socketsHandler = new SocketsHttpHandler + if (OperatingSystem.IsBrowser()) { - // Connection pooling - PooledConnectionIdleTimeout = options.PooledConnectionIdleTimeout ?? DefaultPooledConnectionIdleTimeout, - PooledConnectionLifetime = options.PooledConnectionLifetime ?? DefaultPooledConnectionLifetime, - MaxConnectionsPerServer = options.MaxConnectionsPerServer ?? 10, - - // Enable automatic decompression - AutomaticDecompression = DecompressionMethods.GZip | DecompressionMethods.Deflate, - - // Connection settings - ConnectTimeout = options.ConnectTimeout ?? TimeSpan.FromSeconds(30), - KeepAlivePingPolicy = HttpKeepAlivePingPolicy.WithActiveRequests, - KeepAlivePingTimeout = TimeSpan.FromSeconds(15), - KeepAlivePingDelay = TimeSpan.FromSeconds(30), - - // Enable cookies if needed - UseCookies = options.UseCookies, - }; - - // Enable HTTP/2 if supported - if (options.EnableHttp2) + // The browser fetch handler supports neither SocketsHttpHandler nor the + // HttpClientHandler connection settings, so it is used as is. + handler = new HttpClientHandler(); + } + else { - socketsHandler.EnableMultipleHttp2Connections = true; + handler = CreateSocketsHandler(options); } - - handler = socketsHandler; #else // Fallback for older frameworks var httpHandler = new HttpClientHandler @@ -113,6 +95,40 @@ public static HttpMessageHandler CreateHandler(HttpClientFactoryOptions? options return handler; } + +#if NET5_0_OR_GREATER + private static SocketsHttpHandler CreateSocketsHandler(HttpClientFactoryOptions options) + { + // Use SocketsHttpHandler for best performance + var socketsHandler = new SocketsHttpHandler + { + // Connection pooling + PooledConnectionIdleTimeout = options.PooledConnectionIdleTimeout ?? DefaultPooledConnectionIdleTimeout, + PooledConnectionLifetime = options.PooledConnectionLifetime ?? DefaultPooledConnectionLifetime, + MaxConnectionsPerServer = options.MaxConnectionsPerServer ?? 10, + + // Enable automatic decompression + AutomaticDecompression = DecompressionMethods.GZip | DecompressionMethods.Deflate, + + // Connection settings + ConnectTimeout = options.ConnectTimeout ?? TimeSpan.FromSeconds(30), + KeepAlivePingPolicy = HttpKeepAlivePingPolicy.WithActiveRequests, + KeepAlivePingTimeout = TimeSpan.FromSeconds(15), + KeepAlivePingDelay = TimeSpan.FromSeconds(30), + + // Enable cookies if needed + UseCookies = options.UseCookies, + }; + + // Enable HTTP/2 if supported + if (options.EnableHttp2) + { + socketsHandler.EnableMultipleHttp2Connections = true; + } + + return socketsHandler; + } +#endif } /// diff --git a/src/CarpaNet/Http/ProgressReportingStream.cs b/src/CarpaNet/Http/ProgressReportingStream.cs new file mode 100644 index 0000000..54d321b --- /dev/null +++ b/src/CarpaNet/Http/ProgressReportingStream.cs @@ -0,0 +1,135 @@ +using System; +using System.IO; +using System.Threading; +using System.Threading.Tasks; + +namespace CarpaNet.Http; + +/// +/// A read-only stream wrapper that reports how many bytes have been read, for upload progress. +/// +/// +/// Wrap the content passed to an upload (for example +/// +/// or a generated binary procedure). The reported value is the position in the inner stream, so it +/// drops back when the stream is rewound for a retry. +/// +public sealed class ProgressReportingStream : Stream +{ + private readonly Stream _inner; + private readonly IProgress _progress; + private readonly bool _leaveOpen; + private long _bytesRead; + + /// + /// Creates a progress-reporting wrapper. + /// + /// The stream to read. + /// Receives the total number of bytes read after each read. + /// Whether disposing the wrapper leaves open. + public ProgressReportingStream(Stream inner, IProgress progress, bool leaveOpen = true) + { + _inner = inner ?? throw new ArgumentNullException(nameof(inner)); + _progress = progress ?? throw new ArgumentNullException(nameof(progress)); + _leaveOpen = leaveOpen; + _bytesRead = inner.CanSeek ? inner.Position : 0; + } + + /// + /// Gets the number of bytes read so far. + /// + public long BytesRead => _bytesRead; + + /// + public override bool CanRead => _inner.CanRead; + + /// + public override bool CanSeek => _inner.CanSeek; + + /// + public override bool CanWrite => false; + + /// + public override long Length => _inner.Length; + + /// + public override long Position + { + get => _inner.Position; + set + { + _inner.Position = value; + Report(value); + } + } + + /// + public override int Read(byte[] buffer, int offset, int count) + { + var read = _inner.Read(buffer, offset, count); + Advance(read); + return read; + } + + /// + public override async Task ReadAsync(byte[] buffer, int offset, int count, CancellationToken cancellationToken) + { + var read = await _inner.ReadAsync(buffer, offset, count, cancellationToken).ConfigureAwait(false); + Advance(read); + return read; + } + +#if NET8_0_OR_GREATER + /// + public override async ValueTask ReadAsync(Memory buffer, CancellationToken cancellationToken = default) + { + var read = await _inner.ReadAsync(buffer, cancellationToken).ConfigureAwait(false); + Advance(read); + return read; + } +#endif + + /// + public override long Seek(long offset, SeekOrigin origin) + { + var position = _inner.Seek(offset, origin); + Report(position); + return position; + } + + /// + public override void Flush() + { + } + + /// + public override void SetLength(long value) => throw new NotSupportedException(); + + /// + public override void Write(byte[] buffer, int offset, int count) => throw new NotSupportedException(); + + /// + protected override void Dispose(bool disposing) + { + if (disposing && !_leaveOpen) + { + _inner.Dispose(); + } + + base.Dispose(disposing); + } + + private void Advance(int read) + { + if (read > 0) + { + Report(_bytesRead + read); + } + } + + private void Report(long bytesRead) + { + _bytesRead = bytesRead; + _progress.Report(bytesRead); + } +} diff --git a/src/CarpaNet/Http/RateLimitHandler.cs b/src/CarpaNet/Http/RateLimitHandler.cs index b92a603..50a3f26 100644 --- a/src/CarpaNet/Http/RateLimitHandler.cs +++ b/src/CarpaNet/Http/RateLimitHandler.cs @@ -79,50 +79,118 @@ protected override async Task SendAsync( CancellationToken cancellationToken) { var attempt = 0; - HttpResponseMessage? response = null; + HttpRequestMessage? retryRequest = null; - while (true) + try { - attempt++; - - // Clone request for retry (requests can only be sent once) - using var requestClone = attempt > 1 ? await CloneRequestAsync(request).ConfigureAwait(false) : null; - var requestToSend = requestClone ?? request; - - response = await base.SendAsync(requestToSend, cancellationToken).ConfigureAwait(false); - - // Check for rate limit - if (response.StatusCode != (HttpStatusCode)429) + while (true) { - return response; + attempt++; + + var response = await base.SendAsync(retryRequest ?? request, cancellationToken).ConfigureAwait(false); + + // Check for rate limit + if (response.StatusCode != (HttpStatusCode)429) + { + return response; + } + + // Parse rate limit info + var rateLimitInfo = RateLimitInfo.FromResponse(response); + + // Raise event + RateLimitEncountered?.Invoke(this, new RateLimitEventArgs(rateLimitInfo, attempt)); + + // Check if we should retry + if (!AutoRetryOnRateLimit || attempt >= MaxRetries) + { + _logger.LogWarning("Rate limit max retries exceeded"); + return response; + } + + // A DPoP proof is single use: a signed request can only be retried when the sender + // supplied a callback that signs the copy again. + var prepareRetry = GetRetryPreparer(request); + if (prepareRetry == null && request.Headers.Contains("DPoP")) + { + _logger.LogWarning("Rate limited (429) on a DPoP-signed request without a retry preparer; not retrying"); + return response; + } + + // Clone the request for the retry (requests can only be sent once). A body that + // cannot be read again (a consumed, non-seekable stream) cannot be retried. + HttpRequestMessage nextRequest; + try + { + nextRequest = await CloneRequestAsync(request).ConfigureAwait(false); + } + catch (InvalidOperationException ex) + { + _logger.LogWarning(ex, "Rate limited (429) but the request body cannot be replayed; not retrying"); + return response; + } + + if (prepareRetry != null) + { + await prepareRetry(nextRequest).ConfigureAwait(false); + } + + retryRequest?.Dispose(); + retryRequest = nextRequest; + + _logger.LogWarning("Rate limited (429) on attempt {Attempt}/{MaxRetries}", attempt, MaxRetries); + + // Calculate delay + var delay = CalculateDelay(rateLimitInfo, attempt); + _logger.LogDebug("Rate limit retry after {DelayMs}ms", (int)delay.TotalMilliseconds); + + // Dispose the response before retrying + response.Dispose(); + + // Wait before retrying + await Task.Delay(delay, cancellationToken).ConfigureAwait(false); } + } + finally + { + retryRequest?.Dispose(); + } + } - // Parse rate limit info - var rateLimitInfo = RateLimitInfo.FromResponse(response); - - // Raise event - RateLimitEncountered?.Invoke(this, new RateLimitEventArgs(rateLimitInfo, attempt)); - - // Check if we should retry - if (!AutoRetryOnRateLimit || attempt >= MaxRetries) - { - _logger.LogWarning("Rate limit max retries exceeded"); - return response; - } + /// + /// Attaches a callback that updates the copy of made for each + /// rate-limit retry, before it is sent. Use it for single-use credentials such as DPoP proofs. + /// + /// The original request. + /// Receives the copy to update. + public static void SetRetryPreparer(HttpRequestMessage request, Func prepareRetry) + { + if (request == null) + throw new ArgumentNullException(nameof(request)); + if (prepareRetry == null) + throw new ArgumentNullException(nameof(prepareRetry)); - _logger.LogWarning("Rate limited (429) on attempt {Attempt}/{MaxRetries}", attempt, MaxRetries); +#if NET5_0_OR_GREATER + request.Options.Set(RetryPreparerKey, prepareRetry); +#else + request.Properties[RetryPreparerName] = prepareRetry; +#endif + } - // Calculate delay - var delay = CalculateDelay(rateLimitInfo, attempt); - _logger.LogDebug("Rate limit retry after {DelayMs}ms", (int)delay.TotalMilliseconds); + private static Func? GetRetryPreparer(HttpRequestMessage request) + { +#if NET5_0_OR_GREATER + return request.Options.TryGetValue(RetryPreparerKey, out var prepareRetry) ? prepareRetry : null; +#else + return request.Properties.TryGetValue(RetryPreparerName, out var value) ? value as Func : null; +#endif + } - // Dispose the response before retrying - response.Dispose(); + private const string RetryPreparerName = "CarpaNet.RateLimitHandler.PrepareRetry"; - // Wait before retrying - await Task.Delay(delay, cancellationToken).ConfigureAwait(false); - } - } +#if NET5_0_OR_GREATER + private static readonly HttpRequestOptionsKey> RetryPreparerKey = new(RetryPreparerName); +#endif private TimeSpan CalculateDelay(RateLimitInfo? rateLimitInfo, int attempt) { diff --git a/src/CarpaNet/Http/XrpcHttpHandler.cs b/src/CarpaNet/Http/XrpcHttpHandler.cs index 4ce7aae..3c05848 100644 --- a/src/CarpaNet/Http/XrpcHttpHandler.cs +++ b/src/CarpaNet/Http/XrpcHttpHandler.cs @@ -208,6 +208,43 @@ public static void AddCommonHeaders( } } + /// + /// Adds caller-supplied headers to a request. A header that already exists on the request is replaced. + /// + /// The request. + /// The headers to add, or null. + public static void AddCustomHeaders(HttpRequestMessage request, IReadOnlyDictionary? headers) + { + if (headers == null) + { + return; + } + + foreach (var header in headers) + { + request.Headers.Remove(header.Key); + request.Headers.TryAddWithoutValidation(header.Key, header.Value); + } + } + + /// + /// Returns whether two URLs have the same scheme, host and port. + /// Used to decide whether session credentials may be sent with a request. + /// + /// The request URL. + /// The URL of the server that issued the credentials. + public static bool IsSameOrigin(Uri url, Uri origin) + { + if (url == null) + throw new ArgumentNullException(nameof(url)); + if (origin == null) + throw new ArgumentNullException(nameof(origin)); + + return string.Equals(url.Scheme, origin.Scheme, StringComparison.OrdinalIgnoreCase) + && string.Equals(url.IdnHost, origin.IdnHost, StringComparison.OrdinalIgnoreCase) + && url.Port == origin.Port; + } + /// /// Processes an HTTP response and deserializes the result. /// Throws appropriate exceptions for error responses. diff --git a/src/CarpaNet/IXrpcRequestClient.cs b/src/CarpaNet/IXrpcRequestClient.cs new file mode 100644 index 0000000..ffae59a --- /dev/null +++ b/src/CarpaNet/IXrpcRequestClient.cs @@ -0,0 +1,30 @@ +using System.Net.Http; +using System.Text.Json; +using System.Threading; +using System.Threading.Tasks; + +namespace CarpaNet; + +/// +/// The low-level XRPC send operation shared by the client implementations. +/// +/// +/// , the OAuth client and implement this +/// interface. The extension methods in use it for binary +/// bodies, binary responses, procedures with query parameters and per-request options. +/// +public interface IXrpcRequestClient +{ + /// + /// Gets the JSON options used to serialize request bodies and deserialize responses. + /// + JsonSerializerOptions JsonOptions { get; } + + /// + /// Sends an XRPC request through the client's authentication pipeline and returns the raw response. + /// + /// The request. + /// Cancellation token. + /// The HTTP response, which the caller disposes. Error statuses are not thrown. + Task SendXrpcAsync(XrpcRequest request, CancellationToken cancellationToken = default); +} diff --git a/src/CarpaNet/Identity/DnsOverHttpsResolver.cs b/src/CarpaNet/Identity/DnsOverHttpsResolver.cs new file mode 100644 index 0000000..bb03dff --- /dev/null +++ b/src/CarpaNet/Identity/DnsOverHttpsResolver.cs @@ -0,0 +1,308 @@ +using System; +using System.Collections.Generic; +using System.IO; +using System.Net.Http; +using System.Text; +using System.Text.Json; +using System.Text.Json.Serialization; +using System.Threading; +using System.Threading.Tasks; + +namespace CarpaNet.Identity; + +/// +/// DNS resolver that queries TXT records over HTTPS using the JSON API +/// (application/dns-json) served by public resolvers such as Cloudflare and Google. +/// +/// +/// +/// Use this resolver where raw UDP DNS is unavailable, such as in browsers (WebAssembly), +/// or on networks that block outbound UDP port 53. +/// +/// +/// Each endpoint is queried with +/// GET {endpoint}?name={name}&type=TXT and Accept: application/dns-json. +/// Endpoints are tried in order. The resolver moves to the next endpoint when a request +/// fails, times out, returns a non-success HTTP status, returns malformed JSON, or returns a +/// DNS status other than NOERROR (0) or NXDOMAIN (3). NXDOMAIN and NOERROR are authoritative +/// answers and stop the fallback. When every endpoint fails, an empty list is returned, +/// which matches . +/// +/// +public sealed class DnsOverHttpsResolver : IDnsResolver +{ + private const int DnsStatusNoError = 0; + private const int DnsStatusNxDomain = 3; + private const int DnsTypeTxt = 16; + + private readonly HttpClient _httpClient; + private readonly string[] _endpoints; + private readonly TimeSpan _timeout; + + /// + /// Cloudflare DNS-over-HTTPS JSON endpoint. + /// + public const string CloudflareEndpoint = "https://cloudflare-dns.com/dns-query"; + + /// + /// Google Public DNS JSON endpoint. + /// + public const string GoogleEndpoint = "https://dns.google/resolve"; + + /// + /// Default endpoints, tried in order (Cloudflare, then Google). + /// + public static IReadOnlyList DefaultEndpoints { get; } = new[] { CloudflareEndpoint, GoogleEndpoint }; + + /// + /// Default timeout for each endpoint request. + /// + public static readonly TimeSpan DefaultTimeout = TimeSpan.FromSeconds(5); + + /// + /// Creates a new DnsOverHttpsResolver. + /// + /// The HttpClient used for requests. The resolver does not dispose it. + /// DNS JSON endpoints to query, in fallback order. If null or empty, is used. + /// Timeout for each endpoint request. If null, is used. + public DnsOverHttpsResolver(HttpClient httpClient, IEnumerable? endpoints = null, TimeSpan? timeout = null) + { + _httpClient = httpClient ?? throw new ArgumentNullException(nameof(httpClient)); + + var list = new List(); + if (endpoints != null) + { + foreach (var endpoint in endpoints) + { + if (!string.IsNullOrWhiteSpace(endpoint)) + list.Add(endpoint.Trim()); + } + } + + _endpoints = list.Count > 0 ? list.ToArray() : new List(DefaultEndpoints).ToArray(); + _timeout = timeout ?? DefaultTimeout; + + if (_timeout <= TimeSpan.Zero && _timeout != Timeout.InfiniteTimeSpan) + throw new ArgumentOutOfRangeException(nameof(timeout), "Timeout must be positive or Timeout.InfiniteTimeSpan."); + } + + /// + /// Gets the configured endpoints, in fallback order. + /// + public IReadOnlyList Endpoints => _endpoints; + + /// + /// + /// Multiple character-strings in one TXT record are joined into one string. + /// Cancellation of is propagated as an + /// ; per-endpoint timeouts are not. + /// + public async Task> GetTxtRecordsAsync(string name, CancellationToken cancellationToken = default) + { + if (string.IsNullOrWhiteSpace(name)) + throw new ArgumentException("Name cannot be empty", nameof(name)); + + foreach (var endpoint in _endpoints) + { + cancellationToken.ThrowIfCancellationRequested(); + + var records = await TryQueryEndpointAsync(endpoint, name, cancellationToken).ConfigureAwait(false); + if (records != null) + return records; + } + + return Array.Empty(); + } + + /// + /// Queries one endpoint. Returns null when the next endpoint should be tried. + /// + private async Task?> TryQueryEndpointAsync(string endpoint, string name, CancellationToken cancellationToken) + { + var separator = endpoint.IndexOf('?') >= 0 ? "&" : "?"; + var url = $"{endpoint}{separator}name={Uri.EscapeDataString(name)}&type=TXT"; + + using var cts = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken); + if (_timeout != Timeout.InfiniteTimeSpan) + cts.CancelAfter(_timeout); + + try + { + using var request = new HttpRequestMessage(HttpMethod.Get, url); + request.Headers.TryAddWithoutValidation("Accept", "application/dns-json"); + + using var response = await _httpClient.SendAsync(request, HttpCompletionOption.ResponseHeadersRead, cts.Token).ConfigureAwait(false); + if (!response.IsSuccessStatusCode) + return null; + +#if NET5_0_OR_GREATER + using var stream = await response.Content.ReadAsStreamAsync(cts.Token).ConfigureAwait(false); +#else + using var stream = await response.Content.ReadAsStreamAsync().ConfigureAwait(false); +#endif + var dnsResponse = await JsonSerializer.DeserializeAsync(stream, IdentityJsonContext.Default.DnsJsonResponse, cts.Token).ConfigureAwait(false); + if (dnsResponse == null) + return null; + + return ParseResponse(dnsResponse); + } + catch (OperationCanceledException) when (!cancellationToken.IsCancellationRequested) + { + // Per-endpoint timeout; try next endpoint + return null; + } + catch (HttpRequestException) + { + return null; + } + catch (JsonException) + { + return null; + } + catch (IOException) + { + return null; + } + } + + /// + /// Converts a DNS JSON response to TXT strings. Returns null when the status is not authoritative. + /// + internal static IReadOnlyList? ParseResponse(DnsJsonResponse response) + { + if (response.Status == DnsStatusNxDomain) + return Array.Empty(); + + if (response.Status != DnsStatusNoError) + return null; + + var results = new List(); + if (response.Answer == null) + return results; + + foreach (var answer in response.Answer) + { + if (answer == null || answer.Type != DnsTypeTxt || answer.Data == null) + continue; + + results.Add(ParseTxtData(answer.Data)); + } + + return results; + } + + /// + /// Parses TXT presentation data. Quoted character-strings are unquoted, unescaped, and joined. + /// Unquoted data is returned trimmed. + /// + /// The data value of a TXT answer, for example "\"did=did:plc:abc\"". + /// The TXT record text. + public static string ParseTxtData(string data) + { + if (data == null) + throw new ArgumentNullException(nameof(data)); + + var trimmed = data.Trim(); + if (trimmed.Length == 0 || trimmed[0] != '"') + return trimmed; + + var builder = new StringBuilder(trimmed.Length); + var i = 0; + while (i < trimmed.Length) + { + // Skip whitespace between character-strings + while (i < trimmed.Length && char.IsWhiteSpace(trimmed[i])) + i++; + + if (i >= trimmed.Length) + break; + + if (trimmed[i] != '"') + { + // Unquoted segment: read up to the next whitespace + var start = i; + while (i < trimmed.Length && !char.IsWhiteSpace(trimmed[i])) + i++; + builder.Append(trimmed, start, i - start); + continue; + } + + i++; // Opening quote + while (i < trimmed.Length && trimmed[i] != '"') + { + var c = trimmed[i]; + if (c == '\\' && i + 1 < trimmed.Length) + { + // \DDD decimal escape + if (i + 3 < trimmed.Length + && char.IsDigit(trimmed[i + 1]) && char.IsDigit(trimmed[i + 2]) && char.IsDigit(trimmed[i + 3])) + { + var value = ((trimmed[i + 1] - '0') * 100) + ((trimmed[i + 2] - '0') * 10) + (trimmed[i + 3] - '0'); + builder.Append((char)value); + i += 4; + continue; + } + + builder.Append(trimmed[i + 1]); + i += 2; + continue; + } + + builder.Append(c); + i++; + } + + i++; // Closing quote + } + + return builder.ToString(); + } +} + +/// +/// DNS JSON API response (application/dns-json). +/// +internal sealed class DnsJsonResponse +{ + /// + /// DNS response code (0 = NOERROR, 2 = SERVFAIL, 3 = NXDOMAIN). + /// + [JsonPropertyName("Status")] + public int Status { get; set; } + + /// + /// Answer records. + /// + [JsonPropertyName("Answer")] + public List? Answer { get; set; } +} + +/// +/// One answer record in a DNS JSON API response. +/// +internal sealed class DnsJsonAnswer +{ + /// + /// Record owner name. + /// + [JsonPropertyName("name")] + public string? Name { get; set; } + + /// + /// Record type (16 = TXT). + /// + [JsonPropertyName("type")] + public int Type { get; set; } + + /// + /// Time to live in seconds. + /// + [JsonPropertyName("TTL")] + public int Ttl { get; set; } + + /// + /// Record data in presentation format. + /// + [JsonPropertyName("data")] + public string? Data { get; set; } +} diff --git a/src/CarpaNet/Identity/DnsResolverDefaults.cs b/src/CarpaNet/Identity/DnsResolverDefaults.cs new file mode 100644 index 0000000..43c0a0c --- /dev/null +++ b/src/CarpaNet/Identity/DnsResolverDefaults.cs @@ -0,0 +1,47 @@ +using System; +using System.Net.Http; +using System.Runtime.InteropServices; + +namespace CarpaNet.Identity; + +/// +/// Selects the default for the current platform. +/// +public static class DnsResolverDefaults +{ + /// + /// Gets a value indicating whether the current platform can send raw UDP DNS queries. + /// This is false in browsers (WebAssembly) and under WASI. + /// + public static bool IsUdpDnsSupported { get; } = !IsSandboxedPlatform(); + + /// + /// Creates the default DNS resolver for the current platform. + /// + /// + /// Returns a (with ) + /// when is false, and a (UDP) otherwise. + /// On networks that block UDP port 53, create a directly. + /// + /// The HttpClient for DNS-over-HTTPS requests. It is not used by the UDP resolver, and it is not disposed. + /// The DNS resolver. + public static IDnsResolver CreateDefault(HttpClient httpClient) + { + if (httpClient == null) + throw new ArgumentNullException(nameof(httpClient)); + + return IsUdpDnsSupported + ? new DefaultDnsResolver() + : new DnsOverHttpsResolver(httpClient); + } + + private static bool IsSandboxedPlatform() + { +#if NET8_0_OR_GREATER + return OperatingSystem.IsBrowser() || OperatingSystem.IsWasi(); +#else + return RuntimeInformation.IsOSPlatform(OSPlatform.Create("BROWSER")) + || RuntimeInformation.IsOSPlatform(OSPlatform.Create("WASI")); +#endif + } +} diff --git a/src/CarpaNet/Identity/IdentityJsonContext.cs b/src/CarpaNet/Identity/IdentityJsonContext.cs new file mode 100644 index 0000000..00cf5b7 --- /dev/null +++ b/src/CarpaNet/Identity/IdentityJsonContext.cs @@ -0,0 +1,12 @@ +using System.Text.Json.Serialization; + +namespace CarpaNet.Identity; + +/// +/// JSON serialization context for DNS-over-HTTPS and XRPC handle resolution responses. +/// +[JsonSerializable(typeof(DnsJsonResponse))] +[JsonSerializable(typeof(ResolveHandleResponse))] +internal partial class IdentityJsonContext : JsonSerializerContext +{ +} diff --git a/src/CarpaNet/Identity/IdentityResolver.cs b/src/CarpaNet/Identity/IdentityResolver.cs index 027a42f..f8eacd8 100644 --- a/src/CarpaNet/Identity/IdentityResolver.cs +++ b/src/CarpaNet/Identity/IdentityResolver.cs @@ -12,15 +12,24 @@ namespace CarpaNet.Identity; /// /// Resolves ATProtocol identities (handles and DIDs) to DID documents. -/// Supports did:plc and did:web methods, and handle resolution via DNS TXT and HTTPS. +/// Supports did:plc and did:web methods, and handle resolution via DNS TXT, HTTPS well-known, +/// and (optionally) the com.atproto.identity.resolveHandle XRPC method. /// +/// +/// Handle resolution methods are tried in the order set by +/// (default: DNS, well-known, XRPC). +/// When no DNS resolver is given, selects one for +/// the platform (DNS-over-HTTPS in browsers, UDP elsewhere). +/// public sealed class IdentityResolver : IDisposable { private readonly HttpClient _httpClient; private readonly bool _ownsHttpClient; private readonly string _plcDirectoryUrl; - private readonly IDnsResolver? _dnsResolver; + private readonly IDnsResolver _dnsResolver; private readonly IIdentityCache? _cache; + private readonly HandleResolutionMethod[] _handleResolutionOrder; + private readonly XrpcHandleResolver? _xrpcHandleResolver; private readonly ILogger _logger; /// @@ -33,7 +42,7 @@ public sealed class IdentityResolver : IDisposable /// /// Optional logger factory for diagnostic logging. public IdentityResolver(ILoggerFactory? loggerFactory = null) - : this(new HttpClient(), ownsHttpClient: true, DefaultPlcDirectory, new DefaultDnsResolver(), null, loggerFactory) + : this(new HttpClient(), ownsHttpClient: true, new IdentityResolverOptions(), loggerFactory) { } @@ -42,7 +51,7 @@ public IdentityResolver(ILoggerFactory? loggerFactory = null) /// /// The HttpClient to use for requests. /// The PLC directory URL (default: https://plc.directory). - /// Optional custom DNS resolver for handle resolution. + /// Optional custom DNS resolver for handle resolution. If null, is used. /// Optional identity cache for caching resolved identities. /// Optional logger factory for diagnostic logging. public IdentityResolver( @@ -51,26 +60,69 @@ public IdentityResolver( IDnsResolver? dnsResolver = null, IIdentityCache? cache = null, ILoggerFactory? loggerFactory = null) - : this(httpClient, ownsHttpClient: false, plcDirectoryUrl ?? DefaultPlcDirectory, dnsResolver, cache, loggerFactory) + : this( + httpClient, + ownsHttpClient: false, + new IdentityResolverOptions { PlcDirectoryUrl = plcDirectoryUrl, DnsResolver = dnsResolver, Cache = cache }, + loggerFactory) + { + } + + /// + /// Creates a new IdentityResolver with the specified options. + /// + /// The HttpClient to use for requests. It is not disposed by the resolver. + /// The resolver options. + /// Optional logger factory for diagnostic logging. + public IdentityResolver(HttpClient httpClient, IdentityResolverOptions options, ILoggerFactory? loggerFactory = null) + : this(httpClient, ownsHttpClient: false, options, loggerFactory) { } private IdentityResolver( HttpClient httpClient, bool ownsHttpClient, - string plcDirectoryUrl, - IDnsResolver? dnsResolver, - IIdentityCache? cache, + IdentityResolverOptions options, ILoggerFactory? loggerFactory) { + if (httpClient == null) + throw new ArgumentNullException(nameof(httpClient)); + if (options == null) + throw new ArgumentNullException(nameof(options)); + _httpClient = httpClient; _ownsHttpClient = ownsHttpClient; - _plcDirectoryUrl = plcDirectoryUrl.TrimEnd('/'); - _dnsResolver = dnsResolver; - _cache = cache; + _plcDirectoryUrl = (options.PlcDirectoryUrl ?? DefaultPlcDirectory).TrimEnd('/'); + _dnsResolver = options.DnsResolver ?? DnsResolverDefaults.CreateDefault(httpClient); + _cache = options.Cache; + + var order = options.HandleResolutionOrder; + if (order == null || order.Count == 0) + order = IdentityResolverOptions.DefaultHandleResolutionOrder; + var methods = new List(); + foreach (var method in order) + { + if (!methods.Contains(method)) + methods.Add(method); + } + _handleResolutionOrder = methods.ToArray(); + + if (!string.IsNullOrWhiteSpace(options.HandleResolutionServiceUrl)) + _xrpcHandleResolver = new XrpcHandleResolver(httpClient, options.HandleResolutionServiceUrl!); + _logger = (loggerFactory ?? NullLoggerFactory.Instance).CreateLogger(); } + /// + /// Gets the DNS resolver used for . + /// + public IDnsResolver DnsResolver => _dnsResolver; + + /// + /// Gets the handle resolution methods, in the order they are tried. + /// + public IReadOnlyList HandleResolutionOrder => _handleResolutionOrder; + /// /// Gets the identity cache, if configured. /// @@ -81,7 +133,7 @@ private IdentityResolver( /// /// Optional HttpClient to use for requests. If null, a new one will be created. /// The PLC directory URL (default: https://plc.directory). - /// Optional custom DNS resolver for handle resolution. + /// Optional custom DNS resolver for handle resolution. If null, is used. /// Optional logger factory for diagnostic logging. /// An IdentityResolver with caching enabled. public static IdentityResolver CreateWithCache( @@ -99,7 +151,7 @@ public static IdentityResolver CreateWithCache( /// The identity cache to use. /// Optional HttpClient to use for requests. If null, a new one will be created. /// The PLC directory URL (default: https://plc.directory). - /// Optional custom DNS resolver for handle resolution. + /// Optional custom DNS resolver for handle resolution. If null, is used. /// Optional logger factory for diagnostic logging. /// An IdentityResolver with the specified cache. public static IdentityResolver CreateWithCache( @@ -118,15 +170,19 @@ public static IdentityResolver CreateWithCache( return new IdentityResolver( httpClient, ownsHttpClient, - plcDirectoryUrl ?? DefaultPlcDirectory, - dnsResolver ?? new DefaultDnsResolver(), - cache, + new IdentityResolverOptions { PlcDirectoryUrl = plcDirectoryUrl, DnsResolver = dnsResolver, Cache = cache }, loggerFactory); } /// /// Resolves an identifier (handle or DID) to a DID document. /// + /// + /// For a handle, the DID document's handle claim (alsoKnownAs) must match the handle, + /// or an is thrown. This check runs for every + /// handle resolution method. When the DID came from , + /// the handle-to-DID direction is trusted from the service, not verified by DNS or well-known. + /// /// A handle (e.g., "alice.bsky.social") or DID (e.g., "did:plc:..."). /// Cancellation token. /// The resolved DID document. @@ -293,7 +349,8 @@ public async Task ResolveWebDidAsync(string did, CancellationToken } /// - /// Resolves a handle to a DID using DNS TXT or HTTPS well-known methods. + /// Resolves a handle to a DID using the configured handle resolution methods + /// (by default DNS TXT, then HTTPS well-known, then XRPC if a service is configured). /// /// The handle to resolve (e.g., "alice.bsky.social"). /// Cancellation token. @@ -304,7 +361,8 @@ public Task ResolveHandleAsync(string handle, CancellationToken cancella } /// - /// Resolves a handle to a DID using DNS TXT or HTTPS well-known methods. + /// Resolves a handle to a DID using the configured handle resolution methods + /// (by default DNS TXT, then HTTPS well-known, then XRPC if a service is configured). /// /// The handle to resolve (e.g., "alice.bsky.social"). /// If true, bypasses the cache and forces a fresh resolution. @@ -338,30 +396,28 @@ public async Task ResolveHandleAsync(string handle, bool skipCache, Canc _logger.LogTrace("Cache miss for {Key}", handle); } - // Try DNS TXT first (preferred) - var dnsDid = await TryResolveHandleDnsAsync(handle, cancellationToken).ConfigureAwait(false); - if (!string.IsNullOrEmpty(dnsDid)) + foreach (var method in _handleResolutionOrder) { - _logger.LogDebug("DNS resolution found {Did} for {Handle}", dnsDid, handle); - // Cache the result - if (_cache != null) + var did = method switch { - await _cache.SetHandleDidAsync(handle, dnsDid!, cancellationToken).ConfigureAwait(false); - } - return dnsDid!; - } + HandleResolutionMethod.Dns => await TryResolveHandleDnsAsync(handle, cancellationToken).ConfigureAwait(false), + HandleResolutionMethod.WellKnown => await TryResolveHandleHttpsAsync(handle, cancellationToken).ConfigureAwait(false), + HandleResolutionMethod.Xrpc => await TryResolveHandleXrpcAsync(handle, cancellationToken).ConfigureAwait(false), + _ => null + }; + + if (string.IsNullOrEmpty(did)) + continue; + + _logger.LogDebug("{Method} resolution found {Did} for {Handle}", method, did, handle); - // Fall back to HTTPS well-known - var httpsDid = await TryResolveHandleHttpsAsync(handle, cancellationToken).ConfigureAwait(false); - if (!string.IsNullOrEmpty(httpsDid)) - { - _logger.LogDebug("HTTPS resolution found {Did} for {Handle}", httpsDid, handle); // Cache the result if (_cache != null) { - await _cache.SetHandleDidAsync(handle, httpsDid!, cancellationToken).ConfigureAwait(false); + await _cache.SetHandleDidAsync(handle, did!, cancellationToken).ConfigureAwait(false); } - return httpsDid!; + + return did!; } _logger.LogError("Failed to resolve handle {Handle}", handle); @@ -373,9 +429,6 @@ public async Task ResolveHandleAsync(string handle, bool skipCache, Canc /// private async Task TryResolveHandleDnsAsync(string handle, CancellationToken cancellationToken) { - if (_dnsResolver == null) - return null; // DNS resolution not available - try { var txtRecordName = $"_atproto.{handle}"; @@ -392,9 +445,10 @@ public async Task ResolveHandleAsync(string handle, bool skipCache, Canc } } } - catch + catch (Exception ex) when (!cancellationToken.IsCancellationRequested) { - // DNS resolution failed, will try HTTPS + // DNS resolution failed, will try the next method + _logger.LogDebug(ex, "DNS resolution failed for {Handle}", handle); } return null; @@ -419,9 +473,31 @@ public async Task ResolveHandleAsync(string handle, bool skipCache, Canc if (did.StartsWith("did:", StringComparison.OrdinalIgnoreCase)) return did; } - catch + catch (Exception ex) when (!cancellationToken.IsCancellationRequested) + { + // HTTPS resolution failed, will try the next method + _logger.LogDebug(ex, "Well-known resolution failed for {Handle}", handle); + } + + return null; + } + + /// + /// Tries to resolve a handle via the configured XRPC service. + /// The result is trusted from the service, not verified against DNS or well-known. + /// + private async Task TryResolveHandleXrpcAsync(string handle, CancellationToken cancellationToken) + { + if (_xrpcHandleResolver == null) + return null; // No service configured + + try + { + return await _xrpcHandleResolver.ResolveHandleAsync(handle, cancellationToken).ConfigureAwait(false); + } + catch (Exception ex) when (!cancellationToken.IsCancellationRequested) { - // HTTPS resolution failed + _logger.LogDebug(ex, "XRPC resolution via {Service} failed for {Handle}", _xrpcHandleResolver.ServiceUrl, handle); } return null; diff --git a/src/CarpaNet/Identity/IdentityResolverOptions.cs b/src/CarpaNet/Identity/IdentityResolverOptions.cs new file mode 100644 index 0000000..e788de5 --- /dev/null +++ b/src/CarpaNet/Identity/IdentityResolverOptions.cs @@ -0,0 +1,95 @@ +using System.Collections.Generic; + +namespace CarpaNet.Identity; + +/// +/// A method used to resolve a handle to a DID. +/// +public enum HandleResolutionMethod +{ + /// + /// DNS TXT record at _atproto.{handle}, through the configured . + /// + Dns, + + /// + /// HTTPS request to https://{handle}/.well-known/atproto-did. + /// + WellKnown, + + /// + /// XRPC call to com.atproto.identity.resolveHandle on + /// . + /// The result is only as trustworthy as that service (see ). + /// Skipped when no service URL is configured. + /// + Xrpc, +} + +/// +/// Options for . +/// +public sealed class IdentityResolverOptions +{ + /// + /// Base URL of the public Bluesky AppView, which can be used as + /// . + /// + public const string PublicBlueskyAppViewUrl = "https://public.api.bsky.app"; + + /// + /// Gets or sets the PLC directory URL. If null, is used. + /// + public string? PlcDirectoryUrl { get; set; } + + /// + /// Gets or sets the DNS resolver for . + /// If null, is used. + /// + public IDnsResolver? DnsResolver { get; set; } + + /// + /// Gets or sets the identity cache. If null, results are not cached. + /// + public IIdentityCache? Cache { get; set; } + + /// + /// Gets or sets the handle resolution methods, in the order they are tried. + /// The first method that returns a DID wins. + /// If null or empty, is used. + /// + /// + /// To use only the XRPC service (for example in a browser): + /// + /// new IdentityResolverOptions + /// { + /// HandleResolutionServiceUrl = IdentityResolverOptions.PublicBlueskyAppViewUrl, + /// HandleResolutionOrder = new[] { HandleResolutionMethod.Xrpc }, + /// }; + /// + /// + public IReadOnlyList? HandleResolutionOrder { get; set; } + + /// + /// Gets or sets the base URL of a service (a PDS or AppView, for example + /// ) used for . + /// If null, the XRPC method is skipped. + /// + /// + /// A DID returned by this service is not verified against the handle's DNS record or + /// well-known file. It is only as trustworthy as the service. The DID document's handle + /// claim is still checked by . + /// + public string? HandleResolutionServiceUrl { get; set; } + + /// + /// Gets the default handle resolution order: DNS, then well-known, then XRPC + /// (XRPC only when is set). + /// + public static IReadOnlyList DefaultHandleResolutionOrder { get; } = new[] + { + HandleResolutionMethod.Dns, + HandleResolutionMethod.WellKnown, + HandleResolutionMethod.Xrpc, + }; +} diff --git a/src/CarpaNet/Identity/XrpcHandleResolver.cs b/src/CarpaNet/Identity/XrpcHandleResolver.cs new file mode 100644 index 0000000..6e8f3b4 --- /dev/null +++ b/src/CarpaNet/Identity/XrpcHandleResolver.cs @@ -0,0 +1,122 @@ +using System; +using System.IO; +using System.Net.Http; +using System.Text.Json; +using System.Text.Json.Serialization; +using System.Threading; +using System.Threading.Tasks; + +namespace CarpaNet.Identity; + +/// +/// Resolves handles to DIDs by calling com.atproto.identity.resolveHandle on an +/// ATProtocol service, such as a PDS or an AppView (for example https://public.api.bsky.app). +/// +/// +/// +/// This method works where DNS and /.well-known/atproto-did lookups do not, such as in +/// browsers, where UDP is unavailable and well-known requests are usually blocked by CORS. +/// +/// +/// Trust: the result is only as trustworthy as the service. The client does not check +/// the handle's DNS record or well-known file itself, so a malicious or out-of-date service can +/// return a wrong DID. still +/// checks that the DID document claims the handle, but that check does not prove that the +/// handle's domain points back to the DID. +/// +/// +public sealed class XrpcHandleResolver +{ + private readonly HttpClient _httpClient; + private readonly string _serviceUrl; + + /// + /// The XRPC method used for handle resolution. + /// + public const string ResolveHandleMethod = "com.atproto.identity.resolveHandle"; + + /// + /// Creates a new XrpcHandleResolver. + /// + /// The HttpClient used for requests. The resolver does not dispose it. + /// The base URL of the service (for example https://public.api.bsky.app). + public XrpcHandleResolver(HttpClient httpClient, string serviceUrl) + { + _httpClient = httpClient ?? throw new ArgumentNullException(nameof(httpClient)); + + if (string.IsNullOrWhiteSpace(serviceUrl)) + throw new ArgumentException("Service URL cannot be empty", nameof(serviceUrl)); + + if (!Uri.TryCreate(serviceUrl, UriKind.Absolute, out var uri) + || (uri.Scheme != Uri.UriSchemeHttps && uri.Scheme != Uri.UriSchemeHttp)) + throw new ArgumentException($"Service URL must be an absolute HTTP(S) URL: {serviceUrl}", nameof(serviceUrl)); + + _serviceUrl = serviceUrl.Trim().TrimEnd('/'); + } + + /// + /// Gets the base URL of the service. + /// + public string ServiceUrl => _serviceUrl; + + /// + /// Resolves a handle to a DID through the service. + /// + /// The handle to resolve (e.g., "alice.bsky.social"). + /// Cancellation token. + /// + /// The DID, or null if the service returns an error status (for example 400 when the handle + /// does not resolve) or a response that does not contain a valid DID. + /// + /// The request could not be sent. + public async Task ResolveHandleAsync(string handle, CancellationToken cancellationToken = default) + { + if (string.IsNullOrWhiteSpace(handle)) + throw new ArgumentException("Handle cannot be empty", nameof(handle)); + + var url = $"{_serviceUrl}/xrpc/{ResolveHandleMethod}?handle={Uri.EscapeDataString(handle)}"; + + using var request = new HttpRequestMessage(HttpMethod.Get, url); + request.Headers.TryAddWithoutValidation("Accept", "application/json"); + + using var response = await _httpClient.SendAsync(request, HttpCompletionOption.ResponseHeadersRead, cancellationToken).ConfigureAwait(false); + if (!response.IsSuccessStatusCode) + return null; + + try + { +#if NET5_0_OR_GREATER + using var stream = await response.Content.ReadAsStreamAsync(cancellationToken).ConfigureAwait(false); +#else + using var stream = await response.Content.ReadAsStreamAsync().ConfigureAwait(false); +#endif + var result = await JsonSerializer.DeserializeAsync(stream, IdentityJsonContext.Default.ResolveHandleResponse, cancellationToken).ConfigureAwait(false); + var did = result?.Did?.Trim(); + + if (did != null && IdentityResolver.IsValidDid(did)) + return did; + } + catch (JsonException) + { + // Malformed response + } + catch (IOException) + { + // Truncated response + } + + return null; + } +} + +/// +/// Output of com.atproto.identity.resolveHandle. +/// +internal sealed class ResolveHandleResponse +{ + /// + /// The resolved DID. + /// + [JsonPropertyName("did")] + public string? Did { get; set; } +} diff --git a/src/CarpaNet/README.md b/src/CarpaNet/README.md index 62af926..74a3ce9 100644 --- a/src/CarpaNet/README.md +++ b/src/CarpaNet/README.md @@ -78,7 +78,7 @@ var client = await ATProtoClient.CreateWithSessionAsync( ## Identity Resolution -`IdentityResolver` resolves handles to DIDs and DIDs to DID documents. Supports both `did:plc` (via PLC directory) and `did:web`. Handle resolution uses DNS TXT records with HTTPS fallback. +`IdentityResolver` resolves handles to DIDs and DIDs to DID documents. Supports both `did:plc` (via PLC directory) and `did:web`. Handle resolution uses DNS TXT records, then HTTPS well-known, then (optionally) the `com.atproto.identity.resolveHandle` XRPC method on a service that you configure with `IdentityResolverOptions`. In browsers, DNS uses DNS-over-HTTPS (`DnsOverHttpsResolver`) because UDP is not available. ```csharp var resolver = IdentityResolver.CreateWithCache(); @@ -233,10 +233,12 @@ This differs from `LexiconResolveAuthority` in that it takes a user handle rathe | Property | Description | Default | |---|---|---| -| `CarpaNet_SourceGen_RootNamespace` | Root namespace for generated code | Project's root namespace | +| `CarpaNet_RootNamespace` | Root namespace for generated code | None (namespaces come from the NSID) | | `CarpaNet_JsonContextName` | Name for the generated `JsonSerializerContext` class | `ATProtoJsonContext` | | `CarpaNet_CborContextName` | Name for the generated `CborSerializerContext` class | `ATProtoCborContext` | -| `CarpaNet_EmitValidationAttributes` | Emit validation attributes on generated properties | `false` | +| `CarpaNet_EmitValidationAttributes` | Emit validation attributes on generated properties | `true` | + +The older `CarpaNet_SourceGen_RootNamespace`, `CarpaNet_SourceGen_JsonContextName`, `CarpaNet_SourceGen_CborContextName` and `CarpaNet_SourceGen_EmitValidationAttributes` names are still accepted. If both are set, the `CarpaNet_*` name wins. ### Lexicon Resolution @@ -296,12 +298,25 @@ await client.AppBskyActorGetProfileAsync(parameters); await client.ComAtprotoRepoCreateRecordAsync(input); ``` -Query parameters are gathered into a `*Parameters` class with a `ToQueryParameters()` method returning `IEnumerable>` (so array params like `uris` can emit repeated keys). Procedure inputs use an `*Input` class. Subscriptions generate `SubscribeAsync` extensions returning `IAsyncEnumerable`. +Query parameters are gathered into a `*Parameters` class with a `ToQueryParameters()` method returning `IEnumerable>` (so array params like `uris` can emit repeated keys). Procedures that declare `parameters` get the same class and an extra `parameters` argument. Procedure inputs use an `*Input` class. Subscriptions generate `SubscribeAsync` extensions returning `IAsyncEnumerable`. + +Endpoints that do not use JSON bodies get different signatures: + +```csharp +// Procedure with a non-JSON input (com.atproto.repo.uploadBlob, encoding */*): +// the body is a Stream; contentType defaults to the lexicon encoding, or application/octet-stream for wildcards +var upload = await client.ComAtprotoRepoUploadBlobAsync(stream, "image/png"); + +// Query with a non-JSON output (com.atproto.sync.getBlob, com.atproto.sync.getRepo): returns the raw bytes +byte[] car = await client.ComAtprotoSyncGetRepoAsync(new ComAtproto.Sync.GetRepoParameters { Did = did }); +``` ### Union types (`UnionImplementations.g.cs`) Lexicon `union` references become discriminated union classes with `System.Text.Json` polymorphic serialization via `[JsonPolymorphic]` and `[JsonDerivedType]` attributes. +Open unions (no `closed: true`) can contain members from newer lexicon versions. For each open union interface `I{Name}`, the generator also emits a sealed `Unknown_{Name}` class (for example `AppBsky.Actor.Unknown_DefsPreferences`). A member with an unknown or missing `$type` is read into this class instead of being dropped. It has `Type` (the `$type` value, or an empty string), `Raw` (the whole member as a `JsonElement`) and, when read from CBOR, `RawCbor` (the original bytes). Serialization writes the member back unchanged, so a read-modify-write cycle does not lose data. The underscore keeps the name from colliding with lexicon-derived names, which never contain one after the first character. Closed unions still throw on an unknown `$type`. + ### JSON serialization context (`ATProtoJsonContext.g.cs`) A source-generated `JsonSerializerContext` that registers all generated types. The context name is configurable via `CarpaNet_JsonContextName`. diff --git a/src/CarpaNet/ScopedATProtoClient.cs b/src/CarpaNet/ScopedATProtoClient.cs new file mode 100644 index 0000000..aedbafb --- /dev/null +++ b/src/CarpaNet/ScopedATProtoClient.cs @@ -0,0 +1,120 @@ +using System; +using System.Collections.Generic; +using System.Net.Http; +using System.Text.Json; +using System.Threading; +using System.Threading.Tasks; +using CarpaNet.Auth; +using CarpaNet.Identity; + +namespace CarpaNet; + +/// +/// A view of another client that applies to every request. +/// Created with . +/// +/// +/// The scoped client shares the inner client's session, token provider and . +/// Disposing the inner client makes the scoped client unusable; the scoped client itself owns nothing. +/// +public sealed class ScopedATProtoClient : IATProtoClient, IXrpcRequestClient +{ + private readonly IXrpcRequestClient _xrpc; + + /// + /// Creates a scoped client. + /// + /// The client to send through. It must implement . + /// The options to apply to every request. + public ScopedATProtoClient(IATProtoClient inner, XrpcRequestOptions options) + { + Inner = inner ?? throw new ArgumentNullException(nameof(inner)); + Options = options ?? throw new ArgumentNullException(nameof(options)); + _xrpc = inner as IXrpcRequestClient + ?? throw new ArgumentException( + $"{inner.GetType().Name} does not implement {nameof(IXrpcRequestClient)}.", nameof(inner)); + } + + /// + /// Gets the client requests are sent through. + /// + public IATProtoClient Inner { get; } + + /// + /// Gets the options applied to every request. + /// + public XrpcRequestOptions Options { get; } + + /// + public Uri BaseUrl => Options.ServiceUrl ?? Inner.BaseUrl; + + /// + public bool IsAuthenticated => Inner.IsAuthenticated; + + /// + public string? AuthenticatedDid => Inner.AuthenticatedDid; + + /// + public IdentityResolver? IdentityResolver => Inner.IdentityResolver; + + /// + public ITokenProvider? TokenProvider => Inner.TokenProvider; + + /// + public HttpClient HttpClient => Inner.HttpClient; + + /// + public IReadOnlyList? LabelerDids => Options.AcceptLabelers ?? Inner.LabelerDids; + + /// + public JsonSerializerOptions JsonOptions => _xrpc.JsonOptions; + + /// + public Task SendXrpcAsync(XrpcRequest request, CancellationToken cancellationToken = default) + { + if (request == null) + { + throw new ArgumentNullException(nameof(request)); + } + + request.Options = XrpcRequestOptions.Combine(Options, request.Options); + return _xrpc.SendXrpcAsync(request, cancellationToken); + } + + /// + public Task GetAsync( + string nsid, + IEnumerable>? parameters = null, + CancellationToken cancellationToken = default) + => this.QueryAsync(nsid, parameters, null, cancellationToken); + + /// + public Task GetAsync( + string nsid, + string proxyServiceDid, + IEnumerable>? parameters = null, + CancellationToken cancellationToken = default) + => this.QueryAsync(nsid, parameters, new XrpcRequestOptions { ProxyServiceDid = proxyServiceDid }, cancellationToken); + + /// + public Task PostAsync( + string nsid, + TInput? input, + CancellationToken cancellationToken = default) + => this.ProcedureAsync(nsid, null, input, null, cancellationToken); + + /// + public Task PostAsync( + string nsid, + string proxyServiceDid, + TInput? input, + CancellationToken cancellationToken = default) + => this.ProcedureAsync(nsid, null, input, new XrpcRequestOptions { ProxyServiceDid = proxyServiceDid }, cancellationToken); + + /// + public IAsyncEnumerable SubscribeAsync( + string nsid, + IEnumerable>? parameters = null, + CancellationToken cancellationToken = default) + => Inner.SubscribeAsync(nsid, parameters, cancellationToken); +} diff --git a/src/CarpaNet/XrpcBody.cs b/src/CarpaNet/XrpcBody.cs new file mode 100644 index 0000000..0edbaf3 --- /dev/null +++ b/src/CarpaNet/XrpcBody.cs @@ -0,0 +1,175 @@ +using System; +using System.IO; +using System.Net.Http; +using System.Net.Http.Headers; +using System.Text.Json; +using System.Text.Json.Serialization.Metadata; +using System.Threading; +using System.Threading.Tasks; + +namespace CarpaNet; + +/// +/// The body of an XRPC procedure call. +/// +/// +/// A body can be sent more than once (for a retry after a token refresh) when it is held in +/// memory or comes from a seekable stream. A body from a non-seekable stream is sent once, +/// and a request that fails with 401 is not retried. +/// +public sealed class XrpcBody +{ + private readonly byte[]? _bytes; + private readonly Stream? _stream; + private readonly long _startPosition; + private bool _used; + + private XrpcBody(byte[]? bytes, Stream? stream, string contentType) + { + if (string.IsNullOrEmpty(contentType)) + { + throw new ArgumentException("Content type cannot be null or empty.", nameof(contentType)); + } + + _bytes = bytes; + _stream = stream; + _startPosition = stream != null && stream.CanSeek ? stream.Position : 0; + ContentType = contentType; + } + + /// + /// Gets the MIME type of the body. + /// + public string ContentType { get; } + + /// + /// Gets whether the body can be sent again. + /// + public bool IsReplayable => _bytes != null || (_stream?.CanSeek ?? false); + + /// + /// Creates a body from bytes. + /// + /// The body bytes. + /// The MIME type. + public static XrpcBody FromBytes(byte[] data, string contentType) + { + if (data == null) + { + throw new ArgumentNullException(nameof(data)); + } + + return new XrpcBody(data, null, contentType); + } + + /// + /// Creates a body that streams from , starting at its current position. + /// The stream is not disposed. + /// + /// The stream to read. + /// The MIME type. + public static XrpcBody FromStream(Stream stream, string contentType) + { + if (stream == null) + { + throw new ArgumentNullException(nameof(stream)); + } + + return new XrpcBody(null, stream, contentType); + } + + /// + /// Creates a JSON body. + /// + /// The value type. + /// The value to serialize. + /// The serialization metadata for . + public static XrpcBody FromJson(T value, JsonTypeInfo typeInfo) + { + if (typeInfo == null) + { + throw new ArgumentNullException(nameof(typeInfo)); + } + + return new XrpcBody(JsonSerializer.SerializeToUtf8Bytes(value, typeInfo), null, "application/json"); + } + + /// + /// Creates the HTTP content for one send. A stream body is rewound to its starting position + /// when it is sent again. Used by implementations. + /// + /// The content, or null when the body was already sent and cannot be replayed. + public HttpContent? CreateContent() + { + HttpContent content; + if (_bytes != null) + { + content = new ByteArrayContent(_bytes); + } + else + { + if (_used) + { + if (!_stream!.CanSeek) + { + return null; + } + + _stream.Position = _startPosition; + } + + content = new StreamContent(new NonDisposingStream(_stream!)); + } + + _used = true; + content.Headers.ContentType = MediaTypeHeaderValue.Parse(ContentType); + return content; + } + + /// + /// Wraps a caller-owned stream so that disposing the request does not dispose it. + /// + private sealed class NonDisposingStream : Stream + { + private readonly Stream _inner; + + public NonDisposingStream(Stream inner) + { + _inner = inner; + } + + public override bool CanRead => _inner.CanRead; + + public override bool CanSeek => _inner.CanSeek; + + public override bool CanWrite => false; + + public override long Length => _inner.Length; + + public override long Position + { + get => _inner.Position; + set => _inner.Position = value; + } + + public override void Flush() + { + } + + public override int Read(byte[] buffer, int offset, int count) => _inner.Read(buffer, offset, count); + + public override Task ReadAsync(byte[] buffer, int offset, int count, CancellationToken cancellationToken) + => _inner.ReadAsync(buffer, offset, count, cancellationToken); + +#if NET8_0_OR_GREATER + public override ValueTask ReadAsync(Memory buffer, CancellationToken cancellationToken = default) + => _inner.ReadAsync(buffer, cancellationToken); +#endif + + public override long Seek(long offset, SeekOrigin origin) => _inner.Seek(offset, origin); + + public override void SetLength(long value) => throw new NotSupportedException(); + + public override void Write(byte[] buffer, int offset, int count) => throw new NotSupportedException(); + } +} diff --git a/src/CarpaNet/XrpcRequest.cs b/src/CarpaNet/XrpcRequest.cs new file mode 100644 index 0000000..2c8df73 --- /dev/null +++ b/src/CarpaNet/XrpcRequest.cs @@ -0,0 +1,52 @@ +using System; +using System.Collections.Generic; +using System.Net.Http; + +namespace CarpaNet; + +/// +/// A single XRPC request, sent with . +/// +public sealed class XrpcRequest +{ + /// + /// Creates a request. + /// + /// for a query, for a procedure. + /// The NSID of the method. + public XrpcRequest(HttpMethod method, string nsid) + { + if (string.IsNullOrEmpty(nsid)) + { + throw new ArgumentException("NSID cannot be null or empty.", nameof(nsid)); + } + + Method = method ?? throw new ArgumentNullException(nameof(method)); + Nsid = nsid; + } + + /// + /// Gets the HTTP method. + /// + public HttpMethod Method { get; } + + /// + /// Gets the NSID of the method. + /// + public string Nsid { get; } + + /// + /// Gets or sets the query parameters. + /// + public IEnumerable>? Parameters { get; set; } + + /// + /// Gets or sets the request body (procedures only). + /// + public XrpcBody? Body { get; set; } + + /// + /// Gets or sets the per-request options. + /// + public XrpcRequestOptions? Options { get; set; } +} diff --git a/src/CarpaNet/XrpcRequestOptions.cs b/src/CarpaNet/XrpcRequestOptions.cs new file mode 100644 index 0000000..f945405 --- /dev/null +++ b/src/CarpaNet/XrpcRequestOptions.cs @@ -0,0 +1,151 @@ +using System; +using System.Collections.Generic; + +namespace CarpaNet; + +/// +/// Per-request settings for an XRPC call: service proxying, accepted labelers, extra headers, +/// and an alternate service to send the request to. +/// +/// +/// +/// Options are usually applied to many calls at once with +/// , +/// which returns a client whose generated API methods all use them. +/// +/// +/// Credentials are only attached when a request goes to the authenticated session's own PDS. +/// A request sent to carries no session credentials; put any +/// Authorization it needs (such as a service-auth token) in . +/// +/// +public sealed class XrpcRequestOptions +{ + /// + /// Gets or sets the service DID reference sent in the atproto-proxy header + /// (for example did:web:api.bsky.app#bsky_appview). + /// When set, it replaces the proxy a generated method would use. + /// + public string? ProxyServiceDid { get; set; } + + /// + /// Gets or sets whether the request is sent without an atproto-proxy header, + /// even when a generated method would add one. Takes precedence over . + /// + public bool DisableProxy { get; set; } + + /// + /// Gets or sets the labeler DIDs sent in the atproto-accept-labelers header, + /// replacing the client's own list. Entries may carry parameters such as ;redact + /// (see ). An empty list sends no header. + /// + public IReadOnlyList? AcceptLabelers { get; set; } + + /// + /// Gets or sets additional request headers. These are added after the XRPC headers and + /// may include Authorization, in which case session credentials are not attached. + /// + public IReadOnlyDictionary? Headers { get; set; } + + /// + /// Gets or sets the base URL of a service to send the request to instead of the session's PDS. + /// Requests to this service never carry session credentials and are not re-routed by the + /// repo parameter. + /// + public Uri? ServiceUrl { get; set; } + + /// + /// Gets the proxy to send, taking into account. + /// + public string? EffectiveProxyServiceDid => DisableProxy ? null : ProxyServiceDid; + + /// + /// Gets whether these options set the proxy, either to a service or to none. + /// + internal bool SetsProxy => DisableProxy || ProxyServiceDid != null; + + /// + /// Combines two sets of options. Values set on win; + /// headers are merged, with winning on the same name. + /// + /// The options that take precedence. + /// The options used for anything does not set. + /// The combined options, or null when both are null. + public static XrpcRequestOptions? Combine(XrpcRequestOptions? outer, XrpcRequestOptions? inner) + { + if (outer == null) + { + return inner; + } + + if (inner == null) + { + return outer; + } + + var outerSetsProxy = outer.SetsProxy; + return new XrpcRequestOptions + { + ProxyServiceDid = outerSetsProxy ? outer.ProxyServiceDid : inner.ProxyServiceDid, + DisableProxy = outerSetsProxy ? outer.DisableProxy : inner.DisableProxy, + AcceptLabelers = outer.AcceptLabelers ?? inner.AcceptLabelers, + Headers = MergeHeaders(outer.Headers, inner.Headers), + ServiceUrl = outer.ServiceUrl ?? inner.ServiceUrl, + }; + } + + private static IReadOnlyDictionary? MergeHeaders( + IReadOnlyDictionary? outer, + IReadOnlyDictionary? inner) + { + if (outer == null || outer.Count == 0) + { + return inner; + } + + if (inner == null || inner.Count == 0) + { + return outer; + } + + var merged = new Dictionary(StringComparer.OrdinalIgnoreCase); + foreach (var header in inner) + { + merged[header.Key] = header.Value; + } + + foreach (var header in outer) + { + merged[header.Key] = header.Value; + } + + return merged; + } +} + +/// +/// Helpers for the atproto-accept-labelers header. +/// +public static class AcceptLabelersHeader +{ + /// + /// The header name. + /// + public const string Name = "atproto-accept-labelers"; + + /// + /// Returns the header entry for a labeler whose takedown-level labels should be applied + /// by the AppView (redacting the content), for example did:plc:abc;redact. + /// + /// The labeler DID. + /// The DID with the ;redact parameter. + public static string Redact(string labelerDid) + { + if (string.IsNullOrEmpty(labelerDid)) + { + throw new ArgumentException("Labeler DID cannot be null or empty.", nameof(labelerDid)); + } + + return labelerDid + ";redact"; + } +} diff --git a/src/CarpaNet/build/CarpaNet.targets b/src/CarpaNet/build/CarpaNet.targets index f34bb0d..b28a6db 100644 --- a/src/CarpaNet/build/CarpaNet.targets +++ b/src/CarpaNet/build/CarpaNet.targets @@ -29,6 +29,12 @@ + + + + + + diff --git a/tests/CarpaNet.UnitTests/Auth/SessionInvalidationTests.cs b/tests/CarpaNet.UnitTests/Auth/SessionInvalidationTests.cs new file mode 100644 index 0000000..bc06a67 --- /dev/null +++ b/tests/CarpaNet.UnitTests/Auth/SessionInvalidationTests.cs @@ -0,0 +1,163 @@ +using System; +using System.Collections.Generic; +using System.Linq; +using System.Net; +using System.Net.Http; +using System.Text.Json; +using System.Threading.Tasks; +using CarpaNet.Auth; +using CarpaNet.UnitTests.Http; +using Xunit; + +namespace CarpaNet.UnitTests.Auth; + +/// +/// Tests for on and for +/// refreshing even while the token is still valid. +/// +public class SessionInvalidationTests +{ + private const string Did = "did:plc:user"; + private static readonly Uri Pds = new("https://pds.user.example"); + + [Theory] + [InlineData(HttpStatusCode.BadRequest, "ExpiredToken")] + [InlineData(HttpStatusCode.Unauthorized, "InvalidToken")] + public async Task RejectedRefresh_RaisesSessionInvalidatedOnce_AndClearsTokens(HttpStatusCode status, string error) + { + var handler = new RecordingHandler(_ => XrpcRequestPipelineTests.Json($"{{\"error\":\"{error}\",\"message\":\"x\"}}", status)); + using var provider = CreateProvider(handler); + var events = new List(); + provider.SessionInvalidated += (_, e) => events.Add(e); + + await Assert.ThrowsAnyAsync(() => provider.RefreshAsync()); + await Assert.ThrowsAsync(() => provider.RefreshAsync()); + + var raised = Assert.Single(events); + Assert.Equal(Did, raised.Did); + Assert.Equal(error, raised.Reason); + Assert.Null(provider.AccessJwt); + Assert.Null(provider.RefreshJwt); + Assert.Equal(Did, provider.CurrentDid); + Assert.False(provider.HasValidToken); + } + + [Theory] + [InlineData(HttpStatusCode.InternalServerError)] + [InlineData(HttpStatusCode.BadGateway)] + [InlineData((HttpStatusCode)429)] + public async Task TemporaryFailure_DoesNotInvalidate(HttpStatusCode status) + { + var handler = new RecordingHandler(_ => XrpcRequestPipelineTests.Status(status)); + using var provider = CreateProvider(handler); + var raised = false; + provider.SessionInvalidated += (_, _) => raised = true; + + await Assert.ThrowsAnyAsync(() => provider.RefreshAsync()); + + Assert.False(raised); + Assert.NotNull(provider.RefreshJwt); + } + + [Fact] + public async Task RestoreSession_ResetsInvalidation() + { + var handler = new RecordingHandler(_ => XrpcRequestPipelineTests.Json("{\"error\":\"ExpiredToken\"}", HttpStatusCode.BadRequest)); + using var provider = CreateProvider(handler); + var count = 0; + provider.SessionInvalidated += (_, _) => count++; + + await Assert.ThrowsAnyAsync(() => provider.RefreshAsync()); + Restore(provider); + await Assert.ThrowsAnyAsync(() => provider.RefreshAsync()); + + Assert.Equal(2, count); + } + + [Fact] + public async Task RefreshAsync_WhileTokenStillValid_StillRefreshes() + { + var fresh = XrpcRequestPipelineTests.CreateJwt(Did, DateTimeOffset.UtcNow.AddHours(3), "fresh"); + var handler = new RecordingHandler(_ => XrpcRequestPipelineTests.Json(XrpcRequestPipelineTests.SessionJson(fresh))); + using var provider = CreateProvider(handler); + Assert.True(provider.HasValidToken); + + await provider.RefreshAsync(); + + Assert.Single(handler.Requests); + Assert.Equal(fresh, provider.AccessJwt); + } + + [Fact] + public async Task ConcurrentRefreshes_HitTheServerOnce() + { + var gate = new TaskCompletionSource(); + var fresh = XrpcRequestPipelineTests.CreateJwt(Did, DateTimeOffset.UtcNow.AddHours(3), "fresh"); + var handler = new AsyncHandler(async () => + { + await gate.Task; + return XrpcRequestPipelineTests.Json(XrpcRequestPipelineTests.SessionJson(fresh)); + }); + using var provider = new SessionTokenProvider(new HttpClient(handler)); + Restore(provider); + + var refreshes = Enumerable.Range(0, 5).Select(_ => provider.RefreshAsync()).ToArray(); + gate.SetResult(true); + await Task.WhenAll(refreshes); + + Assert.Equal(1, handler.Calls); + } + + [Fact] + public async Task ATProtoClient_RejectedRefreshOn401_ReturnsOriginalError() + { + var handler = new RecordingHandler(r => r.Uri.AbsolutePath.EndsWith("refreshSession", StringComparison.Ordinal) + ? XrpcRequestPipelineTests.Json("{\"error\":\"ExpiredToken\"}", HttpStatusCode.BadRequest) + : XrpcRequestPipelineTests.Json("{\"error\":\"InvalidToken\"}", HttpStatusCode.Unauthorized)); + using var client = XrpcRequestPipelineTests.CreateSessionClient(handler, out _); + var provider = (INotifySessionInvalidated)client.TokenProvider!; + var raised = false; + provider.SessionInvalidated += (_, _) => raised = true; + + var ex = await Assert.ThrowsAsync(() => client.GetAsync("com.example.get")); + + Assert.Equal("InvalidToken", ex.ErrorCode); + Assert.True(raised); + } + + private static SessionTokenProvider CreateProvider(RecordingHandler handler) + { + var provider = new SessionTokenProvider(new HttpClient(handler)); + Restore(provider); + return provider; + } + + private static void Restore(SessionTokenProvider provider) + { + provider.RestoreSession( + XrpcRequestPipelineTests.CreateJwt(Did, DateTimeOffset.UtcNow.AddHours(2), "access"), + XrpcRequestPipelineTests.CreateJwt(Did, DateTimeOffset.UtcNow.AddDays(30), "refresh"), + Did, + "user.example", + Pds); + } + + private sealed class AsyncHandler : HttpMessageHandler + { + private readonly Func> _respond; + private int _calls; + + public AsyncHandler(Func> respond) + { + _respond = respond; + } + + public int Calls => _calls; + + protected override Task SendAsync(HttpRequestMessage request, System.Threading.CancellationToken cancellationToken) + { + System.Threading.Interlocked.Increment(ref _calls); + return _respond(); + } + } +} diff --git a/tests/CarpaNet.UnitTests/Blob/BlobPipelineTests.cs b/tests/CarpaNet.UnitTests/Blob/BlobPipelineTests.cs new file mode 100644 index 0000000..1c89c05 --- /dev/null +++ b/tests/CarpaNet.UnitTests/Blob/BlobPipelineTests.cs @@ -0,0 +1,102 @@ +using System; +using System.Collections.Generic; +using System.IO; +using System.Linq; +using System.Net; +using System.Net.Http; +using System.Text; +using System.Threading.Tasks; +using CarpaNet.Blob; +using CarpaNet.Identity; +using CarpaNet.UnitTests.Http; +using Xunit; + +namespace CarpaNet.UnitTests.Blob; + +/// +/// Tests for blob upload and download through the client's own request pipeline. +/// +public class BlobPipelineTests +{ + private const string Cid = "bafkreibme22gw2h7y2h7tg2fhqotaqjucnbc24deqo72b6mkl2egezxhvy"; + + [Fact] + public async Task UploadBlobAsync_StreamsThroughClientAuth_AndReturnsBlobRef() + { + var handler = new RecordingHandler(_ => XrpcRequestPipelineTests.Json( + $"{{\"blob\":{{\"$type\":\"blob\",\"ref\":{{\"$link\":\"{Cid}\"}},\"mimeType\":\"image/png\",\"size\":4}}}}")); + using var client = XrpcRequestPipelineTests.CreateSessionClient(handler, out var accessJwt); + var data = new byte[] { 1, 2, 3, 4 }; + + var blob = await client.UploadBlobAsync(new MemoryStream(data), "image/png"); + + var request = Assert.Single(handler.Requests); + Assert.Equal("/xrpc/com.atproto.repo.uploadBlob", request.Uri.AbsolutePath); + Assert.Equal("Bearer " + accessJwt, request.Header("Authorization")); + Assert.Equal("image/png", request.ContentType); + Assert.Equal(data, request.Body); + Assert.Equal(Cid, blob.Ref!.Link); + + var atBlob = blob.ToATBlob(); + Assert.Equal(Cid, atBlob.Ref.ToString()); + Assert.Equal("image/png", atBlob.MimeType); + Assert.Equal(4, atBlob.Size); + } + + [Fact] + public async Task UploadBlobAsync_DoesNotDisposeCallerStream() + { + var handler = new RecordingHandler(_ => XrpcRequestPipelineTests.Json( + $"{{\"blob\":{{\"$type\":\"blob\",\"ref\":{{\"$link\":\"{Cid}\"}},\"mimeType\":\"image/png\",\"size\":1}}}}")); + using var client = XrpcRequestPipelineTests.CreateSessionClient(handler, out _); + using var stream = new MemoryStream(new byte[] { 9 }); + + await client.UploadBlobAsync(stream, "image/png"); + + Assert.True(stream.CanRead); + } + + [Fact] + public void ToATBlob_WithoutCid_Throws() + { + Assert.Throws(() => new BlobRef().ToATBlob()); + } + + [Fact] + public async Task DownloadBlobAsync_OtherAccount_UsesOwnerPdsWithoutCredentials() + { + var bytes = Encoding.UTF8.GetBytes("image"); + var handler = new RecordingHandler(_ => new HttpResponseMessage(HttpStatusCode.OK) { Content = new ByteArrayContent(bytes) }); + var cache = new MemoryIdentityCache(); + await cache.SetDidDocumentAsync("did:plc:other", new DidDocument + { + Id = "did:plc:other", + Service = new List + { + new DidService { Id = "#atproto_pds", Type = "AtprotoPersonalDataServer", ServiceEndpoint = "https://pds.other.example" }, + }, + }); + using var client = XrpcRequestPipelineTests.CreateSessionClient(handler, out _, new IdentityResolver(new HttpClient(handler), cache: cache)); + + var result = await client.DownloadBlobAsync(new ATDid("did:plc:other"), new ATCid(Cid)); + + var request = Assert.Single(handler.Requests); + Assert.Equal("pds.other.example", request.Uri.Host); + Assert.Null(request.Header("Authorization")); + Assert.Contains("did=did%3Aplc%3Aother", request.Uri.Query); + Assert.Equal(bytes, result); + } + + [Fact] + public async Task DownloadBlobAsync_OwnAccount_UsesOwnPdsWithCredentials() + { + var handler = new RecordingHandler(_ => new HttpResponseMessage(HttpStatusCode.OK) { Content = new ByteArrayContent(new byte[] { 1 }) }); + using var client = XrpcRequestPipelineTests.CreateSessionClient(handler, out var accessJwt); + + await client.DownloadBlobAsync(new ATDid("did:plc:user"), new ATCid(Cid)); + + var request = Assert.Single(handler.Requests); + Assert.Equal("pds.user.example", request.Uri.Host); + Assert.Equal("Bearer " + accessJwt, request.Header("Authorization")); + } +} diff --git a/tests/CarpaNet.UnitTests/CarpaNet.UnitTests.csproj b/tests/CarpaNet.UnitTests/CarpaNet.UnitTests.csproj index 1608792..be0b62f 100644 --- a/tests/CarpaNet.UnitTests/CarpaNet.UnitTests.csproj +++ b/tests/CarpaNet.UnitTests/CarpaNet.UnitTests.csproj @@ -19,6 +19,7 @@ + diff --git a/tests/CarpaNet.UnitTests/Generation/BinaryXrpcGenerationTests.cs b/tests/CarpaNet.UnitTests/Generation/BinaryXrpcGenerationTests.cs new file mode 100644 index 0000000..dd39612 --- /dev/null +++ b/tests/CarpaNet.UnitTests/Generation/BinaryXrpcGenerationTests.cs @@ -0,0 +1,454 @@ +using System.Text.RegularExpressions; +using Microsoft.CodeAnalysis.CSharp; +using Xunit; + +namespace CarpaNet.UnitTests.Generation; + +/// +/// Tests the client extension methods generated for procedures with non-JSON bodies or query parameters, +/// and for queries with non-JSON output. Lexicons are copies of the real ones (no network access). +/// +public class BinaryXrpcGenerationTests +{ + private const string UploadBlob = """ + { + "lexicon": 1, + "id": "com.atproto.repo.uploadBlob", + "defs": { + "main": { + "type": "procedure", + "description": "Upload a new blob, to be referenced from a repository record.", + "input": { "encoding": "*/*" }, + "output": { + "encoding": "application/json", + "schema": { + "type": "object", + "required": ["blob"], + "properties": { "blob": { "type": "blob" } } + } + } + } + } + } + """; + + private const string UploadPart = """ + { + "lexicon": 1, + "id": "app.bsky.video.uploadPart", + "defs": { + "main": { + "type": "procedure", + "description": "Upload one part.", + "parameters": { + "type": "params", + "required": ["jobId", "partNumber"], + "properties": { + "jobId": { "type": "string", "minLength": 1, "maxLength": 256 }, + "partNumber": { "type": "integer", "minimum": 1 } + } + }, + "input": { "encoding": "application/octet-stream" }, + "output": { + "encoding": "application/json", + "schema": { + "type": "object", + "required": ["partNumber", "sizeBytes"], + "properties": { + "partNumber": { "type": "integer", "minimum": 1 }, + "sizeBytes": { "type": "integer" } + } + } + }, + "errors": [ { "name": "UploadNotFound" }, { "name": "PartSizeMismatch" } ] + } + } + } + """; + + private const string UploadVideo = """ + { + "lexicon": 1, + "id": "app.bsky.video.uploadVideo", + "defs": { + "main": { + "type": "procedure", + "description": "Upload a video to be processed then stored on the PDS.", + "input": { "encoding": "video/mp4" }, + "output": { + "encoding": "application/json", + "schema": { + "type": "object", + "required": ["jobStatus"], + "properties": { + "jobStatus": { "type": "ref", "ref": "app.bsky.video.defs#jobStatus" } + } + } + } + } + } + } + """; + + private const string VideoDefs = """ + { + "lexicon": 1, + "id": "app.bsky.video.defs", + "defs": { + "jobStatus": { + "type": "object", + "required": ["jobId", "did", "state"], + "properties": { + "jobId": { "type": "string" }, + "did": { "type": "string", "format": "did" }, + "state": { "type": "string" } + } + } + } + } + """; + + private const string GetBlob = """ + { + "lexicon": 1, + "id": "com.atproto.sync.getBlob", + "defs": { + "main": { + "type": "query", + "description": "Get a blob associated with a given account.", + "parameters": { + "type": "params", + "required": ["did", "cid"], + "properties": { + "did": { "type": "string", "format": "did" }, + "cid": { "type": "string", "format": "cid" } + } + }, + "output": { "encoding": "*/*" }, + "errors": [ { "name": "BlobNotFound" } ] + } + } + } + """; + + private const string GetRepo = """ + { + "lexicon": 1, + "id": "com.atproto.sync.getRepo", + "defs": { + "main": { + "type": "query", + "description": "Download a repository export as CAR file.", + "parameters": { + "type": "params", + "required": ["did"], + "properties": { + "did": { "type": "string", "format": "did" }, + "since": { "type": "string", "format": "tid" } + } + }, + "output": { "encoding": "application/vnd.ipld.car" } + } + } + } + """; + + // JSON-bodied procedure that also takes query parameters, plus the unchanged plain cases + private const string JsonProcedures = """ + { + "lexicon": 1, + "id": "com.example.doThing", + "defs": { + "main": { + "type": "procedure", + "parameters": { + "type": "params", + "required": ["mode"], + "properties": { "mode": { "type": "string" }, "dryRun": { "type": "boolean" } } + }, + "input": { + "encoding": "application/json", + "schema": { "type": "object", "required": ["name"], "properties": { "name": { "type": "string" } } } + }, + "output": { + "encoding": "application/json", + "schema": { "type": "object", "properties": { "ok": { "type": "boolean" } } } + } + } + } + } + """; + + private const string PlainProcedure = """ + { + "lexicon": 1, + "id": "com.example.plain", + "defs": { + "main": { + "type": "procedure", + "input": { + "encoding": "application/json", + "schema": { "type": "object", "properties": { "name": { "type": "string" } } } + }, + "output": { + "encoding": "application/json", + "schema": { "type": "object", "properties": { "ok": { "type": "boolean" } } } + } + } + } + } + """; + + private const string PlainQuery = """ + { + "lexicon": 1, + "id": "com.example.getThing", + "defs": { + "main": { + "type": "query", + "parameters": { "type": "params", "properties": { "id": { "type": "string" } } }, + "output": { + "encoding": "application/json", + "schema": { "type": "object", "properties": { "name": { "type": "string" } } } + } + } + } + } + """; + + private const string ChatUpload = """ + { + "lexicon": 1, + "id": "chat.bsky.example.upload", + "defs": { + "main": { + "type": "procedure", + "input": { "encoding": "image/png" } + } + } + } + """; + + private static readonly Lazy Run = new(() => GeneratorTestHarness.Run(new[] + { + UploadBlob, UploadPart, UploadVideo, VideoDefs, GetBlob, GetRepo, JsonProcedures, PlainProcedure, PlainQuery, ChatUpload, + })); + + private static string Extensions => Normalize(Run.Value.Sources["ATProtoExtensions.g.cs"]); + + [Fact] + public void GeneratedCode_Compiles() + { + GeneratorTestHarness.AssertCompiles(Run.Value); + } + + [Fact] + public void UploadBlob_WildcardInput_TakesStreamWithOctetStreamDefault() + { + Assert.Contains( + "public static async System.Threading.Tasks.Task ComAtprotoRepoUploadBlobAsync( " + + "this CarpaNet.IATProtoClient client, System.IO.Stream body, string contentType = \"application/octet-stream\", " + + "System.Threading.CancellationToken cancellationToken = default)", + Extensions); + Assert.Contains( + "return await global::CarpaNet.ATProtoClientXrpcExtensions.PostBinaryAsync( " + + "client, \"com.atproto.repo.uploadBlob\", null, null, body, contentType, cancellationToken);", + Extensions); + } + + [Fact] + public void UploadVideo_ConcreteEncoding_IsDefaultContentType() + { + Assert.Contains( + "AppBskyVideoUploadVideoAsync( this CarpaNet.IATProtoClient client, System.IO.Stream body, string contentType = \"video/mp4\", " + + "System.Threading.CancellationToken cancellationToken = default)", + Extensions); + } + + [Fact] + public void UploadPart_BinaryInputWithParameters_PassesQueryParameters() + { + Assert.Contains( + "public static async System.Threading.Tasks.Task AppBskyVideoUploadPartAsync( " + + "this CarpaNet.IATProtoClient client, System.IO.Stream body, string contentType = \"application/octet-stream\", " + + "AppBsky.Video.UploadPartParameters? parameters = null, System.Threading.CancellationToken cancellationToken = default)", + Extensions); + Assert.Contains( + "PostBinaryAsync( client, \"app.bsky.video.uploadPart\", null, " + + "parameters?.ToQueryParameters(), body, contentType, cancellationToken);", + Extensions); + + // Procedures now get a Parameters class with ToQueryParameters() + var source = Run.Value.AllSource; + Assert.Contains("public partial class UploadPartParameters", source); + Assert.Contains("(\"partNumber\", PartNumber.ToString())", source); + } + + [Fact] + public void JsonProcedureWithParameters_UsesPostWithParameters() + { + Assert.Contains( + "ComExampleDoThingAsync( this CarpaNet.IATProtoClient client, ComExample.DoThingInput input, " + + "ComExample.DoThingParameters? parameters = null, System.Threading.CancellationToken cancellationToken = default)", + Extensions); + Assert.Contains( + "return await global::CarpaNet.ATProtoClientXrpcExtensions.PostWithParametersAsync( " + + "client, \"com.example.doThing\", null, parameters?.ToQueryParameters(), input, cancellationToken);", + Extensions); + } + + [Fact] + public void JsonProcedureWithoutParameters_IsUnchanged() + { + Assert.Contains( + "ComExamplePlainAsync( this CarpaNet.IATProtoClient client, ComExample.PlainInput input, " + + "System.Threading.CancellationToken cancellationToken = default) { " + + "return await client.PostAsync( \"com.example.plain\", input, cancellationToken); }", + Extensions); + } + + [Fact] + public void BinaryProcedure_ForProxiedService_PassesProxyDid() + { + Assert.Contains( + "ChatBskyExampleUploadAsync( this CarpaNet.IATProtoClient client, System.IO.Stream body, string contentType = \"image/png\", " + + "System.Threading.CancellationToken cancellationToken = default)", + Extensions); + Assert.Contains( + "PostBinaryAsync( client, \"chat.bsky.example.upload\", CarpaNet.BlueskyServices.ChatServiceDid, null, body, contentType, cancellationToken);", + Extensions); + } + + [Theory] + [InlineData("ComAtprotoSyncGetBlobAsync", "ComAtproto.Sync.GetBlobParameters", "com.atproto.sync.getBlob")] + [InlineData("ComAtprotoSyncGetRepoAsync", "ComAtproto.Sync.GetRepoParameters", "com.atproto.sync.getRepo")] + public void BinaryOutputQuery_ReturnsBytes(string methodName, string parametersType, string nsid) + { + Assert.Contains( + $"public static async System.Threading.Tasks.Task {methodName}( this CarpaNet.IATProtoClient client, " + + $"{parametersType}? parameters = null, System.Threading.CancellationToken cancellationToken = default)", + Extensions); + Assert.Contains( + $"return await global::CarpaNet.ATProtoClientXrpcExtensions.GetBytesAsync( client, \"{nsid}\", null, parameters?.ToQueryParameters(), cancellationToken);", + Extensions); + } + + [Fact] + public void JsonQuery_IsUnchanged() + { + Assert.Contains( + "return await client.GetAsync( \"com.example.getThing\", parameters?.ToQueryParameters(), cancellationToken);", + Extensions); + } + + [Fact] + public void GeneratedMethods_AreCallableFromConsumerCode() + { + const string usage = """ + using System.IO; + using System.Threading.Tasks; + using CarpaNet; + + internal static class Usage + { + public static async Task CallAll(IATProtoClient client, Stream stream) + { + ComAtproto.Repo.UploadBlobOutput blob = await client.ComAtprotoRepoUploadBlobAsync(stream); + blob = await client.ComAtprotoRepoUploadBlobAsync(stream, "image/jpeg"); + AppBsky.Video.UploadVideoOutput video = await client.AppBskyVideoUploadVideoAsync(stream); + AppBsky.Video.UploadPartOutput part = await client.AppBskyVideoUploadPartAsync( + stream, + parameters: new AppBsky.Video.UploadPartParameters { JobId = "job", PartNumber = 1 }); + byte[] bytes = await client.ComAtprotoSyncGetBlobAsync(new ComAtproto.Sync.GetBlobParameters { Did = default, Cid = "bafy" }); + byte[] car = await client.ComAtprotoSyncGetRepoAsync(); + ComExample.DoThingOutput done = await client.ComExampleDoThingAsync( + new ComExample.DoThingInput { Name = "n" }, + new ComExample.DoThingParameters { Mode = "m", DryRun = true }); + } + } + """; + + var compilation = Run.Value.Compilation.AddSyntaxTrees( + CSharpSyntaxTree.ParseText(usage, new CSharpParseOptions(LanguageVersion.Latest))); + var errors = compilation.GetDiagnostics().Where(d => d.Severity == Microsoft.CodeAnalysis.DiagnosticSeverity.Error).ToList(); + + Assert.True(errors.Count == 0, string.Join("\n", errors)); + } + + [Fact] + public void XrpcEndpoints_StillGenerateForBinaryEndpoints() + { + var run = GeneratorTestHarness.Run( + new[] { UploadBlob, UploadPart, GetBlob }, + new Dictionary { ["CarpaNet_EmitXrpcEndpoints"] = "true" }); + + Assert.DoesNotContain(run.GeneratorDiagnostics, d => d.Id == "ATPG001"); + var controllers = run.Sources["XrpcControllers.g.cs"]; + Assert.Contains("HttpPost(\"/xrpc/com.atproto.repo.uploadBlob\")", controllers); + Assert.Contains("HttpPost(\"/xrpc/app.bsky.video.uploadPart\")", controllers); + Assert.Contains("HttpGet(\"/xrpc/com.atproto.sync.getBlob\")", controllers); + } + + [Fact] + public async Task GeneratedMethods_RunThroughTheClientPipeline() + { + const string driver = """ + using System.IO; + using System.Text.Json; + using System.Threading.Tasks; + using CarpaNet; + + public static class Driver + { + public static JsonSerializerOptions Options => CarpaNet.Json.ATProtoJsonContext.DefaultOptions; + + public static async Task Run(IATProtoClient client) + { + var blob = await client.ComAtprotoRepoUploadBlobAsync(new MemoryStream(new byte[] { 1, 2, 3 }), "image/png"); + var part = await client.AppBskyVideoUploadPartAsync( + new MemoryStream(new byte[] { 4, 5 }), + parameters: new AppBsky.Video.UploadPartParameters { JobId = "job1", PartNumber = 2 }); + var bytes = await client.ComAtprotoSyncGetBlobAsync( + new ComAtproto.Sync.GetBlobParameters { Did = new ATDid("did:plc:user"), Cid = "bafy" }); + return blob.Blob.Ref + "|" + blob.Blob.MimeType + "|" + part.PartNumber + "|" + bytes.Length; + } + } + """; + + var run = Run.Value; + var withDriver = new GeneratorTestHarness.GeneratorRun + { + Compilation = run.Compilation.AddSyntaxTrees(CSharpSyntaxTree.ParseText(driver, new CSharpParseOptions(LanguageVersion.Latest))), + GeneratorDiagnostics = run.GeneratorDiagnostics, + Sources = run.Sources, + }; + var driverType = GeneratorTestHarness.CompileAndLoad(withDriver).GetType("Driver")!; + + var handler = new CarpaNet.UnitTests.Http.RecordingHandler(r => r.Uri.AbsolutePath switch + { + "/xrpc/com.atproto.repo.uploadBlob" => CarpaNet.UnitTests.Http.XrpcRequestPipelineTests.Json( + "{\"blob\":{\"$type\":\"blob\",\"ref\":{\"$link\":\"bafkreibme22gw2h7y2h7tg2fhqotaqjucnbc24deqo72b6mkl2egezxhvy\"},\"mimeType\":\"image/png\",\"size\":3}}"), + "/xrpc/app.bsky.video.uploadPart" => CarpaNet.UnitTests.Http.XrpcRequestPipelineTests.Json("{\"partNumber\":2,\"sizeBytes\":2}"), + _ => new System.Net.Http.HttpResponseMessage(System.Net.HttpStatusCode.OK) { Content = new System.Net.Http.ByteArrayContent(new byte[] { 9, 9, 9, 9 }) }, + }); + var options = CarpaNet.UnitTests.Http.XrpcRequestPipelineTests.CreateOptions(new System.Net.Http.HttpClient(handler)); + options.JsonOptions = (System.Text.Json.JsonSerializerOptions)driverType.GetProperty("Options")!.GetValue(null)!; + using var client = CarpaNet.ATProtoClient.CreateWithRestoredSession( + CarpaNet.UnitTests.Http.XrpcRequestPipelineTests.CreateJwt("did:plc:user", DateTimeOffset.UtcNow.AddHours(1), "a"), + CarpaNet.UnitTests.Http.XrpcRequestPipelineTests.CreateJwt("did:plc:user", DateTimeOffset.UtcNow.AddDays(1), "r"), + "did:plc:user", "user.example", new Uri("https://pds.user.example"), options); + + var result = await (Task)driverType.GetMethod("Run")!.Invoke(null, new object[] { client })!; + + Assert.Equal("bafkreibme22gw2h7y2h7tg2fhqotaqjucnbc24deqo72b6mkl2egezxhvy|image/png|2|4", result); + Assert.Equal(3, handler.Requests.Count); + Assert.Equal("image/png", handler.Requests[0].ContentType); + Assert.Equal(new byte[] { 1, 2, 3 }, handler.Requests[0].Body); + Assert.Equal("application/octet-stream", handler.Requests[1].ContentType); + Assert.Equal("?jobId=job1&partNumber=2", handler.Requests[1].Uri.Query); + Assert.Equal("*/*", handler.Requests[2].Header("Accept")); + Assert.All(handler.Requests, r => Assert.StartsWith("Bearer ", r.Header("Authorization"))); + } + + private static string Normalize(string source) => Regex.Replace(source, @"\s+", " "); +} diff --git a/tests/CarpaNet.UnitTests/Generation/CrossNamespaceAndArrayParameterTests.cs b/tests/CarpaNet.UnitTests/Generation/CrossNamespaceAndArrayParameterTests.cs new file mode 100644 index 0000000..9a4af30 --- /dev/null +++ b/tests/CarpaNet.UnitTests/Generation/CrossNamespaceAndArrayParameterTests.cs @@ -0,0 +1,99 @@ +using System.Collections.Generic; +using System.Linq; +using Xunit; + +namespace CarpaNet.UnitTests.Generation; + +/// +/// Regression tests for two generator bugs found by compiling every atproto lexicon: +/// a ref to an array-of-union def in another namespace emitted an unqualified interface name +/// (tools.ozone.moderation.getAccountPreferences), and integer-array query parameters emitted +/// item?.ToString() on a long (tools.ozone.*.getAssignments). +/// +public class CrossNamespaceAndArrayParameterTests +{ + private const string Defs = """ + { + "lexicon": 1, + "id": "com.example.alpha.defs", + "defs": { + "preferences": { + "type": "array", + "items": { "type": "union", "refs": ["#first", "#second"] } + }, + "first": { "type": "object", "properties": { "a": { "type": "string" } } }, + "second": { "type": "object", "properties": { "b": { "type": "integer" } } } + } + } + """; + + private const string GetPreferences = """ + { + "lexicon": 1, + "id": "org.example.beta.getPreferences", + "defs": { + "main": { + "type": "query", + "output": { + "encoding": "application/json", + "schema": { + "type": "object", + "required": ["preferences"], + "properties": { "preferences": { "type": "ref", "ref": "com.example.alpha.defs#preferences" } } + } + } + } + } + } + """; + + private const string GetAssignments = """ + { + "lexicon": 1, + "id": "org.example.beta.getAssignments", + "defs": { + "main": { + "type": "query", + "parameters": { + "type": "params", + "properties": { + "ids": { "type": "array", "items": { "type": "integer" } }, + "flags": { "type": "array", "items": { "type": "boolean" } } + } + }, + "output": { + "encoding": "application/json", + "schema": { "type": "object", "properties": { "count": { "type": "integer" } } } + } + } + } + } + """; + + [Fact] + public void CrossNamespaceArrayOfUnionRef_Compiles() + { + var run = GeneratorTestHarness.Run(new[] { Defs, GetPreferences }); + + GeneratorTestHarness.AssertCompiles(run); + Assert.Contains("ComExample.Alpha.IDefsPreferences", run.Sources.Values.First(s => s.Contains("GetPreferencesOutput"))); + } + + [Fact] + public void IntegerAndBooleanArrayParameters_FormatInvariantly() + { + var run = GeneratorTestHarness.Run(new[] { GetAssignments }); + var assembly = GeneratorTestHarness.CompileAndLoad(run); + + var parametersType = assembly.GetType("OrgExample.Beta.GetAssignmentsParameters")!; + var parameters = System.Activator.CreateInstance(parametersType)!; + parametersType.GetProperty("Ids")!.SetValue(parameters, new List { 12, -3 }); + parametersType.GetProperty("Flags")!.SetValue(parameters, new List { true, false }); + + var query = ((IEnumerable>)parametersType.GetMethod("ToQueryParameters")!.Invoke(parameters, null)!).ToList(); + + Assert.Equal( + new[] { ("ids", "12"), ("ids", "-3"), ("flags", "true"), ("flags", "false") }, + query.Select(kv => (kv.Key, kv.Value)).ToArray()); + } +} diff --git a/tests/CarpaNet.UnitTests/Generation/GeneratorBuildPropertyTests.cs b/tests/CarpaNet.UnitTests/Generation/GeneratorBuildPropertyTests.cs new file mode 100644 index 0000000..9776a4a --- /dev/null +++ b/tests/CarpaNet.UnitTests/Generation/GeneratorBuildPropertyTests.cs @@ -0,0 +1,152 @@ +using System.Xml.Linq; +using Xunit; + +namespace CarpaNet.UnitTests.Generation; + +/// +/// Verifies that the MSBuild properties exposed by CarpaNet.targets reach the generator, under both +/// the current CarpaNet_* names and the legacy CarpaNet_SourceGen_* names. +/// +public class GeneratorBuildPropertyTests +{ + private const string Lexicon = """ + { + "lexicon": 1, + "id": "com.example.note", + "defs": { + "main": { + "type": "record", + "key": "tid", + "record": { + "type": "object", + "required": ["text"], + "properties": { + "text": { "type": "string", "maxLength": 300 } + } + } + } + } + } + """; + + public static IEnumerable Prefixes => new[] + { + new object[] { "CarpaNet_" }, + new object[] { "CarpaNet_SourceGen_" }, + }; + + [Fact] + public void Defaults_UseDefaultNamesAndEmitValidation() + { + var run = GeneratorTestHarness.Run(new[] { Lexicon }); + + GeneratorTestHarness.AssertCompiles(run); + Assert.Contains("namespace ComExample;", run.AllSource); + Assert.Contains("public sealed class ATProtoJsonContext", run.AllSource); + Assert.Contains("public partial class ATProtoCborContext", run.AllSource); + Assert.Contains("CarpaNet.Validation.ATStringLength(300", run.AllSource); + } + + [Theory] + [MemberData(nameof(Prefixes))] + public void RootNamespace_IsApplied(string prefix) + { + var run = GeneratorTestHarness.Run(new[] { Lexicon }, new Dictionary + { + [prefix + "RootNamespace"] = "My.Lexicons", + }); + + GeneratorTestHarness.AssertCompiles(run); + Assert.Contains("namespace My.Lexicons.ComExample;", run.AllSource); + Assert.Contains("namespace My.Lexicons.Json;", run.AllSource); + Assert.Contains("namespace My.Lexicons.Cbor;", run.AllSource); + Assert.DoesNotContain("namespace ComExample;", run.AllSource); + } + + [Theory] + [MemberData(nameof(Prefixes))] + public void ContextNames_AreApplied(string prefix) + { + var run = GeneratorTestHarness.Run(new[] { Lexicon }, new Dictionary + { + [prefix + "JsonContextName"] = "MyJsonContext", + [prefix + "CborContextName"] = "MyCborContext", + }); + + GeneratorTestHarness.AssertCompiles(run); + Assert.True(run.Sources.ContainsKey("MyJsonContext.g.cs")); + Assert.Contains("public sealed class MyJsonContext", run.AllSource); + Assert.Contains("public partial class MyCborContext", run.AllSource); + Assert.DoesNotContain("class ATProtoJsonContext", run.AllSource); + Assert.DoesNotContain("class ATProtoCborContext", run.AllSource); + } + + [Theory] + [MemberData(nameof(Prefixes))] + public void EmitValidationAttributes_False_OmitsValidationAttributes(string prefix) + { + var run = GeneratorTestHarness.Run(new[] { Lexicon }, new Dictionary + { + [prefix + "EmitValidationAttributes"] = "false", + }); + + GeneratorTestHarness.AssertCompiles(run); + Assert.DoesNotContain("CarpaNet.Validation.AT", run.AllSource); + } + + [Fact] + public void CurrentName_TakesPrecedenceOverLegacyName() + { + var run = GeneratorTestHarness.Run(new[] { Lexicon }, new Dictionary + { + ["CarpaNet_JsonContextName"] = "NewJsonContext", + ["CarpaNet_SourceGen_JsonContextName"] = "OldJsonContext", + }); + + Assert.Contains("public sealed class NewJsonContext", run.AllSource); + Assert.DoesNotContain("OldJsonContext", run.AllSource); + } + + [Fact] + public void EmptyCurrentName_FallsBackToLegacyName() + { + // MSBuild writes "build_property.X = " for unset compiler-visible properties + var run = GeneratorTestHarness.Run(new[] { Lexicon }, new Dictionary + { + ["CarpaNet_RootNamespace"] = "", + ["CarpaNet_SourceGen_RootNamespace"] = "Legacy.Root", + }); + + Assert.Contains("namespace Legacy.Root.ComExample;", run.AllSource); + } + + [Theory] + [InlineData("src/CarpaNet/build/CarpaNet.targets")] + [InlineData("src/CarpaNet.SourceGen/build/CarpaNet.SourceGen.targets")] + public void Targets_ExposeCurrentAndLegacyPropertyNames(string relativePath) + { + var path = Path.Combine(FindRepoRoot(), relativePath); + var visible = XDocument.Load(path) + .Descendants() + .Where(e => e.Name.LocalName == "CompilerVisibleProperty") + .Select(e => (string?)e.Attribute("Include")) + .ToHashSet(); + + foreach (var name in new[] { "RootNamespace", "EmitValidationAttributes", "CborContextName", "JsonContextName" }) + { + Assert.Contains("CarpaNet_" + name, visible); + Assert.Contains("CarpaNet_SourceGen_" + name, visible); + } + } + + private static string FindRepoRoot() + { + var dir = new DirectoryInfo(AppContext.BaseDirectory); + while (dir != null && !File.Exists(Path.Combine(dir.FullName, "CarpaNet.slnx"))) + { + dir = dir.Parent; + } + + return dir?.FullName ?? throw new InvalidOperationException("Could not locate the repository root (CarpaNet.slnx)."); + } +} diff --git a/tests/CarpaNet.UnitTests/Generation/GeneratorTestHarness.cs b/tests/CarpaNet.UnitTests/Generation/GeneratorTestHarness.cs new file mode 100644 index 0000000..93edb1e --- /dev/null +++ b/tests/CarpaNet.UnitTests/Generation/GeneratorTestHarness.cs @@ -0,0 +1,198 @@ +using System.Collections.Immutable; +using System.Reflection; +using System.Runtime.Loader; +using System.Text; +using Microsoft.CodeAnalysis; +using Microsoft.CodeAnalysis.CSharp; +using Microsoft.CodeAnalysis.Diagnostics; +using Microsoft.CodeAnalysis.Text; + +namespace CarpaNet.UnitTests.Generation; + +/// +/// Runs end to end over in-memory lexicon JSON (no MSBuild, no network), +/// optionally compiling and loading the generated code so it can be exercised at runtime. +/// +internal static class GeneratorTestHarness +{ + /// + /// The result of a generator run. + /// + internal sealed class GeneratorRun + { + public required CSharpCompilation Compilation { get; init; } + + public required ImmutableArray GeneratorDiagnostics { get; init; } + + /// Generated sources keyed by hint name (e.g. "ATProtoExtensions.g.cs"). + public required IReadOnlyDictionary Sources { get; init; } + + /// All generated sources concatenated. + public string AllSource => string.Join("\n", Sources.Values); + + /// Errors from compiling the generated code together with any extra sources. + public IReadOnlyList CompilationErrors => + Compilation.GetDiagnostics().Where(d => d.Severity == DiagnosticSeverity.Error).ToList(); + } + + /// + /// Runs the generator over the given lexicon documents. + /// + /// Lexicon documents as JSON text. + /// MSBuild properties exposed to the generator, without the "build_property." prefix. + public static GeneratorRun Run(IEnumerable lexiconJson, IDictionary? buildProperties = null) + { + var additionalTexts = lexiconJson + .Select((json, i) => (AdditionalText)new InMemoryAdditionalText($"/lexicons/lexicon{i}.json", json)) + .ToImmutableArray(); + + var globalOptions = new Dictionary(StringComparer.Ordinal); + if (buildProperties != null) + { + foreach (var kvp in buildProperties) + { + globalOptions["build_property." + kvp.Key] = kvp.Value; + } + } + + var optionsProvider = new TestAnalyzerConfigOptionsProvider(globalOptions); + + var compilation = CSharpCompilation.Create( + "GeneratedLexicons_" + Guid.NewGuid().ToString("N"), + Array.Empty(), + References.Value, + new CSharpCompilationOptions(OutputKind.DynamicallyLinkedLibrary, nullableContextOptions: NullableContextOptions.Enable)); + + var driver = CSharpGeneratorDriver.Create( + new[] { new LexiconGenerator().AsSourceGenerator() }, + additionalTexts, + new CSharpParseOptions(LanguageVersion.Latest), + optionsProvider); + + driver = (CSharpGeneratorDriver)driver.RunGeneratorsAndUpdateCompilation(compilation, out var outputCompilation, out var diagnostics); + + var runResult = driver.GetRunResult(); + var sources = runResult.Results + .SelectMany(r => r.GeneratedSources) + .ToDictionary(s => s.HintName, s => s.SourceText.ToString(), StringComparer.Ordinal); + + return new GeneratorRun + { + Compilation = (CSharpCompilation)outputCompilation, + GeneratorDiagnostics = diagnostics, + Sources = sources, + }; + } + + /// + /// Asserts that the generated code compiles without errors and returns the formatted errors otherwise. + /// + public static void AssertCompiles(GeneratorRun run) + { + var errors = run.CompilationErrors; + if (errors.Count > 0) + { + var message = new StringBuilder("Generated code failed to compile:\n"); + foreach (var error in errors.Take(25)) + { + message.AppendLine(error.ToString()); + } + + throw new Xunit.Sdk.XunitException(message.ToString()); + } + } + + /// + /// Compiles the generated code and loads it into a collectible load context. + /// + public static Assembly CompileAndLoad(GeneratorRun run) + { + AssertCompiles(run); + + using var stream = new MemoryStream(); + var emitResult = run.Compilation.Emit(stream); + if (!emitResult.Success) + { + throw new Xunit.Sdk.XunitException("Emit failed:\n" + string.Join("\n", emitResult.Diagnostics.Where(d => d.Severity == DiagnosticSeverity.Error))); + } + + stream.Position = 0; + var context = new AssemblyLoadContext("GeneratedLexicons", isCollectible: true); + return context.LoadFromStream(stream); + } + + private static readonly Lazy> References = new(() => + { + var tpa = (string?)AppContext.GetData("TRUSTED_PLATFORM_ASSEMBLIES") ?? string.Empty; + var excluded = new HashSet(StringComparer.OrdinalIgnoreCase) + { + // Avoid generator/test types (some share the CarpaNet namespace) leaking into the consumer compilation + "CarpaNet.SourceGen", + "CarpaNet.UnitTests", + }; + + var paths = tpa.Split(Path.PathSeparator, StringSplitOptions.RemoveEmptyEntries) + .Where(p => !excluded.Contains(Path.GetFileNameWithoutExtension(p))) + .ToList(); + + // Make sure the runtime library and CBOR dependency are present even if not yet in the TPA list + foreach (var assembly in new[] { typeof(CarpaNet.IATProtoClient).Assembly, typeof(System.Formats.Cbor.CborReader).Assembly }) + { + if (!paths.Contains(assembly.Location, StringComparer.OrdinalIgnoreCase)) + { + paths.Add(assembly.Location); + } + } + + return paths.Select(p => (MetadataReference)MetadataReference.CreateFromFile(p)).ToImmutableArray(); + }); + + private sealed class InMemoryAdditionalText : AdditionalText + { + private readonly SourceText _text; + + public InMemoryAdditionalText(string path, string text) + { + Path = path; + _text = SourceText.From(text, Encoding.UTF8); + } + + public override string Path { get; } + + public override SourceText GetText(CancellationToken cancellationToken = default) => _text; + } + + private sealed class TestAnalyzerConfigOptionsProvider : AnalyzerConfigOptionsProvider + { + private static readonly TestAnalyzerConfigOptions LexiconFileOptions = new(new Dictionary + { + ["build_metadata.AdditionalFiles.IsATProtoLexicon"] = "true", + }); + + public TestAnalyzerConfigOptionsProvider(Dictionary globalOptions) + { + GlobalOptions = new TestAnalyzerConfigOptions(globalOptions); + } + + public override AnalyzerConfigOptions GlobalOptions { get; } + + public override AnalyzerConfigOptions GetOptions(SyntaxTree tree) => TestAnalyzerConfigOptions.Empty; + + public override AnalyzerConfigOptions GetOptions(AdditionalText textFile) => LexiconFileOptions; + } + + private sealed class TestAnalyzerConfigOptions : AnalyzerConfigOptions + { + public static readonly TestAnalyzerConfigOptions Empty = new(new Dictionary()); + + private readonly Dictionary _values; + + public TestAnalyzerConfigOptions(Dictionary values) + { + _values = values; + } + + public override bool TryGetValue(string key, [System.Diagnostics.CodeAnalysis.NotNullWhen(true)] out string? value) + => _values.TryGetValue(key, out value); + } +} diff --git a/tests/CarpaNet.UnitTests/Generation/OpenUnionGenerationTests.cs b/tests/CarpaNet.UnitTests/Generation/OpenUnionGenerationTests.cs new file mode 100644 index 0000000..299c63c --- /dev/null +++ b/tests/CarpaNet.UnitTests/Generation/OpenUnionGenerationTests.cs @@ -0,0 +1,279 @@ +using System.Collections; +using System.Reflection; +using System.Text.Json; +using System.Text.Json.Nodes; +using CarpaNet.Cbor; +using Xunit; + +namespace CarpaNet.UnitTests.Generation; + +/// +/// End-to-end tests for open unions: unknown members must survive JSON and CBOR round trips. +/// +public class OpenUnionGenerationTests +{ + // Shapes mirror real lexicons: a single open union (post embed), an inline array of open unions + // (facet features), a named array-of-union def (app.bsky.actor.defs#preferences) and a closed union. + private const string Lexicon = """ + { + "lexicon": 1, + "id": "com.example.unions", + "defs": { + "main": { + "type": "record", + "key": "tid", + "record": { + "type": "object", + "properties": { + "embed": { "type": "union", "refs": ["#a", "#b"] }, + "items": { "type": "array", "items": { "type": "union", "refs": ["#a", "#b"] } }, + "prefs": { "type": "ref", "ref": "#preferences" }, + "closedEmbed": { "type": "union", "refs": ["#a", "#b"], "closed": true } + } + } + }, + "a": { + "type": "object", + "required": ["text"], + "properties": { "text": { "type": "string" } } + }, + "b": { + "type": "object", + "properties": { "count": { "type": "integer" } } + }, + "preferences": { + "type": "array", + "items": { "type": "union", "refs": ["#a", "#b"] } + } + } + } + """; + + private const string RecordJson = """ + { + "$type": "com.example.unions", + "embed": { "$type": "com.example.future#thing", "x": 1, "nested": { "y": [1, 2, "three"], "flag": true } }, + "items": [ + { "$type": "com.example.unions#a", "text": "hi" }, + { "$type": "com.example.future#other", "z": "q" }, + { "$type": "com.example.unions#b", "count": 3 }, + { "noType": true } + ], + "prefs": [ + { "$type": "com.example.future#newPref", "enabled": false, "list": ["a", "b"] }, + { "$type": "com.example.unions#a", "text": "known" } + ] + } + """; + + private static readonly Lazy Generated = new(() => + GeneratorTestHarness.CompileAndLoad(GeneratorTestHarness.Run(new[] { Lexicon }))); + + [Fact] + public void OpenUnion_GeneratesUnknownClassImplementingInterface() + { + var assembly = Generated.Value; + + foreach (var (iface, unknown) in new[] + { + ("ComExample.IUnionsEmbed", "ComExample.Unknown_UnionsEmbed"), + ("ComExample.IUnionsItems", "ComExample.Unknown_UnionsItems"), + ("ComExample.IUnionsPreferences", "ComExample.Unknown_UnionsPreferences"), + }) + { + var ifaceType = assembly.GetType(iface, throwOnError: true)!; + var unknownType = assembly.GetType(unknown, throwOnError: true)!; + Assert.True(unknownType.IsSealed); + Assert.True(ifaceType.IsAssignableFrom(unknownType)); + Assert.NotNull(unknownType.GetConstructor(new[] { typeof(string), typeof(JsonElement), typeof(byte[]) })); + } + + // Closed unions keep throwing on unknown members and get no Unknown_ class + Assert.Null(assembly.GetType("ComExample.Unknown_UnionsClosedEmbed")); + } + + [Fact] + public void Json_UnknownMembersArePreservedInOrder() + { + var record = DeserializeRecord(RecordJson); + var recordType = record.GetType(); + + var embed = recordType.GetProperty("Embed")!.GetValue(record)!; + Assert.Equal("Unknown_UnionsEmbed", embed.GetType().Name); + Assert.Equal("com.example.future#thing", GetUnknownType(embed)); + AssertJsonEquivalent( + """{ "$type": "com.example.future#thing", "x": 1, "nested": { "y": [1, 2, "three"], "flag": true } }""", + GetUnknownRaw(embed).GetRawText()); + + var items = ((IEnumerable)recordType.GetProperty("Items")!.GetValue(record)!).Cast().ToList(); + Assert.Equal(4, items.Count); + Assert.Equal("UnionsA", items[0].GetType().Name); + Assert.Equal("Unknown_UnionsItems", items[1].GetType().Name); + Assert.Equal("com.example.future#other", GetUnknownType(items[1])); + Assert.Equal("UnionsB", items[2].GetType().Name); + Assert.Equal("Unknown_UnionsItems", items[3].GetType().Name); + Assert.Equal(string.Empty, GetUnknownType(items[3])); + + var prefs = ((IEnumerable)recordType.GetProperty("Prefs")!.GetValue(record)!).Cast().ToList(); + Assert.Equal(2, prefs.Count); + Assert.Equal("Unknown_UnionsPreferences", prefs[0].GetType().Name); + Assert.Equal("com.example.future#newPref", GetUnknownType(prefs[0])); + Assert.Equal("UnionsA", prefs[1].GetType().Name); + } + + [Fact] + public void Json_RoundTripIsLossless() + { + var record = DeserializeRecord(RecordJson); + + var json = JsonSerializer.Serialize(record, record.GetType(), GetJsonOptions()); + + AssertJsonEquivalent(RecordJson, json); + } + + [Fact] + public void Json_ClosedUnion_UnknownMemberStillThrows() + { + const string json = """ + { "$type": "com.example.unions", "closedEmbed": { "$type": "com.example.future#thing", "x": 1 } } + """; + + Assert.ThrowsAny(() => DeserializeRecord(json)); + } + + [Fact] + public void Json_CallerConstructedUnknownMember_IsWrittenVerbatim() + { + var assembly = Generated.Value; + var recordType = assembly.GetType("ComExample.Unions", throwOnError: true)!; + var unknownType = assembly.GetType("ComExample.Unknown_UnionsEmbed", throwOnError: true)!; + + using var doc = JsonDocument.Parse("""{ "$type": "com.example.custom#x", "value": 42 }"""); + var unknown = Activator.CreateInstance(unknownType, "com.example.custom#x", doc.RootElement, null)!; + + // The constructor clones, so the instance outlives the source document + doc.Dispose(); + + var record = DeserializeRecord("""{ "$type": "com.example.unions" }"""); + recordType.GetProperty("Embed")!.SetValue(record, unknown); + + var json = JsonSerializer.Serialize(record, recordType, GetJsonOptions()); + AssertJsonEquivalent( + """{ "$type": "com.example.unions", "embed": { "$type": "com.example.custom#x", "value": 42 } }""", + json); + + // ToJson/FromJson helpers on the unknown class + var toJson = (JsonElement)unknownType.GetMethod("ToJson")!.Invoke(unknown, null)!; + Assert.Equal(42, toJson.GetProperty("value").GetInt32()); + var fromJson = unknownType.GetMethod("FromJson")!.Invoke(null, new object[] { toJson })!; + Assert.Equal("com.example.custom#x", GetUnknownType(fromJson)); + } + + [Fact] + public void Cbor_UnknownMembersRoundTripWithRawBytes() + { + var record = DeserializeRecord(RecordJson); + var recordType = record.GetType(); + var cborContext = GetCborContext(); + + // JSON-sourced unknowns are converted to CBOR from their JSON value + var firstBytes = CborSerialize(cborContext, recordType, record); + + var decoded = CborDeserialize(cborContext, recordType, firstBytes); + var embed = recordType.GetProperty("Embed")!.GetValue(decoded)!; + Assert.Equal("Unknown_UnionsEmbed", embed.GetType().Name); + Assert.Equal("com.example.future#thing", GetUnknownType(embed)); + Assert.NotNull(embed.GetType().GetProperty("RawCbor")!.GetValue(embed)); + AssertJsonEquivalent( + """{ "$type": "com.example.future#thing", "x": 1, "nested": { "y": [1, 2, "three"], "flag": true } }""", + GetUnknownRaw(embed).GetRawText()); + + var items = ((IEnumerable)recordType.GetProperty("Items")!.GetValue(decoded)!).Cast().ToList(); + Assert.Equal( + new[] { "UnionsA", "Unknown_UnionsItems", "UnionsB", "Unknown_UnionsItems" }, + items.Select(i => i.GetType().Name)); + Assert.Equal("com.example.future#other", GetUnknownType(items[1])); + Assert.Equal(string.Empty, GetUnknownType(items[3])); + + var prefs = ((IEnumerable)recordType.GetProperty("Prefs")!.GetValue(decoded)!).Cast().ToList(); + Assert.Equal(new[] { "Unknown_UnionsPreferences", "UnionsA" }, prefs.Select(p => p.GetType().Name)); + + // Converting back to JSON gives the original document + AssertJsonEquivalent(RecordJson, JsonSerializer.Serialize(decoded, recordType, GetJsonOptions())); + + // CBOR-sourced unknowns are written back from their original bytes + var secondBytes = CborSerialize(cborContext, recordType, decoded); + Assert.Equal(firstBytes, secondBytes); + } + + [Fact] + public void Cbor_ClosedUnion_KnownMemberRoundTrips_UnknownMemberThrows() + { + var recordType = Generated.Value.GetType("ComExample.Unions", throwOnError: true)!; + var cborContext = GetCborContext(); + + var record = DeserializeRecord(""" + { "$type": "com.example.unions", "closedEmbed": { "$type": "com.example.unions#b", "count": 7 }, "items": [] } + """); + var decoded = CborDeserialize(cborContext, recordType, CborSerialize(cborContext, recordType, record)); + Assert.Equal("UnionsB", recordType.GetProperty("ClosedEmbed")!.GetValue(decoded)!.GetType().Name); + + var writer = new DagCborWriter(); + writer.WriteStartMap(2); + writer.WriteTextString("$type"); + writer.WriteTextString("com.example.unions"); + writer.WriteTextString("closedEmbed"); + writer.WriteStartMap(1); + writer.WriteTextString("$type"); + writer.WriteTextString("com.example.future#thing"); + writer.WriteEndMap(); + writer.WriteEndMap(); + var bytes = writer.Encode(); + + var ex = Assert.ThrowsAny(() => CborDeserialize(cborContext, recordType, bytes)); + Assert.IsType(ex is TargetInvocationException tie ? tie.InnerException : ex); + } + + private static object DeserializeRecord(string json) + { + var recordType = Generated.Value.GetType("ComExample.Unions", throwOnError: true)!; + return JsonSerializer.Deserialize(json, recordType, GetJsonOptions())!; + } + + private static JsonSerializerOptions GetJsonOptions() + { + var contextType = Generated.Value.GetType("CarpaNet.Json.ATProtoJsonContext", throwOnError: true)!; + return (JsonSerializerOptions)contextType.GetProperty("DefaultOptions")!.GetValue(null)!; + } + + private static CborSerializerContext GetCborContext() + { + var contextType = Generated.Value.GetType("CarpaNet.Cbor.ATProtoCborContext", throwOnError: true)!; + return (CborSerializerContext)contextType.GetProperty("Default")!.GetValue(null)!; + } + + private static byte[] CborSerialize(CborSerializerContext context, Type type, object value) + { + var method = typeof(CborSerializerContext).GetMethod(nameof(CborSerializerContext.Serialize))!.MakeGenericMethod(type); + return (byte[])method.Invoke(context, new[] { value })!; + } + + private static object CborDeserialize(CborSerializerContext context, Type type, byte[] data) + { + var method = typeof(CborSerializerContext).GetMethod(nameof(CborSerializerContext.Deserialize))!.MakeGenericMethod(type); + return method.Invoke(context, new object[] { new ReadOnlyMemory(data) })!; + } + + private static string GetUnknownType(object unknown) + => (string)unknown.GetType().GetProperty("Type")!.GetValue(unknown)!; + + private static JsonElement GetUnknownRaw(object unknown) + => (JsonElement)unknown.GetType().GetProperty("Raw")!.GetValue(unknown)!; + + private static void AssertJsonEquivalent(string expected, string actual) + { + var expectedNode = JsonNode.Parse(expected); + var actualNode = JsonNode.Parse(actual); + Assert.True(JsonNode.DeepEquals(expectedNode, actualNode), $"JSON differs.\nExpected: {expected}\nActual: {actual}"); + } +} diff --git a/tests/CarpaNet.UnitTests/Http/XrpcRequestPipelineTests.cs b/tests/CarpaNet.UnitTests/Http/XrpcRequestPipelineTests.cs new file mode 100644 index 0000000..e69ae4a --- /dev/null +++ b/tests/CarpaNet.UnitTests/Http/XrpcRequestPipelineTests.cs @@ -0,0 +1,507 @@ +using System; +using System.Collections.Generic; +using System.IO; +using System.Linq; +using System.Net; +using System.Net.Http; +using System.Text; +using System.Text.Json; +using System.Threading; +using System.Threading.Tasks; +using CarpaNet.Http; +using CarpaNet.Identity; +using Xunit; + +namespace CarpaNet.UnitTests.Http; + +/// +/// Tests for the shared XRPC send pipeline: credential scoping, repo routing, per-request options, +/// binary bodies and responses, and body replay on token refresh. +/// +public class XrpcRequestPipelineTests +{ + private const string UserDid = "did:plc:user"; + private const string OtherDid = "did:plc:other"; + private static readonly Uri UserPds = new("https://pds.user.example"); + private static readonly Uri OtherPds = new("https://pds.other.example"); + + [Fact] + public async Task Get_OwnPds_AttachesBearerToken() + { + var handler = new RecordingHandler(_ => Json("{\"value\":1}")); + using var client = CreateSessionClient(handler, out var accessJwt); + + await client.GetAsync("com.example.get"); + + var request = Assert.Single(handler.Requests); + Assert.Equal(UserPds.Host, request.Uri.Host); + Assert.Equal("Bearer " + accessJwt, request.Header("Authorization")); + } + + [Fact] + public async Task Get_ForeignRepo_RoutesToOwnerPdsWithoutCredentials() + { + var handler = new RecordingHandler(_ => Json("{\"value\":1}")); + using var client = CreateSessionClient(handler, out _, await CreateResolverAsync(handler)); + + await client.GetAsync("com.atproto.repo.getRecord", Params(("repo", OtherDid), ("collection", "app.bsky.feed.post"), ("rkey", "abc"))); + + var request = Assert.Single(handler.Requests); + Assert.Equal(OtherPds.Host, request.Uri.Host); + Assert.Null(request.Header("Authorization")); + Assert.Null(request.Header("DPoP")); + } + + [Fact] + public async Task Get_OwnRepo_KeepsCredentials() + { + var handler = new RecordingHandler(_ => Json("{\"value\":1}")); + using var client = CreateSessionClient(handler, out var accessJwt, await CreateResolverAsync(handler)); + + await client.GetAsync("com.atproto.repo.getRecord", Params(("repo", UserDid))); + + var request = Assert.Single(handler.Requests); + Assert.Equal(UserPds.Host, request.Uri.Host); + Assert.Equal("Bearer " + accessJwt, request.Header("Authorization")); + } + + [Fact] + public async Task Get_WithProxy_IsNotReroutedByRepo() + { + var handler = new RecordingHandler(_ => Json("{\"value\":1}")); + using var client = CreateSessionClient(handler, out _, await CreateResolverAsync(handler)); + + await client.GetAsync("app.bsky.example.get", "did:web:api.bsky.app#bsky_appview", Params(("repo", OtherDid))); + + var request = Assert.Single(handler.Requests); + Assert.Equal(UserPds.Host, request.Uri.Host); + Assert.Equal("did:web:api.bsky.app#bsky_appview", request.Header("atproto-proxy")); + Assert.NotNull(request.Header("Authorization")); + } + + [Fact] + public async Task ServiceUrl_SendsNoSessionCredentials_ButKeepsCallerAuthorization() + { + var handler = new RecordingHandler(_ => Json("{\"value\":1}")); + using var client = CreateSessionClient(handler, out _); + var video = client + .WithServiceUrl(new Uri("https://video.example")) + .WithHeader("Authorization", "Bearer service-auth-token"); + + await video.GetAsync("app.bsky.video.getUploadLimits"); + + var request = Assert.Single(handler.Requests); + Assert.Equal("video.example", request.Uri.Host); + Assert.Equal("Bearer service-auth-token", request.Header("Authorization")); + } + + [Fact] + public async Task ServiceUrl_WithoutCallerAuthorization_SendsNone() + { + var handler = new RecordingHandler(_ => Json("{\"value\":1}")); + using var client = CreateSessionClient(handler, out _); + + await client.WithServiceUrl(new Uri("https://public.api.example")).GetAsync("app.bsky.feed.getFeed"); + + Assert.Null(Assert.Single(handler.Requests).Header("Authorization")); + } + + [Fact] + public async Task WithoutProxy_OverridesGeneratedProxy() + { + var handler = new RecordingHandler(_ => Json("{\"value\":1}")); + using var client = CreateSessionClient(handler, out _); + + // Generated chat methods call the proxy overload. + await client.WithoutProxy().GetAsync("chat.bsky.convo.getLog", BlueskyServices.ChatServiceDid); + + Assert.Null(Assert.Single(handler.Requests).Header("atproto-proxy")); + } + + [Fact] + public async Task WithProxy_AppliesToEveryCall() + { + var handler = new RecordingHandler(_ => Json("{\"value\":1}")); + using var client = CreateSessionClient(handler, out _); + var appview = client.WithProxy("did:web:api.bsky.app#bsky_appview"); + + await appview.GetAsync("app.bsky.feed.getTimeline"); + await appview.PostAsync("app.bsky.graph.muteActor", new { actor = OtherDid }); + + Assert.All(handler.Requests, r => Assert.Equal("did:web:api.bsky.app#bsky_appview", r.Header("atproto-proxy"))); + } + + [Fact] + public async Task NestedScopes_OuterScopeWins() + { + var handler = new RecordingHandler(_ => Json("{\"value\":1}")); + using var client = CreateSessionClient(handler, out _); + var scoped = client + .WithProxy("did:web:inner.example#svc") + .WithHeader("x-test", "inner") + .WithProxy("did:web:outer.example#svc") + .WithHeader("x-test", "outer"); + + await scoped.GetAsync("com.example.get"); + + var request = Assert.Single(handler.Requests); + Assert.Equal("did:web:outer.example#svc", request.Header("atproto-proxy")); + Assert.Equal("outer", request.Header("x-test")); + Assert.Same(client, scoped.Inner); + } + + [Fact] + public async Task SetLabelerDids_ChangesHeaderOnLaterRequests_AndKeepsRedact() + { + var handler = new RecordingHandler(_ => Json("{\"value\":1}")); + using var client = CreateSessionClient(handler, out _); + + await client.GetAsync("com.example.get"); + client.SetLabelerDids(new[] { AcceptLabelersHeader.Redact("did:plc:mod"), "did:plc:custom" }); + await client.GetAsync("com.example.get"); + + Assert.Null(handler.Requests[0].Header("atproto-accept-labelers")); + Assert.Equal("did:plc:mod;redact,did:plc:custom", handler.Requests[1].Header("atproto-accept-labelers")); + } + + [Fact] + public async Task WithAcceptLabelers_ReplacesClientList() + { + var handler = new RecordingHandler(_ => Json("{\"value\":1}")); + using var client = CreateSessionClient(handler, out _); + client.SetLabelerDids(new[] { "did:plc:client" }); + + var scoped = client.WithAcceptLabelers(new[] { "did:plc:scoped" }); + await scoped.GetAsync("com.example.get"); + + Assert.Equal("did:plc:scoped", Assert.Single(handler.Requests).Header("atproto-accept-labelers")); + Assert.Equal(new[] { "did:plc:scoped" }, scoped.LabelerDids); + } + + [Fact] + public async Task PostBinaryAsync_SendsStreamWithContentTypeAndParameters() + { + var handler = new RecordingHandler(_ => Json("{\"ok\":true}")); + using var client = CreateSessionClient(handler, out _); + var data = Encoding.UTF8.GetBytes("part-bytes"); + + var result = await client.PostBinaryAsync( + "app.bsky.video.uploadPart", null, Params(("jobId", "job1"), ("partNumber", "2")), + new MemoryStream(data), "application/octet-stream"); + + var request = Assert.Single(handler.Requests); + Assert.Equal(HttpMethod.Post, request.Method); + Assert.Equal("?jobId=job1&partNumber=2", request.Uri.Query); + Assert.Equal("application/octet-stream", request.ContentType); + Assert.Equal(data, request.Body); + Assert.True(result.GetProperty("ok").GetBoolean()); + } + + [Fact] + public async Task PostBinaryAsync_SeekableStream_IsReplayedAfterTokenRefresh() + { + var calls = 0; + var handler = new RecordingHandler(r => + { + if (r.Uri.AbsolutePath.EndsWith("refreshSession", StringComparison.Ordinal)) + { + return Json(SessionJson(CreateJwt(UserDid, DateTimeOffset.UtcNow.AddHours(2), "fresh"))); + } + + return ++calls == 1 ? Status(HttpStatusCode.Unauthorized) : Json("{\"ok\":true}"); + }); + using var client = CreateSessionClient(handler, out _); + var data = Encoding.UTF8.GetBytes("blob-data"); + + await client.PostBinaryAsync("com.atproto.repo.uploadBlob", null, null, new MemoryStream(data), "image/png"); + + var uploads = handler.Requests.Where(r => r.Uri.AbsolutePath.EndsWith("uploadBlob", StringComparison.Ordinal)).ToList(); + Assert.Equal(2, uploads.Count); + Assert.Equal(data, uploads[0].Body); + Assert.Equal(data, uploads[1].Body); + } + + [Fact] + public async Task PostBinaryAsync_NonSeekableStream_IsNotRetried() + { + var handler = new RecordingHandler(_ => Status(HttpStatusCode.Unauthorized)); + using var client = CreateSessionClient(handler, out _); + + await Assert.ThrowsAsync(() => + client.PostBinaryAsync("com.atproto.repo.uploadBlob", null, null, new NonSeekableStream(new byte[] { 1, 2, 3 }), "image/png")); + + Assert.Single(handler.Requests); + } + + [Fact] + public async Task GetBytesAsync_ReturnsBodyAndAcceptsAnyType() + { + var bytes = new byte[] { 0x89, 0x50, 0x4E, 0x47 }; + var handler = new RecordingHandler(_ => new HttpResponseMessage(HttpStatusCode.OK) { Content = new ByteArrayContent(bytes) }); + using var client = CreateSessionClient(handler, out _); + + var result = await client.GetBytesAsync("com.atproto.sync.getBlob", null, Params(("did", UserDid), ("cid", "bafy"))); + + Assert.Equal(bytes, result); + Assert.Equal("*/*", Assert.Single(handler.Requests).Header("Accept")); + } + + [Fact] + public async Task GetBytesAsync_ErrorStatus_Throws() + { + var handler = new RecordingHandler(_ => Json("{\"error\":\"BlobNotFound\",\"message\":\"nope\"}", HttpStatusCode.BadRequest)); + using var client = CreateSessionClient(handler, out _); + + var ex = await Assert.ThrowsAsync(() => client.GetBytesAsync("com.atproto.sync.getBlob", null, null)); + Assert.Equal("BlobNotFound", ex.ErrorCode); + } + + [Fact] + public async Task PostWithParametersAsync_SendsParametersAndJsonBody() + { + var handler = new RecordingHandler(_ => Json("{\"ok\":true}")); + using var client = CreateSessionClient(handler, out _); + + await client.PostWithParametersAsync( + "com.example.procedure", "did:web:svc.example#svc", Params(("mode", "fast")), new { name = "x" }); + + var request = Assert.Single(handler.Requests); + Assert.Equal("?mode=fast", request.Uri.Query); + Assert.Equal("application/json", request.ContentType); + Assert.Equal("{\"name\":\"x\"}", Encoding.UTF8.GetString(request.Body!)); + Assert.Equal("did:web:svc.example#svc", request.Header("atproto-proxy")); + } + + [Fact] + public async Task UserAgent_IsAddedPerRequest_WhenCallerSuppliesHttpClient() + { + var handler = new RecordingHandler(_ => Json("{\"value\":1}")); + var options = CreateOptions(new HttpClient(handler)); + options.UserAgent = "CarpaNetTests/1.0"; + using var client = ATProtoClient.Create(options); + + await client.GetAsync("com.example.get"); + + Assert.Equal("CarpaNetTests/1.0", Assert.Single(handler.Requests).Header("User-Agent")); + } + + [Fact] + public void UserAgent_IsSetOnOwnedHttpClient() + { + var options = CreateOptions(null); + options.UserAgent = "CarpaNetTests/1.0"; + using var client = ATProtoClient.Create(options); + + Assert.Equal("CarpaNetTests/1.0", client.HttpClient.DefaultRequestHeaders.UserAgent.ToString()); + } + + [Fact] + public void Combine_OuterProxyAndDisableWin_HeadersMerge() + { + var inner = new XrpcRequestOptions + { + ProxyServiceDid = "did:web:inner#svc", + AcceptLabelers = new[] { "did:plc:inner" }, + Headers = new Dictionary { ["a"] = "inner", ["b"] = "inner" }, + }; + var outer = new XrpcRequestOptions + { + DisableProxy = true, + Headers = new Dictionary { ["A"] = "outer" }, + }; + + var combined = XrpcRequestOptions.Combine(outer, inner)!; + + Assert.Null(combined.EffectiveProxyServiceDid); + Assert.Equal(new[] { "did:plc:inner" }, combined.AcceptLabelers); + Assert.Equal("outer", combined.Headers!["a"]); + Assert.Equal("inner", combined.Headers!["b"]); + Assert.Same(inner, XrpcRequestOptions.Combine(null, inner)); + } + + [Fact] + public void Combine_OuterWithoutProxy_KeepsInnerProxy() + { + var combined = XrpcRequestOptions.Combine( + new XrpcRequestOptions { Headers = new Dictionary { ["x"] = "1" } }, + new XrpcRequestOptions { ProxyServiceDid = BlueskyServices.ChatServiceDid })!; + + Assert.Equal(BlueskyServices.ChatServiceDid, combined.EffectiveProxyServiceDid); + } + + [Theory] + [InlineData("https://pds.example", "https://pds.example/xrpc/a", true)] + [InlineData("https://pds.example", "https://PDS.example:443/xrpc/a", true)] + [InlineData("https://pds.example", "http://pds.example/xrpc/a", false)] + [InlineData("https://pds.example", "https://pds.example:8443/xrpc/a", false)] + [InlineData("https://pds.example", "https://evil.pds.example/xrpc/a", false)] + public void IsSameOrigin_ComparesSchemeHostPort(string origin, string url, bool expected) + { + Assert.Equal(expected, XrpcHttpHandler.IsSameOrigin(new Uri(url), new Uri(origin))); + } + + [Fact] + public async Task ProgressReportingStream_ReportsBytesRead() + { + var reports = new List(); + var progress = new SyncProgress(reports.Add); + using var stream = new ProgressReportingStream(new MemoryStream(new byte[10]), progress); + var buffer = new byte[4]; + + while (await stream.ReadAsync(buffer, 0, buffer.Length) > 0) + { + } + + Assert.Equal(new long[] { 4, 8, 10 }, reports); + Assert.Equal(10, stream.BytesRead); + } + + [Fact] + public async Task RateLimitHandler_UnreplayableBody_ReturnsRateLimitResponse() + { + var inner = new RecordingHandler(_ => Status((HttpStatusCode)429)); + using var rateLimit = new RateLimitHandler(inner) { AutoRetryOnRateLimit = true, MaxRetries = 3 }; + using var http = new HttpClient(rateLimit); + using var request = new HttpRequestMessage(HttpMethod.Post, "https://pds.example/xrpc/com.atproto.repo.uploadBlob") + { + Content = new StreamContent(new NonSeekableStream(new byte[] { 1, 2, 3 })), + }; + + using var response = await http.SendAsync(request); + + Assert.Equal((HttpStatusCode)429, response.StatusCode); + Assert.Single(inner.Requests); + } + + #region Helpers + + internal static ATProtoClientOptions CreateOptions(HttpClient? httpClient) + { + return new ATProtoClientOptions + { + HttpClient = httpClient, + JsonOptions = TestHelpers.CreateJsonOptions(), + CborContext = TestHelpers.CreateCborContext(), + CreateIdentityResolver = false, + }; + } + + internal static ATProtoClient CreateSessionClient(RecordingHandler handler, out string accessJwt, IdentityResolver? resolver = null) + { + accessJwt = CreateJwt(UserDid, DateTimeOffset.UtcNow.AddHours(2), "access"); + var refreshJwt = CreateJwt(UserDid, DateTimeOffset.UtcNow.AddDays(30), "refresh"); + var options = CreateOptions(new HttpClient(handler)); + options.IdentityResolver = resolver; + return ATProtoClient.CreateWithRestoredSession(accessJwt, refreshJwt, UserDid, "user.example", UserPds, options); + } + + private static async Task CreateResolverAsync(RecordingHandler handler) + { + var cache = new MemoryIdentityCache(); + await cache.SetDidDocumentAsync(UserDid, DidDoc(UserDid, UserPds)); + await cache.SetDidDocumentAsync(OtherDid, DidDoc(OtherDid, OtherPds)); + return new IdentityResolver(new HttpClient(handler), cache: cache); + } + + private static DidDocument DidDoc(string did, Uri pds) + { + return new DidDocument + { + Id = did, + Service = new List + { + new DidService { Id = "#atproto_pds", Type = "AtprotoPersonalDataServer", ServiceEndpoint = pds.ToString().TrimEnd('/') }, + }, + }; + } + + internal static string CreateJwt(string sub, DateTimeOffset expires, string nonce) + { + static string B64(string json) => Convert.ToBase64String(Encoding.UTF8.GetBytes(json)).TrimEnd('=').Replace('+', '-').Replace('/', '_'); + return B64("{\"alg\":\"none\",\"typ\":\"JWT\"}") + "." + + B64($"{{\"sub\":\"{sub}\",\"exp\":{expires.ToUnixTimeSeconds()},\"jti\":\"{nonce}\"}}") + ".sig"; + } + + internal static string SessionJson(string accessJwt) + { + var refresh = CreateJwt(UserDid, DateTimeOffset.UtcNow.AddDays(30), "refresh2"); + return $"{{\"accessJwt\":\"{accessJwt}\",\"refreshJwt\":\"{refresh}\",\"handle\":\"user.example\",\"did\":\"{UserDid}\"}}"; + } + + internal static IEnumerable> Params(params (string Key, string Value)[] values) + => values.Select(v => new KeyValuePair(v.Key, v.Value)).ToList(); + + internal static HttpResponseMessage Json(string json, HttpStatusCode status = HttpStatusCode.OK) + => new(status) { Content = new StringContent(json, Encoding.UTF8, "application/json") }; + + internal static HttpResponseMessage Status(HttpStatusCode status) + => new(status) { Content = new StringContent(string.Empty) }; + + #endregion +} + +/// +/// Records each request (with its body read at send time) and answers through a callback. +/// +internal sealed class RecordingHandler : HttpMessageHandler +{ + private readonly Func _respond; + + public RecordingHandler(Func respond) + { + _respond = respond; + } + + public List Requests { get; } = new(); + + protected override async Task SendAsync(HttpRequestMessage request, CancellationToken cancellationToken) + { + byte[]? body = null; + string? contentType = null; + if (request.Content != null) + { + // Copy without buffering, as a socket handler does, so an unreplayable body stays unreplayable. + using var copy = new MemoryStream(); + await request.Content.CopyToAsync(copy, cancellationToken); + body = copy.ToArray(); + contentType = request.Content.Headers.ContentType?.MediaType; + } + + var headers = request.Headers.ToDictionary(h => h.Key, h => string.Join(",", h.Value), StringComparer.OrdinalIgnoreCase); + var recorded = new RecordedRequest(request.Method, request.RequestUri!, headers, body, contentType); + Requests.Add(recorded); + return _respond(recorded); + } +} + +internal sealed record RecordedRequest( + HttpMethod Method, + Uri Uri, + Dictionary Headers, + byte[]? Body, + string? ContentType) +{ + public string? Header(string name) => Headers.TryGetValue(name, out var value) ? value : null; +} + +internal sealed class NonSeekableStream : MemoryStream +{ + public NonSeekableStream(byte[] data) + : base(data) + { + } + + public override bool CanSeek => false; +} + +internal sealed class SyncProgress : IProgress +{ + private readonly Action _report; + + public SyncProgress(Action report) + { + _report = report; + } + + public void Report(long value) => _report(value); +} diff --git a/tests/CarpaNet.UnitTests/Identity/DnsOverHttpsResolverTests.cs b/tests/CarpaNet.UnitTests/Identity/DnsOverHttpsResolverTests.cs new file mode 100644 index 0000000..218584e --- /dev/null +++ b/tests/CarpaNet.UnitTests/Identity/DnsOverHttpsResolverTests.cs @@ -0,0 +1,301 @@ +using System; +using System.Collections.Generic; +using System.Linq; +using System.Net; +using System.Net.Http; +using System.Text; +using System.Threading; +using System.Threading.Tasks; +using CarpaNet.Identity; +using Xunit; + +namespace CarpaNet.UnitTests.Identity; + +/// +/// HttpMessageHandler that answers requests with a delegate and records them. No network access. +/// +internal sealed class FakeHttpHandler : HttpMessageHandler +{ + private readonly Func> _respond; + + public FakeHttpHandler(Func respond) + { + _respond = (request, _) => Task.FromResult(respond(request)); + } + + public FakeHttpHandler(Func> respond) + { + _respond = respond; + } + + public List Requests { get; } = new(); + + protected override Task SendAsync(HttpRequestMessage request, CancellationToken cancellationToken) + { + Requests.Add(request); + return _respond(request, cancellationToken); + } + + public static HttpResponseMessage Json(string json, HttpStatusCode status = HttpStatusCode.OK) + { + return new HttpResponseMessage(status) + { + Content = new StringContent(json, Encoding.UTF8, "application/json"), + }; + } + + public static HttpResponseMessage Text(string text, HttpStatusCode status = HttpStatusCode.OK) + { + return new HttpResponseMessage(status) + { + Content = new StringContent(text, Encoding.UTF8, "text/plain"), + }; + } +} + +public class DnsOverHttpsResolverTests +{ + private const string Endpoint1 = "https://doh1.test/dns-query"; + private const string Endpoint2 = "https://doh2.test/resolve"; + + private static string TxtResponse(params string[] data) + { + var answers = string.Join(",", data.Select(d => + $"{{\"name\":\"_atproto.example.com\",\"type\":16,\"TTL\":300,\"data\":{System.Text.Json.JsonSerializer.Serialize(d)}}}")); + return $"{{\"Status\":0,\"TC\":false,\"RD\":true,\"RA\":true,\"AD\":false,\"CD\":false,\"Answer\":[{answers}]}}"; + } + + private static (DnsOverHttpsResolver Resolver, FakeHttpHandler Handler) Create( + Func respond, TimeSpan? timeout = null) + { + var handler = new FakeHttpHandler(respond); + var resolver = new DnsOverHttpsResolver(new HttpClient(handler), new[] { Endpoint1, Endpoint2 }, timeout); + return (resolver, handler); + } + + [Fact] + public async Task GetTxtRecords_SendsJsonQuery() + { + var (resolver, handler) = Create(_ => FakeHttpHandler.Json(TxtResponse("\"did=did:plc:abc\""))); + + await resolver.GetTxtRecordsAsync("_atproto.example.com"); + + var request = Assert.Single(handler.Requests); + Assert.Equal(HttpMethod.Get, request.Method); + Assert.Equal("https://doh1.test/dns-query?name=_atproto.example.com&type=TXT", request.RequestUri!.ToString()); + Assert.Contains(request.Headers.Accept, h => h.MediaType == "application/dns-json"); + } + + [Fact] + public async Task GetTxtRecords_SingleQuotedTxt_IsUnquoted() + { + var (resolver, _) = Create(_ => FakeHttpHandler.Json(TxtResponse("\"did=did:plc:abc123\""))); + + var records = await resolver.GetTxtRecordsAsync("_atproto.example.com"); + + Assert.Equal(new[] { "did=did:plc:abc123" }, records); + } + + [Fact] + public async Task GetTxtRecords_UnquotedTxt_IsReturnedAsIs() + { + var (resolver, _) = Create(_ => FakeHttpHandler.Json(TxtResponse("did=did:plc:abc123"))); + + var records = await resolver.GetTxtRecordsAsync("_atproto.example.com"); + + Assert.Equal(new[] { "did=did:plc:abc123" }, records); + } + + [Fact] + public async Task GetTxtRecords_SplitCharacterStrings_AreJoined() + { + var (resolver, _) = Create(_ => FakeHttpHandler.Json(TxtResponse("\"did=did:plc:\" \"abc123\""))); + + var records = await resolver.GetTxtRecordsAsync("_atproto.example.com"); + + Assert.Equal(new[] { "did=did:plc:abc123" }, records); + } + + [Fact] + public async Task GetTxtRecords_MultipleAnswers_IgnoresNonTxt() + { + const string json = "{\"Status\":0,\"Answer\":[" + + "{\"name\":\"_atproto.example.com\",\"type\":5,\"TTL\":300,\"data\":\"target.example.com.\"}," + + "{\"name\":\"target.example.com\",\"type\":16,\"TTL\":300,\"data\":\"\\\"v=spf1 -all\\\"\"}," + + "{\"name\":\"target.example.com\",\"type\":16,\"TTL\":300,\"data\":\"\\\"did=did:plc:xyz\\\"\"}]}"; + var (resolver, _) = Create(_ => FakeHttpHandler.Json(json)); + + var records = await resolver.GetTxtRecordsAsync("_atproto.example.com"); + + Assert.Equal(new[] { "v=spf1 -all", "did=did:plc:xyz" }, records); + } + + [Fact] + public async Task GetTxtRecords_NoErrorWithoutAnswer_ReturnsEmptyWithoutFallback() + { + var (resolver, handler) = Create(_ => FakeHttpHandler.Json("{\"Status\":0}")); + + var records = await resolver.GetTxtRecordsAsync("_atproto.example.com"); + + Assert.Empty(records); + Assert.Single(handler.Requests); + } + + [Fact] + public async Task GetTxtRecords_NxDomain_ReturnsEmptyWithoutFallback() + { + var (resolver, handler) = Create(_ => FakeHttpHandler.Json("{\"Status\":3,\"Authority\":[]}")); + + var records = await resolver.GetTxtRecordsAsync("_atproto.example.com"); + + Assert.Empty(records); + Assert.Single(handler.Requests); + } + + [Fact] + public async Task GetTxtRecords_ServFail_FallsBackToNextEndpoint() + { + var (resolver, handler) = Create(request => + request.RequestUri!.Host == "doh1.test" + ? FakeHttpHandler.Json("{\"Status\":2}") + : FakeHttpHandler.Json(TxtResponse("\"did=did:plc:fallback\""))); + + var records = await resolver.GetTxtRecordsAsync("_atproto.example.com"); + + Assert.Equal(new[] { "did=did:plc:fallback" }, records); + Assert.Equal(2, handler.Requests.Count); + Assert.StartsWith(Endpoint2, handler.Requests[1].RequestUri!.ToString()); + } + + [Fact] + public async Task GetTxtRecords_HttpError_FallsBackToNextEndpoint() + { + var (resolver, _) = Create(request => + request.RequestUri!.Host == "doh1.test" + ? FakeHttpHandler.Text("bad gateway", HttpStatusCode.BadGateway) + : FakeHttpHandler.Json(TxtResponse("\"did=did:plc:fallback\""))); + + var records = await resolver.GetTxtRecordsAsync("_atproto.example.com"); + + Assert.Equal(new[] { "did=did:plc:fallback" }, records); + } + + [Fact] + public async Task GetTxtRecords_MalformedJson_FallsBackToNextEndpoint() + { + var (resolver, _) = Create(request => + request.RequestUri!.Host == "doh1.test" + ? FakeHttpHandler.Json("{\"Status\":0,\"Answer\":[{") + : FakeHttpHandler.Json(TxtResponse("\"did=did:plc:fallback\""))); + + var records = await resolver.GetTxtRecordsAsync("_atproto.example.com"); + + Assert.Equal(new[] { "did=did:plc:fallback" }, records); + } + + [Fact] + public async Task GetTxtRecords_TransportError_FallsBackToNextEndpoint() + { + var (resolver, _) = Create(request => + request.RequestUri!.Host == "doh1.test" + ? throw new HttpRequestException("connection refused") + : FakeHttpHandler.Json(TxtResponse("\"did=did:plc:fallback\""))); + + var records = await resolver.GetTxtRecordsAsync("_atproto.example.com"); + + Assert.Equal(new[] { "did=did:plc:fallback" }, records); + } + + [Fact] + public async Task GetTxtRecords_AllEndpointsFail_ReturnsEmpty() + { + var (resolver, handler) = Create(_ => FakeHttpHandler.Json("not json")); + + var records = await resolver.GetTxtRecordsAsync("_atproto.example.com"); + + Assert.Empty(records); + Assert.Equal(2, handler.Requests.Count); + } + + [Fact] + public async Task GetTxtRecords_EndpointTimeout_FallsBackToNextEndpoint() + { + var handler = new FakeHttpHandler(async (request, ct) => + { + if (request.RequestUri!.Host == "doh1.test") + await Task.Delay(Timeout.Infinite, ct); + return FakeHttpHandler.Json(TxtResponse("\"did=did:plc:fallback\"")); + }); + var resolver = new DnsOverHttpsResolver(new HttpClient(handler), new[] { Endpoint1, Endpoint2 }, TimeSpan.FromMilliseconds(50)); + + var records = await resolver.GetTxtRecordsAsync("_atproto.example.com"); + + Assert.Equal(new[] { "did=did:plc:fallback" }, records); + } + + [Fact] + public async Task GetTxtRecords_CallerCancellation_Throws() + { + var (resolver, _) = Create(_ => FakeHttpHandler.Json(TxtResponse("\"did=did:plc:abc\""))); + using var cts = new CancellationTokenSource(); + cts.Cancel(); + + await Assert.ThrowsAnyAsync( + () => resolver.GetTxtRecordsAsync("_atproto.example.com", cts.Token)); + } + + [Fact] + public void Constructor_NoEndpoints_UsesDefaults() + { + var resolver = new DnsOverHttpsResolver(new HttpClient(new FakeHttpHandler(_ => FakeHttpHandler.Text("")))); + + Assert.Equal(DnsOverHttpsResolver.DefaultEndpoints, resolver.Endpoints); + Assert.Equal(DnsOverHttpsResolver.CloudflareEndpoint, resolver.Endpoints[0]); + } + + [Fact] + public void Constructor_NullHttpClient_Throws() + { + Assert.Throws(() => new DnsOverHttpsResolver(null!)); + } + + [Theory] + [InlineData("\"did=did:plc:abc\"", "did=did:plc:abc")] + [InlineData("did=did:plc:abc", "did=did:plc:abc")] + [InlineData("\"a\" \"b\" \"c\"", "abc")] + [InlineData("\"say \\\"hi\\\"\"", "say \"hi\"")] + [InlineData("\"a\\059b\"", "a;b")] + [InlineData("\"\"", "")] + public void ParseTxtData_HandlesPresentationFormat(string data, string expected) + { + Assert.Equal(expected, DnsOverHttpsResolver.ParseTxtData(data)); + } +} + +public class DnsResolverDefaultsTests +{ + [Fact] + public void CreateDefault_OnNonBrowserHost_ReturnsUdpResolver() + { + // Unit tests run on a desktop/server runtime, which supports UDP. + Assert.True(DnsResolverDefaults.IsUdpDnsSupported); + + var resolver = DnsResolverDefaults.CreateDefault(new HttpClient(new FakeHttpHandler(_ => FakeHttpHandler.Text("")))); + + Assert.IsType(resolver); + } + + [Fact] + public void CreateDefault_NullHttpClient_Throws() + { + Assert.Throws(() => DnsResolverDefaults.CreateDefault(null!)); + } + + [Fact] + public void IdentityResolver_WithoutDnsResolver_UsesPlatformDefault() + { + using var resolver = new IdentityResolver(new HttpClient(new FakeHttpHandler(_ => FakeHttpHandler.Text("")))); + + Assert.IsType(resolver.DnsResolver); + } +} diff --git a/tests/CarpaNet.UnitTests/Identity/HandleResolutionOrderTests.cs b/tests/CarpaNet.UnitTests/Identity/HandleResolutionOrderTests.cs new file mode 100644 index 0000000..2680a22 --- /dev/null +++ b/tests/CarpaNet.UnitTests/Identity/HandleResolutionOrderTests.cs @@ -0,0 +1,266 @@ +using System; +using System.Collections.Generic; +using System.Net; +using System.Net.Http; +using System.Threading; +using System.Threading.Tasks; +using CarpaNet.Identity; +using Xunit; + +namespace CarpaNet.UnitTests.Identity; + +public class HandleResolutionOrderTests +{ + private const string Handle = "alice.example.com"; + private const string Service = "https://appview.test"; + + private sealed class FakeDnsResolver : IDnsResolver + { + private readonly Func> _answer; + + public FakeDnsResolver(Func> answer) => _answer = answer; + + public int Calls { get; private set; } + + public Task> GetTxtRecordsAsync(string name, CancellationToken cancellationToken = default) + { + Calls++; + return Task.FromResult(_answer(name)); + } + } + + private static FakeDnsResolver NoDns() => new(_ => Array.Empty()); + + private static bool IsWellKnown(HttpRequestMessage r) => r.RequestUri!.AbsolutePath == "/.well-known/atproto-did"; + + private static bool IsXrpc(HttpRequestMessage r) => r.RequestUri!.AbsolutePath == "/xrpc/com.atproto.identity.resolveHandle"; + + [Fact] + public async Task DefaultOrder_DnsWins_NoHttpRequests() + { + var handler = new FakeHttpHandler(_ => FakeHttpHandler.Text("", HttpStatusCode.NotFound)); + var dns = new FakeDnsResolver(_ => new[] { "did=did:plc:fromdns" }); + using var resolver = new IdentityResolver(new HttpClient(handler), new IdentityResolverOptions + { + DnsResolver = dns, + HandleResolutionServiceUrl = Service, + }); + + var did = await resolver.ResolveHandleAsync(Handle); + + Assert.Equal("did:plc:fromdns", did); + Assert.Empty(handler.Requests); + } + + [Fact] + public async Task DefaultOrder_WellKnownBeforeXrpc() + { + var handler = new FakeHttpHandler(r => + IsWellKnown(r) ? FakeHttpHandler.Text("did:plc:fromwellknown\n") + : FakeHttpHandler.Json("{\"did\":\"did:plc:fromxrpc\"}")); + using var resolver = new IdentityResolver(new HttpClient(handler), new IdentityResolverOptions + { + DnsResolver = NoDns(), + HandleResolutionServiceUrl = Service, + }); + + var did = await resolver.ResolveHandleAsync(Handle); + + Assert.Equal("did:plc:fromwellknown", did); + Assert.DoesNotContain(handler.Requests, IsXrpc); + } + + [Fact] + public async Task DefaultOrder_FallsBackToXrpc() + { + var handler = new FakeHttpHandler(r => + IsWellKnown(r) ? FakeHttpHandler.Text("", HttpStatusCode.NotFound) + : FakeHttpHandler.Json("{\"did\":\"did:plc:fromxrpc\"}")); + using var resolver = new IdentityResolver(new HttpClient(handler), new IdentityResolverOptions + { + DnsResolver = NoDns(), + HandleResolutionServiceUrl = Service + "/", + }); + + var did = await resolver.ResolveHandleAsync(Handle); + + Assert.Equal("did:plc:fromxrpc", did); + var xrpc = Assert.Single(handler.Requests, IsXrpc); + Assert.Equal("https://appview.test/xrpc/com.atproto.identity.resolveHandle?handle=alice.example.com", xrpc.RequestUri!.ToString()); + } + + [Fact] + public async Task XrpcOnly_SkipsDnsAndWellKnown() + { + var handler = new FakeHttpHandler(_ => FakeHttpHandler.Json("{\"did\":\"did:plc:fromxrpc\"}")); + var dns = new FakeDnsResolver(_ => new[] { "did=did:plc:fromdns" }); + using var resolver = new IdentityResolver(new HttpClient(handler), new IdentityResolverOptions + { + DnsResolver = dns, + HandleResolutionServiceUrl = Service, + HandleResolutionOrder = new[] { HandleResolutionMethod.Xrpc }, + }); + + var did = await resolver.ResolveHandleAsync(Handle); + + Assert.Equal("did:plc:fromxrpc", did); + Assert.Equal(0, dns.Calls); + Assert.All(handler.Requests, r => Assert.True(IsXrpc(r))); + } + + [Fact] + public async Task CustomOrder_XrpcBeforeDns() + { + var handler = new FakeHttpHandler(_ => FakeHttpHandler.Json("{\"error\":\"InvalidRequest\",\"message\":\"Unable to resolve handle\"}", HttpStatusCode.BadRequest)); + var dns = new FakeDnsResolver(_ => new[] { "did=did:plc:fromdns" }); + using var resolver = new IdentityResolver(new HttpClient(handler), new IdentityResolverOptions + { + DnsResolver = dns, + HandleResolutionServiceUrl = Service, + HandleResolutionOrder = new[] { HandleResolutionMethod.Xrpc, HandleResolutionMethod.Dns }, + }); + + var did = await resolver.ResolveHandleAsync(Handle); + + Assert.Equal("did:plc:fromdns", did); + Assert.Single(handler.Requests, IsXrpc); + Assert.Equal(1, dns.Calls); + } + + [Fact] + public async Task Xrpc_WithoutServiceUrl_IsSkipped() + { + var handler = new FakeHttpHandler(_ => FakeHttpHandler.Text("", HttpStatusCode.NotFound)); + using var resolver = new IdentityResolver(new HttpClient(handler), new IdentityResolverOptions + { + DnsResolver = NoDns(), + }); + + await Assert.ThrowsAsync(() => resolver.ResolveHandleAsync(Handle)); + Assert.DoesNotContain(handler.Requests, IsXrpc); + } + + [Theory] + [InlineData("{\"did\":\"not-a-did\"}", HttpStatusCode.OK)] + [InlineData("{\"did\":", HttpStatusCode.OK)] + [InlineData("{}", HttpStatusCode.OK)] + [InlineData("{\"error\":\"InvalidRequest\"}", HttpStatusCode.BadRequest)] + [InlineData("oops", HttpStatusCode.InternalServerError)] + public async Task Xrpc_BadResponse_FailsResolution(string body, HttpStatusCode status) + { + var handler = new FakeHttpHandler(_ => FakeHttpHandler.Json(body, status)); + using var resolver = new IdentityResolver(new HttpClient(handler), new IdentityResolverOptions + { + HandleResolutionServiceUrl = Service, + HandleResolutionOrder = new[] { HandleResolutionMethod.Xrpc }, + }); + + await Assert.ThrowsAsync(() => resolver.ResolveHandleAsync(Handle)); + } + + [Fact] + public async Task Xrpc_TransportError_FailsResolution() + { + var handler = new FakeHttpHandler(_ => throw new HttpRequestException("connection refused")); + using var resolver = new IdentityResolver(new HttpClient(handler), new IdentityResolverOptions + { + HandleResolutionServiceUrl = Service, + HandleResolutionOrder = new[] { HandleResolutionMethod.Xrpc }, + }); + + await Assert.ThrowsAsync(() => resolver.ResolveHandleAsync(Handle)); + } + + [Fact] + public async Task Xrpc_Result_IsCached() + { + var handler = new FakeHttpHandler(_ => FakeHttpHandler.Json("{\"did\":\"did:plc:fromxrpc\"}")); + var cache = new MemoryIdentityCache(); + using var resolver = new IdentityResolver(new HttpClient(handler), new IdentityResolverOptions + { + Cache = cache, + HandleResolutionServiceUrl = Service, + HandleResolutionOrder = new[] { HandleResolutionMethod.Xrpc }, + }); + + await resolver.ResolveHandleAsync(Handle); + var did = await resolver.ResolveHandleAsync(Handle); + + Assert.Equal("did:plc:fromxrpc", did); + Assert.Single(handler.Requests); + Assert.Equal("did:plc:fromxrpc", await cache.GetHandleDidAsync(Handle)); + } + + [Fact] + public async Task ResolveAsync_XrpcResolvedHandle_StillChecksDidDocumentHandle() + { + const string didDoc = "{\"id\":\"did:plc:fromxrpc\",\"alsoKnownAs\":[\"at://someone-else.example.com\"]," + + "\"service\":[{\"id\":\"#atproto_pds\",\"type\":\"AtprotoPersonalDataServer\",\"serviceEndpoint\":\"https://pds.test\"}]}"; + var handler = new FakeHttpHandler(r => + IsXrpc(r) ? FakeHttpHandler.Json("{\"did\":\"did:plc:fromxrpc\"}") + : FakeHttpHandler.Json(didDoc)); + using var resolver = new IdentityResolver(new HttpClient(handler), new IdentityResolverOptions + { + PlcDirectoryUrl = "https://plc.test", + HandleResolutionServiceUrl = Service, + HandleResolutionOrder = new[] { HandleResolutionMethod.Xrpc }, + }); + + await Assert.ThrowsAsync(() => resolver.ResolveAsync(Handle)); + } + + [Fact] + public async Task ResolveAsync_XrpcResolvedHandle_MatchingDidDocument_Succeeds() + { + const string didDoc = "{\"id\":\"did:plc:fromxrpc\",\"alsoKnownAs\":[\"at://alice.example.com\"]," + + "\"service\":[{\"id\":\"#atproto_pds\",\"type\":\"AtprotoPersonalDataServer\",\"serviceEndpoint\":\"https://pds.test\"}]}"; + var handler = new FakeHttpHandler(r => + IsXrpc(r) ? FakeHttpHandler.Json("{\"did\":\"did:plc:fromxrpc\"}") + : FakeHttpHandler.Json(didDoc)); + using var resolver = new IdentityResolver(new HttpClient(handler), new IdentityResolverOptions + { + PlcDirectoryUrl = "https://plc.test", + HandleResolutionServiceUrl = Service, + HandleResolutionOrder = new[] { HandleResolutionMethod.Xrpc }, + }); + + var doc = await resolver.ResolveAsync(Handle); + + Assert.Equal("did:plc:fromxrpc", doc.Id); + Assert.Contains(handler.Requests, r => r.RequestUri!.ToString() == "https://plc.test/did:plc:fromxrpc"); + } + + [Fact] + public void Options_DuplicateMethods_AreRemoved() + { + using var resolver = new IdentityResolver(new HttpClient(new FakeHttpHandler(_ => FakeHttpHandler.Text(""))), new IdentityResolverOptions + { + DnsResolver = NoDns(), + HandleResolutionOrder = new[] { HandleResolutionMethod.WellKnown, HandleResolutionMethod.Dns, HandleResolutionMethod.WellKnown }, + }); + + Assert.Equal(new[] { HandleResolutionMethod.WellKnown, HandleResolutionMethod.Dns }, resolver.HandleResolutionOrder); + } + + [Fact] + public void Options_EmptyOrder_UsesDefault() + { + using var resolver = new IdentityResolver(new HttpClient(new FakeHttpHandler(_ => FakeHttpHandler.Text(""))), new IdentityResolverOptions + { + DnsResolver = NoDns(), + HandleResolutionOrder = Array.Empty(), + }); + + Assert.Equal(IdentityResolverOptions.DefaultHandleResolutionOrder, resolver.HandleResolutionOrder); + } + + [Fact] + public void XrpcHandleResolver_InvalidServiceUrl_Throws() + { + var http = new HttpClient(new FakeHttpHandler(_ => FakeHttpHandler.Text(""))); + + Assert.Throws(() => new XrpcHandleResolver(http, "not a url")); + Assert.Throws(() => new XrpcHandleResolver(http, "ftp://example.com")); + Assert.Throws(() => new IdentityResolver(http, new IdentityResolverOptions { HandleResolutionServiceUrl = "relative/path" })); + } +} diff --git a/tests/CarpaNet.UnitTests/OAuth/OAuthCallbackValidationTests.cs b/tests/CarpaNet.UnitTests/OAuth/OAuthCallbackValidationTests.cs new file mode 100644 index 0000000..04aa116 --- /dev/null +++ b/tests/CarpaNet.UnitTests/OAuth/OAuthCallbackValidationTests.cs @@ -0,0 +1,373 @@ +using System; +using System.Collections.Generic; +using System.Linq; +using System.Net; +using System.Net.Http; +using System.Text; +using System.Threading; +using System.Threading.Tasks; +using CarpaNet.Identity; +using CarpaNet.OAuth; +using CarpaNet.OAuth.Storage; +using Xunit; + +namespace CarpaNet.UnitTests.OAuth; + +/// +/// Tests for the validation done by : the RFC 9207 +/// iss parameter, and verification of the token response's sub. +/// No network access: every HTTP call is answered by . +/// +public class OAuthCallbackValidationTests +{ + private const string Issuer = "https://auth.example.com"; + private const string OtherIssuer = "https://evil.example.com"; + private const string PdsUrl = "https://pds.example.com"; + private const string EntrywayUrl = "https://entryway.example.com"; + private const string PlcDirectory = "https://plc.test"; + private const string UserDid = "did:plc:abcdefghijklmnopqrstuvwx"; + private const string OtherDid = "did:plc:zyxwvutsrqponmlkjihgfedc"; + private const string RedirectUri = "http://127.0.0.1:8080/callback"; + + [Fact] + public async Task Callback_IssMismatch_ThrowsAndDoesNotRequestToken() + { + var server = new FakeOAuthServer(); + using var session = server.CreateSession(); + + var state = await StartAsync(session, UserDid); + var ex = await Assert.ThrowsAsync( + () => session.CallbackAsync(CallbackUrl(state, OtherIssuer))); + + Assert.Equal("issuer_mismatch", ex.ErrorCode); + Assert.Equal(0, server.TokenRequests); + } + + [Fact] + public async Task Callback_IssMismatch_ThrowsEvenWhenServerDoesNotAdvertiseIss() + { + var server = new FakeOAuthServer { IssParameterSupported = false }; + using var session = server.CreateSession(); + + var state = await StartAsync(session, UserDid); + var ex = await Assert.ThrowsAsync( + () => session.CallbackAsync(CallbackUrl(state, OtherIssuer))); + + Assert.Equal("issuer_mismatch", ex.ErrorCode); + Assert.Equal(0, server.TokenRequests); + } + + [Fact] + public async Task Callback_IssMissing_WhenRequired_ThrowsAndDoesNotRequestToken() + { + var server = new FakeOAuthServer { IssParameterSupported = true }; + using var session = server.CreateSession(); + + var state = await StartAsync(session, UserDid, appState: "my-app-state"); + var ex = await Assert.ThrowsAsync( + () => session.CallbackAsync(CallbackUrl(state, iss: null))); + + Assert.Equal("missing_iss", ex.ErrorCode); + Assert.Equal("my-app-state", ex.AppState); + Assert.Equal(0, server.TokenRequests); + } + + [Fact] + public async Task Callback_IssMissing_WhenNotRequired_Succeeds() + { + var server = new FakeOAuthServer { IssParameterSupported = false }; + using var session = server.CreateSession(); + + var state = await StartAsync(session, UserDid); + using var client = await session.CallbackAsync(CallbackUrl(state, iss: null)); + + Assert.Equal(UserDid, client.Did); + Assert.Equal(1, server.TokenRequests); + } + + [Fact] + public async Task Callback_IssMatches_Succeeds() + { + var server = new FakeOAuthServer { IssParameterSupported = true }; + using var session = server.CreateSession(); + + var state = await StartAsync(session, UserDid, appState: "app"); + using var client = await session.CallbackAsync(CallbackUrl(state, Issuer)); + + Assert.Equal(UserDid, client.Did); + Assert.Equal("app", client.AppState); + Assert.Equal(new Uri(PdsUrl), client.TokenProvider.PdsUrl); + Assert.NotNull(await server.SessionStore.GetAsync(UserDid)); + Assert.Equal(0, server.RevocationRequests); + } + + [Fact] + public async Task Callback_SubNotADid_RevokesAndDoesNotStoreSession() + { + var server = new FakeOAuthServer { TokenSub = "alice.example.com" }; + using var session = server.CreateSession(); + + var state = await StartAsync(session, PdsUrl); + var ex = await Assert.ThrowsAsync( + () => session.CallbackAsync(CallbackUrl(state, Issuer))); + + Assert.Equal("invalid_sub", ex.ErrorCode); + Assert.Equal(1, server.RevocationRequests); + Assert.Contains("token=refresh-token", server.LastRevocationBody); + Assert.Null(await server.SessionStore.GetAsync("alice.example.com")); + } + + [Fact] + public async Task Callback_SubPdsNotProtectedByIssuer_RevokesAndDoesNotStoreSession() + { + var server = new FakeOAuthServer(); + server.PdsAuthorizationServers = new[] { OtherIssuer }; + using var session = server.CreateSession(); + + // Start from the entryway, which (correctly) points at the issuer + var state = await StartAsync(session, EntrywayUrl); + var ex = await Assert.ThrowsAsync( + () => session.CallbackAsync(CallbackUrl(state, Issuer))); + + Assert.Equal("sub_issuer_mismatch", ex.ErrorCode); + Assert.Equal(1, server.TokenRequests); + Assert.Equal(1, server.RevocationRequests); + Assert.Null(await server.SessionStore.GetAsync(UserDid)); + } + + [Fact] + public async Task Callback_SubDidResolutionFails_RevokesAndThrows() + { + var server = new FakeOAuthServer { TokenSub = OtherDid }; // No DID document for OtherDid + using var session = server.CreateSession(); + + var state = await StartAsync(session, EntrywayUrl); + var ex = await Assert.ThrowsAsync( + () => session.CallbackAsync(CallbackUrl(state, Issuer))); + + Assert.Equal("sub_verification_failed", ex.ErrorCode); + Assert.Equal(1, server.RevocationRequests); + Assert.Null(await server.SessionStore.GetAsync(OtherDid)); + } + + [Fact] + public async Task Callback_SubDiffersFromExpectedDid_RevokesAndThrows() + { + var server = new FakeOAuthServer { TokenSub = OtherDid }; + server.AddDidDocument(OtherDid, PdsUrl); + using var session = server.CreateSession(); + + var state = await StartAsync(session, UserDid); + var ex = await Assert.ThrowsAsync( + () => session.CallbackAsync(CallbackUrl(state, Issuer))); + + Assert.Equal("sub_mismatch", ex.ErrorCode); + Assert.Equal(1, server.RevocationRequests); + Assert.Null(await server.SessionStore.GetAsync(OtherDid)); + } + + [Fact] + public async Task Callback_StartedFromEntryway_UsesResolvedPdsAsAudience() + { + var server = new FakeOAuthServer(); + using var session = server.CreateSession(); + + var state = await StartAsync(session, EntrywayUrl); + using var client = await session.CallbackAsync(CallbackUrl(state, Issuer)); + + Assert.Equal(UserDid, client.Did); + Assert.Equal(new Uri(PdsUrl), client.BaseUrl); + Assert.Equal(new Uri(PdsUrl), client.TokenProvider.PdsUrl); + + var stored = await server.SessionStore.GetAsync(UserDid); + Assert.NotNull(stored); + Assert.Equal(PdsUrl, stored!.TokenSet.Audience); + Assert.Equal(Issuer, stored.TokenSet.Issuer); + } + + [Fact] + public async Task Callback_StartedFromEntryway_AnyAccountMaySignIn() + { + var server = new FakeOAuthServer { TokenSub = OtherDid }; + server.AddDidDocument(OtherDid, "https://other-pds.example.com/"); + server.AddProtectedResource("https://other-pds.example.com", Issuer); + using var session = server.CreateSession(); + + var state = await StartAsync(session, EntrywayUrl); + using var client = await session.CallbackAsync(CallbackUrl(state, Issuer)); + + Assert.Equal(OtherDid, client.Did); + Assert.Equal(new Uri("https://other-pds.example.com"), client.TokenProvider.PdsUrl); + } + + [Fact] + public async Task RestoreSession_AfterEntrywayCallback_KeepsResolvedPds() + { + var server = new FakeOAuthServer(); + + using (var session = server.CreateSession()) + { + var state = await StartAsync(session, EntrywayUrl); + using var client = await session.CallbackAsync(CallbackUrl(state, Issuer)); + } + + using var restoredSession = server.CreateSession(); + using var restored = await restoredSession.RestoreSessionAsync(UserDid); + + Assert.NotNull(restored); + Assert.Equal(UserDid, restored!.Did); + Assert.Equal(new Uri(PdsUrl), restored.TokenProvider.PdsUrl); + Assert.True(restored.IsAuthenticated); + } + + private static async Task StartAsync(OAuthSession session, string input, string? appState = null) + { + var url = await session.AuthorizeAsync(input, appState); + var query = System.Web.HttpUtility.ParseQueryString(new Uri(url).Query); + return query["state"] ?? throw new InvalidOperationException("No state in authorization URL."); + } + + private static string CallbackUrl(string state, string? iss) + { + var url = $"{RedirectUri}?code=auth-code&state={Uri.EscapeDataString(state)}"; + if (iss != null) + { + url += $"&iss={Uri.EscapeDataString(iss)}"; + } + + return url; + } + + /// + /// In-memory authorization server, PDS, entryway and PLC directory. + /// + private sealed class FakeOAuthServer : HttpMessageHandler + { + private readonly Dictionary _didDocuments = new(StringComparer.Ordinal); + private readonly Dictionary _protectedResources = new(StringComparer.OrdinalIgnoreCase); + + public FakeOAuthServer() + { + AddDidDocument(UserDid, PdsUrl); + AddProtectedResource(EntrywayUrl, Issuer); + } + + public bool IssParameterSupported { get; set; } = true; + + public string TokenSub { get; set; } = UserDid; + + public string[] PdsAuthorizationServers + { + set => _protectedResources[PdsUrl] = value; + } + + public MemoryOAuthSessionStore SessionStore { get; } = new(); + + public MemoryOAuthStateStore StateStore { get; } = new(); + + public int TokenRequests { get; private set; } + + public int RevocationRequests { get; private set; } + + public string LastRevocationBody { get; private set; } = string.Empty; + + public void AddDidDocument(string did, string pdsEndpoint) + { + _didDocuments[did] = $$""" + { + "id": "{{did}}", + "alsoKnownAs": ["at://user.example.com"], + "service": [ + { "id": "#atproto_pds", "type": "AtprotoPersonalDataServer", "serviceEndpoint": "{{pdsEndpoint}}" } + ] + } + """; + _protectedResources.TryAdd(pdsEndpoint.TrimEnd('/'), new[] { Issuer }); + } + + public void AddProtectedResource(string resourceUrl, params string[] authorizationServers) + { + _protectedResources[resourceUrl.TrimEnd('/')] = authorizationServers; + } + + public OAuthSession CreateSession() + { + var httpClient = new HttpClient(this, disposeHandler: false); + return new OAuthSession(new OAuthClientConfig + { + ClientId = OAuthClientConfig.CreateLoopbackClientId(8080), + RedirectUri = RedirectUri, + HttpClient = httpClient, + StateStore = StateStore, + SessionStore = SessionStore, + IdentityResolver = new IdentityResolver(httpClient, PlcDirectory), + }); + } + + protected override async Task SendAsync(HttpRequestMessage request, CancellationToken cancellationToken) + { + var uri = request.RequestUri!; + var origin = $"{uri.Scheme}://{uri.Authority}"; + var path = uri.AbsolutePath; + + if (request.Method == HttpMethod.Get && path == "/.well-known/oauth-protected-resource" && + _protectedResources.TryGetValue(origin, out var servers)) + { + var list = string.Join(",", servers.Select(s => $"\"{s}\"")); + return Json($$"""{ "resource": "{{origin}}", "authorization_servers": [{{list}}] }"""); + } + + if (request.Method == HttpMethod.Get && origin == Issuer && path == "/.well-known/oauth-authorization-server") + { + return Json($$""" + { + "issuer": "{{Issuer}}", + "authorization_endpoint": "{{Issuer}}/oauth/authorize", + "token_endpoint": "{{Issuer}}/oauth/token", + "revocation_endpoint": "{{Issuer}}/oauth/revoke", + "authorization_response_iss_parameter_supported": {{(IssParameterSupported ? "true" : "false")}}, + "dpop_signing_alg_values_supported": ["ES256"] + } + """); + } + + if (request.Method == HttpMethod.Get && origin == PlcDirectory) + { + var did = Uri.UnescapeDataString(path.TrimStart('/')); + return _didDocuments.TryGetValue(did, out var doc) + ? Json(doc) + : new HttpResponseMessage(HttpStatusCode.NotFound); + } + + if (request.Method == HttpMethod.Post && uri.ToString() == $"{Issuer}/oauth/token") + { + TokenRequests++; + Assert.True(request.Headers.Contains("DPoP")); + return Json($$""" + { + "access_token": "access-token", + "token_type": "DPoP", + "expires_in": 3600, + "refresh_token": "refresh-token", + "scope": "atproto", + "sub": "{{TokenSub}}" + } + """); + } + + if (request.Method == HttpMethod.Post && uri.ToString() == $"{Issuer}/oauth/revoke") + { + RevocationRequests++; + LastRevocationBody = await request.Content!.ReadAsStringAsync(cancellationToken); + return new HttpResponseMessage(HttpStatusCode.OK); + } + + return new HttpResponseMessage(HttpStatusCode.NotFound); + } + + private static HttpResponseMessage Json(string json) => new(HttpStatusCode.OK) + { + Content = new StringContent(json, Encoding.UTF8, "application/json"), + }; + } +} diff --git a/tests/CarpaNet.UnitTests/OAuth/OAuthClientPipelineTests.cs b/tests/CarpaNet.UnitTests/OAuth/OAuthClientPipelineTests.cs new file mode 100644 index 0000000..2b1448b --- /dev/null +++ b/tests/CarpaNet.UnitTests/OAuth/OAuthClientPipelineTests.cs @@ -0,0 +1,343 @@ +using System; +using System.Collections.Generic; +using System.IO; +using System.Linq; +using System.Net; +using System.Net.Http; +using System.Text; +using System.Text.Json; +using System.Threading.Tasks; +using CarpaNet.Auth; +using CarpaNet.Blob; +using CarpaNet.Identity; +using CarpaNet.OAuth; +using CarpaNet.OAuth.Crypto; +using CarpaNet.OAuth.Storage; +using CarpaNet.UnitTests.Http; +using Xunit; + +namespace CarpaNet.UnitTests.OAuth; + +/// +/// Tests for the OAuth client's request pipeline: DPoP credentials only for the session's PDS, +/// nonce retries, blob upload through DPoP, and session invalidation on a rejected refresh. +/// +public class OAuthClientPipelineTests +{ + private const string Did = "did:plc:oauthuser"; + private const string Pds = "https://pds.oauth.example"; + private const string Issuer = "https://auth.oauth.example"; + private const string TokenEndpoint = "https://auth.oauth.example/oauth/token"; + + [Fact] + public async Task Get_OwnPds_SendsDPoPProofAndToken() + { + var handler = new RecordingHandler(_ => XrpcRequestPipelineTests.Json("{\"value\":1}")); + using var client = await CreateClientAsync(handler); + + await client.GetAsync("com.example.get"); + + var request = Assert.Single(handler.Requests); + Assert.Equal("DPoP access-1", request.Header("Authorization")); + var proof = DecodePayload(request.Header("DPoP")!); + Assert.Equal("GET", proof.GetProperty("htm").GetString()); + Assert.StartsWith(Pds + "/xrpc/com.example.get", proof.GetProperty("htu").GetString()); + } + + [Fact] + public async Task ServiceUrl_SendsNoDPoPCredentials() + { + var handler = new RecordingHandler(_ => XrpcRequestPipelineTests.Json("{\"value\":1}")); + using var client = await CreateClientAsync(handler); + + await client.WithServiceUrl(new Uri("https://video.example")).GetAsync("app.bsky.video.getJobStatus"); + + var request = Assert.Single(handler.Requests); + Assert.Null(request.Header("Authorization")); + Assert.Null(request.Header("DPoP")); + } + + [Fact] + public async Task UseDPoPNonceChallenge_IsRetriedOnceWithTheNewNonce() + { + var calls = 0; + var handler = new RecordingHandler(_ => + { + if (++calls == 1) + { + var challenge = XrpcRequestPipelineTests.Json("{\"error\":\"use_dpop_nonce\"}", HttpStatusCode.Unauthorized); + challenge.Headers.TryAddWithoutValidation("WWW-Authenticate", "DPoP error=\"use_dpop_nonce\", error_description=\"Resource server requires nonce in DPoP proof\""); + challenge.Headers.TryAddWithoutValidation("DPoP-Nonce", "nonce-123"); + return challenge; + } + + return XrpcRequestPipelineTests.Json("{\"value\":1}"); + }); + using var client = await CreateClientAsync(handler); + + await client.GetAsync("com.example.get"); + + Assert.Equal(2, handler.Requests.Count); + Assert.False(DecodePayload(handler.Requests[0].Header("DPoP")!).TryGetProperty("nonce", out _)); + Assert.Equal("nonce-123", DecodePayload(handler.Requests[1].Header("DPoP")!).GetProperty("nonce").GetString()); + Assert.DoesNotContain(handler.Requests, r => r.Uri.ToString() == TokenEndpoint); + } + + [Fact] + public async Task UploadBlob_UsesDPoP() + { + var handler = new RecordingHandler(_ => XrpcRequestPipelineTests.Json( + "{\"blob\":{\"$type\":\"blob\",\"ref\":{\"$link\":\"bafkreibme22gw2h7y2h7tg2fhqotaqjucnbc24deqo72b6mkl2egezxhvy\"},\"mimeType\":\"image/jpeg\",\"size\":3}}")); + using var client = await CreateClientAsync(handler); + + await client.UploadBlobAsync(new MemoryStream(new byte[] { 1, 2, 3 }), "image/jpeg"); + + var request = Assert.Single(handler.Requests); + Assert.Equal("DPoP access-1", request.Header("Authorization")); + Assert.Equal("POST", DecodePayload(request.Header("DPoP")!).GetProperty("htm").GetString()); + Assert.Equal(new byte[] { 1, 2, 3 }, request.Body); + } + + [Fact] + public async Task SetLabelerDids_ChangesHeader() + { + var handler = new RecordingHandler(_ => XrpcRequestPipelineTests.Json("{\"value\":1}")); + using var client = await CreateClientAsync(handler); + + client.SetLabelerDids(new[] { AcceptLabelersHeader.Redact("did:plc:mod") }); + await client.GetAsync("com.example.get"); + + Assert.Equal("did:plc:mod;redact", Assert.Single(handler.Requests).Header("atproto-accept-labelers")); + } + + [Fact] + public async Task RejectedRefresh_RaisesSessionInvalidated() + { + var handler = new RecordingHandler(_ => XrpcRequestPipelineTests.Json( + "{\"error\":\"invalid_grant\",\"error_description\":\"refresh token revoked\"}", HttpStatusCode.BadRequest)); + var provider = await CreateProviderAsync(handler); + var events = new List(); + provider.SessionInvalidated += (_, e) => events.Add(e); + + await Assert.ThrowsAsync(() => provider.RefreshAsync()); + + var raised = Assert.Single(events); + Assert.Equal(Did, raised.Did); + Assert.Equal("invalid_grant", raised.Reason); + Assert.False(provider.HasValidToken); + await Assert.ThrowsAsync(() => provider.RefreshAsync()); + Assert.Single(events); + } + + [Fact] + public async Task TemporaryRefreshFailure_DoesNotInvalidate() + { + var handler = new RecordingHandler(_ => XrpcRequestPipelineTests.Json( + "{\"error\":\"server_error\"}", HttpStatusCode.InternalServerError)); + var provider = await CreateProviderAsync(handler); + var raised = false; + provider.SessionInvalidated += (_, _) => raised = true; + + await Assert.ThrowsAsync(() => provider.RefreshAsync()); + + Assert.False(raised); + Assert.NotNull(provider.RefreshToken); + } + + [Fact] + public async Task Refresh_IssuerStillAuthoritative_RefreshesAndKeepsAudienceWithoutSlash() + { + var handler = new RecordingHandler(r => IdentityAndTokenResponses(r, Issuer)); + var provider = await CreateProviderAsync(handler, withResolver: true); + + await provider.RefreshAsync(); + + Assert.Equal("access-2", provider.AccessToken); + Assert.Equal(new Uri(Pds), provider.PdsUrl); + Assert.Contains(handler.Requests, r => r.Uri.ToString() == TokenEndpoint); + Assert.Contains(handler.Requests, r => r.Uri.AbsolutePath == "/.well-known/oauth-protected-resource"); + } + + [Fact] + public async Task Refresh_IssuerNoLongerAuthoritative_InvalidatesWithoutTokenRequest() + { + var handler = new RecordingHandler(r => IdentityAndTokenResponses(r, "https://other-auth.example")); + var provider = await CreateProviderAsync(handler, withResolver: true); + var events = new List(); + provider.SessionInvalidated += (_, e) => events.Add(e); + + await Assert.ThrowsAsync(() => provider.RefreshAsync()); + + Assert.Equal("issuer_mismatch", Assert.Single(events).Reason); + Assert.DoesNotContain(handler.Requests, r => r.Uri.ToString() == TokenEndpoint); + Assert.False(provider.HasValidToken); + } + + [Fact] + public async Task Refresh_IdentityResolutionFails_IsTemporary() + { + var handler = new RecordingHandler(r => r.Uri.Host == "plc.directory" + ? XrpcRequestPipelineTests.Status(HttpStatusCode.ServiceUnavailable) + : IdentityAndTokenResponses(r, Issuer)); + var provider = await CreateProviderAsync(handler, withResolver: true); + var raised = false; + provider.SessionInvalidated += (_, _) => raised = true; + + await Assert.ThrowsAsync(() => provider.RefreshAsync()); + + Assert.False(raised); + Assert.Equal("refresh-1", provider.RefreshToken); + Assert.DoesNotContain(handler.Requests, r => r.Uri.ToString() == TokenEndpoint); + } + + [Fact] + public async Task RateLimitedRequest_IsRetriedWithAFreshDPoPProof() + { + var calls = 0; + var inner = new RecordingHandler(_ => + { + if (++calls == 1) + { + var limited = XrpcRequestPipelineTests.Status((HttpStatusCode)429); + limited.Headers.Add("Retry-After", "0"); + return limited; + } + + return XrpcRequestPipelineTests.Json("{\"value\":1}"); + }); + using var rateLimit = new CarpaNet.Http.RateLimitHandler(inner) { AutoRetryOnRateLimit = true, MaxRetries = 3 }; + using var client = await CreateClientAsync(inner, new HttpClient(rateLimit)); + + await client.GetAsync("com.example.get"); + + Assert.Equal(2, inner.Requests.Count); + var first = DecodePayload(inner.Requests[0].Header("DPoP")!).GetProperty("jti").GetString(); + var second = DecodePayload(inner.Requests[1].Header("DPoP")!).GetProperty("jti").GetString(); + Assert.NotEqual(first, second); + Assert.Equal("DPoP access-1", inner.Requests[1].Header("Authorization")); + } + + [Fact] + public async Task RateLimitHandler_DPoPRequestWithoutPreparer_IsNotRetried() + { + var inner = new RecordingHandler(_ => + { + var limited = XrpcRequestPipelineTests.Status((HttpStatusCode)429); + limited.Headers.Add("Retry-After", "0"); + return limited; + }); + using var rateLimit = new CarpaNet.Http.RateLimitHandler(inner) { AutoRetryOnRateLimit = true, MaxRetries = 3 }; + using var http = new HttpClient(rateLimit); + using var request = new HttpRequestMessage(HttpMethod.Get, Pds + "/xrpc/com.example.get"); + request.Headers.Add("DPoP", "proof"); + + using var response = await http.SendAsync(request); + + Assert.Equal((HttpStatusCode)429, response.StatusCode); + Assert.Single(inner.Requests); + } + + [Fact] + public async Task DisposingClientOrSession_DoesNotDisposeCallerIdentityResolver() + { + var handler = new RecordingHandler(_ => XrpcRequestPipelineTests.Json("{}")); + var resolver = new IdentityResolver(); // owns its HttpClient + var session = new OAuthSession(new OAuthClientConfig + { + ClientId = "https://app.example/client-metadata.json", + RedirectUri = "https://app.example/callback", + HttpClient = new HttpClient(handler), + IdentityResolver = resolver, + }); + var client = await CreateClientAsync(handler, identityResolver: resolver); + + client.Dispose(); + session.Dispose(); + + var resolverHttp = (HttpClient)typeof(IdentityResolver) + .GetField("_httpClient", System.Reflection.BindingFlags.NonPublic | System.Reflection.BindingFlags.Instance)! + .GetValue(resolver)!; + resolverHttp.Timeout = TimeSpan.FromSeconds(5); // throws ObjectDisposedException if disposed + resolver.Dispose(); + Assert.Throws(() => resolverHttp.Timeout = TimeSpan.FromSeconds(6)); + } + + private static HttpResponseMessage IdentityAndTokenResponses(RecordedRequest request, string protectingIssuer) + { + if (request.Uri.Host == "plc.directory") + { + return XrpcRequestPipelineTests.Json( + $"{{\"id\":\"{Did}\",\"alsoKnownAs\":[\"at://oauthuser.example\"],\"service\":[{{\"id\":\"#atproto_pds\",\"type\":\"AtprotoPersonalDataServer\",\"serviceEndpoint\":\"{Pds}/\"}}]}}"); + } + + if (request.Uri.AbsolutePath == "/.well-known/oauth-protected-resource") + { + return XrpcRequestPipelineTests.Json( + $"{{\"resource\":\"{Pds}\",\"authorization_servers\":[\"{protectingIssuer}\"]}}"); + } + + if (request.Uri.ToString() == TokenEndpoint) + { + return XrpcRequestPipelineTests.Json( + $"{{\"access_token\":\"access-2\",\"token_type\":\"DPoP\",\"expires_in\":3600,\"refresh_token\":\"refresh-2\",\"scope\":\"atproto\",\"sub\":\"{Did}\"}}"); + } + + return XrpcRequestPipelineTests.Status(HttpStatusCode.NotFound); + } + + private static async Task CreateProviderAsync(RecordingHandler handler, bool withResolver = false) + { + var provider = new DPoPTokenProvider( + new HttpClient(handler), + new MemoryOAuthSessionStore(), + clientId: "https://app.example/client-metadata.json", + identityResolver: withResolver ? new IdentityResolver(new HttpClient(handler), cache: new MemoryIdentityCache()) : null); + var tokenSet = new TokenSet + { + Issuer = Issuer, + Sub = Did, + Audience = Pds, + Scope = "atproto", + AccessToken = "access-1", + RefreshToken = "refresh-1", + ExpiresAt = DateTimeOffset.UtcNow.AddHours(1), + }; + var metadata = new OAuthAuthorizationServerMetadata + { + Issuer = Issuer, + TokenEndpoint = TokenEndpoint, + AuthorizationEndpoint = Issuer + "/oauth/authorize", + }; + + await provider.SetupAsync(Did, tokenSet, DPoPKeyPair.Generate(), metadata); + return provider; + } + + private static async Task CreateClientAsync(RecordingHandler handler, HttpClient? httpClient = null, IdentityResolver? identityResolver = null) + { + var provider = await CreateProviderAsync(handler); + var session = new OAuthSession(new OAuthClientConfig + { + ClientId = "https://app.example/client-metadata.json", + RedirectUri = "https://app.example/callback", + HttpClient = new HttpClient(handler), + }); + + return new ATProtoOAuthClient( + Did, + Pds, + provider, + session, + appState: null, + identityResolver: identityResolver ?? new IdentityResolver(new HttpClient(handler), cache: new MemoryIdentityCache()), + jsonOptions: TestHelpers.CreateJsonOptions(), + httpClient: httpClient ?? new HttpClient(handler)); + } + + private static JsonElement DecodePayload(string jwt) + { + var payload = jwt.Split('.')[1].Replace('-', '+').Replace('_', '/'); + payload = payload.PadRight(payload.Length + ((4 - (payload.Length % 4)) % 4), '='); + return JsonDocument.Parse(Encoding.UTF8.GetString(Convert.FromBase64String(payload))).RootElement.Clone(); + } +} diff --git a/tests/CarpaNet.UnitTests/OAuth/Scopes/AccountPermissionTests.cs b/tests/CarpaNet.UnitTests/OAuth/Scopes/AccountPermissionTests.cs new file mode 100644 index 0000000..66fc03c --- /dev/null +++ b/tests/CarpaNet.UnitTests/OAuth/Scopes/AccountPermissionTests.cs @@ -0,0 +1,115 @@ +using System; +using CarpaNet.OAuth.Scopes; +using Xunit; + +namespace CarpaNet.UnitTests.OAuth.Scopes; + +/// +/// Ported from oauth-scopes scopes/account-permission.test.ts. +/// +public class AccountPermissionTests +{ + [Fact] + public void Parse_ValidScopes() + { + Assert.True(AccountPermission.TryParse("account:email?action=read", out var scope1)); + Assert.Equal(AccountAttribute.Email, scope1!.Attribute); + Assert.Equal(AccountActions.Read, scope1.Actions); + + Assert.True(AccountPermission.TryParse("account:repo?action=manage", out var scope2)); + Assert.Equal(AccountAttribute.Repo, scope2!.Attribute); + Assert.Equal(AccountActions.Manage, scope2.Actions); + } + + [Fact] + public void Parse_WithoutAction_DefaultsToRead() + { + var scope = AccountPermission.Parse("account:status"); + Assert.Equal(AccountAttribute.Status, scope.Attribute); + Assert.Equal(AccountActions.Read, scope.Actions); + } + + [Theory] + [InlineData("account:invalid")] + [InlineData("account:email?action=invalid")] + [InlineData("invalid:email")] + [InlineData("account")] + [InlineData("")] + [InlineData("account:")] + [InlineData("account:email?attr=repo")] + [InlineData("account:email?unknown=1")] + public void Parse_Invalid_ReturnsFalse(string scope) + { + Assert.False(AccountPermission.TryParse(scope, out var parsed)); + Assert.Null(parsed); + Assert.Throws(() => AccountPermission.Parse(scope)); + } + + [Theory] + [InlineData(AccountAttribute.Email, AccountActions.Read, "account:email")] + [InlineData(AccountAttribute.Repo, AccountActions.Read, "account:repo")] + [InlineData(AccountAttribute.Status, AccountActions.Read, "account:status")] + [InlineData(AccountAttribute.Email, AccountActions.Manage, "account:email?action=manage")] + [InlineData(AccountAttribute.Repo, AccountActions.Manage, "account:repo?action=manage")] + [InlineData(AccountAttribute.Status, AccountActions.Manage, "account:status?action=manage")] + public void ScopeNeededFor(AccountAttribute attribute, AccountActions action, string expected) + { + Assert.Equal(expected, AccountPermission.ScopeNeededFor(attribute, action)); + } + + [Fact] + public void Matches() + { + Assert.True(AccountPermission.Parse("account:email?action=read").Matches(AccountAttribute.Email, AccountActions.Read)); + Assert.True(AccountPermission.Parse("account:repo?action=manage").Matches(AccountAttribute.Repo, AccountActions.Manage)); + Assert.False(AccountPermission.Parse("account:email?action=read").Matches(AccountAttribute.Email, AccountActions.Manage)); + Assert.False(AccountPermission.Parse("account:email?action=read").Matches(AccountAttribute.Repo, AccountActions.Read)); + + var defaulted = AccountPermission.Parse("account:email"); + Assert.True(defaulted.Matches(AccountAttribute.Email, AccountActions.Read)); + Assert.False(defaulted.Matches(AccountAttribute.Email, AccountActions.Manage)); + + Assert.True(AccountPermission.Parse("account:status?action=read").Matches(AccountAttribute.Status, AccountActions.Read)); + + // "manage" implies "read" + Assert.True(AccountPermission.Parse("account:email?action=manage").Matches(AccountAttribute.Email, AccountActions.Read)); + } + + [Fact] + public void Format() + { + Assert.Equal("account:email?action=manage", new AccountPermission(AccountAttribute.Email, AccountActions.Manage).ToString()); + Assert.Equal("account:repo", new AccountPermission(AccountAttribute.Repo, AccountActions.Read).ToString()); + Assert.Equal("account:email", new AccountPermission(AccountAttribute.Email).ToString()); + Assert.Equal("account:status", new AccountPermission(AccountAttribute.Status).ToString()); + Assert.Equal( + "account:email?action=read&action=manage", + new AccountPermission(AccountAttribute.Email, AccountActions.Read | AccountActions.Manage).ToString()); + } + + [Fact] + public void Constructor_RejectsInvalidActions() + { + Assert.Throws(() => new AccountPermission(AccountAttribute.Email, AccountActions.None)); + Assert.Throws(() => new AccountPermission((AccountAttribute)42)); + } + + [Theory] + [InlineData("account:email")] + [InlineData("account:email?action=manage")] + [InlineData("account:repo")] + [InlineData("account:repo?action=manage")] + [InlineData("account:status")] + [InlineData("account:status?action=manage")] + public void RoundTrip(string scope) + { + Assert.Equal(scope, AccountPermission.Parse(scope).ToString()); + } + + [Fact] + public void QueryForm_IsNormalizedToPositional() + { + Assert.Equal("account:email?action=manage", AccountPermission.Parse("account?attr=email&action=manage").ToString()); + Assert.Equal("account:email", AccountPermission.Parse("account:email?action=read").ToString()); + } +} diff --git a/tests/CarpaNet.UnitTests/OAuth/Scopes/BlobPermissionTests.cs b/tests/CarpaNet.UnitTests/OAuth/Scopes/BlobPermissionTests.cs new file mode 100644 index 0000000..cce3204 --- /dev/null +++ b/tests/CarpaNet.UnitTests/OAuth/Scopes/BlobPermissionTests.cs @@ -0,0 +1,98 @@ +using System; +using CarpaNet.OAuth.Scopes; +using Xunit; + +namespace CarpaNet.UnitTests.OAuth.Scopes; + +/// +/// Ported from oauth-scopes scopes/blob-permission.test.ts. +/// +public class BlobPermissionTests +{ + [Fact] + public void Parse_Positional() + { + Assert.Equal(new[] { "image/png" }, BlobPermission.Parse("blob:image/png").Accept); + } + + [Fact] + public void Parse_MultipleAccept() + { + Assert.Equal( + new[] { "image/png", "image/jpeg" }, + BlobPermission.Parse("blob?accept=image/png&accept=image/jpeg").Accept); + } + + [Theory] + [InlineData("blob")] + [InlineData("invalid")] + [InlineData("scope")] + [InlineData("blob:invalid")] + [InlineData("blob?accept=invalid-mime")] + [InlineData("blob?accept=invalid")] + [InlineData("blob:*/**")] + [InlineData("blob:*/png")] + [InlineData("blob:image/png?accept=image/jpeg")] + [InlineData("blob:image/png?mime=image/jpeg")] + public void Parse_Invalid_ReturnsFalse(string scope) + { + Assert.False(BlobPermission.TryParse(scope, out var parsed)); + Assert.Null(parsed); + } + + [Fact] + public void ScopeNeededFor() + { + Assert.Equal("blob:image/png", BlobPermission.ScopeNeededFor("image/png")); + Assert.Equal("blob:application/json", BlobPermission.ScopeNeededFor("application/json")); + } + + [Fact] + public void Matches() + { + Assert.True(BlobPermission.Parse("blob:image/png").Matches("image/png")); + Assert.False(BlobPermission.Parse("blob:image/png").Matches("image/jpeg")); + + var any = BlobPermission.Parse("blob:*/*"); + Assert.True(any.Matches("image/jpeg")); + Assert.True(any.Matches("application/json")); + + Assert.True(BlobPermission.Parse("blob:image/*").Matches("image/gif")); + + var multiple = BlobPermission.Parse("blob?accept=image/png&accept=image/jpeg"); + Assert.True(multiple.Matches("image/png")); + Assert.True(multiple.Matches("image/jpeg")); + Assert.False(multiple.Matches("image/gif")); + } + + [Fact] + public void Format_MultipleAccept_UsesSortedQuery() + { + Assert.Equal("blob?accept=image/jpeg&accept=image/png", new BlobPermission("image/png", "image/jpeg").ToString()); + } + + [Fact] + public void Format_StripsRedundantAccept() + { + Assert.Equal("blob:*/*", new BlobPermission("*/*", "image/*").ToString()); + Assert.Equal("blob:*/*", new BlobPermission("*/*", "image/png").ToString()); + Assert.Equal("blob:image/*", new BlobPermission("image/*", "image/png").ToString()); + } + + [Fact] + public void Format_SingleAccept_UsesPositional() + { + Assert.Equal("blob:image/png", new BlobPermission("image/png").ToString()); + Assert.Equal("blob:image/*", new BlobPermission("image/*").ToString()); + Assert.Equal("blob:*/*", new BlobPermission("*/*").ToString()); + Assert.Equal("blob:image/png", new BlobPermission("IMAGE/PNG").ToString()); + } + + [Fact] + public void Constructor_RejectsInvalidAccept() + { + Assert.Throws(() => new BlobPermission()); + Assert.Throws(() => new BlobPermission("image")); + Assert.Throws(() => new BlobPermission("*/png")); + } +} diff --git a/tests/CarpaNet.UnitTests/OAuth/Scopes/IdentityPermissionTests.cs b/tests/CarpaNet.UnitTests/OAuth/Scopes/IdentityPermissionTests.cs new file mode 100644 index 0000000..ff45d4b --- /dev/null +++ b/tests/CarpaNet.UnitTests/OAuth/Scopes/IdentityPermissionTests.cs @@ -0,0 +1,60 @@ +using CarpaNet.OAuth.Scopes; +using Xunit; + +namespace CarpaNet.UnitTests.OAuth.Scopes; + +/// +/// Ported from oauth-scopes scopes/identity-permission.test.ts. +/// +public class IdentityPermissionTests +{ + [Fact] + public void Parse_Positional() + { + Assert.Equal(IdentityAttribute.Handle, IdentityPermission.Parse("identity:handle").Attribute); + Assert.Equal(IdentityAttribute.All, IdentityPermission.Parse("identity:*").Attribute); + Assert.Equal(IdentityAttribute.Handle, IdentityPermission.Parse("identity?attr=handle").Attribute); + } + + [Theory] + [InlineData("invalid")] + [InlineData("identity:invalid")] + [InlineData("identity:*?action=*")] + [InlineData("identity:*?action=manage")] + [InlineData("identity:*?action=submit")] + [InlineData("identity:handle?action=invalid")] + [InlineData("identity?attribute=invalid&action=invalid")] + [InlineData("identity")] + [InlineData("identity:handle?attr=handle")] + public void Parse_Invalid_ReturnsFalse(string scope) + { + Assert.False(IdentityPermission.TryParse(scope, out var parsed)); + Assert.Null(parsed); + } + + [Fact] + public void ScopeNeededFor() + { + Assert.Equal("identity:handle", IdentityPermission.ScopeNeededFor(IdentityAttribute.Handle)); + Assert.Equal("identity:*", IdentityPermission.ScopeNeededFor(IdentityAttribute.All)); + } + + [Fact] + public void Matches() + { + var handle = IdentityPermission.Parse("identity:handle"); + Assert.True(handle.Matches(IdentityAttribute.Handle)); + Assert.False(handle.Matches(IdentityAttribute.All)); + + var all = IdentityPermission.Parse("identity:*"); + Assert.True(all.Matches(IdentityAttribute.All)); + Assert.True(all.Matches(IdentityAttribute.Handle)); + } + + [Fact] + public void Format() + { + Assert.Equal("identity:handle", new IdentityPermission(IdentityAttribute.Handle).ToString()); + Assert.Equal("identity:*", new IdentityPermission(IdentityAttribute.All).ToString()); + } +} diff --git a/tests/CarpaNet.UnitTests/OAuth/Scopes/IncludeScopeTests.cs b/tests/CarpaNet.UnitTests/OAuth/Scopes/IncludeScopeTests.cs new file mode 100644 index 0000000..c298015 --- /dev/null +++ b/tests/CarpaNet.UnitTests/OAuth/Scopes/IncludeScopeTests.cs @@ -0,0 +1,93 @@ +using System; +using CarpaNet.OAuth.Scopes; +using Xunit; + +namespace CarpaNet.UnitTests.OAuth.Scopes; + +/// +/// Ported from oauth-scopes scopes/include-scope.test.ts (parsing, formatting and +/// isParentAuthorityOf; permission-set expansion is not implemented). +/// +public class IncludeScopeTests +{ + [Theory] + [InlineData("include:com.example.bar", "com.example.bar", null)] + [InlineData("include:com.example.baz?aud=did:web:example.com%23my_service", "com.example.baz", "did:web:example.com#my_service")] + [InlineData("include:com.example.baz?aud=did:web:example.com#my_service", "com.example.baz", "did:web:example.com#my_service")] + [InlineData("include?nsid=com.example.baz", "com.example.baz", null)] + [InlineData("include?aud=did:web:example.com%23my_service&nsid=com.example.baz", "com.example.baz", "did:web:example.com#my_service")] + public void Parse_Valid(string scope, string nsid, string? aud) + { + var include = IncludeScope.Parse(scope); + Assert.Equal(nsid, include.Nsid); + Assert.Equal(aud, include.Aud); + } + + [Theory] + [InlineData("")] + [InlineData("repo:com.example.baz")] + [InlineData("include")] + [InlineData("include#")] + // Invalid NSID + [InlineData("include:")] + [InlineData("include:#")] + [InlineData("include:&")] + [InlineData("include:com..example")] + [InlineData("include:com")] + [InlineData("include:com.example")] + [InlineData("include:9com.example.foo")] + [InlineData("include:com.example.-bar")] + [InlineData("include:invalid^nsid")] + [InlineData("include:nsid")] + // Invalid AUD + [InlineData("include:com.example.baz?aud=")] + [InlineData("include:com.example.baz?aud=did:web:example.com")] + [InlineData("include:com.example.baz?aud=invalid^did")] + // Duplicate or unknown params + [InlineData("include:com.example.baz?nsid=com.example.baz")] + [InlineData("include:com.example.baz?aud=did:web:a.com%23x&aud=did:web:b.com%23x")] + [InlineData("include:com.example.baz?lxm=com.example.baz")] + public void Parse_Invalid_ReturnsFalse(string scope) + { + Assert.False(IncludeScope.TryParse(scope, out var parsed)); + Assert.Null(parsed); + } + + [Fact] + public void Format() + { + Assert.Equal("include:com.example.foo", new IncludeScope("com.example.foo").ToString()); + Assert.Equal( + "include:com.example.foo?aud=did:web:example.com%23my_service", + new IncludeScope("com.example.foo", "did:web:example.com#my_service").ToString()); + Assert.Equal( + "include:com.example.baz?aud=did:web:example.com%23my_service", + IncludeScope.Parse("include?aud=did:web:example.com%23my_service&nsid=com.example.baz").ToString()); + } + + [Fact] + public void Constructor_RejectsInvalidInput() + { + Assert.Throws(() => new IncludeScope("nsid")); + Assert.Throws(() => new IncludeScope("com.example.foo", "did:web:example.com")); + } + + [Theory] + [InlineData("com.example.foo.identifier", true)] + [InlineData("com.example.foo.bar.baz", true)] + [InlineData("com.example.foo.bar.baz.quz", true)] + [InlineData("com", false)] + [InlineData("com.example", false)] + [InlineData("com.example.bar", false)] + [InlineData("com.example.bar.foo", false)] + [InlineData("com.example.bar.qux", false)] + [InlineData("com.atproto.foo", false)] + [InlineData("com.atproto.foo.auth", false)] + [InlineData("com.atproto.foo.bar", false)] + [InlineData("*", false)] + public void IsParentAuthorityOf(string nsid, bool expected) + { + var scope = new IncludeScope("com.example.foo.auth"); + Assert.Equal(expected, scope.IsParentAuthorityOf(nsid)); + } +} diff --git a/tests/CarpaNet.UnitTests/OAuth/Scopes/MimeTests.cs b/tests/CarpaNet.UnitTests/OAuth/Scopes/MimeTests.cs new file mode 100644 index 0000000..3fa8b3a --- /dev/null +++ b/tests/CarpaNet.UnitTests/OAuth/Scopes/MimeTests.cs @@ -0,0 +1,75 @@ +using CarpaNet.OAuth.Scopes; +using Xunit; + +namespace CarpaNet.UnitTests.OAuth.Scopes; + +/// +/// Ported from oauth-scopes lib/mime.test.ts. +/// +public class MimeTests +{ + [Theory] + [InlineData("image/png", true)] + [InlineData("application/json", true)] + [InlineData("text/html", true)] + [InlineData("image/*", true)] + [InlineData("*/*", true)] + [InlineData("image//png", false)] + [InlineData("/png", false)] + [InlineData("image/", false)] + [InlineData("image/**", false)] + [InlineData("*/png", false)] + [InlineData("*", false)] + [InlineData("image/png/extra", false)] + public void IsAccept(string value, bool expected) + { + Assert.Equal(expected, BlobPermission.IsAccept(value)); + } + + [Theory] + [InlineData("image/png", true)] + [InlineData("application/json", true)] + [InlineData("image/*", false)] + [InlineData("*/*", false)] + [InlineData("image/png/extra", false)] + [InlineData("*/mime", false)] + [InlineData("/png", false)] + [InlineData("image/", false)] + [InlineData("image", false)] + [InlineData("image/ png", false)] + [InlineData("image//png", false)] + public void IsMime(string value, bool expected) + { + Assert.Equal(expected, BlobPermission.IsMime(value)); + } + + [Theory] + [InlineData("image/png", "image/png", true)] + [InlineData("image/*", "image/jpeg", true)] + [InlineData("image/*", "image/gif", true)] + [InlineData("image/png", "image/jpeg", false)] + [InlineData("image/*", "text/html", false)] + [InlineData("*/*", "application/json", true)] + [InlineData("image/png", "*/mime", false)] + [InlineData("image/png", "image", false)] + [InlineData("image/*", "image//png", false)] + [InlineData("image/*", "image/ png", false)] + [InlineData("*/*", "image/", false)] + [InlineData("*/*", "/mime", false)] + public void MatchesAccept(string accept, string mime, bool expected) + { + Assert.Equal(expected, BlobPermission.MatchesAccept(accept, mime)); + } + + [Fact] + public void MatchesAnyAccept() + { + var accepts = new[] { "image/png", "application/json" }; + Assert.True(BlobPermission.MatchesAnyAccept(accepts, "image/png")); + Assert.True(BlobPermission.MatchesAnyAccept(accepts, "application/json")); + Assert.False(BlobPermission.MatchesAnyAccept(accepts, "text/html")); + Assert.False(BlobPermission.MatchesAnyAccept(System.Array.Empty(), "image/png")); + Assert.True(BlobPermission.MatchesAnyAccept(new[] { "image/*" }, "image/jpeg")); + Assert.False(BlobPermission.MatchesAnyAccept(new[] { "image/*" }, "text/html")); + } +} diff --git a/tests/CarpaNet.UnitTests/OAuth/Scopes/RepoPermissionTests.cs b/tests/CarpaNet.UnitTests/OAuth/Scopes/RepoPermissionTests.cs new file mode 100644 index 0000000..362f51b --- /dev/null +++ b/tests/CarpaNet.UnitTests/OAuth/Scopes/RepoPermissionTests.cs @@ -0,0 +1,142 @@ +using System; +using CarpaNet.OAuth.Scopes; +using Xunit; + +namespace CarpaNet.UnitTests.OAuth.Scopes; + +/// +/// Ported from oauth-scopes scopes/repo-permission.test.ts. +/// +public class RepoPermissionTests +{ + [Fact] + public void Parse_Positional_DefaultsToAllActions() + { + var scope = RepoPermission.Parse("repo:com.example.foo"); + Assert.Equal(new[] { "com.example.foo" }, scope.Collections); + Assert.Equal(RepoActions.All, scope.Actions); + } + + [Fact] + public void Parse_MultipleActions() + { + var scope = RepoPermission.Parse("repo:com.example.foo?action=create&action=update"); + Assert.Equal(new[] { "com.example.foo" }, scope.Collections); + Assert.Equal(RepoActions.Create | RepoActions.Update, scope.Actions); + } + + [Fact] + public void Parse_WildcardCollection_WithAction() + { + var scope = RepoPermission.Parse("repo:*?action=create"); + Assert.Equal(new[] { "*" }, scope.Collections); + Assert.Equal(RepoActions.Create, scope.Actions); + Assert.True(scope.Matches("any.collection", RepoActions.Create)); + Assert.False(scope.Matches("any.collection", RepoActions.Update)); + } + + [Fact] + public void Parse_WildcardCollection_WithoutActions() + { + var scope = RepoPermission.Parse("repo:*"); + Assert.Equal(new[] { "*" }, scope.Collections); + Assert.Equal(RepoActions.All, scope.Actions); + Assert.True(scope.Matches("any.collection", RepoActions.Create)); + Assert.True(scope.Matches("any.collection", RepoActions.Update)); + Assert.True(scope.Matches("any.collection", RepoActions.Delete)); + } + + [Theory] + [InlineData("repo:foo bar")] + [InlineData("repo:.foo")] + [InlineData("repo:bar.")] + [InlineData("repo:com.example.foo?action=invalid")] + [InlineData("invalid")] + [InlineData("scope")] + [InlineData("repo:*?action=*")] + [InlineData("repo:invalid")] + [InlineData("repo?collection=invalid&action=invalid")] + [InlineData("repo")] + [InlineData("repo:")] + [InlineData("repo:com.example.foo?collection=com.example.bar")] + [InlineData("repo:com.example.foo?unknown=x")] + public void Parse_Invalid_ReturnsFalse(string scope) + { + Assert.False(RepoPermission.TryParse(scope, out var parsed)); + Assert.Null(parsed); + } + + [Fact] + public void ScopeNeededFor() + { + Assert.Equal("repo:com.example.foo?action=create", RepoPermission.ScopeNeededFor("com.example.foo", RepoActions.Create)); + Assert.Equal("repo:*?action=create", RepoPermission.ScopeNeededFor("*", RepoActions.Create)); + + // scopeNeededFor assumes valid input and does not validate + Assert.Equal("repo:invalid?action=create", RepoPermission.ScopeNeededFor("invalid", RepoActions.Create)); + } + + [Fact] + public void Matches() + { + var create = RepoPermission.Parse("repo:com.example.foo?action=create"); + Assert.True(create.Matches("com.example.foo", RepoActions.Create)); + Assert.False(create.Matches("com.example.foo", RepoActions.Update)); + + var wildcard = RepoPermission.Parse("repo:*?action=create"); + Assert.True(wildcard.Matches("com.example.bar", RepoActions.Create)); + Assert.False(wildcard.Matches("com.example.bar", RepoActions.Delete)); + + var multiple = RepoPermission.Parse("repo:com.example.foo?action=create&action=update"); + Assert.True(multiple.Matches("com.example.foo", RepoActions.Create)); + Assert.True(multiple.Matches("com.example.foo", RepoActions.Update)); + Assert.False(multiple.Matches("com.example.foo", RepoActions.Delete)); + + var defaulted = RepoPermission.Parse("repo:com.example.foo"); + Assert.True(defaulted.Matches("com.example.foo", RepoActions.Create)); + Assert.True(defaulted.Matches("com.example.foo", RepoActions.Update)); + Assert.True(defaulted.Matches("com.example.foo", RepoActions.Delete)); + Assert.False(defaulted.Matches("com.example.bar", RepoActions.Create)); + } + + [Fact] + public void Format() + { + Assert.Equal( + "repo:com.example.foo?action=create&action=update", + new RepoPermission("com.example.foo", RepoActions.Create | RepoActions.Update).ToString()); + Assert.Equal("repo:com.example.foo", new RepoPermission("com.example.foo").ToString()); + } + + [Fact] + public void Constructor_RejectsInvalidInput() + { + Assert.Throws(() => new RepoPermission("invalid")); + Assert.Throws(() => new RepoPermission("com.example.foo", RepoActions.None)); + Assert.Throws(() => new RepoPermission(Array.Empty())); + } + + [Theory] + [InlineData("repo:com.example.foo", "repo:com.example.foo")] + [InlineData("repo:com.example.foo?action=create", "repo:com.example.foo?action=create")] + [InlineData("repo:com.example.foo?action=create&action=update", "repo:com.example.foo?action=create&action=update")] + [InlineData("repo:*?action=create&action=update&action=delete", "repo:*")] + [InlineData("repo:com.example.foo?action=create&action=update&action=delete", "repo:com.example.foo")] + [InlineData("repo:*?action=create", "repo:*?action=create")] + [InlineData("repo:*?action=update", "repo:*?action=update")] + [InlineData("repo?collection=*&action=update", "repo:*?action=update")] + [InlineData("repo?collection=*&collection=com.example.foo&action=update", "repo:*?action=update")] + [InlineData("repo?collection=*", "repo:*")] + [InlineData("repo?collection=*&action=create&action=update&action=delete", "repo:*")] + [InlineData("repo?collection=*&collection=com.example.foo", "repo:*")] + [InlineData("repo?action=create&collection=com.example.foo", "repo:com.example.foo?action=create")] + [InlineData("repo?collection=com.example.foo&action=create&action=update&action=delete", "repo:com.example.foo")] + [InlineData( + "repo?action=create&collection=com.example.foo&collection=com.example.bar", + "repo?collection=com.example.bar&collection=com.example.foo&action=create")] + [InlineData("repo:com.example.foo?action=delete&action=create", "repo:com.example.foo?action=create&action=delete")] + public void Reformat(string input, string expected) + { + Assert.Equal(expected, RepoPermission.Parse(input).ToString()); + } +} diff --git a/tests/CarpaNet.UnitTests/OAuth/Scopes/RpcPermissionTests.cs b/tests/CarpaNet.UnitTests/OAuth/Scopes/RpcPermissionTests.cs new file mode 100644 index 0000000..3dc078f --- /dev/null +++ b/tests/CarpaNet.UnitTests/OAuth/Scopes/RpcPermissionTests.cs @@ -0,0 +1,164 @@ +using System; +using CarpaNet.OAuth.Scopes; +using Xunit; + +namespace CarpaNet.UnitTests.OAuth.Scopes; + +/// +/// Ported from oauth-scopes scopes/rpc-permission.test.ts. +/// +public class RpcPermissionTests +{ + [Fact] + public void Parse_Positional() + { + var scope = RpcPermission.Parse("rpc:com.example.service?aud=did:web:example.com%23service_id"); + Assert.Equal("did:web:example.com#service_id", scope.Aud); + Assert.Equal(new[] { "com.example.service" }, scope.Lxm); + } + + [Fact] + public void Parse_QueryAndPositionalForms() + { + var named = RpcPermission.Parse("rpc?lxm=com.example.method1&aud=*"); + Assert.Equal("*", named.Aud); + Assert.Equal(new[] { "com.example.method1" }, named.Lxm); + + var positional = RpcPermission.Parse("rpc:com.example.method1?aud=*"); + Assert.Equal("*", positional.Aud); + Assert.Equal(new[] { "com.example.method1" }, positional.Lxm); + } + + [Fact] + public void Parse_MultipleLxm() + { + var scope = RpcPermission.Parse("rpc?aud=*&lxm=com.example.method1&lxm=com.example.method2"); + Assert.Equal("*", scope.Aud); + Assert.Equal(new[] { "com.example.method1", "com.example.method2" }, scope.Lxm); + } + + [Theory] + // Missing lxm + [InlineData("rpc?aud=did:web:example.com%23service_id")] + [InlineData("rpc:?aud=did:web:example.com%23service_id")] + [InlineData("rpc?aud=did:web:example.com")] + // Missing aud + [InlineData("rpc?lxm=com.example.method1")] + [InlineData("rpc:com.example.method1")] + [InlineData("rpc:com.example.service")] + // lxm in both positional and query form + [InlineData("rpc:com.example.method1?aud=did:web:example.com&lxm=com.example.method2")] + // Any aud and any lxm + [InlineData("rpc?aud=*&lxm=*")] + [InlineData("rpc:*?aud=*")] + // Invalid aud / lxm + [InlineData("rpc:com.example.service?aud=invalid")] + [InlineData("rpc:invalid")] + [InlineData("rpc?lxm=invalid")] + [InlineData("rpc:*")] + [InlineData("invalid")] + [InlineData("rpc:invalid?aud=did:web:example.com")] + [InlineData("rpc:invalid?aud=did:web:example.com%23service_id")] + [InlineData("rpc:foo.bar")] + [InlineData("rpc:com.example.service?aud=did:web:example.com%23service_id&invalid=param")] + [InlineData("rpc:foo.bar.baz?aud=did:web")] + [InlineData("rpc:foo.bar.baz?aud=did:web%23service_id")] + [InlineData("rpc:foo.bar.baz?aud=did:plc:111")] + [InlineData("rpc:foo.bar.baz?aud=did:plc:111%23service_id")] + [InlineData("rpc:foo.bar.baz?aud=did:foo:bar")] + [InlineData("rpc:foo.bar.baz?aud=did:foo:bar%23service_id")] + [InlineData("rpc:foo.bar.baz?aud=did:web:example.com%23service_id&lxm=foo.bar.baz")] + [InlineData("rpc:foo.bar.baz?aud=invalid")] + [InlineData("notrpc:com.example.service?aud=did:web:example.com%23service_id")] + [InlineData("rpc?lxm=invalid&aud=invalid")] + // Extra atproto DID rules: aud with a path, a port, an empty fragment, two auds + [InlineData("rpc:foo.bar.baz?aud=did:web:example.com:path%23svc")] + [InlineData("rpc:foo.bar.baz?aud=did:web:example.com%253A8080%23svc")] + [InlineData("rpc:foo.bar.baz?aud=did:web:example.com%23")] + [InlineData("rpc:foo.bar.baz?aud=*&aud=*")] + public void Parse_Invalid_ReturnsFalse(string scope) + { + Assert.False(RpcPermission.TryParse(scope, out var parsed)); + Assert.Null(parsed); + } + + [Fact] + public void Parse_ValidAtprotoDidAudiences() + { + Assert.True(RpcPermission.TryParse("rpc:foo.bar.baz?aud=did:plc:abcdefghijklmnopqrstuvwx%23svc", out _)); + Assert.True(RpcPermission.TryParse("rpc:foo.bar.baz?aud=did:web:localhost%253A2583%23svc", out _)); + Assert.True(RpcPermission.TryParse("rpc:app.bsky.feed.getTimeline?aud=did:web:api.bsky.app%23bsky_appview", out _)); + } + + [Fact] + public void ScopeNeededFor() + { + Assert.Equal( + "rpc:com.example.service?aud=did:web:example.com%23service_id", + RpcPermission.ScopeNeededFor("com.example.service", "did:web:example.com#service_id")); + Assert.Equal("rpc:com.example.method1?aud=*", RpcPermission.ScopeNeededFor("com.example.method1", "*")); + } + + [Fact] + public void Matches() + { + var exact = RpcPermission.Parse("rpc:com.example.service?aud=did:web:example.com%23service_id"); + Assert.True(exact.Matches("com.example.service", "did:web:example.com#service_id")); + Assert.False(exact.Matches("com.example.OtherService", "did:web:example.com#service_id")); + Assert.False(exact.Matches("com.example.service", "did:example:456#service_id")); + + Assert.True(RpcPermission.Parse("rpc:com.example.method1?aud=*") + .Matches("com.example.method1", "did:web:example.com#service_id")); + + var anyLxm = RpcPermission.Parse("rpc:*?aud=did:web:example.com%23service_id"); + Assert.True(anyLxm.Matches("com.example.method1", "did:web:example.com#service_id")); + Assert.True(anyLxm.Matches("com.example.anyMethod", "did:web:example.com#service_id")); + } + + [Fact] + public void Format() + { + Assert.Equal( + "rpc:com.example.service?aud=did:web:example.com%23service_id", + new RpcPermission("did:web:example.com#service_id", "com.example.service").ToString()); + Assert.Equal("rpc:com.example.method1?aud=*", new RpcPermission("*", "com.example.method1").ToString()); + Assert.Equal( + "rpc?lxm=com.example.method1&lxm=com.example.method2&aud=did:web:example.com%23service_id", + new RpcPermission("did:web:example.com#service_id", new[] { "com.example.method1", "com.example.method2" }).ToString()); + Assert.Equal( + "rpc:*?aud=did:web:example.com%23service_id", + new RpcPermission("did:web:example.com#service_id", "*").ToString()); + + // Simplifies lxm if one of them is "*" + Assert.Equal( + "rpc:*?aud=did:web:example.com%23service_id", + new RpcPermission("did:web:example.com#service_id", new[] { "*", "com.example.method1" }).ToString()); + } + + [Fact] + public void Constructor_RejectsInvalidInput() + { + Assert.Throws(() => new RpcPermission("*", "*")); + Assert.Throws(() => new RpcPermission("did:web:example.com", "com.example.foo")); + Assert.Throws(() => new RpcPermission("*", "invalid")); + } + + [Theory] + [InlineData("rpc:com.example.service?aud=did:web:example.com%23service_id", "rpc:com.example.service?aud=did:web:example.com%23service_id")] + [InlineData("rpc:com.example.service?aud=did:web:example.com#service_id", "rpc:com.example.service?aud=did:web:example.com%23service_id")] + [InlineData("rpc?lxm=com.example.method1&lxm=com.example.method2&aud=*", "rpc?lxm=com.example.method1&lxm=com.example.method2&aud=*")] + [InlineData( + "rpc?lxm=com.example.method1&lxm=com.example.method2&lxm=*&aud=did:web:example.com%23service_id", + "rpc:*?aud=did:web:example.com%23service_id")] + [InlineData("rpc?aud=did:web:example.com%23foo&lxm=com.example.service", "rpc:com.example.service?aud=did:web:example.com%23foo")] + [InlineData("rpc?lxm=com.example.method1&aud=did:web:example.com#foo", "rpc:com.example.method1?aud=did:web:example.com%23foo")] + [InlineData("rpc?lxm=com.example.method1&aud=did:web:example.com%23bar", "rpc:com.example.method1?aud=did:web:example.com%23bar")] + [InlineData("rpc:com.example.method1?&aud=*", "rpc:com.example.method1?aud=*")] + [InlineData( + "rpc?lxm=com.example.b&lxm=com.example.a&lxm=com.example.b&aud=*", + "rpc?lxm=com.example.a&lxm=com.example.b&aud=*")] + public void Reformat(string input, string expected) + { + Assert.Equal(expected, RpcPermission.Parse(input).ToString()); + } +} diff --git a/tests/CarpaNet.UnitTests/OAuth/Scopes/ScopeSetTests.cs b/tests/CarpaNet.UnitTests/OAuth/Scopes/ScopeSetTests.cs new file mode 100644 index 0000000..e00719a --- /dev/null +++ b/tests/CarpaNet.UnitTests/OAuth/Scopes/ScopeSetTests.cs @@ -0,0 +1,118 @@ +using System; +using CarpaNet.OAuth; +using CarpaNet.OAuth.Scopes; +using Xunit; + +namespace CarpaNet.UnitTests.OAuth.Scopes; + +/// +/// Ported from oauth-scopes scopes-set.test.ts, plus builder tests. +/// +public class ScopeSetTests +{ + [Fact] + public void NewSet_IsEmpty() + { + var set = new ScopeSet(); + Assert.Empty(set); + Assert.Equal(string.Empty, set.ToString()); + } + + [Fact] + public void Add_And_Remove() + { + var set = new ScopeSet().Add("repo:read"); + Assert.Single(set); + Assert.True(set.Contains("repo:read")); + Assert.False(set.Contains("repo:write")); + + Assert.True(set.Remove("repo:read")); + Assert.Empty(set); + Assert.False(set.Contains("repo:read")); + Assert.False(set.Remove("repo:read")); + } + + [Fact] + public void MatchesRepo() + { + var set = new ScopeSet(new[] { "repo:com.example.foo" }); + Assert.True(set.MatchesRepo("com.example.foo", RepoActions.Create)); + Assert.False(set.MatchesRepo("com.example.bar", RepoActions.Create)); + + var createOnly = new ScopeSet(new[] { "repo:com.example.foo?action=create" }); + Assert.False(createOnly.MatchesRepo("com.example.foo", RepoActions.Delete)); + + var invalid = new ScopeSet(new[] { "repo:not-a-valid-nsid" }); + Assert.False(invalid.MatchesRepo("not-a-valid-nsid", RepoActions.Create)); + } + + [Fact] + public void Matches_OtherResources() + { + var set = ScopeSet.Parse( + "atproto rpc:app.bsky.actor.getProfile?aud=did:web:api.bsky.app%23bsky_appview blob:image/* account:email identity:handle"); + + Assert.True(set.MatchesRpc("app.bsky.actor.getProfile", "did:web:api.bsky.app#bsky_appview")); + Assert.False(set.MatchesRpc("app.bsky.actor.getProfiles", "did:web:api.bsky.app#bsky_appview")); + Assert.True(set.MatchesBlob("image/png")); + Assert.False(set.MatchesBlob("video/mp4")); + Assert.True(set.MatchesAccount(AccountAttribute.Email, AccountActions.Read)); + Assert.False(set.MatchesAccount(AccountAttribute.Email, AccountActions.Manage)); + Assert.True(set.MatchesIdentity(IdentityAttribute.Handle)); + Assert.False(set.MatchesIdentity(IdentityAttribute.All)); + } + + [Fact] + public void Builder_ProducesSpaceSeparatedString() + { + var set = new ScopeSet() + .AddAtproto() + .AddRepo("app.bsky.feed.post", RepoActions.Create) + .AddRepo("app.bsky.feed.like") + .AddBlob("image/*", "video/mp4") + .AddRpc("app.bsky.actor.getProfile", "did:web:api.bsky.app#bsky_appview") + .AddAccount(AccountAttribute.Email) + .AddIdentity(IdentityAttribute.Handle) + .AddInclude("com.example.authBasic", "did:web:example.com#svc") + .AddTransitionGeneric() + .AddAtproto(); // duplicate is ignored + + Assert.Equal( + "atproto repo:app.bsky.feed.post?action=create repo:app.bsky.feed.like blob?accept=image/*&accept=video/mp4 " + + "rpc:app.bsky.actor.getProfile?aud=did:web:api.bsky.app%23bsky_appview account:email identity:handle " + + "include:com.example.authBasic?aud=did:web:example.com%23svc transition:generic", + set.ToString()); + + foreach (var value in set) + { + Assert.True(AtprotoScope.IsValid(value), value); + } + } + + [Fact] + public void Parse_KeepsUnknownValues_AndNormalizedStringDropsThem() + { + var set = ScopeSet.Parse("transition:generic atproto future:thing"); + Assert.Equal(3, set.Count); + Assert.True(set.Contains("future:thing")); + Assert.Equal("transition:generic atproto future:thing", set.ToString()); + Assert.Equal("atproto transition:generic", set.ToNormalizedString()); + } + + [Fact] + public void Add_RejectsWhitespaceAndEmpty() + { + Assert.Throws(() => new ScopeSet().Add("atproto transition:generic")); + Assert.Throws(() => new ScopeSet().Add(string.Empty)); + } + + [Fact] + public void OAuthClientConfig_SetScope() + { + var config = new OAuthClientConfig() + .SetScope(new ScopeSet().AddAtproto().AddTransitionGeneric()); + Assert.Equal("atproto transition:generic", config.Scope); + + Assert.Throws(() => new OAuthClientConfig().SetScope(new ScopeSet().AddTransitionGeneric())); + } +} diff --git a/tests/CarpaNet.UnitTests/OAuth/Scopes/ScopeSyntaxTests.cs b/tests/CarpaNet.UnitTests/OAuth/Scopes/ScopeSyntaxTests.cs new file mode 100644 index 0000000..8a100c8 --- /dev/null +++ b/tests/CarpaNet.UnitTests/OAuth/Scopes/ScopeSyntaxTests.cs @@ -0,0 +1,111 @@ +using CarpaNet.OAuth.Scopes; +using Xunit; + +namespace CarpaNet.UnitTests.OAuth.Scopes; + +/// +/// Ported from oauth-scopes lib/syntax.test.ts, lib/syntax-string.test.ts and atproto-oauth-scope.ts. +/// +public class ScopeSyntaxTests +{ + [Theory] + [InlineData("prefix", "prefix", true)] + [InlineData("prefix", "differentResource", false)] + [InlineData("prefix:positional", "prefix", true)] + [InlineData("differentResource:positional", "prefix", false)] + [InlineData("prefix?param=value", "prefix", true)] + [InlineData("prefix", "prefi", false)] + [InlineData("prefix:pos", "prefi", false)] + [InlineData("prefix?param=value", "prefi", false)] + [InlineData("prefix", "fix", false)] + [InlineData("prefix:pos", "fix", false)] + [InlineData("prefix?param=value", "fix", false)] + [InlineData("differentResource?param=value", "prefix", false)] + public void IsScopeStringFor(string value, string prefix, bool expected) + { + Assert.Equal(expected, AtprotoScope.IsScopeStringFor(value, prefix)); + } + + [Theory] + [InlineData("atproto")] + [InlineData("transition:generic")] + [InlineData("transition:chat.bsky")] + [InlineData("transition:email")] + public void StaticScopes_AreValid(string scope) + { + Assert.True(AtprotoScope.IsStaticScope(scope)); + Assert.True(AtprotoScope.IsValid(scope)); + Assert.Equal(scope, AtprotoScope.NormalizeValue(scope)); + } + + [Theory] + [InlineData("account:email")] + [InlineData("blob:image/png")] + [InlineData("identity:handle")] + [InlineData("include:com.example.foo")] + [InlineData("repo:com.example.foo")] + [InlineData("rpc:com.example.foo?aud=*")] + public void PermissionScopes_AreValid(string scope) + { + Assert.True(AtprotoScope.IsValid(scope)); + Assert.True(AtprotoScope.TryParse(scope, out var parsed)); + Assert.Equal(scope, parsed!.ToString()); + } + + [Theory] + [InlineData("")] + [InlineData("transition:unknown")] + [InlineData("atproto2")] + [InlineData("unknown:foo")] + [InlineData("repo:invalid")] + [InlineData("rpc:com.example.foo")] + [InlineData("include:com.example.foo?aud=did:web:example.com")] + public void InvalidScopes_AreNotValid(string scope) + { + Assert.False(AtprotoScope.IsValid(scope)); + Assert.Null(AtprotoScope.NormalizeValue(scope)); + } + + [Fact] + public void Normalize_NormalizesSortsDropsInvalidAndDuplicates() + { + var normalized = AtprotoScope.Normalize( + "repo:com.example.foo?action=create&action=update&action=delete atproto bogus " + + "rpc?lxm=com.example.b&lxm=com.example.a&aud=did:web:example.com#svc atproto blob?accept=image/png&accept=image/*"); + + Assert.Equal( + "atproto blob:image/* repo:com.example.foo rpc?lxm=com.example.a&lxm=com.example.b&aud=did:web:example.com%23svc", + normalized); + } + + [Fact] + public void Positional_UrlEncoding_IsDecoded() + { + // "my-res:my%20pos" => "my pos"; checked through a permission whose positional value matters + Assert.True(RepoPermission.TryParse("repo:com.example.f%6Fo", out var repo)); + Assert.Equal("com.example.foo", repo!.Collections[0]); + } + + [Fact] + public void Named_UrlEncoding_IsDecoded() + { + Assert.True(RpcPermission.TryParse("rpc:com.example.foo?aud=did%3Aweb%3Aexample.com%23svc", out var rpc)); + Assert.Equal("did:web:example.com#svc", rpc!.Aud); + } + + [Fact] + public void Positional_MalformedEncoding_IsInvalid() + { + Assert.False(RepoPermission.TryParse("repo:com.example.foo%zz", out _)); + Assert.False(RepoPermission.TryParse("repo:com.example.foo%2", out _)); + } + + [Fact] + public void Query_QuestionMarkInsideParameterValue_IsPartOfValue() + { + // "rpc:foo.bar?aud=did:foo:bar?lxm=bar.baz" => aud is "did:foo:bar?lxm=bar.baz" + Assert.True(RpcPermission.TryParse("rpc:foo.bar.baz?aud=did:web:example.com%23svc?lxm=x", out var rpc)); + Assert.Equal(new[] { "foo.bar.baz" }, rpc!.Lxm); + Assert.Equal("did:web:example.com#svc?lxm=x", rpc!.Aud); + } +} diff --git a/tests/CarpaNet.UnitTests/TypeRegistryTests.cs b/tests/CarpaNet.UnitTests/TypeRegistryTests.cs index cd0b9c6..433c222 100644 --- a/tests/CarpaNet.UnitTests/TypeRegistryTests.cs +++ b/tests/CarpaNet.UnitTests/TypeRegistryTests.cs @@ -202,8 +202,10 @@ public void ResolveToCSharpType_ResolvesUnionAsInterface() registry.RegisterDocument(doc); + // Qualified with the union's namespace: the ref crosses namespaces (app.bsky.feed -> app.bsky), + // and an unqualified interface name would not compile there. var result = registry.ResolveToCSharpType("app.bsky.embed#embedUnion", "app.bsky.feed.post"); - Assert.Equal("IEmbedEmbedUnion", result); + Assert.Equal("AppBsky.IEmbedEmbedUnion", result); } [Fact]