diff --git a/docs/concepts/authentication/oauth-proxy.md b/docs/concepts/authentication/oauth-proxy.md new file mode 100644 index 000000000..0434ed133 --- /dev/null +++ b/docs/concepts/authentication/oauth-proxy.md @@ -0,0 +1,158 @@ +--- +title: OAuth Proxy +description: Bridge MCP clients to OAuth providers that do not support Dynamic Client Registration. +uid: oauth-proxy +--- + +# OAuth proxy + +MCP clients normally discover an authorization server and register themselves dynamically. Many +enterprise identity providers require applications to be registered ahead of time instead. The +ASP.NET Core OAuth proxy bridges these models by exposing an MCP-compatible authorization server +while using one fixed upstream application registration. + +The proxy: + +- accepts Dynamic Client Registration from MCP clients; +- keeps each client's callback URI and PKCE challenge separate from the upstream flow; +- uses an independent PKCE verifier and a fixed callback URI with the upstream provider; +- stores upstream tokens encrypted on the server; +- returns only application-minted, short-lived access tokens to MCP clients; and +- rotates opaque proxy refresh handles without exposing upstream refresh tokens. + +## Configure the proxy + +Register an `IMcpOAuthProxyStore`, configure the proxy, and map its endpoints. The issuer path must +match the path passed to `MapMcpOAuthProxy`. + +```csharp +var proxyIssuer = new Uri("https://mcp.example.com/oauth"); +var tokenIssuer = new ProxyTokenIssuer(/* signing credentials */); + +// The application store must atomically consume one-time records. For example, a Redis-backed +// implementation can use GETDEL. See Storage requirements below. +builder.Services.AddSingleton(); +builder.Services.AddRateLimiter(options => +{ + options.AddFixedWindowLimiter("oauth-proxy", limiter => + { + limiter.PermitLimit = 60; + limiter.Window = TimeSpan.FromMinutes(1); + limiter.QueueLimit = 0; + }); +}); + +builder.Services.AddMcpOAuthProxy(options => +{ + options.Issuer = proxyIssuer; + + // The proxy discovers authorization and token endpoints from this issuer. + options.UpstreamIssuer = new Uri( + "https://login.microsoftonline.com/contoso.onmicrosoft.com/v2.0"); + options.UpstreamClientId = builder.Configuration["OAuth:ClientId"]!; + options.UpstreamClientSecret = builder.Configuration["OAuth:ClientSecret"]!; + options.UpstreamRedirectUri = new Uri("https://mcp.example.com/oauth/callback"); + + options.AllowedScopes.Add("openid"); + options.AllowedScopes.Add("profile"); + options.AllowedScopes.Add("offline_access"); + options.AllowedScopes.Add("api://mcp-server/mcp:tools"); + options.AllowedResources.Add("https://mcp.example.com"); + + options.ClientRegistrationValidator = (context, cancellationToken) => + context.HttpContext.RequestServices + .GetRequiredService() + .ValidateAsync(context, cancellationToken); + options.AuthorizationValidator = (context, cancellationToken) => + context.HttpContext.RequestServices + .GetRequiredService() + .ValidateAsync(context, cancellationToken); + + options.TokenFactory = tokenIssuer.IssueAsync; + options.JwksUri = new Uri("https://mcp.example.com/.well-known/jwks.json"); +}); + +var app = builder.Build(); +app.MapMcpOAuthProxy("/oauth").RequireRateLimiting("oauth-proxy"); +``` + +`ProxyTokenIssuer` is application code. Its `IssueAsync` method receives an + and returns an +. It should validate +the upstream identity, mint a new JWT for the proxy issuer and MCP resource, and return the JWT's +subject, audience, and unique token identifier for the structured audit event. + +The proxy rejects token-factory results that: + +- equal an upstream access, refresh, or ID token; +- are not Bearer tokens; +- have empty subject, audience, or token identifiers; or +- exceed `MaximumAccessTokenLifetime`. + +Configure the MCP server's JWT Bearer handler to validate the proxy-issued token, not the upstream +provider token. The MCP protected-resource metadata should advertise `proxyIssuer` as its +authorization server. + +Both policy callbacks are required. `ClientRegistrationValidator` should verify an initial access +token, trusted redirect URI, or equivalent registration policy. `AuthorizationValidator` should +enforce user and client approval before the proxy uses its upstream application registration. The +OAuth endpoints allow anonymous HTTP access so they continue to work with ASP.NET Core fallback +authorization policies; the callbacks are therefore the application security boundary. + +Every accepted `resource` value must appear exactly in `AllowedResources`. This prevents a dynamic +client from choosing an arbitrary audience for the proxy-issued token. The proxy also constrains +minted scopes to the set returned by the upstream authorization server. + +## Explicit upstream endpoints + +For an upstream provider without OpenID Connect discovery, configure both endpoints explicitly: + +```csharp +options.UpstreamAuthorizationEndpoint = new Uri("https://idp.example.com/authorize"); +options.UpstreamTokenEndpoint = new Uri("https://idp.example.com/token"); +``` + +The proxy supports `client_secret_basic`, `client_secret_post`, and public upstream clients through +`UpstreamClientAuthenticationMethod`. Provider-specific parameters can be supplied through +`AdditionalAuthorizationParameters` and `AdditionalTokenParameters`; standard OAuth parameters +cannot be overridden. + +## Storage requirements + +The proxy protects every stored record with ASP.NET Core Data Protection before passing it to +. Persist the Data +Protection key ring and restrict access to both the key ring and the cache. Losing or deleting keys +that still protect active records invalidates registrations and refresh handles. Store keys contain +only one-way digests, not bearer credentials or client callback URIs. + +`IDistributedCache` does not define an atomic get-and-delete operation. The provided + prevents +replay within one process, but it is not registered automatically because it cannot provide atomic +consumption across multiple application instances. A single-instance application can register it +explicitly. Multi-instance deployments must provide an `IMcpOAuthProxyStore` whose `TakeAsync` +operation is atomic in the backing store. + +Do not use an in-memory distributed cache in production. Registrations and refresh mappings must +survive application restarts, and every instance must share the same protected records. + +Opaque refresh handles rotate after every successful use. Reuse of a consumed handle revokes the +current descendant handle for that token family. Retryable upstream failures restore the handle for +only the remainder of its original absolute lifetime. + +## Security requirements + +- Serve the issuer, callback, and MCP resource over HTTPS. HTTP is accepted only for loopback + development addresses. +- Register the exact fixed `UpstreamRedirectUri` with the upstream provider. +- Keep upstream client credentials and Data Protection keys outside source control. +- Keep `CookieSecurePolicy.Always`, the default, outside loopback development. +- Apply request-rate limits to the public registration, authorization, and token endpoints. +- Keep both proxy policy callbacks default-deny and audit their approval decisions. +- Treat token-mint audit records as security data. They contain subject, audience, scopes, client + ID, token ID, and mint time, but never token values. +- Keep access-token lifetimes short and configure refresh-token and registration lifetimes for the + application's revocation requirements. + +The proxy requires S256 PKCE on both the client-to-proxy and proxy-to-upstream legs. Authorization +callbacks are additionally bound to the browser that initiated the transaction through a +short-lived, HTTP-only cookie. diff --git a/docs/concepts/toc.yml b/docs/concepts/toc.yml index 276415e13..bd4bc4522 100644 --- a/docs/concepts/toc.yml +++ b/docs/concepts/toc.yml @@ -33,6 +33,8 @@ items: uid: elicitation - name: Server Features items: + - name: OAuth Proxy + uid: oauth-proxy - name: Tools uid: tools - name: Resources diff --git a/src/ModelContextProtocol.AspNetCore/Authentication/DistributedCacheMcpOAuthProxyStore.cs b/src/ModelContextProtocol.AspNetCore/Authentication/DistributedCacheMcpOAuthProxyStore.cs new file mode 100644 index 000000000..7b50f8611 --- /dev/null +++ b/src/ModelContextProtocol.AspNetCore/Authentication/DistributedCacheMcpOAuthProxyStore.cs @@ -0,0 +1,71 @@ +using Microsoft.Extensions.Caching.Distributed; +namespace ModelContextProtocol.AspNetCore.Authentication; + +/// +/// Stores protected OAuth proxy records in an . +/// +/// +/// does not expose an atomic get-and-remove operation. This +/// implementation serializes calls within one process. Multi-instance +/// deployments must replace it with a store whose +/// operation is atomic in the backing store. +/// +public sealed class DistributedCacheMcpOAuthProxyStore(IDistributedCache cache) : IMcpOAuthProxyStore +{ + private readonly SemaphoreSlim[] _locks = Enumerable.Range(0, 64) + .Select(static _ => new SemaphoreSlim(1, 1)) + .ToArray(); + + /// + public async ValueTask SetAsync(string key, ReadOnlyMemory value, TimeSpan lifetime, CancellationToken cancellationToken = default) + { + ArgumentException.ThrowIfNullOrEmpty(key); + if (lifetime <= TimeSpan.Zero) + { + throw new ArgumentOutOfRangeException(nameof(lifetime)); + } + + await cache.SetAsync( + key, + value.ToArray(), + new DistributedCacheEntryOptions { AbsoluteExpirationRelativeToNow = lifetime }, + cancellationToken).ConfigureAwait(false); + } + + /// + public async ValueTask?> GetAsync(string key, CancellationToken cancellationToken = default) + { + ArgumentException.ThrowIfNullOrEmpty(key); + return await cache.GetAsync(key, cancellationToken).ConfigureAwait(false); + } + + /// + public async ValueTask?> TakeAsync(string key, CancellationToken cancellationToken = default) + { + ArgumentException.ThrowIfNullOrEmpty(key); + var gate = _locks[(StringComparer.Ordinal.GetHashCode(key) & int.MaxValue) % _locks.Length]; + await gate.WaitAsync(cancellationToken).ConfigureAwait(false); + + try + { + var value = await cache.GetAsync(key, cancellationToken).ConfigureAwait(false); + if (value is not null) + { + await cache.RemoveAsync(key, cancellationToken).ConfigureAwait(false); + } + + return value; + } + finally + { + gate.Release(); + } + } + + /// + public ValueTask RemoveAsync(string key, CancellationToken cancellationToken = default) + { + ArgumentException.ThrowIfNullOrEmpty(key); + return new(cache.RemoveAsync(key, cancellationToken)); + } +} diff --git a/src/ModelContextProtocol.AspNetCore/Authentication/IMcpOAuthProxyStore.cs b/src/ModelContextProtocol.AspNetCore/Authentication/IMcpOAuthProxyStore.cs new file mode 100644 index 000000000..e59e56334 --- /dev/null +++ b/src/ModelContextProtocol.AspNetCore/Authentication/IMcpOAuthProxyStore.cs @@ -0,0 +1,31 @@ +namespace ModelContextProtocol.AspNetCore.Authentication; + +/// +/// Stores protected OAuth proxy records. +/// +/// +/// Values supplied to this interface are already encrypted and authenticated. Implementations used +/// by multiple application instances must make atomic across those instances. +/// +public interface IMcpOAuthProxyStore +{ + /// + /// Stores a protected record. + /// + ValueTask SetAsync(string key, ReadOnlyMemory value, TimeSpan lifetime, CancellationToken cancellationToken = default); + + /// + /// Gets a protected record without consuming it. + /// + ValueTask?> GetAsync(string key, CancellationToken cancellationToken = default); + + /// + /// Atomically gets and removes a protected record. + /// + ValueTask?> TakeAsync(string key, CancellationToken cancellationToken = default); + + /// + /// Removes a protected record. + /// + ValueTask RemoveAsync(string key, CancellationToken cancellationToken = default); +} diff --git a/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyAuthorizationContext.cs b/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyAuthorizationContext.cs new file mode 100644 index 000000000..d4edc6701 --- /dev/null +++ b/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyAuthorizationContext.cs @@ -0,0 +1,39 @@ +using Microsoft.AspNetCore.Http; + +namespace ModelContextProtocol.AspNetCore.Authentication; + +/// +/// Describes a validated proxy authorization request for application policy evaluation. +/// +public sealed class McpOAuthProxyAuthorizationContext +{ + /// + /// Gets the HTTP context for the authorization request. + /// + public required HttpContext HttpContext { get; init; } + + /// + /// Gets the registered dynamic client identifier. + /// + public required string ClientId { get; init; } + + /// + /// Gets the registered client name, when present. + /// + public string? ClientName { get; init; } + + /// + /// Gets the exact registered redirect URI selected by the client. + /// + public required string RedirectUri { get; init; } + + /// + /// Gets the requested scopes. + /// + public required IReadOnlyList Scopes { get; init; } + + /// + /// Gets the validated resource indicator, when present. + /// + public string? Resource { get; init; } +} diff --git a/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyClientRegistrationContext.cs b/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyClientRegistrationContext.cs new file mode 100644 index 000000000..0e96423e5 --- /dev/null +++ b/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyClientRegistrationContext.cs @@ -0,0 +1,44 @@ +using Microsoft.AspNetCore.Http; + +namespace ModelContextProtocol.AspNetCore.Authentication; + +/// +/// Describes a validated dynamic client registration for application policy evaluation. +/// +public sealed class McpOAuthProxyClientRegistrationContext +{ + /// + /// Gets the HTTP request that submitted the registration. + /// + public required HttpContext HttpContext { get; init; } + + /// + /// Gets the requested redirect URIs. + /// + public required IReadOnlyList RedirectUris { get; init; } + + /// + /// Gets the requested scopes. + /// + public required IReadOnlyList Scopes { get; init; } + + /// + /// Gets the normalized grant types requested by the dynamic client. + /// + public required IReadOnlyList GrantTypes { get; init; } + + /// + /// Gets the optional client name. + /// + public string? ClientName { get; init; } + + /// + /// Gets the optional client information URI. + /// + public string? ClientUri { get; init; } + + /// + /// Gets the requested OIDC application type, when present. + /// + public string? ApplicationType { get; init; } +} diff --git a/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyEndpointRouteBuilderExtensions.cs b/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyEndpointRouteBuilderExtensions.cs new file mode 100644 index 000000000..5b20a2137 --- /dev/null +++ b/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyEndpointRouteBuilderExtensions.cs @@ -0,0 +1,60 @@ +using Microsoft.AspNetCore.Http; +using Microsoft.AspNetCore.Routing; +using Microsoft.Extensions.DependencyInjection; +using Microsoft.Extensions.Options; +using ModelContextProtocol.AspNetCore.Authentication; + +namespace Microsoft.AspNetCore.Builder; + +/// +/// Extension methods for mapping MCP OAuth proxy endpoints. +/// +public static class McpOAuthProxyEndpointRouteBuilderExtensions +{ + /// + /// Maps OAuth proxy discovery, registration, authorization, callback, and token endpoints. + /// + /// The endpoint route builder. + /// The route prefix for operational OAuth endpoints. + /// The operational endpoint route group. + public static RouteGroupBuilder MapMcpOAuthProxy(this IEndpointRouteBuilder endpoints, string pattern = "/oauth") + { + ArgumentNullException.ThrowIfNull(endpoints); + ArgumentException.ThrowIfNullOrEmpty(pattern); + if (!pattern.StartsWith('/') || pattern.Contains('{') || pattern.Contains('}')) + { + throw new ArgumentException("The OAuth proxy pattern must be a static absolute path.", nameof(pattern)); + } + + pattern = pattern.TrimEnd('/'); + var services = endpoints.ServiceProvider; + var options = services.GetRequiredService>().Value; + var service = services.GetRequiredService(); + if (!string.Equals(options.Issuer.AbsolutePath.TrimEnd('/'), pattern, StringComparison.Ordinal)) + { + throw new InvalidOperationException("McpOAuthProxyOptions.Issuer path must match the MapMcpOAuthProxy route pattern."); + } + + var issuerPath = pattern.Trim('/'); + var metadataPath = string.IsNullOrEmpty(issuerPath) + ? "/.well-known/oauth-authorization-server" + : $"/.well-known/oauth-authorization-server/{issuerPath}"; + RequestDelegate metadataHandler = context => service.HandleMetadataRequest().ExecuteAsync(context); + endpoints.MapGet(metadataPath, metadataHandler).AllowAnonymous(); + + var group = endpoints.MapGroup(pattern).AllowAnonymous(); + group.MapGet("/.well-known/openid-configuration", metadataHandler); + group.MapPost("/register", ToRequestDelegate(service.HandleRegistrationAsync)); + group.MapGet("/authorize", ToRequestDelegate(service.HandleAuthorizationAsync)); + group.MapGet("/callback", ToRequestDelegate(service.HandleCallbackAsync)); + group.MapPost("/token", ToRequestDelegate(service.HandleTokenAsync)); + return group; + } + + private static RequestDelegate ToRequestDelegate(Func> handler) => + async context => + { + var result = await handler(context).ConfigureAwait(false); + await result.ExecuteAsync(context).ConfigureAwait(false); + }; +} diff --git a/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyModels.cs b/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyModels.cs new file mode 100644 index 000000000..ce5cf747e --- /dev/null +++ b/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyModels.cs @@ -0,0 +1,248 @@ +using System.Text.Json.Serialization; + +namespace ModelContextProtocol.AspNetCore.Authentication; + +internal sealed class OAuthProxyClientRegistrationRequest +{ + [JsonPropertyName("redirect_uris")] + public string[]? RedirectUris { get; set; } + + [JsonPropertyName("token_endpoint_auth_method")] + public string? TokenEndpointAuthMethod { get; set; } + + [JsonPropertyName("grant_types")] + public string[]? GrantTypes { get; set; } + + [JsonPropertyName("response_types")] + public string[]? ResponseTypes { get; set; } + + [JsonPropertyName("client_name")] + public string? ClientName { get; set; } + + [JsonPropertyName("client_uri")] + public string? ClientUri { get; set; } + + [JsonPropertyName("scope")] + public string? Scope { get; set; } + + [JsonPropertyName("application_type")] + public string? ApplicationType { get; set; } +} + +internal sealed class OAuthProxyClientRegistrationResponse +{ + [JsonPropertyName("client_id")] + public required string ClientId { get; init; } + + [JsonPropertyName("client_id_issued_at")] + public required long ClientIdIssuedAt { get; init; } + + [JsonPropertyName("redirect_uris")] + public required string[] RedirectUris { get; init; } + + [JsonPropertyName("token_endpoint_auth_method")] + public string TokenEndpointAuthMethod { get; init; } = "none"; + + [JsonPropertyName("grant_types")] + public string[] GrantTypes { get; init; } = ["authorization_code", "refresh_token"]; + + [JsonPropertyName("response_types")] + public string[] ResponseTypes { get; init; } = ["code"]; + + [JsonPropertyName("client_name")] + [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)] + public string? ClientName { get; init; } + + [JsonPropertyName("client_uri")] + [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)] + public string? ClientUri { get; init; } + + [JsonPropertyName("scope")] + [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)] + public string? Scope { get; init; } + + [JsonPropertyName("application_type")] + [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)] + public string? ApplicationType { get; init; } +} + +internal sealed class OAuthProxyAuthorizationServerMetadata +{ + [JsonPropertyName("issuer")] + public required string Issuer { get; init; } + + [JsonPropertyName("authorization_endpoint")] + public required string AuthorizationEndpoint { get; init; } + + [JsonPropertyName("token_endpoint")] + public required string TokenEndpoint { get; init; } + + [JsonPropertyName("registration_endpoint")] + public required string RegistrationEndpoint { get; init; } + + [JsonPropertyName("jwks_uri")] + [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)] + public string? JwksUri { get; init; } + + [JsonPropertyName("scopes_supported")] + public required string[] ScopesSupported { get; init; } + + [JsonPropertyName("response_types_supported")] + public string[] ResponseTypesSupported { get; init; } = ["code"]; + + [JsonPropertyName("grant_types_supported")] + public string[] GrantTypesSupported { get; init; } = ["authorization_code", "refresh_token"]; + + [JsonPropertyName("token_endpoint_auth_methods_supported")] + public string[] TokenEndpointAuthMethodsSupported { get; init; } = ["none"]; + + [JsonPropertyName("code_challenge_methods_supported")] + public string[] CodeChallengeMethodsSupported { get; init; } = ["S256"]; + + [JsonPropertyName("authorization_response_iss_parameter_supported")] + public bool AuthorizationResponseIssuerParameterSupported { get; init; } = true; +} + +internal sealed class OAuthProxyErrorResponse +{ + [JsonPropertyName("error")] + public required string Error { get; init; } + + [JsonPropertyName("error_description")] + [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)] + public string? ErrorDescription { get; init; } +} + +internal sealed class OAuthProxyTokenResponse +{ + [JsonPropertyName("access_token")] + public required string AccessToken { get; init; } + + [JsonPropertyName("token_type")] + public required string TokenType { get; init; } + + [JsonPropertyName("expires_in")] + public required long ExpiresIn { get; init; } + + [JsonPropertyName("refresh_token")] + [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)] + public string? RefreshToken { get; init; } + + [JsonPropertyName("scope")] + [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)] + public string? Scope { get; init; } +} + +internal sealed class OAuthProxyUpstreamMetadata +{ + [JsonPropertyName("issuer")] + public string? Issuer { get; set; } + + [JsonPropertyName("authorization_endpoint")] + public string? AuthorizationEndpoint { get; set; } + + [JsonPropertyName("token_endpoint")] + public string? TokenEndpoint { get; set; } + + [JsonPropertyName("code_challenge_methods_supported")] + public string[]? CodeChallengeMethodsSupported { get; set; } +} + +internal sealed class OAuthProxyUpstreamTokenResponse +{ + [JsonPropertyName("access_token")] + public string? AccessToken { get; set; } + + [JsonPropertyName("refresh_token")] + public string? RefreshToken { get; set; } + + [JsonPropertyName("id_token")] + public string? IdToken { get; set; } + + [JsonPropertyName("token_type")] + public string? TokenType { get; set; } + + [JsonPropertyName("expires_in")] + public long? ExpiresIn { get; set; } + + [JsonPropertyName("scope")] + public string? Scope { get; set; } + + [JsonPropertyName("error")] + public string? Error { get; set; } + + [JsonPropertyName("error_description")] + public string? ErrorDescription { get; set; } +} + +internal sealed class OAuthProxyClientRecord +{ + public required string ClientId { get; set; } + public required string[] RedirectUris { get; set; } + public required string[] Scopes { get; set; } + public string? ClientName { get; set; } + public string? ClientUri { get; set; } + public string? ApplicationType { get; set; } + public bool SupportsRefreshTokens { get; set; } +} + +internal sealed class OAuthProxyAuthorizationTransaction +{ + public required string ClientId { get; set; } + public required string RedirectUri { get; set; } + public required string CodeChallenge { get; set; } + public required string UpstreamCodeVerifier { get; set; } + public required string BrowserBinding { get; set; } + public required string[] Scopes { get; set; } + public string? ClientState { get; set; } + public string? Resource { get; set; } + public bool SupportsRefreshTokens { get; set; } +} + +internal sealed class OAuthProxyAuthorizationCode +{ + public required string ClientId { get; set; } + public required string RedirectUri { get; set; } + public required string CodeChallenge { get; set; } + public required string[] Scopes { get; set; } + public string? Resource { get; set; } + public bool SupportsRefreshTokens { get; set; } + public required OAuthProxyUpstreamTokenResponse UpstreamToken { get; set; } +} + +internal sealed class OAuthProxyRefreshRecord +{ + public required string ClientId { get; set; } + public required string[] Scopes { get; set; } + public string? Resource { get; set; } + public required OAuthProxyUpstreamTokenResponse UpstreamToken { get; set; } + public required string FamilyId { get; set; } + public required DateTimeOffset ExpiresAt { get; set; } +} + +internal sealed class OAuthProxyRefreshFamilyRecord +{ + public required string CurrentRefreshToken { get; set; } + public required DateTimeOffset ExpiresAt { get; set; } +} + +internal sealed class OAuthProxyConsumedRefreshRecord +{ + public required string FamilyId { get; set; } + public required DateTimeOffset ExpiresAt { get; set; } +} + +[JsonSerializable(typeof(OAuthProxyClientRegistrationRequest))] +[JsonSerializable(typeof(OAuthProxyClientRegistrationResponse))] +[JsonSerializable(typeof(OAuthProxyAuthorizationServerMetadata))] +[JsonSerializable(typeof(OAuthProxyErrorResponse))] +[JsonSerializable(typeof(OAuthProxyTokenResponse))] +[JsonSerializable(typeof(OAuthProxyUpstreamMetadata))] +[JsonSerializable(typeof(OAuthProxyUpstreamTokenResponse))] +[JsonSerializable(typeof(OAuthProxyClientRecord))] +[JsonSerializable(typeof(OAuthProxyAuthorizationTransaction))] +[JsonSerializable(typeof(OAuthProxyAuthorizationCode))] +[JsonSerializable(typeof(OAuthProxyRefreshRecord))] +[JsonSerializable(typeof(OAuthProxyRefreshFamilyRecord))] +[JsonSerializable(typeof(OAuthProxyConsumedRefreshRecord))] +internal sealed partial class McpOAuthProxyJsonContext : JsonSerializerContext; diff --git a/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyOptions.cs b/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyOptions.cs new file mode 100644 index 000000000..edffb2305 --- /dev/null +++ b/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyOptions.cs @@ -0,0 +1,178 @@ +using Microsoft.AspNetCore.Http; + +namespace ModelContextProtocol.AspNetCore.Authentication; + +/// +/// Configures an OAuth authorization-server facade for upstream providers that do not support +/// Dynamic Client Registration. +/// +public sealed class McpOAuthProxyOptions +{ + /// + /// Gets or sets the public issuer URI advertised by the proxy. + /// + public Uri Issuer { get; set; } = null!; + + /// + /// Gets or sets the upstream OpenID Connect issuer used for endpoint discovery. + /// + /// + /// Set this when and + /// are not configured explicitly. + /// + public Uri? UpstreamIssuer { get; set; } + + /// + /// Gets or sets the upstream authorization endpoint. + /// + public Uri? UpstreamAuthorizationEndpoint { get; set; } + + /// + /// Gets or sets the upstream token endpoint. + /// + public Uri? UpstreamTokenEndpoint { get; set; } + + /// + /// Gets or sets the pre-registered upstream client identifier. + /// + public string UpstreamClientId { get; set; } = string.Empty; + + /// + /// Gets or sets the pre-registered upstream client secret, when required. + /// + public string? UpstreamClientSecret { get; set; } + + /// + /// Gets or sets how the proxy authenticates at the upstream token endpoint. + /// + public McpOAuthProxyClientAuthenticationMethod UpstreamClientAuthenticationMethod { get; set; } = + McpOAuthProxyClientAuthenticationMethod.ClientSecretPost; + + /// + /// Gets or sets the fixed proxy callback URI registered with the upstream provider. + /// + public Uri UpstreamRedirectUri { get; set; } = null!; + + /// + /// Gets the complete set of scopes that dynamic clients may request. + /// + public ISet AllowedScopes { get; } = new HashSet(StringComparer.Ordinal); + + /// + /// Gets the exact resource indicator URIs that dynamic clients may request. + /// + /// + /// A requested resource is rejected unless it appears in this set. This prevents an untrusted + /// dynamic client from choosing the audience of a proxy-issued token. + /// + public ISet AllowedResources { get; } = new HashSet(StringComparer.Ordinal); + + /// + /// Gets additional parameters included in upstream authorization requests. + /// + public IDictionary AdditionalAuthorizationParameters { get; } = + new Dictionary(StringComparer.Ordinal); + + /// + /// Gets additional parameters included in upstream token requests. + /// + public IDictionary AdditionalTokenParameters { get; } = + new Dictionary(StringComparer.Ordinal); + + /// + /// Gets or sets whether a downstream RFC 8707 resource indicator is forwarded upstream. + /// + public bool ForwardResourceIndicator { get; set; } + + /// + /// Gets or sets an optional JSON Web Key Set URI advertised by the proxy. + /// + /// + /// Configure this when issues JWT access tokens that clients or + /// resource servers validate through the proxy metadata. + /// + public Uri? JwksUri { get; set; } + + /// + /// Gets or sets the callback that mints a client-facing access token. + /// + /// + /// The callback must issue a new token for the proxy issuer and protected resource. The proxy + /// rejects a result that equals any upstream access, refresh, or ID token. + /// + public Func> TokenFactory { get; set; } = null!; + + /// + /// Gets or sets the policy that approves dynamic client registrations. + /// + /// + /// This policy is required and should authenticate an initial access token, enforce a trusted + /// redirect-URI policy, or otherwise establish that the dynamic client may use the proxy's + /// upstream application registration. + /// + public Func> ClientRegistrationValidator { get; set; } = null!; + + /// + /// Gets or sets the policy that approves an authorization request before redirecting upstream. + /// + /// + /// This policy is required. Applications should use it to enforce downstream client consent, + /// authenticated-user policy, or an equivalent approval boundary for the registered client. + /// + public Func> AuthorizationValidator { get; set; } = null!; + + /// + /// Gets or sets the maximum accepted lifetime of a token returned by . + /// + public TimeSpan MaximumAccessTokenLifetime { get; set; } = TimeSpan.FromMinutes(15); + + /// + /// Gets or sets the lifetime of authorization transactions. + /// + public TimeSpan AuthorizationTransactionLifetime { get; set; } = TimeSpan.FromMinutes(10); + + /// + /// Gets or sets the lifetime of proxy authorization codes. + /// + public TimeSpan AuthorizationCodeLifetime { get; set; } = TimeSpan.FromMinutes(1); + + /// + /// Gets or sets the lifetime of dynamic client registrations. + /// + public TimeSpan ClientRegistrationLifetime { get; set; } = TimeSpan.FromDays(90); + + /// + /// Gets or sets the maximum lifetime of proxy refresh-token mappings. + /// + public TimeSpan RefreshTokenLifetime { get; set; } = TimeSpan.FromDays(30); + + /// + /// Gets or sets the secure policy for the browser-binding cookie. + /// + /// + /// Keep the default in production. may be used + /// for loopback-only development servers. + /// + public CookieSecurePolicy CookieSecurePolicy { get; set; } = CookieSecurePolicy.Always; +} + +/// +/// Specifies how the proxy authenticates to the upstream token endpoint. +/// +public enum McpOAuthProxyClientAuthenticationMethod +{ + /// + /// Sends the client credentials using HTTP Basic authentication. + /// + ClientSecretBasic, + + /// + /// Sends the client credentials in the form body. + /// + ClientSecretPost, + + /// + /// Does not send a client secret. + /// + None, +} diff --git a/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyOptionsValidator.cs b/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyOptionsValidator.cs new file mode 100644 index 000000000..20a6bdf5e --- /dev/null +++ b/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyOptionsValidator.cs @@ -0,0 +1,142 @@ +using Microsoft.Extensions.Options; + +namespace ModelContextProtocol.AspNetCore.Authentication; + +internal sealed class McpOAuthProxyOptionsValidator : IValidateOptions +{ + private static readonly HashSet ReservedAuthorizationParameters = new(StringComparer.Ordinal) + { + "client_id", "redirect_uri", "response_type", "code_challenge", "code_challenge_method", + "scope", "state", "resource" + }; + + private static readonly HashSet ReservedTokenParameters = new(StringComparer.Ordinal) + { + "client_id", "client_secret", "redirect_uri", "grant_type", "code", "code_verifier", + "refresh_token", "resource", "scope" + }; + + public ValidateOptionsResult Validate(string? name, McpOAuthProxyOptions options) + { + List failures = []; + + ValidateSecureUri(options.Issuer, nameof(options.Issuer), failures); + if (options.Issuer is not null && (!string.IsNullOrEmpty(options.Issuer.Query) || !string.IsNullOrEmpty(options.Issuer.Fragment))) + { + failures.Add($"{nameof(options.Issuer)} must not contain a query or fragment."); + } + + ValidateSecureUri(options.UpstreamRedirectUri, nameof(options.UpstreamRedirectUri), failures); + + var hasExplicitAuthorizationEndpoint = options.UpstreamAuthorizationEndpoint is not null; + var hasExplicitTokenEndpoint = options.UpstreamTokenEndpoint is not null; + if (hasExplicitAuthorizationEndpoint != hasExplicitTokenEndpoint) + { + failures.Add($"{nameof(options.UpstreamAuthorizationEndpoint)} and {nameof(options.UpstreamTokenEndpoint)} must be configured together."); + } + else if (hasExplicitAuthorizationEndpoint) + { + ValidateSecureUri(options.UpstreamAuthorizationEndpoint, nameof(options.UpstreamAuthorizationEndpoint), failures); + ValidateSecureUri(options.UpstreamTokenEndpoint, nameof(options.UpstreamTokenEndpoint), failures); + } + else + { + ValidateSecureUri(options.UpstreamIssuer, nameof(options.UpstreamIssuer), failures); + } + + if (string.IsNullOrWhiteSpace(options.UpstreamClientId)) + { + failures.Add($"{nameof(options.UpstreamClientId)} is required."); + } + + if (options.UpstreamClientAuthenticationMethod is not McpOAuthProxyClientAuthenticationMethod.None && + string.IsNullOrEmpty(options.UpstreamClientSecret)) + { + failures.Add($"{nameof(options.UpstreamClientSecret)} is required for the selected upstream client authentication method."); + } + + if (options.TokenFactory is null) + { + failures.Add($"{nameof(options.TokenFactory)} is required."); + } + + if (options.ClientRegistrationValidator is null) + { + failures.Add($"{nameof(options.ClientRegistrationValidator)} is required."); + } + + if (options.AuthorizationValidator is null) + { + failures.Add($"{nameof(options.AuthorizationValidator)} is required."); + } + + ValidatePositiveLifetime(options.MaximumAccessTokenLifetime, nameof(options.MaximumAccessTokenLifetime), failures); + ValidatePositiveLifetime(options.AuthorizationTransactionLifetime, nameof(options.AuthorizationTransactionLifetime), failures); + ValidatePositiveLifetime(options.AuthorizationCodeLifetime, nameof(options.AuthorizationCodeLifetime), failures); + ValidatePositiveLifetime(options.ClientRegistrationLifetime, nameof(options.ClientRegistrationLifetime), failures); + ValidatePositiveLifetime(options.RefreshTokenLifetime, nameof(options.RefreshTokenLifetime), failures); + + foreach (var scope in options.AllowedScopes) + { + if (string.IsNullOrWhiteSpace(scope) || scope.Any(char.IsWhiteSpace)) + { + failures.Add($"{nameof(options.AllowedScopes)} contains an invalid scope value."); + break; + } + } + + foreach (var resource in options.AllowedResources) + { + if (!Uri.TryCreate(resource, UriKind.Absolute, out var resourceUri) || + !McpOAuthProxyUtilities.IsSecureEndpoint(resourceUri) || + !string.IsNullOrEmpty(resourceUri.Fragment)) + { + failures.Add($"{nameof(options.AllowedResources)} contains an invalid resource URI."); + break; + } + } + + if (options.AdditionalAuthorizationParameters.Keys.Any(ReservedAuthorizationParameters.Contains)) + { + failures.Add($"{nameof(options.AdditionalAuthorizationParameters)} cannot override standard OAuth parameters."); + } + + if (options.AdditionalTokenParameters.Keys.Any(ReservedTokenParameters.Contains)) + { + failures.Add($"{nameof(options.AdditionalTokenParameters)} cannot override standard OAuth parameters."); + } + + if (options.AdditionalAuthorizationParameters.Keys.Any(string.IsNullOrWhiteSpace) || + options.AdditionalTokenParameters.Keys.Any(string.IsNullOrWhiteSpace)) + { + failures.Add("Additional OAuth parameter names cannot be empty."); + } + + if (options.JwksUri is not null) + { + ValidateSecureUri(options.JwksUri, nameof(options.JwksUri), failures); + } + + return failures.Count == 0 ? ValidateOptionsResult.Success : ValidateOptionsResult.Fail(failures); + } + + private static void ValidateSecureUri(Uri? uri, string name, List failures) + { + if (uri is null) + { + failures.Add($"{name} is required."); + } + else if (!McpOAuthProxyUtilities.IsSecureEndpoint(uri)) + { + failures.Add($"{name} must be an absolute HTTPS URI, or an HTTP loopback URI for development."); + } + } + + private static void ValidatePositiveLifetime(TimeSpan lifetime, string name, List failures) + { + if (lifetime <= TimeSpan.Zero) + { + failures.Add($"{name} must be positive."); + } + } +} diff --git a/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyProtectedStore.cs b/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyProtectedStore.cs new file mode 100644 index 000000000..e2d5b3493 --- /dev/null +++ b/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyProtectedStore.cs @@ -0,0 +1,73 @@ +using Microsoft.AspNetCore.DataProtection; +using Microsoft.AspNetCore.WebUtilities; +using System.Security.Cryptography; +using System.Text; +using System.Text.Json; +using System.Text.Json.Serialization.Metadata; + +namespace ModelContextProtocol.AspNetCore.Authentication; + +internal sealed class McpOAuthProxyProtectedStore +{ + private const string KeyPrefix = "mcp-oauth-proxy:"; + private readonly IMcpOAuthProxyStore _store; + private readonly IDataProtector _protector; + + public McpOAuthProxyProtectedStore(IMcpOAuthProxyStore store, IDataProtectionProvider dataProtectionProvider) + { + _store = store; + _protector = dataProtectionProvider.CreateProtector("ModelContextProtocol.AspNetCore.Authentication.McpOAuthProxy", "v1"); + } + + public async ValueTask SetAsync( + string category, + string identifier, + T value, + JsonTypeInfo typeInfo, + TimeSpan lifetime, + CancellationToken cancellationToken) + { + var serialized = JsonSerializer.SerializeToUtf8Bytes(value, typeInfo); + var protectedValue = _protector.Protect(serialized); + await _store.SetAsync(CreateKey(category, identifier), protectedValue, lifetime, cancellationToken).ConfigureAwait(false); + } + + public async ValueTask GetAsync( + string category, + string identifier, + JsonTypeInfo typeInfo, + CancellationToken cancellationToken) + { + var protectedValue = await _store.GetAsync(CreateKey(category, identifier), cancellationToken).ConfigureAwait(false); + return protectedValue is null ? default : Deserialize(protectedValue.Value, typeInfo); + } + + public async ValueTask TakeAsync( + string category, + string identifier, + JsonTypeInfo typeInfo, + CancellationToken cancellationToken) + { + var protectedValue = await _store.TakeAsync(CreateKey(category, identifier), cancellationToken).ConfigureAwait(false); + return protectedValue is null ? default : Deserialize(protectedValue.Value, typeInfo); + } + + public ValueTask RemoveAsync( + string category, + string identifier, + CancellationToken cancellationToken) => + _store.RemoveAsync(CreateKey(category, identifier), cancellationToken); + + private T Deserialize(ReadOnlyMemory protectedValue, JsonTypeInfo typeInfo) + { + var serialized = _protector.Unprotect(protectedValue.ToArray()); + return JsonSerializer.Deserialize(serialized, typeInfo) ?? + throw new InvalidOperationException("The OAuth proxy store contained an invalid record."); + } + + private static string CreateKey(string category, string identifier) + { + var input = Encoding.UTF8.GetBytes($"{category}\0{identifier}"); + return $"{KeyPrefix}{category}:{WebEncoders.Base64UrlEncode(SHA256.HashData(input))}"; + } +} diff --git a/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyService.cs b/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyService.cs new file mode 100644 index 000000000..2ecb434bb --- /dev/null +++ b/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyService.cs @@ -0,0 +1,1225 @@ +using Microsoft.AspNetCore.Http; +using Microsoft.AspNetCore.WebUtilities; +using Microsoft.Extensions.DependencyInjection; +using Microsoft.Extensions.Logging; +using Microsoft.Extensions.Options; +using System.Buffers; +using System.Net; +using System.Net.Http.Headers; +using System.Text; +using System.Text.Json; + +namespace ModelContextProtocol.AspNetCore.Authentication; + +internal sealed partial class McpOAuthProxyService +{ + private const int MaxRequestBodySize = 32 * 1024; + private const string BrowserBindingCookieName = "mcp-oauth-proxy-binding"; + private readonly McpOAuthProxyOptions _options; + private readonly McpOAuthProxyProtectedStore _store; + private readonly IHttpClientFactory _httpClientFactory; + private readonly TimeProvider _timeProvider; + private readonly ILogger _logger; + private readonly SemaphoreSlim _metadataLock = new(1, 1); + private OAuthProxyUpstreamMetadata? _upstreamMetadata; + + public McpOAuthProxyService( + IOptions options, + McpOAuthProxyProtectedStore store, + IHttpClientFactory httpClientFactory, + TimeProvider timeProvider, + ILogger logger) + { + _options = options.Value; + _store = store; + _httpClientFactory = httpClientFactory; + _timeProvider = timeProvider; + _logger = logger; + } + + public IResult HandleMetadataRequest() + { + var endpointBase = _options.Issuer; + var metadata = new OAuthProxyAuthorizationServerMetadata + { + Issuer = _options.Issuer.AbsoluteUri.TrimEnd('/'), + AuthorizationEndpoint = McpOAuthProxyUtilities.AppendPath(endpointBase, "authorize").AbsoluteUri, + TokenEndpoint = McpOAuthProxyUtilities.AppendPath(endpointBase, "token").AbsoluteUri, + RegistrationEndpoint = McpOAuthProxyUtilities.AppendPath(endpointBase, "register").AbsoluteUri, + JwksUri = _options.JwksUri?.AbsoluteUri, + ScopesSupported = _options.AllowedScopes.Order(StringComparer.Ordinal).ToArray(), + }; + + return Results.Json(metadata, McpOAuthProxyJsonContext.Default.OAuthProxyAuthorizationServerMetadata); + } + + public async Task HandleRegistrationAsync(HttpContext context) + { + OAuthProxyClientRegistrationRequest? request; + try + { + var body = await ReadRequestBodyAsync(context.Request, context.RequestAborted).ConfigureAwait(false); + request = JsonSerializer.Deserialize(body, McpOAuthProxyJsonContext.Default.OAuthProxyClientRegistrationRequest); + } + catch (Exception exception) when (exception is JsonException or InvalidDataException) + { + return Error(StatusCodes.Status400BadRequest, "invalid_client_metadata", "The registration document is not valid JSON."); + } + + if (request?.RedirectUris is not { Length: > 0 } redirectUris || redirectUris.Length > 20) + { + return Error(StatusCodes.Status400BadRequest, "invalid_redirect_uri", "At least one and no more than 20 redirect URIs are required."); + } + + if (redirectUris.Any(static uri => !McpOAuthProxyUtilities.IsValidRedirectUri(uri)) || + redirectUris.Distinct(StringComparer.Ordinal).Count() != redirectUris.Length) + { + return Error(StatusCodes.Status400BadRequest, "invalid_redirect_uri", "Every redirect URI must be a unique HTTPS URI or HTTP loopback URI without user information or a fragment."); + } + + if (request.TokenEndpointAuthMethod is not null and not "none" and not "client_secret_post") + { + return Error(StatusCodes.Status400BadRequest, "invalid_client_metadata", "Only public clients are supported."); + } + + var grantTypes = request.GrantTypes ?? ["authorization_code"]; + if (!grantTypes.Contains("authorization_code", StringComparer.Ordinal) || + grantTypes.Any(static value => value is not "authorization_code" and not "refresh_token")) + { + return Error(StatusCodes.Status400BadRequest, "invalid_client_metadata", "Only authorization_code and refresh_token grants are supported."); + } + + var responseTypes = request.ResponseTypes ?? ["code"]; + if (responseTypes.Length != 1 || responseTypes[0] != "code") + { + return Error(StatusCodes.Status400BadRequest, "invalid_client_metadata", "Only the code response type is supported."); + } + + if (request.ApplicationType is not null and not "native" and not "web") + { + return Error(StatusCodes.Status400BadRequest, "invalid_client_metadata", "application_type must be 'native' or 'web'."); + } + + var scopes = McpOAuthProxyUtilities.ParseScopes(request.Scope); + if (scopes.Length > 32 || !AreScopesAllowed(scopes)) + { + return Error(StatusCodes.Status400BadRequest, "invalid_client_metadata", "The registration requests an unsupported scope."); + } + + if (request.ClientName?.Length > 200 || + request.ClientUri?.Length > 2048 || + (request.ClientUri is not null && + (!Uri.TryCreate(request.ClientUri, UriKind.Absolute, out var clientUri) || + !McpOAuthProxyUtilities.IsSecureEndpoint(clientUri)))) + { + return Error(StatusCodes.Status400BadRequest, "invalid_client_metadata", "The client metadata is invalid."); + } + + bool registrationAllowed; + try + { + registrationAllowed = await _options.ClientRegistrationValidator( + new McpOAuthProxyClientRegistrationContext + { + HttpContext = context, + RedirectUris = redirectUris, + Scopes = scopes, + GrantTypes = grantTypes, + ClientName = request.ClientName, + ClientUri = request.ClientUri, + ApplicationType = request.ApplicationType, + }, + context.RequestAborted).ConfigureAwait(false); + } + catch (Exception exception) when (exception is not OperationCanceledException) + { + LogRegistrationPolicyFailed(_logger, exception); + return Error(StatusCodes.Status500InternalServerError, "server_error", "The registration policy could not evaluate the client."); + } + + if (!registrationAllowed) + { + return Error(StatusCodes.Status403Forbidden, "access_denied", "The dynamic client registration was not approved."); + } + + var clientId = $"mcp_{McpOAuthProxyUtilities.CreateRandomToken()}"; + await _store.SetAsync( + "client", + clientId, + new OAuthProxyClientRecord + { + ClientId = clientId, + RedirectUris = redirectUris, + Scopes = scopes, + ClientName = request.ClientName, + ClientUri = request.ClientUri, + ApplicationType = request.ApplicationType, + SupportsRefreshTokens = grantTypes.Contains("refresh_token", StringComparer.Ordinal), + }, + McpOAuthProxyJsonContext.Default.OAuthProxyClientRecord, + _options.ClientRegistrationLifetime, + context.RequestAborted).ConfigureAwait(false); + + var response = new OAuthProxyClientRegistrationResponse + { + ClientId = clientId, + ClientIdIssuedAt = _timeProvider.GetUtcNow().ToUnixTimeSeconds(), + RedirectUris = redirectUris, + GrantTypes = grantTypes, + ResponseTypes = responseTypes, + ClientName = request.ClientName, + ClientUri = request.ClientUri, + Scope = scopes.Length == 0 ? null : string.Join(' ', scopes), + ApplicationType = request.ApplicationType, + }; + + return Results.Json( + response, + McpOAuthProxyJsonContext.Default.OAuthProxyClientRegistrationResponse, + statusCode: StatusCodes.Status201Created); + } + + public async Task HandleAuthorizationAsync(HttpContext context) + { + var query = context.Request.Query; + var clientId = query["client_id"].ToString(); + var redirectUri = query["redirect_uri"].ToString(); + + if (!McpOAuthProxyUtilities.IsValidClientId(clientId) || string.IsNullOrEmpty(redirectUri)) + { + return Error(StatusCodes.Status400BadRequest, "invalid_request", "client_id and redirect_uri are required."); + } + + var client = await _store.GetAsync( + "client", + clientId, + McpOAuthProxyJsonContext.Default.OAuthProxyClientRecord, + context.RequestAborted).ConfigureAwait(false); + + if (client is null) + { + return Error(StatusCodes.Status400BadRequest, "invalid_request", "The client is not registered."); + } + + if (!client.RedirectUris.Contains(redirectUri, StringComparer.Ordinal)) + { + return Error(StatusCodes.Status400BadRequest, "invalid_request", "redirect_uri does not match the registered client."); + } + + if (query["response_type"].ToString() != "code") + { + return RedirectError(redirectUri, query["state"], "unsupported_response_type", "Only the code response type is supported."); + } + + var codeChallenge = query["code_challenge"].ToString(); + if (query["code_challenge_method"].ToString() != "S256" || + !McpOAuthProxyUtilities.IsValidCodeChallenge(codeChallenge)) + { + return RedirectError(redirectUri, query["state"], "invalid_request", "A valid S256 PKCE challenge is required."); + } + + var scopes = string.IsNullOrWhiteSpace(query["scope"]) + ? client.Scopes + : McpOAuthProxyUtilities.ParseScopes(query["scope"]); + if (!AreScopesAllowed(scopes) || scopes.Any(scope => !client.Scopes.Contains(scope, StringComparer.Ordinal))) + { + return RedirectError(redirectUri, query["state"], "invalid_scope", "The request includes a scope that was not registered."); + } + + var resource = query["resource"].ToString(); + if (!string.IsNullOrEmpty(resource) && !_options.AllowedResources.Contains(resource)) + { + return RedirectError(redirectUri, query["state"], "invalid_target", "The requested resource is not allowed by the proxy."); + } + + bool authorizationAllowed; + try + { + authorizationAllowed = await _options.AuthorizationValidator( + new McpOAuthProxyAuthorizationContext + { + HttpContext = context, + ClientId = clientId, + ClientName = client.ClientName, + RedirectUri = redirectUri, + Scopes = scopes, + Resource = string.IsNullOrEmpty(resource) ? null : resource, + }, + context.RequestAborted).ConfigureAwait(false); + } + catch (Exception exception) when (exception is not OperationCanceledException) + { + LogAuthorizationPolicyFailed(_logger, exception); + return RedirectError(redirectUri, query["state"], "server_error", "The authorization policy could not evaluate the request."); + } + + if (!authorizationAllowed) + { + return RedirectError(redirectUri, query["state"], "access_denied", "The authorization request was not approved."); + } + + OAuthProxyUpstreamMetadata upstreamMetadata; + try + { + upstreamMetadata = await GetUpstreamMetadataAsync(context.RequestAborted).ConfigureAwait(false); + } + catch (OperationCanceledException exception) when (!context.RequestAborted.IsCancellationRequested) + { + LogDiscoveryFailed(_logger, exception); + return RedirectError(redirectUri, query["state"], "server_error", "The upstream authorization server is unavailable."); + } + catch (Exception exception) when (exception is HttpRequestException or JsonException or InvalidOperationException) + { + LogDiscoveryFailed(_logger, exception); + return RedirectError(redirectUri, query["state"], "server_error", "The upstream authorization server is unavailable."); + } + + var transactionId = McpOAuthProxyUtilities.CreateRandomToken(); + var browserBinding = McpOAuthProxyUtilities.CreateRandomToken(); + var upstreamCodeVerifier = McpOAuthProxyUtilities.CreateRandomToken(48); + await _store.SetAsync( + "transaction", + transactionId, + new OAuthProxyAuthorizationTransaction + { + ClientId = clientId, + RedirectUri = redirectUri, + CodeChallenge = codeChallenge, + UpstreamCodeVerifier = upstreamCodeVerifier, + BrowserBinding = browserBinding, + Scopes = scopes, + ClientState = query["state"], + Resource = string.IsNullOrEmpty(resource) ? null : resource, + SupportsRefreshTokens = client.SupportsRefreshTokens, + }, + McpOAuthProxyJsonContext.Default.OAuthProxyAuthorizationTransaction, + _options.AuthorizationTransactionLifetime, + context.RequestAborted).ConfigureAwait(false); + + context.Response.Cookies.Append( + GetBrowserBindingCookieName(transactionId), + browserBinding, + new CookieOptions + { + HttpOnly = true, + Secure = _options.CookieSecurePolicy is CookieSecurePolicy.Always || + (_options.CookieSecurePolicy is CookieSecurePolicy.SameAsRequest && context.Request.IsHttps), + SameSite = Microsoft.AspNetCore.Http.SameSiteMode.Lax, + IsEssential = true, + Path = GetCallbackPath(), + MaxAge = _options.AuthorizationTransactionLifetime, + }); + + Dictionary parameters = new(StringComparer.Ordinal) + { + ["client_id"] = _options.UpstreamClientId, + ["redirect_uri"] = _options.UpstreamRedirectUri.AbsoluteUri, + ["response_type"] = "code", + ["code_challenge"] = McpOAuthProxyUtilities.CreateCodeChallenge(upstreamCodeVerifier), + ["code_challenge_method"] = "S256", + ["scope"] = scopes.Length == 0 ? null : string.Join(' ', scopes), + ["state"] = transactionId, + }; + + if (_options.ForwardResourceIndicator && !string.IsNullOrEmpty(resource)) + { + parameters["resource"] = resource; + } + + foreach (var parameter in _options.AdditionalAuthorizationParameters) + { + parameters[parameter.Key] = parameter.Value; + } + + return Results.Redirect(QueryHelpers.AddQueryString(upstreamMetadata.AuthorizationEndpoint!, parameters)); + } + + public async Task HandleCallbackAsync(HttpContext context) + { + var state = context.Request.Query["state"].ToString(); + if (!McpOAuthProxyUtilities.IsValidOpaqueToken(state)) + { + return Error(StatusCodes.Status400BadRequest, "invalid_request", "state is required."); + } + + var transaction = await _store.GetAsync( + "transaction", + state, + McpOAuthProxyJsonContext.Default.OAuthProxyAuthorizationTransaction, + context.RequestAborted).ConfigureAwait(false); + + if (transaction is null) + { + return Error(StatusCodes.Status400BadRequest, "invalid_request", "The authorization transaction is invalid or expired."); + } + + var browserBindingCookieName = GetBrowserBindingCookieName(state); + if (!context.Request.Cookies.TryGetValue(browserBindingCookieName, out var browserBinding) || + !McpOAuthProxyUtilities.FixedTimeEquals(transaction.BrowserBinding, browserBinding)) + { + return Error(StatusCodes.Status400BadRequest, "invalid_request", "The authorization response is not bound to the initiating browser."); + } + + transaction = await _store.TakeAsync( + "transaction", + state, + McpOAuthProxyJsonContext.Default.OAuthProxyAuthorizationTransaction, + context.RequestAborted).ConfigureAwait(false); + if (transaction is null) + { + return Error(StatusCodes.Status400BadRequest, "invalid_request", "The authorization transaction was already consumed."); + } + + context.Response.Cookies.Delete(browserBindingCookieName, new CookieOptions { Path = GetCallbackPath() }); + + var upstreamError = context.Request.Query["error"].ToString(); + if (!string.IsNullOrEmpty(upstreamError)) + { + return RedirectError(transaction.RedirectUri, transaction.ClientState, upstreamError, "The upstream authorization request was rejected."); + } + + var code = context.Request.Query["code"].ToString(); + if (string.IsNullOrEmpty(code)) + { + return RedirectError(transaction.RedirectUri, transaction.ClientState, "invalid_request", "The upstream authorization response did not include a code."); + } + + var exchange = await ExchangeAuthorizationCodeAsync( + code, + transaction.UpstreamCodeVerifier, + transaction.Resource, + context.RequestAborted).ConfigureAwait(false); + if (exchange.Token is null) + { + return RedirectError(transaction.RedirectUri, transaction.ClientState, "server_error", "The upstream token exchange failed."); + } + + var grantedScopes = ResolveGrantedScopes(transaction.Scopes, exchange.Token.Scope); + if (grantedScopes is null) + { + return RedirectError(transaction.RedirectUri, transaction.ClientState, "server_error", "The upstream authorization server returned an invalid scope grant."); + } + + var proxyCode = McpOAuthProxyUtilities.CreateRandomToken(); + await _store.SetAsync( + "code", + proxyCode, + new OAuthProxyAuthorizationCode + { + ClientId = transaction.ClientId, + RedirectUri = transaction.RedirectUri, + CodeChallenge = transaction.CodeChallenge, + Scopes = grantedScopes, + Resource = transaction.Resource, + SupportsRefreshTokens = transaction.SupportsRefreshTokens, + UpstreamToken = exchange.Token, + }, + McpOAuthProxyJsonContext.Default.OAuthProxyAuthorizationCode, + _options.AuthorizationCodeLifetime, + context.RequestAborted).ConfigureAwait(false); + + Dictionary parameters = new(StringComparer.Ordinal) + { + ["code"] = proxyCode, + ["state"] = transaction.ClientState, + ["iss"] = _options.Issuer.AbsoluteUri.TrimEnd('/'), + }; + return Results.Redirect(QueryHelpers.AddQueryString(transaction.RedirectUri, parameters)); + } + + public async Task HandleTokenAsync(HttpContext context) + { + if (!context.Request.HasFormContentType) + { + return Error(StatusCodes.Status400BadRequest, "invalid_request", "The token request must use form encoding."); + } + + Dictionary form; + try + { + var body = await ReadRequestBodyAsync(context.Request, context.RequestAborted).ConfigureAwait(false); + var parsed = QueryHelpers.ParseQuery(Encoding.UTF8.GetString(body)); + if (parsed.Any(static field => field.Value.Count != 1)) + { + return Error(StatusCodes.Status400BadRequest, "invalid_request", "Token request parameters must occur exactly once."); + } + + form = parsed.ToDictionary( + static field => field.Key, + static field => field.Value.ToString(), + StringComparer.Ordinal); + } + catch (InvalidDataException) + { + return Error(StatusCodes.Status400BadRequest, "invalid_request", "The token request body is too large."); + } + + var grantType = form.GetValueOrDefault("grant_type") ?? string.Empty; + return grantType switch + { + "authorization_code" => await ExchangeProxyAuthorizationCodeAsync(context, form).ConfigureAwait(false), + "refresh_token" => await ExchangeProxyRefreshTokenAsync(context, form).ConfigureAwait(false), + _ => Error(StatusCodes.Status400BadRequest, "unsupported_grant_type", "Only authorization_code and refresh_token grants are supported."), + }; + } + + private async Task ExchangeProxyAuthorizationCodeAsync(HttpContext context, IReadOnlyDictionary form) + { + var code = form.GetValueOrDefault("code") ?? string.Empty; + var clientId = form.GetValueOrDefault("client_id") ?? string.Empty; + var redirectUri = form.GetValueOrDefault("redirect_uri") ?? string.Empty; + var codeVerifier = form.GetValueOrDefault("code_verifier") ?? string.Empty; + if (!McpOAuthProxyUtilities.IsValidOpaqueToken(code) || + !McpOAuthProxyUtilities.IsValidClientId(clientId) || + string.IsNullOrEmpty(redirectUri) || + !McpOAuthProxyUtilities.IsValidCodeVerifier(codeVerifier)) + { + return Error(StatusCodes.Status400BadRequest, "invalid_request", "code, client_id, redirect_uri, and code_verifier are required."); + } + + var authorizationCode = await _store.GetAsync( + "code", + code, + McpOAuthProxyJsonContext.Default.OAuthProxyAuthorizationCode, + context.RequestAborted).ConfigureAwait(false); + if (authorizationCode is null || + !McpOAuthProxyUtilities.FixedTimeEquals(authorizationCode.ClientId, clientId) || + !McpOAuthProxyUtilities.FixedTimeEquals(authorizationCode.RedirectUri, redirectUri) || + !McpOAuthProxyUtilities.FixedTimeEquals( + authorizationCode.CodeChallenge, + McpOAuthProxyUtilities.CreateCodeChallenge(codeVerifier))) + { + return Error(StatusCodes.Status400BadRequest, "invalid_grant", "The authorization code is invalid, expired, or does not match the request."); + } + + + authorizationCode = await _store.TakeAsync( + "code", + code, + McpOAuthProxyJsonContext.Default.OAuthProxyAuthorizationCode, + context.RequestAborted).ConfigureAwait(false); + if (authorizationCode is null) + { + return Error(StatusCodes.Status400BadRequest, "invalid_grant", "The authorization code was already consumed."); + } + + var token = await MintTokenAsync( + authorizationCode.ClientId, + authorizationCode.Resource, + authorizationCode.Scopes, + authorizationCode.UpstreamToken, + isRefresh: false, + context.RequestAborted).ConfigureAwait(false); + if (token is null) + { + return Error(StatusCodes.Status500InternalServerError, "server_error", "The proxy could not issue an access token."); + } + + string? refreshToken = null; + if (authorizationCode.SupportsRefreshTokens && !string.IsNullOrEmpty(authorizationCode.UpstreamToken.RefreshToken)) + { + refreshToken = McpOAuthProxyUtilities.CreateRandomToken(); + var familyId = McpOAuthProxyUtilities.CreateRandomToken(); + var expiresAt = _timeProvider.GetUtcNow() + _options.RefreshTokenLifetime; + await StoreRefreshRecordAsync( + refreshToken, + authorizationCode.ClientId, + authorizationCode.Resource, + authorizationCode.Scopes, + authorizationCode.UpstreamToken, + familyId, + expiresAt, + context.RequestAborted).ConfigureAwait(false); + _ = await StoreRefreshFamilyAsync(familyId, refreshToken, expiresAt, context.RequestAborted).ConfigureAwait(false); + } + + return TokenResponse(token, refreshToken, authorizationCode.Scopes); + } + + private async Task ExchangeProxyRefreshTokenAsync(HttpContext context, IReadOnlyDictionary form) + { + var refreshToken = form.GetValueOrDefault("refresh_token") ?? string.Empty; + var clientId = form.GetValueOrDefault("client_id") ?? string.Empty; + if (!McpOAuthProxyUtilities.IsValidOpaqueToken(refreshToken) || + !McpOAuthProxyUtilities.IsValidClientId(clientId)) + { + return Error(StatusCodes.Status400BadRequest, "invalid_request", "refresh_token and client_id are required."); + } + + var refreshRecord = await _store.GetAsync( + "refresh", + refreshToken, + McpOAuthProxyJsonContext.Default.OAuthProxyRefreshRecord, + context.RequestAborted).ConfigureAwait(false); + if (refreshRecord is null) + { + await RevokeRefreshFamilyOnReuseAsync(refreshToken, context.RequestAborted).ConfigureAwait(false); + return Error(StatusCodes.Status400BadRequest, "invalid_grant", "The refresh token is invalid or expired."); + } + + if (!McpOAuthProxyUtilities.FixedTimeEquals(refreshRecord.ClientId, clientId) || + string.IsNullOrEmpty(refreshRecord.UpstreamToken.RefreshToken)) + { + return Error(StatusCodes.Status400BadRequest, "invalid_grant", "The refresh token is invalid or expired."); + } + + var remainingLifetime = refreshRecord.ExpiresAt - _timeProvider.GetUtcNow(); + if (remainingLifetime <= TimeSpan.Zero) + { + return Error(StatusCodes.Status400BadRequest, "invalid_grant", "The refresh token is expired."); + } + + await _store.SetAsync( + "consumed-refresh", + refreshToken, + new OAuthProxyConsumedRefreshRecord + { + FamilyId = refreshRecord.FamilyId, + ExpiresAt = refreshRecord.ExpiresAt, + }, + McpOAuthProxyJsonContext.Default.OAuthProxyConsumedRefreshRecord, + remainingLifetime, + context.RequestAborted).ConfigureAwait(false); + + refreshRecord = await _store.TakeAsync( + "refresh", + refreshToken, + McpOAuthProxyJsonContext.Default.OAuthProxyRefreshRecord, + context.RequestAborted).ConfigureAwait(false); + if (refreshRecord is null) + { + await RevokeRefreshFamilyOnReuseAsync(refreshToken, context.RequestAborted).ConfigureAwait(false); + return Error(StatusCodes.Status400BadRequest, "invalid_grant", "The refresh token was already consumed."); + } + + var family = await _store.TakeAsync( + "refresh-family", + refreshRecord.FamilyId, + McpOAuthProxyJsonContext.Default.OAuthProxyRefreshFamilyRecord, + context.RequestAborted).ConfigureAwait(false); + if (family is null || !McpOAuthProxyUtilities.FixedTimeEquals(family.CurrentRefreshToken, refreshToken)) + { + if (family is not null) + { + await _store.RemoveAsync("refresh", family.CurrentRefreshToken, context.RequestAborted).ConfigureAwait(false); + } + + return Error(StatusCodes.Status400BadRequest, "invalid_grant", "Refresh-token reuse was detected and the token family was revoked."); + } + + var requestedScope = form.GetValueOrDefault("scope") ?? string.Empty; + var requestedScopes = string.IsNullOrWhiteSpace(requestedScope) + ? refreshRecord.Scopes + : McpOAuthProxyUtilities.ParseScopes(requestedScope); + if (requestedScopes.Length > 32 || + requestedScopes.Any(scope => !refreshRecord.Scopes.Contains(scope, StringComparer.Ordinal))) + { + await RestoreRefreshFamilyAsync(refreshToken, refreshRecord, context.RequestAborted).ConfigureAwait(false); + return Error(StatusCodes.Status400BadRequest, "invalid_scope", "A refresh request cannot expand the original scope grant."); + } + + var exchange = await ExchangeRefreshTokenAsync( + refreshRecord.UpstreamToken.RefreshToken!, + refreshRecord.Resource, + requestedScopes, + context.RequestAborted).ConfigureAwait(false); + if (exchange.Token is null) + { + if (exchange.IsTransient) + { + await RestoreRefreshFamilyAsync(refreshToken, refreshRecord, context.RequestAborted).ConfigureAwait(false); + } + + return Error( + exchange.IsTransient ? StatusCodes.Status503ServiceUnavailable : StatusCodes.Status400BadRequest, + exchange.IsTransient ? "temporarily_unavailable" : "invalid_grant", + "The upstream token refresh failed."); + } + + exchange.Token.RefreshToken ??= refreshRecord.UpstreamToken.RefreshToken; + var grantedScopes = ResolveGrantedScopes(requestedScopes, exchange.Token.Scope); + if (grantedScopes is null) + { + return Error(StatusCodes.Status400BadRequest, "invalid_grant", "The upstream authorization server returned an invalid scope grant."); + } + + var token = await MintTokenAsync( + refreshRecord.ClientId, + refreshRecord.Resource, + grantedScopes, + exchange.Token, + isRefresh: true, + context.RequestAborted).ConfigureAwait(false); + if (token is null) + { + refreshRecord.Scopes = grantedScopes; + refreshRecord.UpstreamToken = exchange.Token; + await RestoreRefreshFamilyAsync(refreshToken, refreshRecord, context.RequestAborted).ConfigureAwait(false); + return Error(StatusCodes.Status500InternalServerError, "server_error", "The proxy could not issue an access token."); + } + + var rotatedRefreshToken = McpOAuthProxyUtilities.CreateRandomToken(); + await StoreRefreshRecordAsync( + rotatedRefreshToken, + refreshRecord.ClientId, + refreshRecord.Resource, + grantedScopes, + exchange.Token, + refreshRecord.FamilyId, + refreshRecord.ExpiresAt, + context.RequestAborted).ConfigureAwait(false); + var familyPublished = await StoreRefreshFamilyAsync( + refreshRecord.FamilyId, + rotatedRefreshToken, + refreshRecord.ExpiresAt, + context.RequestAborted).ConfigureAwait(false); + if (!familyPublished) + { + return Error(StatusCodes.Status400BadRequest, "invalid_grant", "Refresh-token reuse was detected and the token family was revoked."); + } + + return TokenResponse(token, rotatedRefreshToken, grantedScopes); + } + + private async Task StoreRefreshRecordAsync( + string refreshToken, + string clientId, + string? resource, + string[] scopes, + OAuthProxyUpstreamTokenResponse upstreamToken, + string familyId, + DateTimeOffset expiresAt, + CancellationToken cancellationToken) + { + var remainingLifetime = expiresAt - _timeProvider.GetUtcNow(); + if (remainingLifetime <= TimeSpan.Zero) + { + return; + } + + await _store.SetAsync( + "refresh", + refreshToken, + new OAuthProxyRefreshRecord + { + ClientId = clientId, + Resource = resource, + Scopes = scopes, + UpstreamToken = upstreamToken, + FamilyId = familyId, + ExpiresAt = expiresAt, + }, + McpOAuthProxyJsonContext.Default.OAuthProxyRefreshRecord, + remainingLifetime, + cancellationToken).ConfigureAwait(false); + } + + private async Task StoreRefreshFamilyAsync( + string familyId, + string currentRefreshToken, + DateTimeOffset expiresAt, + CancellationToken cancellationToken) + { + var remainingLifetime = expiresAt - _timeProvider.GetUtcNow(); + if (remainingLifetime <= TimeSpan.Zero) + { + return false; + } + + await _store.SetAsync( + "refresh-family", + familyId, + new OAuthProxyRefreshFamilyRecord + { + CurrentRefreshToken = currentRefreshToken, + ExpiresAt = expiresAt, + }, + McpOAuthProxyJsonContext.Default.OAuthProxyRefreshFamilyRecord, + remainingLifetime, + cancellationToken).ConfigureAwait(false); + + var revoked = await _store.GetAsync( + "revoked-refresh-family", + familyId, + McpOAuthProxyJsonContext.Default.OAuthProxyConsumedRefreshRecord, + cancellationToken).ConfigureAwait(false); + if (revoked is null) + { + return true; + } + + var publishedFamily = await _store.TakeAsync( + "refresh-family", + familyId, + McpOAuthProxyJsonContext.Default.OAuthProxyRefreshFamilyRecord, + cancellationToken).ConfigureAwait(false); + if (publishedFamily is not null) + { + await _store.RemoveAsync("refresh", publishedFamily.CurrentRefreshToken, cancellationToken).ConfigureAwait(false); + } + + return false; + } + + private async Task RestoreRefreshFamilyAsync( + string refreshToken, + OAuthProxyRefreshRecord refreshRecord, + CancellationToken cancellationToken) + { + await StoreRefreshRecordAsync( + refreshToken, + refreshRecord.ClientId, + refreshRecord.Resource, + refreshRecord.Scopes, + refreshRecord.UpstreamToken, + refreshRecord.FamilyId, + refreshRecord.ExpiresAt, + cancellationToken).ConfigureAwait(false); + _ = await StoreRefreshFamilyAsync( + refreshRecord.FamilyId, + refreshToken, + refreshRecord.ExpiresAt, + cancellationToken).ConfigureAwait(false); + } + + private async Task RevokeRefreshFamilyOnReuseAsync(string refreshToken, CancellationToken cancellationToken) + { + var consumed = await _store.GetAsync( + "consumed-refresh", + refreshToken, + McpOAuthProxyJsonContext.Default.OAuthProxyConsumedRefreshRecord, + cancellationToken).ConfigureAwait(false); + if (consumed is null) + { + return; + } + + var remainingLifetime = consumed.ExpiresAt - _timeProvider.GetUtcNow(); + if (remainingLifetime <= TimeSpan.Zero) + { + return; + } + + await _store.SetAsync( + "revoked-refresh-family", + consumed.FamilyId, + consumed, + McpOAuthProxyJsonContext.Default.OAuthProxyConsumedRefreshRecord, + remainingLifetime, + cancellationToken).ConfigureAwait(false); + + var family = await _store.TakeAsync( + "refresh-family", + consumed.FamilyId, + McpOAuthProxyJsonContext.Default.OAuthProxyRefreshFamilyRecord, + cancellationToken).ConfigureAwait(false); + if (family is not null) + { + await _store.RemoveAsync("refresh", family.CurrentRefreshToken, cancellationToken).ConfigureAwait(false); + } + } + + private async Task MintTokenAsync( + string clientId, + string? resource, + string[] scopes, + OAuthProxyUpstreamTokenResponse upstreamToken, + bool isRefresh, + CancellationToken cancellationToken) + { + McpOAuthProxyTokenResult token; + try + { + token = await _options.TokenFactory( + new McpOAuthProxyTokenContext + { + ClientId = clientId, + Resource = resource, + Scopes = scopes, + UpstreamAccessToken = upstreamToken.AccessToken!, + UpstreamIdToken = upstreamToken.IdToken, + UpstreamTokenType = upstreamToken.TokenType, + UpstreamExpiresIn = upstreamToken.ExpiresIn is long expiresIn ? TimeSpan.FromSeconds(expiresIn) : null, + IsRefresh = isRefresh, + }, + cancellationToken).ConfigureAwait(false); + } + catch (Exception exception) when (exception is not OperationCanceledException) + { + LogTokenFactoryFailed(_logger, exception); + return null; + } + + if (string.IsNullOrWhiteSpace(token.AccessToken) || + string.IsNullOrWhiteSpace(token.Subject) || + string.IsNullOrWhiteSpace(token.Audience) || + string.IsNullOrWhiteSpace(token.TokenId) || + !string.Equals(token.TokenType, "Bearer", StringComparison.OrdinalIgnoreCase) || + token.ExpiresIn < TimeSpan.FromSeconds(1) || + token.ExpiresIn > _options.MaximumAccessTokenLifetime || + (!string.IsNullOrEmpty(resource) && !string.Equals(token.Audience, resource, StringComparison.Ordinal)) || + IsUpstreamToken(token.AccessToken, upstreamToken)) + { + LogInvalidTokenFactoryResult(_logger); + return null; + } + + LogTokenMinted( + _logger, + token.Subject, + token.Audience, + string.Join(' ', scopes), + clientId, + token.TokenId, + _timeProvider.GetUtcNow()); + return token; + } + + private IResult TokenResponse(McpOAuthProxyTokenResult token, string? refreshToken, string[] scopes) => + new NoStoreResult(Results.Json( + new OAuthProxyTokenResponse + { + AccessToken = token.AccessToken, + TokenType = token.TokenType, + ExpiresIn = checked((long)token.ExpiresIn.TotalSeconds), + RefreshToken = refreshToken, + Scope = scopes.Length == 0 ? null : string.Join(' ', scopes), + }, + McpOAuthProxyJsonContext.Default.OAuthProxyTokenResponse)); + + private async Task ExchangeAuthorizationCodeAsync( + string code, + string codeVerifier, + string? resource, + CancellationToken cancellationToken) + { + Dictionary parameters = new(StringComparer.Ordinal) + { + ["grant_type"] = "authorization_code", + ["code"] = code, + ["redirect_uri"] = _options.UpstreamRedirectUri.AbsoluteUri, + ["code_verifier"] = codeVerifier, + }; + if (_options.ForwardResourceIndicator && !string.IsNullOrEmpty(resource)) + { + parameters["resource"] = resource; + } + + return await ExchangeUpstreamTokenAsync(parameters, cancellationToken).ConfigureAwait(false); + } + + private async Task ExchangeRefreshTokenAsync( + string refreshToken, + string? resource, + string[] scopes, + CancellationToken cancellationToken) + { + Dictionary parameters = new(StringComparer.Ordinal) + { + ["grant_type"] = "refresh_token", + ["refresh_token"] = refreshToken, + }; + if (scopes.Length > 0) + { + parameters["scope"] = string.Join(' ', scopes); + } + if (_options.ForwardResourceIndicator && !string.IsNullOrEmpty(resource)) + { + parameters["resource"] = resource; + } + + return await ExchangeUpstreamTokenAsync(parameters, cancellationToken).ConfigureAwait(false); + } + + private async Task ExchangeUpstreamTokenAsync( + Dictionary parameters, + CancellationToken cancellationToken) + { + OAuthProxyUpstreamMetadata metadata; + try + { + metadata = await GetUpstreamMetadataAsync(cancellationToken).ConfigureAwait(false); + } + catch (OperationCanceledException exception) when (!cancellationToken.IsCancellationRequested) + { + LogDiscoveryFailed(_logger, exception); + return new(null, IsTransient: true); + } + catch (Exception exception) when (exception is HttpRequestException or JsonException or InvalidOperationException) + { + LogDiscoveryFailed(_logger, exception); + return new(null, IsTransient: true); + } + if (_options.UpstreamClientAuthenticationMethod is not McpOAuthProxyClientAuthenticationMethod.ClientSecretBasic) + { + parameters["client_id"] = _options.UpstreamClientId; + } + + if (_options.UpstreamClientAuthenticationMethod is McpOAuthProxyClientAuthenticationMethod.ClientSecretPost) + { + parameters["client_secret"] = _options.UpstreamClientSecret!; + } + + foreach (var parameter in _options.AdditionalTokenParameters) + { + parameters[parameter.Key] = parameter.Value; + } + + using var request = new HttpRequestMessage(HttpMethod.Post, metadata.TokenEndpoint) + { + Content = new FormUrlEncodedContent(parameters), + }; + request.Headers.Accept.Add(new MediaTypeWithQualityHeaderValue("application/json")); + if (_options.UpstreamClientAuthenticationMethod is McpOAuthProxyClientAuthenticationMethod.ClientSecretBasic) + { + var encodedClient = WebUtility.UrlEncode(_options.UpstreamClientId); + var encodedSecret = WebUtility.UrlEncode(_options.UpstreamClientSecret!); + request.Headers.Authorization = new AuthenticationHeaderValue( + "Basic", + Convert.ToBase64String(Encoding.UTF8.GetBytes($"{encodedClient}:{encodedSecret}"))); + } + + HttpResponseMessage response; + try + { + response = await _httpClientFactory.CreateClient(McpOAuthProxyServiceCollectionExtensions.HttpClientName) + .SendAsync(request, HttpCompletionOption.ResponseHeadersRead, cancellationToken) + .ConfigureAwait(false); + } + catch (OperationCanceledException exception) when (!cancellationToken.IsCancellationRequested) + { + LogUpstreamTokenRequestException(_logger, exception); + return new(null, IsTransient: true); + } + catch (HttpRequestException exception) + { + LogUpstreamTokenRequestException(_logger, exception); + return new(null, IsTransient: true); + } + + using (response) + { + OAuthProxyUpstreamTokenResponse? token; + try + { + await using var stream = await response.Content.ReadAsStreamAsync(cancellationToken).ConfigureAwait(false); + token = await JsonSerializer.DeserializeAsync( + stream, + McpOAuthProxyJsonContext.Default.OAuthProxyUpstreamTokenResponse, + cancellationToken).ConfigureAwait(false); + } + catch (OperationCanceledException exception) when (!cancellationToken.IsCancellationRequested) + { + LogInvalidUpstreamTokenResponse(_logger, response.StatusCode, exception); + return new(null, IsTransient: true); + } + catch (Exception exception) when (exception is JsonException or HttpRequestException) + { + LogInvalidUpstreamTokenResponse(_logger, response.StatusCode, exception); + return new( + null, + IsTransient: IsTransientUpstreamStatus(response.StatusCode)); + } + + if (!response.IsSuccessStatusCode || string.IsNullOrEmpty(token?.AccessToken)) + { + LogUpstreamTokenRequestFailed(_logger, response.StatusCode, token?.Error); + return new(null, IsTransient: IsTransientUpstreamStatus(response.StatusCode)); + } + + return new(token, IsTransient: false); + } + } + + private async Task GetUpstreamMetadataAsync(CancellationToken cancellationToken) + { + if (_upstreamMetadata is not null) + { + return _upstreamMetadata; + } + + await _metadataLock.WaitAsync(cancellationToken).ConfigureAwait(false); + try + { + if (_upstreamMetadata is not null) + { + return _upstreamMetadata; + } + + if (_options.UpstreamAuthorizationEndpoint is not null && _options.UpstreamTokenEndpoint is not null) + { + return _upstreamMetadata = new OAuthProxyUpstreamMetadata + { + Issuer = _options.UpstreamIssuer?.AbsoluteUri, + AuthorizationEndpoint = _options.UpstreamAuthorizationEndpoint.AbsoluteUri, + TokenEndpoint = _options.UpstreamTokenEndpoint.AbsoluteUri, + CodeChallengeMethodsSupported = ["S256"], + }; + } + + var discoveryUri = McpOAuthProxyUtilities.AppendPath(_options.UpstreamIssuer!, "/.well-known/openid-configuration"); + using var response = await _httpClientFactory.CreateClient(McpOAuthProxyServiceCollectionExtensions.HttpClientName) + .GetAsync(discoveryUri, cancellationToken) + .ConfigureAwait(false); + response.EnsureSuccessStatusCode(); + await using var stream = await response.Content.ReadAsStreamAsync(cancellationToken).ConfigureAwait(false); + var metadata = await JsonSerializer.DeserializeAsync( + stream, + McpOAuthProxyJsonContext.Default.OAuthProxyUpstreamMetadata, + cancellationToken).ConfigureAwait(false) ?? + throw new InvalidOperationException("The upstream discovery document is empty."); + + if (!SameIssuer(metadata.Issuer, _options.UpstreamIssuer!) || + !Uri.TryCreate(metadata.AuthorizationEndpoint, UriKind.Absolute, out var authorizationEndpoint) || + !Uri.TryCreate(metadata.TokenEndpoint, UriKind.Absolute, out var tokenEndpoint) || + !McpOAuthProxyUtilities.IsSecureEndpoint(authorizationEndpoint) || + !McpOAuthProxyUtilities.IsSecureEndpoint(tokenEndpoint) || + metadata.CodeChallengeMethodsSupported?.Contains("S256", StringComparer.Ordinal) is not true) + { + throw new InvalidOperationException("The upstream discovery document has an invalid issuer, endpoint, or PKCE configuration."); + } + + return _upstreamMetadata = metadata; + } + finally + { + _metadataLock.Release(); + } + } + + private bool AreScopesAllowed(IEnumerable scopes) => + scopes.All(_options.AllowedScopes.Contains); + + private static string[]? ResolveGrantedScopes(string[] requestedScopes, string? upstreamScope) + { + if (string.IsNullOrWhiteSpace(upstreamScope)) + { + return requestedScopes; + } + + var grantedScopes = McpOAuthProxyUtilities.ParseScopes(upstreamScope); + return grantedScopes.All(scope => requestedScopes.Contains(scope, StringComparer.Ordinal)) + ? grantedScopes + : null; + } + + private static bool IsTransientUpstreamStatus(HttpStatusCode statusCode) => + statusCode is HttpStatusCode.RequestTimeout or HttpStatusCode.TooManyRequests || + (int)statusCode >= 500; + + private static async Task ReadRequestBodyAsync(HttpRequest request, CancellationToken cancellationToken) + { + if (request.ContentLength > MaxRequestBodySize) + { + throw new InvalidDataException("The request body exceeds the configured limit."); + } + + var buffer = ArrayPool.Shared.Rent(8192); + try + { + using var body = new MemoryStream(); + while (true) + { + var read = await request.Body.ReadAsync(buffer, cancellationToken).ConfigureAwait(false); + if (read == 0) + { + return body.ToArray(); + } + + if (body.Length + read > MaxRequestBodySize) + { + throw new InvalidDataException("The request body exceeds the configured limit."); + } + + body.Write(buffer, 0, read); + } + } + finally + { + ArrayPool.Shared.Return(buffer); + } + } + + private static string GetBrowserBindingCookieName(string transactionId) => + $"{BrowserBindingCookieName}-{transactionId}"; + + private string GetCallbackPath() => + $"{_options.Issuer.AbsolutePath.TrimEnd('/')}/callback"; + + private static bool SameIssuer(string? discoveredIssuer, Uri configuredIssuer) => + !string.IsNullOrEmpty(discoveredIssuer) && + string.Equals(discoveredIssuer.TrimEnd('/'), configuredIssuer.AbsoluteUri.TrimEnd('/'), StringComparison.Ordinal); + + private static bool IsUpstreamToken(string candidate, OAuthProxyUpstreamTokenResponse upstreamToken) => + McpOAuthProxyUtilities.FixedTimeEquals(candidate, upstreamToken.AccessToken!) || + (!string.IsNullOrEmpty(upstreamToken.RefreshToken) && McpOAuthProxyUtilities.FixedTimeEquals(candidate, upstreamToken.RefreshToken)) || + (!string.IsNullOrEmpty(upstreamToken.IdToken) && McpOAuthProxyUtilities.FixedTimeEquals(candidate, upstreamToken.IdToken)); + + private IResult RedirectError( + string redirectUri, + string? state, + string error, + string description) + { + Dictionary parameters = new(StringComparer.Ordinal) + { + ["error"] = error, + ["error_description"] = description, + ["state"] = state, + ["iss"] = _options.Issuer.AbsoluteUri.TrimEnd('/'), + }; + return Results.Redirect(QueryHelpers.AddQueryString(redirectUri, parameters)); + } + + private static IResult Error(int statusCode, string error, string description) => + new NoStoreResult(Results.Json( + new OAuthProxyErrorResponse { Error = error, ErrorDescription = description }, + McpOAuthProxyJsonContext.Default.OAuthProxyErrorResponse, + statusCode: statusCode)); + + private readonly record struct UpstreamExchangeResult(OAuthProxyUpstreamTokenResponse? Token, bool IsTransient); + + private sealed class NoStoreResult(IResult inner) : IResult + { + public Task ExecuteAsync(HttpContext httpContext) + { + httpContext.Response.Headers.CacheControl = "no-store"; + httpContext.Response.Headers.Pragma = "no-cache"; + return inner.ExecuteAsync(httpContext); + } + } + + [LoggerMessage(Level = LogLevel.Error, Message = "OAuth proxy upstream discovery failed.")] + private static partial void LogDiscoveryFailed(ILogger logger, Exception exception); + + [LoggerMessage(Level = LogLevel.Error, Message = "OAuth proxy token factory failed.")] + private static partial void LogTokenFactoryFailed(ILogger logger, Exception exception); + + [LoggerMessage(Level = LogLevel.Error, Message = "OAuth proxy client registration policy failed.")] + private static partial void LogRegistrationPolicyFailed(ILogger logger, Exception exception); + + [LoggerMessage(Level = LogLevel.Error, Message = "OAuth proxy authorization policy failed.")] + private static partial void LogAuthorizationPolicyFailed(ILogger logger, Exception exception); + + [LoggerMessage(Level = LogLevel.Error, Message = "OAuth proxy token factory returned an invalid or unsafe token result.")] + private static partial void LogInvalidTokenFactoryResult(ILogger logger); + + [LoggerMessage(Level = LogLevel.Information, Message = "OAuth proxy minted token: subject={Subject}, audience={Audience}, scopes={Scopes}, client_id={ClientId}, jti={TokenId}, minted_at={MintedAt}.")] + private static partial void LogTokenMinted( + ILogger logger, + string subject, + string audience, + string scopes, + string clientId, + string tokenId, + DateTimeOffset mintedAt); + + [LoggerMessage(Level = LogLevel.Warning, Message = "OAuth proxy could not reach the upstream token endpoint.")] + private static partial void LogUpstreamTokenRequestException(ILogger logger, Exception exception); + + [LoggerMessage(Level = LogLevel.Warning, Message = "OAuth proxy received an invalid token response with status {StatusCode} from the upstream provider.")] + private static partial void LogInvalidUpstreamTokenResponse(ILogger logger, HttpStatusCode statusCode, Exception exception); + + [LoggerMessage(Level = LogLevel.Warning, Message = "OAuth proxy upstream token request failed with status {StatusCode} and error {Error}.")] + private static partial void LogUpstreamTokenRequestFailed(ILogger logger, HttpStatusCode statusCode, string? error); +} diff --git a/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyServiceCollectionExtensions.cs b/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyServiceCollectionExtensions.cs new file mode 100644 index 000000000..167fb2b70 --- /dev/null +++ b/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyServiceCollectionExtensions.cs @@ -0,0 +1,42 @@ +using Microsoft.Extensions.DependencyInjection.Extensions; +using Microsoft.Extensions.Options; +using ModelContextProtocol.AspNetCore.Authentication; + +namespace Microsoft.Extensions.DependencyInjection; + +/// +/// Extension methods for adding an MCP OAuth proxy to an ASP.NET Core application. +/// +public static class McpOAuthProxyServiceCollectionExtensions +{ + internal const string HttpClientName = "ModelContextProtocol.AspNetCore.Authentication.McpOAuthProxy"; + + /// + /// Adds an OAuth authorization-server facade for an upstream provider that does not support + /// Dynamic Client Registration. + /// + /// + /// An must be registered. Single-instance applications may + /// explicitly use . Multi-instance deployments + /// must provide a store whose consume operation is atomic in its backing store. + /// + public static IServiceCollection AddMcpOAuthProxy( + this IServiceCollection services, + Action configureOptions) + { + ArgumentNullException.ThrowIfNull(services); + ArgumentNullException.ThrowIfNull(configureOptions); + + services.AddDataProtection(); + services.AddHttpClient(HttpClientName); + services.AddOptions() + .Configure(configureOptions) + .ValidateOnStart(); + services.TryAddEnumerable( + ServiceDescriptor.Singleton, McpOAuthProxyOptionsValidator>()); + services.TryAddSingleton(TimeProvider.System); + services.TryAddSingleton(); + services.TryAddSingleton(); + return services; + } +} diff --git a/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyTokenContext.cs b/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyTokenContext.cs new file mode 100644 index 000000000..9224fe1f8 --- /dev/null +++ b/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyTokenContext.cs @@ -0,0 +1,83 @@ +namespace ModelContextProtocol.AspNetCore.Authentication; + +/// +/// Provides the validated OAuth transaction and upstream token response used to mint a proxy token. +/// +public sealed class McpOAuthProxyTokenContext +{ + /// + /// Gets the dynamic proxy client identifier. + /// + public required string ClientId { get; init; } + + /// + /// Gets the resource requested by the MCP client, when present. + /// + public string? Resource { get; init; } + + /// + /// Gets the scopes granted to the MCP client. + /// + public required IReadOnlyList Scopes { get; init; } + + /// + /// Gets the upstream access token. This value must never be returned as the proxy access token. + /// + public required string UpstreamAccessToken { get; init; } + + /// + /// Gets the upstream ID token, when present. + /// + public string? UpstreamIdToken { get; init; } + + /// + /// Gets the upstream token type, when present. + /// + public string? UpstreamTokenType { get; init; } + + /// + /// Gets the upstream token lifetime, when present. + /// + public TimeSpan? UpstreamExpiresIn { get; init; } + + /// + /// Gets a value indicating whether the token is being minted during a refresh. + /// + public bool IsRefresh { get; init; } +} + +/// +/// Describes a client-facing access token minted by an OAuth proxy token factory. +/// +public sealed class McpOAuthProxyTokenResult +{ + /// + /// Gets the newly minted access token. + /// + public required string AccessToken { get; init; } + + /// + /// Gets the token type returned to the client. + /// + public string TokenType { get; init; } = "Bearer"; + + /// + /// Gets the access-token lifetime. + /// + public required TimeSpan ExpiresIn { get; init; } + + /// + /// Gets the authenticated subject recorded in the token-mint audit event. + /// + public required string Subject { get; init; } + + /// + /// Gets the token audience recorded in the token-mint audit event. + /// + public required string Audience { get; init; } + + /// + /// Gets the unique token identifier recorded in the token-mint audit event. + /// + public required string TokenId { get; init; } +} diff --git a/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyUtilities.cs b/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyUtilities.cs new file mode 100644 index 000000000..946d221bd --- /dev/null +++ b/src/ModelContextProtocol.AspNetCore/Authentication/McpOAuthProxyUtilities.cs @@ -0,0 +1,75 @@ +using Microsoft.AspNetCore.WebUtilities; +using System.Security.Cryptography; +using System.Text; + +namespace ModelContextProtocol.AspNetCore.Authentication; + +internal static class McpOAuthProxyUtilities +{ + public static string CreateRandomToken(int byteCount = 32) => + WebEncoders.Base64UrlEncode(RandomNumberGenerator.GetBytes(byteCount)); + + public static string CreateCodeChallenge(string verifier) => + WebEncoders.Base64UrlEncode(SHA256.HashData(Encoding.ASCII.GetBytes(verifier))); + + public static bool FixedTimeEquals(string left, string right) + { + var leftBytes = Encoding.UTF8.GetBytes(left); + var rightBytes = Encoding.UTF8.GetBytes(right); + return leftBytes.Length == rightBytes.Length && + CryptographicOperations.FixedTimeEquals(leftBytes, rightBytes); + } + + public static bool IsSecureEndpoint(Uri uri) => + uri.IsAbsoluteUri && + (uri.Scheme.Equals(Uri.UriSchemeHttps, StringComparison.OrdinalIgnoreCase) || + (uri.Scheme.Equals(Uri.UriSchemeHttp, StringComparison.OrdinalIgnoreCase) && uri.IsLoopback)); + + public static bool IsValidRedirectUri(string value) + { + if (value.Length > 2048 || + !Uri.TryCreate(value, UriKind.Absolute, out var uri) || + !string.IsNullOrEmpty(uri.Fragment) || + !string.IsNullOrEmpty(uri.UserInfo)) + { + return false; + } + + return uri.Scheme.Equals(Uri.UriSchemeHttps, StringComparison.OrdinalIgnoreCase) || + (uri.Scheme.Equals(Uri.UriSchemeHttp, StringComparison.OrdinalIgnoreCase) && uri.IsLoopback); + } + + public static string[] ParseScopes(string? scope) => + string.IsNullOrWhiteSpace(scope) + ? [] + : scope.Split(' ', StringSplitOptions.RemoveEmptyEntries | StringSplitOptions.TrimEntries) + .Distinct(StringComparer.Ordinal) + .ToArray(); + + public static bool IsValidCodeChallenge(string value) => + value.Length == 43 && value.All(IsBase64UrlCharacter); + + public static bool IsValidCodeVerifier(string value) => + value.Length is >= 43 and <= 128 && value.All(IsPkceVerifierCharacter); + + public static bool IsValidOpaqueToken(string value) => + value.Length == 43 && value.All(IsBase64UrlCharacter); + + public static bool IsValidClientId(string value) => + value.StartsWith("mcp_", StringComparison.Ordinal) && IsValidOpaqueToken(value[4..]); + + public static Uri AppendPath(Uri baseUri, string path) + { + var builder = new UriBuilder(baseUri) + { + Path = $"{baseUri.AbsolutePath.TrimEnd('/')}/{path.TrimStart('/')}" + }; + return builder.Uri; + } + + private static bool IsBase64UrlCharacter(char value) => + char.IsAsciiLetterOrDigit(value) || value is '-' or '_'; + + private static bool IsPkceVerifierCharacter(char value) => + char.IsAsciiLetterOrDigit(value) || value is '-' or '.' or '_' or '~'; +} diff --git a/tests/ModelContextProtocol.AspNetCore.Tests/OAuth/McpOAuthProxyTests.cs b/tests/ModelContextProtocol.AspNetCore.Tests/OAuth/McpOAuthProxyTests.cs new file mode 100644 index 000000000..de723d609 --- /dev/null +++ b/tests/ModelContextProtocol.AspNetCore.Tests/OAuth/McpOAuthProxyTests.cs @@ -0,0 +1,782 @@ +using Microsoft.AspNetCore.Builder; +using Microsoft.AspNetCore.Authorization; +using Microsoft.AspNetCore.Http; +using Microsoft.AspNetCore.WebUtilities; +using Microsoft.Extensions.DependencyInjection; +using ModelContextProtocol.AspNetCore.Authentication; +using ModelContextProtocol.AspNetCore.Tests.Utils; +using System.Collections.Concurrent; +using System.Net; +using System.Net.Http.Json; +using System.Security.Cryptography; +using System.Text; +using System.Text.Json; +using System.Text.Json.Serialization; + +namespace ModelContextProtocol.AspNetCore.Tests.OAuth; + +public sealed class McpOAuthProxyTests : KestrelInMemoryTest +{ + private const string ClientRedirectUri = "http://127.0.0.1:43123/callback"; + private const string Resource = "https://mcp.example.com/server"; + private readonly RecordingStore _store = new(); + private readonly StubHttpClientFactory _upstream = new(); + private bool _allowRegistration = true; + private bool _returnUpstreamAccessToken; + private bool _useDiscovery; + private string[] _lastRegistrationGrantTypes = []; + private string[] _lastMintedScopes = []; + private int _tokenNumber; + + public McpOAuthProxyTests(ITestOutputHelper testOutputHelper) + : base(testOutputHelper) + { + SocketsHttpHandler.AllowAutoRedirect = false; + Builder.Services.AddMcpOAuthProxy(options => + { + options.Issuer = new Uri("http://localhost:5000/oauth"); + if (_useDiscovery) + { + options.UpstreamIssuer = new Uri("http://localhost:5000/upstream"); + } + else + { + options.UpstreamAuthorizationEndpoint = new Uri("http://localhost:5000/upstream/authorize"); + options.UpstreamTokenEndpoint = new Uri("http://localhost:5000/upstream/token"); + } + options.UpstreamClientId = "upstream-client"; + options.UpstreamClientSecret = "upstream-secret"; + options.UpstreamRedirectUri = new Uri("http://localhost:5000/oauth/callback"); + options.AllowedScopes.Add("mcp:tools"); + options.AllowedScopes.Add("mcp:read"); + options.AllowedResources.Add(Resource); + options.ForwardResourceIndicator = true; + options.CookieSecurePolicy = CookieSecurePolicy.SameAsRequest; + options.ClientRegistrationValidator = (context, _) => + { + _lastRegistrationGrantTypes = context.GrantTypes.ToArray(); + return ValueTask.FromResult(_allowRegistration); + }; + options.AuthorizationValidator = static (_, _) => ValueTask.FromResult(true); + options.TokenFactory = (context, _) => + { + _lastMintedScopes = context.Scopes.ToArray(); + var tokenNumber = Interlocked.Increment(ref _tokenNumber); + return ValueTask.FromResult(new McpOAuthProxyTokenResult + { + AccessToken = _returnUpstreamAccessToken ? context.UpstreamAccessToken : $"proxy-access-{tokenNumber}", + TokenType = "Bearer", + ExpiresIn = TimeSpan.FromMinutes(5), + Subject = "user-123", + Audience = context.Resource ?? Resource, + TokenId = $"jti-{tokenNumber}", + }); + }; + }); + Builder.Services.AddSingleton(_store); + Builder.Services.AddSingleton(_upstream); + Builder.Services.AddAuthorizationBuilder() + .SetFallbackPolicy(new AuthorizationPolicyBuilder().RequireAuthenticatedUser().Build()); + } + + [Fact] + public async Task MetadataAdvertisesPathAwareProxyEndpoints() + { + await using var app = await StartServerAsync(); + + using var response = await HttpClient.GetAsync( + "/.well-known/oauth-authorization-server/oauth", + TestContext.Current.CancellationToken); + + Assert.Equal(HttpStatusCode.OK, response.StatusCode); + using var document = JsonDocument.Parse(await response.Content.ReadAsStreamAsync(TestContext.Current.CancellationToken)); + var root = document.RootElement; + Assert.Equal("http://localhost:5000/oauth", root.GetProperty("issuer").GetString()); + Assert.Equal("http://localhost:5000/oauth/authorize", root.GetProperty("authorization_endpoint").GetString()); + Assert.Equal("http://localhost:5000/oauth/token", root.GetProperty("token_endpoint").GetString()); + Assert.Equal("http://localhost:5000/oauth/register", root.GetProperty("registration_endpoint").GetString()); + Assert.Equal("S256", root.GetProperty("code_challenge_methods_supported")[0].GetString()); + } + + [Fact] + public async Task RegistrationRejectsUnsafeRedirectUri() + { + await using var app = await StartServerAsync(); + + using var response = await HttpClient.PostAsJsonAsync( + "/oauth/register", + new OAuthProxyRegistrationTestRequest + { + RedirectUris = ["https://client.example/callback#fragment"], + TokenEndpointAuthMethod = "client_secret_post", + }, + OAuthProxyTestJsonContext.Default.OAuthProxyRegistrationTestRequest, + TestContext.Current.CancellationToken); + + Assert.Equal(HttpStatusCode.BadRequest, response.StatusCode); + using var document = JsonDocument.Parse(await response.Content.ReadAsStreamAsync(TestContext.Current.CancellationToken)); + Assert.Equal("invalid_redirect_uri", document.RootElement.GetProperty("error").GetString()); + } + + [Fact] + public async Task RegistrationPolicyCanRejectDynamicClient() + { + _allowRegistration = false; + await using var app = await StartServerAsync(); + + using var response = await HttpClient.PostAsJsonAsync( + "/oauth/register", + new OAuthProxyRegistrationTestRequest + { + RedirectUris = [ClientRedirectUri], + TokenEndpointAuthMethod = "none", + }, + OAuthProxyTestJsonContext.Default.OAuthProxyRegistrationTestRequest, + TestContext.Current.CancellationToken); + + Assert.Equal(HttpStatusCode.Forbidden, response.StatusCode); + } + + [Fact] + public async Task RegistrationDefaultsToAuthorizationCodeGrant() + { + await using var app = await StartServerAsync(); + + using var response = await HttpClient.PostAsJsonAsync( + "/oauth/register", + new OAuthProxyRegistrationTestRequest + { + RedirectUris = [ClientRedirectUri], + TokenEndpointAuthMethod = "none", + }, + OAuthProxyTestJsonContext.Default.OAuthProxyRegistrationTestRequest, + TestContext.Current.CancellationToken); + + Assert.Equal(HttpStatusCode.Created, response.StatusCode); + using var document = JsonDocument.Parse(await response.Content.ReadAsStreamAsync(TestContext.Current.CancellationToken)); + Assert.Equal(["authorization_code"], _lastRegistrationGrantTypes); + Assert.Equal("authorization_code", document.RootElement.GetProperty("grant_types")[0].GetString()); + Assert.Equal(1, document.RootElement.GetProperty("grant_types").GetArrayLength()); + } + + [Fact] + public async Task AuthorizationErrorRedirectIncludesIssuer() + { + await using var app = await StartServerAsync(); + var registration = await RegisterClientAsync(); + var verifier = WebEncoders.Base64UrlEncode(RandomNumberGenerator.GetBytes(48)); + var challenge = WebEncoders.Base64UrlEncode(SHA256.HashData(Encoding.ASCII.GetBytes(verifier))); + var url = QueryHelpers.AddQueryString( + "/oauth/authorize", + new Dictionary + { + ["client_id"] = registration.ClientId, + ["redirect_uri"] = ClientRedirectUri, + ["response_type"] = "token", + ["code_challenge"] = challenge, + ["code_challenge_method"] = "S256", + ["state"] = "client-state", + }); + + using var response = await HttpClient.GetAsync(url, TestContext.Current.CancellationToken); + + Assert.Equal(HttpStatusCode.Redirect, response.StatusCode); + var query = QueryHelpers.ParseQuery(Assert.IsType(response.Headers.Location).Query); + Assert.Equal("unsupported_response_type", query["error"]); + Assert.Equal("http://localhost:5000/oauth", query["iss"]); + } + + [Fact] + public async Task AuthorizationCodeFlowSeparatesTokensAndRotatesRefreshHandle() + { + await using var app = await StartServerAsync(); + var registration = await RegisterClientAsync(); + + var authorization = await BeginAuthorizationAsync(registration.ClientId); + Assert.NotEqual(authorization.DownstreamCodeChallenge, authorization.UpstreamCodeChallenge); + Assert.Equal(Resource, authorization.UpstreamResource); + + var proxyCode = await CompleteCallbackAsync(authorization.TransactionId); + var firstToken = await ExchangeCodeAsync(registration.ClientId, proxyCode, authorization.CodeVerifier); + + Assert.Equal("proxy-access-1", firstToken.AccessToken); + Assert.NotEqual("upstream-access-1", firstToken.AccessToken); + Assert.False(string.IsNullOrEmpty(firstToken.RefreshToken)); + Assert.Equal("no-store", firstToken.CacheControl); + + using var codeReplayResponse = await PostTokenAsync(new Dictionary + { + ["grant_type"] = "authorization_code", + ["code"] = proxyCode, + ["client_id"] = registration.ClientId, + ["redirect_uri"] = ClientRedirectUri, + ["code_verifier"] = authorization.CodeVerifier, + }); + Assert.Equal(HttpStatusCode.BadRequest, codeReplayResponse.StatusCode); + + var secondToken = await RefreshAsync(registration.ClientId, firstToken.RefreshToken!); + Assert.Equal("proxy-access-2", secondToken.AccessToken); + Assert.NotEqual(firstToken.RefreshToken, secondToken.RefreshToken); + Assert.Equal(2, _upstream.TokenRequests.Count); + Assert.Equal("authorization_code", _upstream.TokenRequests[0]["grant_type"]); + Assert.Equal(Resource, _upstream.TokenRequests[0]["resource"]); + Assert.Equal("refresh_token", _upstream.TokenRequests[1]["grant_type"]); + Assert.Equal(Resource, _upstream.TokenRequests[1]["resource"]); + + using var replayResponse = await PostTokenAsync(new Dictionary + { + ["grant_type"] = "refresh_token", + ["refresh_token"] = firstToken.RefreshToken!, + ["client_id"] = registration.ClientId, + }); + Assert.Equal(HttpStatusCode.BadRequest, replayResponse.StatusCode); + + using var revokedDescendantResponse = await PostTokenAsync(new Dictionary + { + ["grant_type"] = "refresh_token", + ["refresh_token"] = secondToken.RefreshToken!, + ["client_id"] = registration.ClientId, + }); + Assert.Equal(HttpStatusCode.BadRequest, revokedDescendantResponse.StatusCode); + + Assert.All(_store.Values, value => + { + var serialized = Encoding.UTF8.GetString(value.Span); + Assert.DoesNotContain("upstream-access", serialized, StringComparison.Ordinal); + Assert.DoesNotContain("upstream-refresh", serialized, StringComparison.Ordinal); + Assert.DoesNotContain(ClientRedirectUri, serialized, StringComparison.Ordinal); + }); + Assert.All(_store.Keys, key => + { + Assert.DoesNotContain(firstToken.RefreshToken!, key, StringComparison.Ordinal); + Assert.DoesNotContain(secondToken.RefreshToken!, key, StringComparison.Ordinal); + }); + } + + [Fact] + public async Task TokenFactoryReceivesOnlyScopesGrantedByUpstream() + { + await using var app = await StartServerAsync(); + var registration = await RegisterClientAsync("mcp:tools mcp:read"); + var authorization = await BeginAuthorizationAsync(registration.ClientId, "mcp:tools mcp:read"); + var proxyCode = await CompleteCallbackAsync(authorization.TransactionId); + + await ExchangeCodeAsync(registration.ClientId, proxyCode, authorization.CodeVerifier); + + Assert.Equal(["mcp:tools"], _lastMintedScopes); + } + + [Fact] + public async Task ClientWithoutRefreshGrantDoesNotReceiveRefreshToken() + { + await using var app = await StartServerAsync(); + var registration = await RegisterClientAsync(includeRefreshGrant: false); + var authorization = await BeginAuthorizationAsync(registration.ClientId); + var proxyCode = await CompleteCallbackAsync(authorization.TransactionId); + + var token = await ExchangeCodeAsync(registration.ClientId, proxyCode, authorization.CodeVerifier); + + Assert.Null(token.RefreshToken); + } + + [Fact] + public async Task AuthorizationUsesValidatedOidcDiscovery() + { + _useDiscovery = true; + await using var app = await StartServerAsync(); + var registration = await RegisterClientAsync(); + + var authorization = await BeginAuthorizationAsync(registration.ClientId); + + Assert.False(string.IsNullOrEmpty(authorization.TransactionId)); + Assert.Equal(1, _upstream.DiscoveryRequests); + } + + [Fact] + public async Task TokenEndpointRejectsUpstreamTokenPassthrough() + { + _returnUpstreamAccessToken = true; + await using var app = await StartServerAsync(); + var registration = await RegisterClientAsync(); + var authorization = await BeginAuthorizationAsync(registration.ClientId); + var proxyCode = await CompleteCallbackAsync(authorization.TransactionId); + + using var response = await PostTokenAsync(new Dictionary + { + ["grant_type"] = "authorization_code", + ["code"] = proxyCode, + ["client_id"] = registration.ClientId, + ["redirect_uri"] = ClientRedirectUri, + ["code_verifier"] = authorization.CodeVerifier, + }); + + Assert.Equal(HttpStatusCode.InternalServerError, response.StatusCode); + using var document = JsonDocument.Parse(await response.Content.ReadAsStreamAsync(TestContext.Current.CancellationToken)); + Assert.Equal("server_error", document.RootElement.GetProperty("error").GetString()); + } + + [Fact] + public async Task TransientRefreshFailurePreservesProxyRefreshHandleForRetry() + { + await using var app = await StartServerAsync(); + var registration = await RegisterClientAsync(); + var authorization = await BeginAuthorizationAsync(registration.ClientId); + var proxyCode = await CompleteCallbackAsync(authorization.TransactionId); + var firstToken = await ExchangeCodeAsync(registration.ClientId, proxyCode, authorization.CodeVerifier); + + _upstream.FailNextRefreshRequest = true; + using var failedResponse = await PostTokenAsync(new Dictionary + { + ["grant_type"] = "refresh_token", + ["refresh_token"] = firstToken.RefreshToken!, + ["client_id"] = registration.ClientId, + }); + + Assert.Equal(HttpStatusCode.ServiceUnavailable, failedResponse.StatusCode); + using var error = JsonDocument.Parse(await failedResponse.Content.ReadAsStreamAsync(TestContext.Current.CancellationToken)); + Assert.Equal("temporarily_unavailable", error.RootElement.GetProperty("error").GetString()); + + var retriedToken = await RefreshAsync(registration.ClientId, firstToken.RefreshToken!); + Assert.Equal("proxy-access-2", retriedToken.AccessToken); + Assert.NotEqual(firstToken.RefreshToken, retriedToken.RefreshToken); + } + + [Fact] + public async Task ConcurrentRefreshReplayRevokesInFlightTokenFamily() + { + await using var app = await StartServerAsync(); + var registration = await RegisterClientAsync(); + var authorization = await BeginAuthorizationAsync(registration.ClientId); + var proxyCode = await CompleteCallbackAsync(authorization.TransactionId); + var firstToken = await ExchangeCodeAsync(registration.ClientId, proxyCode, authorization.CodeVerifier); + + _store.BlockNextRefreshTake = true; + var refreshParameters = new Dictionary + { + ["grant_type"] = "refresh_token", + ["refresh_token"] = firstToken.RefreshToken!, + ["client_id"] = registration.ClientId, + }; + var delayedRefresh = PostTokenAsync(refreshParameters); + await _store.RefreshTakeStarted.Task.WaitAsync(TestContext.Current.CancellationToken); + + using var replayResponse = await PostTokenAsync(refreshParameters); + Assert.Equal(HttpStatusCode.BadRequest, replayResponse.StatusCode); + + _store.ReleaseRefreshTake.TrySetResult(true); + using var delayedResponse = await delayedRefresh; + Assert.Equal(HttpStatusCode.BadRequest, delayedResponse.StatusCode); + } + + [Fact] + public async Task RefreshReplayDuringUpstreamExchangePreventsDescendantPublication() + { + await using var app = await StartServerAsync(); + var registration = await RegisterClientAsync(); + var authorization = await BeginAuthorizationAsync(registration.ClientId); + var proxyCode = await CompleteCallbackAsync(authorization.TransactionId); + var firstToken = await ExchangeCodeAsync(registration.ClientId, proxyCode, authorization.CodeVerifier); + + _upstream.BlockNextRefreshRequest = true; + var refreshParameters = new Dictionary + { + ["grant_type"] = "refresh_token", + ["refresh_token"] = firstToken.RefreshToken!, + ["client_id"] = registration.ClientId, + }; + var inFlightRefresh = PostTokenAsync(refreshParameters); + await _upstream.RefreshRequestStarted.Task.WaitAsync(TestContext.Current.CancellationToken); + + using var replayResponse = await PostTokenAsync(refreshParameters); + Assert.Equal(HttpStatusCode.BadRequest, replayResponse.StatusCode); + + _upstream.ReleaseRefreshRequest.TrySetResult(true); + using var inFlightResponse = await inFlightRefresh; + Assert.Equal(HttpStatusCode.BadRequest, inFlightResponse.StatusCode); + } + + [Fact] + public async Task UpstreamRefreshTimeoutPreservesProxyRefreshHandleForRetry() + { + await using var app = await StartServerAsync(); + var registration = await RegisterClientAsync(); + var authorization = await BeginAuthorizationAsync(registration.ClientId); + var proxyCode = await CompleteCallbackAsync(authorization.TransactionId); + var firstToken = await ExchangeCodeAsync(registration.ClientId, proxyCode, authorization.CodeVerifier); + + _upstream.ThrowNextRefreshTimeout = true; + using var timeoutResponse = await PostTokenAsync(new Dictionary + { + ["grant_type"] = "refresh_token", + ["refresh_token"] = firstToken.RefreshToken!, + ["client_id"] = registration.ClientId, + }); + Assert.Equal(HttpStatusCode.ServiceUnavailable, timeoutResponse.StatusCode); + + var retriedToken = await RefreshAsync(registration.ClientId, firstToken.RefreshToken!); + Assert.Equal("proxy-access-2", retriedToken.AccessToken); + } + + [Theory] + [InlineData(HttpStatusCode.RequestTimeout)] + [InlineData(HttpStatusCode.TooManyRequests)] + public async Task MalformedTransientResponsePreservesProxyRefreshHandleForRetry(HttpStatusCode statusCode) + { + await using var app = await StartServerAsync(); + var registration = await RegisterClientAsync(); + var authorization = await BeginAuthorizationAsync(registration.ClientId); + var proxyCode = await CompleteCallbackAsync(authorization.TransactionId); + var firstToken = await ExchangeCodeAsync(registration.ClientId, proxyCode, authorization.CodeVerifier); + + _upstream.MalformedRefreshFailureStatus = statusCode; + using var throttledResponse = await PostTokenAsync(new Dictionary + { + ["grant_type"] = "refresh_token", + ["refresh_token"] = firstToken.RefreshToken!, + ["client_id"] = registration.ClientId, + }); + Assert.Equal(HttpStatusCode.ServiceUnavailable, throttledResponse.StatusCode); + + var retriedToken = await RefreshAsync(registration.ClientId, firstToken.RefreshToken!); + Assert.Equal("proxy-access-2", retriedToken.AccessToken); + } + + [Fact] + public async Task CallbackRequiresInitiatingBrowserCookie() + { + await using var app = await StartServerAsync(); + var registration = await RegisterClientAsync(); + var authorization = await BeginAuthorizationAsync(registration.ClientId); + + using var clientWithoutCookies = new HttpClient(new SocketsHttpHandler + { + AllowAutoRedirect = false, + ConnectCallback = SocketsHttpHandler.ConnectCallback, + }) + { + BaseAddress = HttpClient.BaseAddress, + }; + using var response = await clientWithoutCookies.GetAsync( + $"/oauth/callback?code=upstream-code&state={Uri.EscapeDataString(authorization.TransactionId)}", + TestContext.Current.CancellationToken); + + Assert.Equal(HttpStatusCode.BadRequest, response.StatusCode); + Assert.Empty(_upstream.TokenRequests); + } + + private async Task StartServerAsync() + { + var app = Builder.Build(); + app.UseAuthorization(); + app.MapMcpOAuthProxy(); + await app.StartAsync(TestContext.Current.CancellationToken); + return app; + } + + private async Task RegisterClientAsync( + string scope = "mcp:tools", + bool includeRefreshGrant = true) + { + using var response = await HttpClient.PostAsJsonAsync( + "/oauth/register", + new OAuthProxyRegistrationTestRequest + { + RedirectUris = [ClientRedirectUri], + TokenEndpointAuthMethod = "client_secret_post", + GrantTypes = includeRefreshGrant + ? new[] { "authorization_code", "refresh_token" } + : new[] { "authorization_code" }, + ResponseTypes = ["code"], + Scope = scope, + ClientName = "Test MCP Client", + ApplicationType = "native", + }, + OAuthProxyTestJsonContext.Default.OAuthProxyRegistrationTestRequest, + TestContext.Current.CancellationToken); + Assert.Equal(HttpStatusCode.Created, response.StatusCode); + using var document = JsonDocument.Parse(await response.Content.ReadAsStreamAsync(TestContext.Current.CancellationToken)); + Assert.Equal("none", document.RootElement.GetProperty("token_endpoint_auth_method").GetString()); + Assert.Equal(scope, document.RootElement.GetProperty("scope").GetString()); + Assert.Equal("Test MCP Client", document.RootElement.GetProperty("client_name").GetString()); + Assert.Equal("native", document.RootElement.GetProperty("application_type").GetString()); + return new(document.RootElement.GetProperty("client_id").GetString()!); + } + + private async Task BeginAuthorizationAsync(string clientId, string scope = "mcp:tools") + { + var verifier = WebEncoders.Base64UrlEncode(RandomNumberGenerator.GetBytes(48)); + var challenge = WebEncoders.Base64UrlEncode(SHA256.HashData(Encoding.ASCII.GetBytes(verifier))); + var url = QueryHelpers.AddQueryString( + "/oauth/authorize", + new Dictionary + { + ["client_id"] = clientId, + ["redirect_uri"] = ClientRedirectUri, + ["response_type"] = "code", + ["code_challenge"] = challenge, + ["code_challenge_method"] = "S256", + ["scope"] = scope, + ["state"] = "client-state", + ["resource"] = Resource, + }); + + using var response = await HttpClient.GetAsync(url, TestContext.Current.CancellationToken); + Assert.Equal(HttpStatusCode.Redirect, response.StatusCode); + var location = Assert.IsType(response.Headers.Location); + Assert.Equal("/upstream/authorize", location.AbsolutePath); + var query = QueryHelpers.ParseQuery(location.Query); + Assert.Equal("S256", query["code_challenge_method"]); + return new( + verifier, + challenge, + query["code_challenge"].ToString(), + query["state"].ToString(), + query["resource"].ToString()); + } + + private async Task CompleteCallbackAsync(string transactionId) + { + using var response = await HttpClient.GetAsync( + $"/oauth/callback?code=upstream-code&state={Uri.EscapeDataString(transactionId)}", + TestContext.Current.CancellationToken); + Assert.Equal(HttpStatusCode.Redirect, response.StatusCode); + var location = Assert.IsType(response.Headers.Location); + Assert.Equal(ClientRedirectUri, location.GetLeftPart(UriPartial.Path)); + var query = QueryHelpers.ParseQuery(location.Query); + Assert.Equal("client-state", query["state"]); + Assert.Equal("http://localhost:5000/oauth", query["iss"]); + return query["code"].ToString(); + } + + private async Task ExchangeCodeAsync(string clientId, string code, string verifier) + { + using var response = await PostTokenAsync(new Dictionary + { + ["grant_type"] = "authorization_code", + ["code"] = code, + ["client_id"] = clientId, + ["redirect_uri"] = ClientRedirectUri, + ["code_verifier"] = verifier, + }); + Assert.Equal(HttpStatusCode.OK, response.StatusCode); + return await ReadTokenAsync(response); + } + + private async Task RefreshAsync(string clientId, string refreshToken) + { + using var response = await PostTokenAsync(new Dictionary + { + ["grant_type"] = "refresh_token", + ["refresh_token"] = refreshToken, + ["client_id"] = clientId, + }); + Assert.Equal(HttpStatusCode.OK, response.StatusCode); + return await ReadTokenAsync(response); + } + + private Task PostTokenAsync(Dictionary parameters) => + HttpClient.PostAsync( + "/oauth/token", + new FormUrlEncodedContent(parameters), + TestContext.Current.CancellationToken); + + private static async Task ReadTokenAsync(HttpResponseMessage response) + { + using var document = JsonDocument.Parse(await response.Content.ReadAsStreamAsync(TestContext.Current.CancellationToken)); + var root = document.RootElement; + return new( + root.GetProperty("access_token").GetString()!, + root.TryGetProperty("refresh_token", out var refreshToken) ? refreshToken.GetString() : null, + response.Headers.CacheControl?.ToString()); + } + + private sealed record ClientRegistration(string ClientId); + + private sealed record AuthorizationStart( + string CodeVerifier, + string DownstreamCodeChallenge, + string UpstreamCodeChallenge, + string TransactionId, + string UpstreamResource); + + private sealed record TokenResult(string AccessToken, string? RefreshToken, string? CacheControl); + + private sealed class RecordingStore : IMcpOAuthProxyStore + { + private readonly ConcurrentDictionary> _values = new(StringComparer.Ordinal); + + public IEnumerable> Values => _values.Values; + + public IEnumerable Keys => _values.Keys; + + public bool BlockNextRefreshTake { get; set; } + + public TaskCompletionSource RefreshTakeStarted { get; } = + new(TaskCreationOptions.RunContinuationsAsynchronously); + + public TaskCompletionSource ReleaseRefreshTake { get; } = + new(TaskCreationOptions.RunContinuationsAsynchronously); + + public ValueTask SetAsync(string key, ReadOnlyMemory value, TimeSpan lifetime, CancellationToken cancellationToken = default) + { + _values[key] = value.ToArray(); + return ValueTask.CompletedTask; + } + + public ValueTask?> GetAsync(string key, CancellationToken cancellationToken = default) => + ValueTask.FromResult(_values.TryGetValue(key, out var value) ? (ReadOnlyMemory?)value : null); + + public async ValueTask?> TakeAsync( + string key, + CancellationToken cancellationToken = default) + { + var result = _values.TryRemove(key, out var value) ? value : (ReadOnlyMemory?)null; + if (BlockNextRefreshTake && key.Contains(":refresh:", StringComparison.Ordinal)) + { + BlockNextRefreshTake = false; + RefreshTakeStarted.TrySetResult(true); + await ReleaseRefreshTake.Task.WaitAsync(cancellationToken); + } + + return result; + } + + public ValueTask RemoveAsync(string key, CancellationToken cancellationToken = default) + { + _values.TryRemove(key, out _); + return ValueTask.CompletedTask; + } + } + + private sealed class StubHttpClientFactory : IHttpClientFactory + { + private readonly HttpClient _client; + + public StubHttpClientFactory() + { + _client = new HttpClient(new StubHandler(HandleAsync)); + } + + public List> TokenRequests { get; } = []; + + public int DiscoveryRequests { get; private set; } + + public bool FailNextRefreshRequest { get; set; } + + public bool BlockNextRefreshRequest { get; set; } + + public bool ThrowNextRefreshTimeout { get; set; } + + public HttpStatusCode? MalformedRefreshFailureStatus { get; set; } + + public TaskCompletionSource RefreshRequestStarted { get; } = + new(TaskCreationOptions.RunContinuationsAsynchronously); + + public TaskCompletionSource ReleaseRefreshRequest { get; } = + new(TaskCreationOptions.RunContinuationsAsynchronously); + + public HttpClient CreateClient(string name) => _client; + + private async Task HandleAsync(HttpRequestMessage request, CancellationToken cancellationToken) + { + if (request.Method == HttpMethod.Get) + { + Assert.Equal("/upstream/.well-known/openid-configuration", request.RequestUri?.AbsolutePath); + DiscoveryRequests++; + return new HttpResponseMessage(HttpStatusCode.OK) + { + Content = new StringContent( + """ + {"issuer":"http://localhost:5000/upstream","authorization_endpoint":"http://localhost:5000/upstream/authorize","token_endpoint":"http://localhost:5000/upstream/token","code_challenge_methods_supported":["S256"]} + """, + Encoding.UTF8, + "application/json"), + }; + } + + Assert.Equal("/upstream/token", request.RequestUri?.AbsolutePath); + var body = await request.Content!.ReadAsStringAsync(cancellationToken); + var parsed = QueryHelpers.ParseQuery(body); + var fields = parsed.ToDictionary(static pair => pair.Key, static pair => pair.Value.ToString(), StringComparer.Ordinal); + TokenRequests.Add(fields); + + var isRefresh = fields["grant_type"] == "refresh_token"; + if (isRefresh && BlockNextRefreshRequest) + { + BlockNextRefreshRequest = false; + RefreshRequestStarted.TrySetResult(true); + await ReleaseRefreshRequest.Task.WaitAsync(cancellationToken); + } + + if (isRefresh && ThrowNextRefreshTimeout) + { + ThrowNextRefreshTimeout = false; + throw new TaskCanceledException("The simulated upstream request timed out."); + } + + if (isRefresh && MalformedRefreshFailureStatus is HttpStatusCode statusCode) + { + MalformedRefreshFailureStatus = null; + return new HttpResponseMessage(statusCode) + { + Content = new StringContent("not-json", Encoding.UTF8, "application/json"), + }; + } + + if (isRefresh && FailNextRefreshRequest) + { + FailNextRefreshRequest = false; + return new HttpResponseMessage(HttpStatusCode.ServiceUnavailable) + { + Content = new StringContent( + """{"error":"temporarily_unavailable"}""", + Encoding.UTF8, + "application/json"), + }; + } + + var suffix = isRefresh ? "2" : "1"; + return new HttpResponseMessage(HttpStatusCode.OK) + { + Content = new StringContent( + $$"""{"access_token":"upstream-access-{{suffix}}","refresh_token":"upstream-refresh-{{suffix}}","id_token":"upstream-id-{{suffix}}","token_type":"Bearer","expires_in":3600,"scope":"mcp:tools"}""", + Encoding.UTF8, + "application/json"), + }; + } + } + + private sealed class StubHandler( + Func> handler) : HttpMessageHandler + { + protected override Task SendAsync(HttpRequestMessage request, CancellationToken cancellationToken) => + handler(request, cancellationToken); + } +} + +internal sealed class OAuthProxyRegistrationTestRequest +{ + [JsonPropertyName("redirect_uris")] + public required string[] RedirectUris { get; init; } + + [JsonPropertyName("token_endpoint_auth_method")] + public string? TokenEndpointAuthMethod { get; init; } + + [JsonPropertyName("grant_types")] + public string[]? GrantTypes { get; init; } + + [JsonPropertyName("response_types")] + public string[]? ResponseTypes { get; init; } + + [JsonPropertyName("scope")] + public string? Scope { get; init; } + + [JsonPropertyName("client_name")] + public string? ClientName { get; init; } + + [JsonPropertyName("application_type")] + public string? ApplicationType { get; init; } +} + +[JsonSerializable(typeof(OAuthProxyRegistrationTestRequest))] +internal sealed partial class OAuthProxyTestJsonContext : JsonSerializerContext;