diff --git a/Directory.Packages.props b/Directory.Packages.props index 565d01689..9de34f11d 100644 --- a/Directory.Packages.props +++ b/Directory.Packages.props @@ -25,6 +25,7 @@ + diff --git a/docs/concepts/transports/transports.md b/docs/concepts/transports/transports.md index bb4e155f2..3917f645b 100644 --- a/docs/concepts/transports/transports.md +++ b/docs/concepts/transports/transports.md @@ -145,6 +145,21 @@ var transport = new HttpClientTransport(new HttpClientTransportOptions }); ``` +#### Controlling client timeout timers in tests + + controls + and +. It defaults to +`TimeProvider.System`. Tests can set it to a `FakeTimeProvider` from the +`Microsoft.Extensions.TimeProvider.Testing` package and advance time explicitly instead +of waiting for real deadlines. + +Use synchronization signals to wait until the request or authorization phase being tested +has started before advancing the clock. Fake time does not control HTTP processing or +task scheduling, so keep a real-time outer test deadline as a safety bound. The client +time provider does not change `HttpClient.Timeout`, OAuth token expiration, or +. + #### Resuming sessions Streamable HTTP supports session resumption. Save the session ID, server capabilities, and server info from the original session, then use to reconnect: diff --git a/src/ModelContextProtocol.Core/Authentication/ClientOAuthProvider.cs b/src/ModelContextProtocol.Core/Authentication/ClientOAuthProvider.cs index 785e3cc2e..872d0572b 100644 --- a/src/ModelContextProtocol.Core/Authentication/ClientOAuthProvider.cs +++ b/src/ModelContextProtocol.Core/Authentication/ClientOAuthProvider.cs @@ -198,7 +198,11 @@ internal override async Task SendAsync(HttpRequestMessage r if (request.Headers.Authorization is null && request.RequestUri is not null) { string? accessToken; - (accessToken, attemptedRefresh) = await GetAccessTokenSilentAsync(request.RequestUri, cancellationToken).ConfigureAwait(false); + using (message?.Context?.RequestTimeout?.Suspend()) + { + cancellationToken.ThrowIfCancellationRequested(); + (accessToken, attemptedRefresh) = await GetAccessTokenSilentAsync(request.RequestUri, cancellationToken).ConfigureAwait(false); + } if (!string.IsNullOrEmpty(accessToken)) { @@ -308,7 +312,12 @@ private async Task HandleUnauthorizedResponseAsync( throw new McpException($"The server does not support the '{BearerScheme}' authentication scheme. Server supports: [{serverSchemes}]."); } - var accessToken = await GetAccessTokenAsync(response, attemptedRefresh, usedAccessToken, cancellationToken).ConfigureAwait(false); + string accessToken; + using (originalJsonRpcMessage?.Context?.RequestTimeout?.Suspend()) + { + cancellationToken.ThrowIfCancellationRequested(); + accessToken = await GetAccessTokenAsync(response, attemptedRefresh, usedAccessToken, cancellationToken).ConfigureAwait(false); + } using var retryRequest = new HttpRequestMessage(originalRequest.Method, originalRequest.RequestUri); diff --git a/src/ModelContextProtocol.Core/Client/AutoDetectingClientSessionTransport.cs b/src/ModelContextProtocol.Core/Client/AutoDetectingClientSessionTransport.cs index 7563acd10..aa36e27f1 100644 --- a/src/ModelContextProtocol.Core/Client/AutoDetectingClientSessionTransport.cs +++ b/src/ModelContextProtocol.Core/Client/AutoDetectingClientSessionTransport.cs @@ -161,8 +161,15 @@ private async Task InitializeSseTransportAsync(JsonRpcMessage message, HttpReque try { LogAttemptingSSE(_name); + // Discovery has been abandoned. Stop its timer rather than restarting it after + // the legacy GET; caller/initialization cancellation and ConnectionTimeout still apply. + message.Context?.RequestTimeout?.Stop(); await sseTransport.ConnectAsync(cancellationToken).ConfigureAwait(false); - await sseTransport.SendMessageAsync(message, cancellationToken).ConfigureAwait(false); + + if (message is not JsonRpcRequest { Method: RequestMethods.ServerDiscover }) + { + await sseTransport.SendMessageAsync(message, cancellationToken).ConfigureAwait(false); + } LogUsingSSE(_name); ActiveTransport = sseTransport; @@ -186,6 +193,12 @@ private async Task InitializeSseTransportAsync(JsonRpcMessage message, HttpReque await sseTransport.DisposeAsync().ConfigureAwait(false); throw; } + + if (message is JsonRpcRequest { Method: RequestMethods.ServerDiscover }) + { + // Let the client apply its initialization and minimum-version policy; never send discover over SSE. + throw new ServerDiscoverSkippedForSseException(); + } } public async ValueTask DisposeAsync() diff --git a/src/ModelContextProtocol.Core/Client/McpClientImpl.cs b/src/ModelContextProtocol.Core/Client/McpClientImpl.cs index d1f2a9d7a..4dece667f 100644 --- a/src/ModelContextProtocol.Core/Client/McpClientImpl.cs +++ b/src/ModelContextProtocol.Core/Client/McpClientImpl.cs @@ -286,8 +286,9 @@ public async Task ConnectAsync(CancellationToken cancellationToken = default) _ = _sessionHandler.ProcessMessagesAsync(CancellationToken.None); // Perform initialization sequence - using var initializationCts = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken); - initializationCts.CancelAfter(_options.InitializationTimeout); + var timeProvider = _options.TimeProvider; + using var initializationTimeout = new RequestTimeout(_options.InitializationTimeout, timeProvider, cancellationToken); + var initializationToken = initializationTimeout.Token; try { @@ -296,31 +297,39 @@ public async Task ConnectAsync(CancellationToken cancellationToken = default) // capabilities and then begins sending normal RPCs that carry protocolVersion / // clientInfo / clientCapabilities in their per-request _meta. A null ProtocolVersion // prefers the 2026-07-28 revision and automatically falls back to the initialize - // handshake when the server doesn't support it. The initialize branch below runs only when - // the caller explicitly pins a version that still supports Streamable HTTP sessions (opting out of the default). + // handshake when the server doesn't support it. HTTP+SSE defaults to the initialize handshake, + // including when AutoDetect selects it while sending the discovery probe. if (_options.ProtocolVersion is null || McpProtocolVersions.RequiresPerRequestMetadata(_options.ProtocolVersion)) { string preferredVersion = _options.ProtocolVersion ?? McpProtocolVersions.July2026ProtocolVersion; DiscoverResult? discoverResult = null; - bool fallbackToInitialize = false; + // Modern-over-SSE is unusual, but honor an explicit version choice instead of forcing initialize. + bool fallbackToInitialize = _transport is SseClientSessionTransport && _options.ProtocolVersion is null; IList? serverSupportedVersions = null; string discoverVersion = preferredVersion; // Apply a probe timeout so dual-path clients don't block forever waiting for an // initialize-handshake server that silently drops unknown methods (per stdio.mdx fallback rules). // The probe timeout is configurable via McpClientOptions.DiscoverProbeTimeout and is - // always bounded by InitializationTimeout (only applied when it is the tighter bound). + // always bounded by InitializationTimeout. OAuth can suspend only the probe timer. var probeTimeout = _options.DiscoverProbeTimeout; - using var probeCts = CancellationTokenSource.CreateLinkedTokenSource(initializationCts.Token); - if (_options.InitializationTimeout > probeTimeout) - { - probeCts.CancelAfter(probeTimeout); - } + using var probeTimeoutController = !fallbackToInitialize && probeTimeout != Timeout.InfiniteTimeSpan && + (_options.InitializationTimeout == Timeout.InfiniteTimeSpan || probeTimeout < _options.InitializationTimeout) + ? new RequestTimeout(probeTimeout, timeProvider, initializationToken) + : null; + var probeToken = probeTimeoutController?.Token ?? initializationToken; try { - discoverResult = await SendDiscoverAsync(discoverVersion, probeCts.Token).ConfigureAwait(false); + if (!fallbackToInitialize) + { + discoverResult = await SendDiscoverAsync(discoverVersion, probeToken).ConfigureAwait(false); + } + } + catch (ServerDiscoverSkippedForSseException) + { + fallbackToInitialize = true; } catch (UnsupportedProtocolVersionException ex) { @@ -346,7 +355,7 @@ public async Task ConnectAsync(CancellationToken cancellationToken = default) } discoverVersion = retryVersion; - discoverResult = await SendDiscoverAsync(discoverVersion, probeCts.Token).ConfigureAwait(false); + discoverResult = await SendDiscoverAsync(discoverVersion, probeToken).ConfigureAwait(false); } else { @@ -391,7 +400,7 @@ public async Task ConnectAsync(CancellationToken cancellationToken = default) // server, so fall back. Other statuses stay uncaught and surface to the caller. fallbackToInitialize = true; } - catch (OperationCanceledException) when (probeCts.IsCancellationRequested && !initializationCts.IsCancellationRequested) + catch (OperationCanceledException) when (probeToken.IsCancellationRequested && !initializationToken.IsCancellationRequested) { // Probe timeout elapsed without a response. Per stdio.mdx fallback rules, no // response within a reasonable timeout means the server requires initialize. Fall back. @@ -433,7 +442,7 @@ public async Task ConnectAsync(CancellationToken cancellationToken = default) : $"Server-supported versions: {string.Join(", ", serverSupportedVersions)}.")); } - await PerformInitializeHandshakeAsync(fallbackVersion, initializationCts.Token).ConfigureAwait(false); + await PerformInitializeHandshakeAsync(fallbackVersion, initializationToken).ConfigureAwait(false); } else { @@ -465,6 +474,7 @@ async Task SendDiscoverAsync(string protocolVersion, Cancellatio new DiscoverRequestParams(), McpJsonUtilities.JsonContext.Default.DiscoverRequestParams, McpJsonUtilities.JsonContext.Default.DiscoverResult, + context: probeTimeoutController is null ? null : new JsonRpcMessageContext { RequestTimeout = probeTimeoutController }, cancellationToken: cancellationToken).ConfigureAwait(false); } } @@ -474,10 +484,10 @@ async Task SendDiscoverAsync(string protocolVersion, Cancellatio // ProtocolVersion that still supports Streamable HTTP sessions (opting out of the default), so // _options.ProtocolVersion is non-null here. string requestProtocol = _options.ProtocolVersion ?? McpProtocolVersions.November2025ProtocolVersion; - await PerformInitializeHandshakeAsync(requestProtocol, initializationCts.Token).ConfigureAwait(false); + await PerformInitializeHandshakeAsync(requestProtocol, initializationToken).ConfigureAwait(false); } } - catch (OperationCanceledException oce) when (initializationCts.IsCancellationRequested && !cancellationToken.IsCancellationRequested) + catch (OperationCanceledException oce) when (initializationToken.IsCancellationRequested && !cancellationToken.IsCancellationRequested) { LogClientInitializationTimeout(_endpointName); throw new TimeoutException("Initialization timed out", oce); diff --git a/src/ModelContextProtocol.Core/Client/McpClientOptions.cs b/src/ModelContextProtocol.Core/Client/McpClientOptions.cs index 61a0613df..7d9c701cc 100644 --- a/src/ModelContextProtocol.Core/Client/McpClientOptions.cs +++ b/src/ModelContextProtocol.Core/Client/McpClientOptions.cs @@ -70,6 +70,10 @@ public sealed class McpClientOptions /// negotiates a different version. To try more than one version, leave this unset for automatic fallback /// or retry the connection with a different value. /// + /// + /// HTTP+SSE connections use the initialize handshake by default. + /// An explicit protocol version is attempted when is selected. + /// /// public string? ProtocolVersion { get; set; } @@ -86,12 +90,36 @@ public sealed class McpClientOptions /// an exception is thrown. /// /// + /// This timeout includes OAuth token acquisition performed during the handshake. Neither this timeout nor + /// caller cancellation is suspended while authenticating. Transport connection establishment that precedes + /// the handshake, such as an explicitly selected SSE connection, retains its transport-specific timeout. + /// + /// /// Setting an appropriate timeout prevents the client from hanging indefinitely when /// connecting to unresponsive servers. /// /// public TimeSpan InitializationTimeout { get; set; } = TimeSpan.FromSeconds(60); + /// + /// Gets or sets the time provider used for and . + /// + /// The time provider. The default is . + /// + /// This provider does not control HTTP client timeouts, OAuth token expiration, or transport-specific + /// deadlines such as . + /// + /// The value is . + public TimeProvider TimeProvider + { + get; + set + { + Throw.IfNull(value); + field = value; + } + } = TimeProvider.System; + /// /// Gets or sets the timeout applied to the server/discover probe that the client issues /// before falling back to the initialize handshake. @@ -121,6 +149,12 @@ public sealed class McpClientOptions /// greater than or equal to , the probe is effectively bounded by /// alone. /// + /// + /// SDK OAuth token acquisition, including metadata discovery, registration, interactive authorization, + /// and token refresh or exchange, is excluded from the probe timeout. After token acquisition, the + /// HTTP request gets a fresh full probe budget, covering both response headers and body processing. + /// and caller cancellation continue to apply during authentication. + /// /// /// /// The value is not positive and is not . diff --git a/src/ModelContextProtocol.Core/Client/ServerDiscoverSkippedForSseException.cs b/src/ModelContextProtocol.Core/Client/ServerDiscoverSkippedForSseException.cs new file mode 100644 index 000000000..b6cb1e448 --- /dev/null +++ b/src/ModelContextProtocol.Core/Client/ServerDiscoverSkippedForSseException.cs @@ -0,0 +1,5 @@ +namespace ModelContextProtocol.Client; + +/// Signals that AutoDetect selected SSE and the client must initialize instead of discovering. +internal sealed class ServerDiscoverSkippedForSseException() + : Exception("AutoDetect selected HTTP+SSE. Use initialize instead of server/discover."); diff --git a/src/ModelContextProtocol.Core/McpSession.Methods.cs b/src/ModelContextProtocol.Core/McpSession.Methods.cs index 9ad210fbb..0bd5368c6 100644 --- a/src/ModelContextProtocol.Core/McpSession.Methods.cs +++ b/src/ModelContextProtocol.Core/McpSession.Methods.cs @@ -38,7 +38,7 @@ public ValueTask SendRequestAsync( serializerOptions.GetTypeInfo(), serializerOptions.GetTypeInfo(), requestId, - cancellationToken); + cancellationToken: cancellationToken); } /// @@ -51,6 +51,7 @@ public ValueTask SendRequestAsync( /// The type information for request parameter serialization. /// The type information for result deserialization. /// The request ID for the request. + /// Non-serialized runtime context for the request. /// The to monitor for cancellation requests. The default is . /// A task that represents the asynchronous operation. The task result contains the deserialized result. internal async ValueTask SendRequestAsync( @@ -59,6 +60,7 @@ internal async ValueTask SendRequestAsync( JsonTypeInfo parametersTypeInfo, JsonTypeInfo resultTypeInfo, RequestId requestId = default, + JsonRpcMessageContext? context = null, CancellationToken cancellationToken = default) where TResult : notnull { @@ -71,6 +73,7 @@ internal async ValueTask SendRequestAsync( Id = requestId, Method = method, Params = JsonSerializer.SerializeToNode(parameters, parametersTypeInfo), + Context = context, }; JsonRpcResponse response = await SendRequestAsync(jsonRpcRequest, cancellationToken).ConfigureAwait(false); diff --git a/src/ModelContextProtocol.Core/ModelContextProtocol.Core.csproj b/src/ModelContextProtocol.Core/ModelContextProtocol.Core.csproj index 3fbef0377..f3eec2350 100644 --- a/src/ModelContextProtocol.Core/ModelContextProtocol.Core.csproj +++ b/src/ModelContextProtocol.Core/ModelContextProtocol.Core.csproj @@ -42,6 +42,7 @@ + diff --git a/src/ModelContextProtocol.Core/Protocol/JsonRpcMessageContext.cs b/src/ModelContextProtocol.Core/Protocol/JsonRpcMessageContext.cs index 0b56caa29..676c89d1c 100644 --- a/src/ModelContextProtocol.Core/Protocol/JsonRpcMessageContext.cs +++ b/src/ModelContextProtocol.Core/Protocol/JsonRpcMessageContext.cs @@ -130,4 +130,9 @@ public sealed class JsonRpcMessageContext /// log notifications for the request. Legacy requests continue to use their negotiated logging behavior. /// public LoggingLevel? LogLevel { get; set; } + + /// + /// Gets or sets the discovery-owned timer, allowing awaited OAuth work to suspend only the probe deadline. + /// + internal RequestTimeout? RequestTimeout { get; set; } } diff --git a/src/ModelContextProtocol.Core/RequestTimeout.cs b/src/ModelContextProtocol.Core/RequestTimeout.cs new file mode 100644 index 000000000..9fcde4e88 --- /dev/null +++ b/src/ModelContextProtocol.Core/RequestTimeout.cs @@ -0,0 +1,67 @@ +namespace ModelContextProtocol; + +/// A request-local timer that can be suspended without suspending linked cancellation. +/// +/// Owned by one awaited initialization or discovery operation, linked to its caller's cancellation. +/// Suspension scopes must be sequential and disposed before their owner. +/// Cancellation may race with suspension, but an expired timer cannot be restarted. +/// +internal sealed class RequestTimeout : IDisposable +{ + private readonly CancellationTokenSource _source; + private readonly ITimer _timer; + private readonly TimeSpan _timeout; + + public RequestTimeout(TimeSpan timeout, TimeProvider timeProvider, CancellationToken cancellationToken) + { + _timeout = timeout; + _source = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken); + Token = _source.Token; + try + { + _timer = timeProvider.CreateTimer(static state => + { + try + { + ((CancellationTokenSource)state!).Cancel(); + } + catch (ObjectDisposedException) + { + // A timer callback already queued when Dispose ran can outlive the source. + } + }, _source, timeout, Timeout.InfiniteTimeSpan); + } + catch + { + _source.Dispose(); + throw; + } + } + + public CancellationToken Token { get; } + + public void Stop() => _timer.Change(Timeout.InfiniteTimeSpan, Timeout.InfiniteTimeSpan); + + public Suspension Suspend() + { + Stop(); + return new Suspension(this); + } + + public void Dispose() + { + _timer.Dispose(); + _source.Dispose(); + } + + public readonly struct Suspension(RequestTimeout owner) : IDisposable + { + public void Dispose() + { + if (!owner.Token.IsCancellationRequested) + { + owner._timer.Change(owner._timeout, Timeout.InfiniteTimeSpan); + } + } + } +} diff --git a/tests/ModelContextProtocol.AspNetCore.Tests/July2026ProtocolHttpFallbackTests.cs b/tests/ModelContextProtocol.AspNetCore.Tests/July2026ProtocolHttpFallbackTests.cs index 9126331de..22cc36ad7 100644 --- a/tests/ModelContextProtocol.AspNetCore.Tests/July2026ProtocolHttpFallbackTests.cs +++ b/tests/ModelContextProtocol.AspNetCore.Tests/July2026ProtocolHttpFallbackTests.cs @@ -2,6 +2,7 @@ using Microsoft.AspNetCore.Http; using Microsoft.AspNetCore.Http.Json; using Microsoft.Extensions.DependencyInjection; +using Microsoft.Extensions.Time.Testing; using ModelContextProtocol.AspNetCore.Tests.Utils; using ModelContextProtocol.Client; using ModelContextProtocol.Protocol; @@ -78,6 +79,60 @@ private async Task StartServerAsync(RequestDelegate handler, bool acceptGet = fa private static JsonTypeInfo GetJsonTypeInfo() => (JsonTypeInfo)McpJsonUtilities.DefaultOptions.GetTypeInfo(typeof(T)); + [Theory] + [InlineData(null, 200)] + [InlineData("application/json", 200)] + [InlineData("text/event-stream", 200)] + [InlineData("application/json", 400)] + public async Task SilentDiscoverHeadersOrBody_UseProbeBudget(string? contentType, int statusCode) + { + var timeProvider = new FakeTimeProvider(); + var probeBudget = TimeSpan.FromSeconds(5); + var stalled = new AsyncGate(); + var methods = new List(); + await StartServerAsync(async context => + { + var message = await JsonSerializer.DeserializeAsync(context.Request.Body, GetJsonTypeInfo(), context.RequestAborted); + if (message is not JsonRpcRequest request) + { + context.Response.StatusCode = StatusCodes.Status202Accepted; + return; + } + methods.Add(request.Method); + if (request.Method == RequestMethods.ServerDiscover) + { + if (contentType is not null) + { + context.Response.StatusCode = statusCode; + context.Response.ContentType = contentType; + await context.Response.WriteAsync(contentType == "text/event-stream" ? ": waiting\n\n" : "{", context.RequestAborted); + await context.Response.Body.FlushAsync(context.RequestAborted); + } + await stalled.WaitAsync(context.RequestAborted); + return; + } + var response = new JsonRpcResponse + { + Id = request.Id, + Result = JsonSerializer.SerializeToNode(new InitializeResult + { + ProtocolVersion = McpProtocolVersions.November2025ProtocolVersion, + Capabilities = new(), + ServerInfo = new() { Name = "legacy", Version = "1" }, + }, McpJsonUtilities.DefaultOptions), + }; + context.Response.ContentType = "application/json"; + await JsonSerializer.SerializeAsync(context.Response.Body, response, GetJsonTypeInfo(), context.RequestAborted); + }); + await using var transport = new HttpClientTransport(new() { Endpoint = new("http://localhost:5000/mcp") }, HttpClient, LoggerFactory); + var connecting = McpClient.CreateAsync(transport, new() { TimeProvider = timeProvider, DiscoverProbeTimeout = probeBudget }, LoggerFactory, TestContext.Current.CancellationToken); + await stalled.WaitUntilEnteredAsync(connecting); + timeProvider.Advance(probeBudget); + await using var client = await connecting.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + Assert.Equal(McpProtocolVersions.November2025ProtocolVersion, client.NegotiatedProtocolVersion); + Assert.Equal([RequestMethods.ServerDiscover, RequestMethods.Initialize], methods); + } + private static async Task WriteJsonRpcErrorAsync(HttpContext context, HttpStatusCode statusCode, int code, string message) { var rpcError = new JsonRpcError diff --git a/tests/ModelContextProtocol.AspNetCore.Tests/MapMcpTests.Mrtr.cs b/tests/ModelContextProtocol.AspNetCore.Tests/MapMcpTests.Mrtr.cs index 03af131b4..cf8a5ab33 100644 --- a/tests/ModelContextProtocol.AspNetCore.Tests/MapMcpTests.Mrtr.cs +++ b/tests/ModelContextProtocol.AspNetCore.Tests/MapMcpTests.Mrtr.cs @@ -10,12 +10,10 @@ namespace ModelContextProtocol.AspNetCore.Tests; public abstract partial class MapMcpTests { - // Starting with the 2026-07-28 protocol revision, Streamable HTTP no longer supports sessions (SEP-2567): - // the handler refuses a request when the server opted into sessions (SessionMode = HttpServerSessionMode.Stateful), so a client pinned - // to that revision downgrades to legacy instead of negotiating 2026-07-28. These MRTR tests therefore can't - // run on the stateful Streamable HTTP fixture; the same coverage runs on the stateless and legacy-SSE fixtures. + // This fixture's strict stateful Streamable HTTP mode rejects the modern revision. + // Stateless and hybrid HTTP servers, and explicitly selected SSE, can serve it. private const string July2026StatefulStreamableHttpSkipReason = - "Starting with the 2026-07-28 protocol revision, Streamable HTTP no longer supports sessions (SEP-2567); stateful Streamable HTTP refuses it. Covered by the stateless and SSE fixtures."; + "The strict stateful Streamable HTTP fixture rejects 2026-07-28. Covered by the stateless and SSE fixtures."; private ServerMessageTracker ConfigureServer(params Delegate[] tools) { diff --git a/tests/ModelContextProtocol.AspNetCore.Tests/MapMcpTests.cs b/tests/ModelContextProtocol.AspNetCore.Tests/MapMcpTests.cs index 43a3c12b5..d413bf528 100644 --- a/tests/ModelContextProtocol.AspNetCore.Tests/MapMcpTests.cs +++ b/tests/ModelContextProtocol.AspNetCore.Tests/MapMcpTests.cs @@ -329,12 +329,10 @@ await client.CallToolAsync("echo_with_user_name", new Dictionary { ["message"] = "hi" }, cancellationToken: TestContext.Current.CancellationToken); - // The client now defaults to the 2026-07-28 protocol revision, whose handshake is server/discover - // rather than the legacy initialize request. On the stateful Streamable HTTP fixture the - // request is refused, so the client downgrades to the legacy initialize. - var expectedHandshakeMethod = UseStreamableHttp && !Stateless - ? RequestMethods.Initialize - : RequestMethods.ServerDiscover; + // With default client options, only the stateless HTTP fixture uses discovery. + var expectedHandshakeMethod = UseStreamableHttp && Stateless + ? RequestMethods.ServerDiscover + : RequestMethods.Initialize; Assert.Contains(expectedHandshakeMethod, observedMethods); Assert.Contains(RequestMethods.ToolsList, observedMethods); Assert.Contains(RequestMethods.ToolsCall, observedMethods); diff --git a/tests/ModelContextProtocol.AspNetCore.Tests/OAuth/AuthTests.cs b/tests/ModelContextProtocol.AspNetCore.Tests/OAuth/AuthTests.cs index 693c77943..476a0ccf3 100644 --- a/tests/ModelContextProtocol.AspNetCore.Tests/OAuth/AuthTests.cs +++ b/tests/ModelContextProtocol.AspNetCore.Tests/OAuth/AuthTests.cs @@ -5,8 +5,10 @@ using Microsoft.AspNetCore.Http; using Microsoft.AspNetCore.WebUtilities; using Microsoft.Extensions.DependencyInjection; +using Microsoft.Extensions.Time.Testing; using ModelContextProtocol; using ModelContextProtocol.AspNetCore.Authentication; +using ModelContextProtocol.AspNetCore.Tests.Utils; using ModelContextProtocol.Authentication; using ModelContextProtocol.Client; using ModelContextProtocol.Protocol; @@ -1506,8 +1508,12 @@ public async Task CanAuthenticate_WithResourceMetadataPathFallbacks() { const string resourcePath = "/mcp"; List wellKnownRequests = []; + var metadataGate = new AsyncGate(); + var timeProvider = new FakeTimeProvider(); + var probeBudget = TimeSpan.FromSeconds(5); Builder.Services.Configure(options => options.DefaultChallengeScheme = JwtBearerDefaults.AuthenticationScheme); + Builder.Services.Configure(options => options.Stateless = true); await using var app = Builder.Build(); var metadata = new ProtectedResourceMetadata @@ -1523,6 +1529,7 @@ public async Task CanAuthenticate_WithResourceMetadataPathFallbacks() wellKnownRequests.Add(context.Request.Path); if (remaining.HasValue) { + await metadataGate.WaitAsync(context.RequestAborted); context.Response.StatusCode = StatusCodes.Status404NotFound; return; } @@ -1552,9 +1559,16 @@ public async Task CanAuthenticate_WithResourceMetadataPathFallbacks() }, }, HttpClient, LoggerFactory); - await using var client = await McpClient.CreateAsync( - transport, loggerFactory: LoggerFactory, cancellationToken: TestContext.Current.CancellationToken); + var connecting = McpClient.CreateAsync( + transport, new() { TimeProvider = timeProvider, DiscoverProbeTimeout = probeBudget }, loggerFactory: LoggerFactory, cancellationToken: TestContext.Current.CancellationToken); + await metadataGate.WaitUntilEnteredAsync(connecting); + timeProvider.Advance(probeBudget * 2); + Assert.False(metadataGate.Token.IsCancellationRequested); + metadataGate.Release.SetResult(); + await using var client = await connecting.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + Assert.Equal(McpProtocolVersions.July2026ProtocolVersion, client.NegotiatedProtocolVersion); + Assert.Equal(1, TestOAuthServer.AuthorizationCodeTokenRequestCount); Assert.Equal( [ $"/.well-known/oauth-protected-resource{resourcePath}", diff --git a/tests/ModelContextProtocol.AspNetCore.Tests/OAuth/DiscoveryTimeoutTests.cs b/tests/ModelContextProtocol.AspNetCore.Tests/OAuth/DiscoveryTimeoutTests.cs new file mode 100644 index 000000000..287084628 --- /dev/null +++ b/tests/ModelContextProtocol.AspNetCore.Tests/OAuth/DiscoveryTimeoutTests.cs @@ -0,0 +1,259 @@ +using Microsoft.AspNetCore.Builder; +using Microsoft.AspNetCore.Http; +using Microsoft.Extensions.DependencyInjection; +using Microsoft.Extensions.Time.Testing; +using ModelContextProtocol.AspNetCore.Tests.Utils; +using ModelContextProtocol.Authentication; +using ModelContextProtocol.Client; +using ModelContextProtocol.Protocol; +using ModelContextProtocol.Tests.Utils; +using System.Collections.Concurrent; + +namespace ModelContextProtocol.AspNetCore.Tests.OAuth; + +public class DiscoveryTimeoutTests(ITestOutputHelper outputHelper) : OAuthTestBase(outputHelper) +{ + private static readonly TimeSpan ProbeBudget = TimeSpan.FromSeconds(5); + private readonly FakeTimeProvider _timeProvider = new(); + private readonly ConcurrentQueue _methods = new(); + private readonly AsyncGate _authorization = new(); + private int _callbackCount; + + [Fact] + public async Task SlowSilentAcquisition_IsExcludedBeforeTheInitialPost() + { + ConfigureModernServer(); + await using var app = await StartMcpServerAsync(); + var cache = new GatedCache(); + await using var transport = CreateTransport(cache); + _authorization.Release.SetResult(); + var connecting = McpClient.CreateAsync(transport, Options(), LoggerFactory, TestContext.Current.CancellationToken); + await cache.Gate.WaitUntilEnteredAsync(connecting); + _timeProvider.Advance(ProbeBudget * 2); + Assert.False(cache.Gate.Token.IsCancellationRequested); + Assert.Empty(_methods); + cache.Gate.Release.SetResult(); + + await using var client = await connecting.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + Assert.Equal(McpProtocolVersions.July2026ProtocolVersion, client.NegotiatedProtocolVersion); + Assert.Equal([RequestMethods.ServerDiscover], _methods); + Assert.Equal(1, _callbackCount); + } + + [Theory] + [InlineData(false)] + [InlineData(true)] + public async Task Authorization_ObservesCallerAndInitializationCancellation(bool initializationTimeout) + { + ConfigureModernServer(); + await using var app = await StartMcpServerAsync(); + await using var transport = CreateTransport(); + using var caller = CancellationTokenSource.CreateLinkedTokenSource(TestContext.Current.CancellationToken); + var options = Options(); + options.InitializationTimeout = initializationTimeout ? ProbeBudget * 4 : TestConstants.DefaultTimeout; + var connecting = McpClient.CreateAsync(transport, options, LoggerFactory, caller.Token); + await _authorization.WaitUntilEnteredAsync(connecting); + if (initializationTimeout) + { + _timeProvider.Advance(options.InitializationTimeout); + var error = await Assert.ThrowsAsync(() => connecting.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken)); + Assert.Equal("Initialization timed out", error.Message); + } + else + { + caller.Cancel(); + await Assert.ThrowsAnyAsync(() => connecting.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken)); + } + await _authorization.Canceled.Task.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + Assert.Equal(1, _callbackCount); + Assert.Empty(_methods); + } + + [Fact] + public async Task SlowAuthorization_PreservesModernProtocol_WithoutSuspendingAnotherClientsProbe() + { + ConfigureModernServer(); + int posts = 0; + var secondPost = new AsyncGate(); + await using var app = await StartMcpServerAsync(configureMiddleware: app => + { + app.Use(async (context, next) => + { + if (context.Request.Method == HttpMethods.Post && Interlocked.Increment(ref posts) == 2) + { + await secondPost.WaitAsync(context.RequestAborted); + } + await next(); + }); + app.UseAuthentication(); + app.UseAuthorization(); + }); + await using var transport = CreateTransport(); + var first = McpClient.CreateAsync(transport, Options(), LoggerFactory, TestContext.Current.CancellationToken); + await _authorization.WaitUntilEnteredAsync(first); + var secondOptions = Options(pinned: true); + secondOptions.DiscoverProbeTimeout = ProbeBudget * 2; + var second = McpClient.CreateAsync(transport, secondOptions, LoggerFactory, TestContext.Current.CancellationToken); + await secondPost.WaitUntilEnteredAsync(second); + // The second probe expiring proves the first authorization survived more than its own budget. + _timeProvider.Advance(secondOptions.DiscoverProbeTimeout); + await Assert.ThrowsAsync(() => second.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken)); + Assert.False(_authorization.Token.IsCancellationRequested); + _authorization.Release.SetResult(); + await using var client = await first.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + Assert.Equal(McpProtocolVersions.July2026ProtocolVersion, client.NegotiatedProtocolVersion); + Assert.Equal(1, _callbackCount); + Assert.Equal([RequestMethods.ServerDiscover], _methods); + } + + [Fact] + public async Task Authentication_RestartsProbeBudgetBeforeRetryHeaders() + { + ConfigureModernServer(); + var initialHeaders = new AsyncGate(); + var retryHeaders = new AsyncGate(); + await using var app = await StartMcpServerAsync(configureMiddleware: app => + { + app.Use(async (context, next) => + { + if (context.Request.Method == HttpMethods.Post) + { + await (context.Request.Headers.Authorization.Count == 0 ? initialHeaders : retryHeaders).WaitAsync(context.RequestAborted); + } + await next(); + }); + app.UseAuthentication(); + app.UseAuthorization(); + }); + await using var transport = CreateTransport(); + var options = Options(pinned: true); + var connecting = McpClient.CreateAsync(transport, options, LoggerFactory, TestContext.Current.CancellationToken); + await initialHeaders.WaitUntilEnteredAsync(connecting); + _timeProvider.Advance(ProbeBudget * 0.6); + initialHeaders.Release.SetResult(); + await _authorization.WaitUntilEnteredAsync(connecting); + _timeProvider.Advance(ProbeBudget * 2); + Assert.False(_authorization.Token.IsCancellationRequested); + _authorization.Release.SetResult(); + await retryHeaders.WaitUntilEnteredAsync(connecting); + _timeProvider.Advance(ProbeBudget * 0.6); + retryHeaders.Release.SetResult(); + + await using var client = await connecting.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + Assert.Equal(McpProtocolVersions.July2026ProtocolVersion, client.NegotiatedProtocolVersion); + Assert.Equal(1, _callbackCount); + Assert.Equal([RequestMethods.ServerDiscover], _methods); + } + + [Fact] + public async Task AuthenticatedRetryHeaders_RemainProbeBounded() + { + ConfigureModernServer(); + var headers = new AsyncGate(); + await using var app = await StartMcpServerAsync(configureMiddleware: app => app.Use(async (context, next) => + { + if (context.Request.Method == HttpMethods.Post && context.Request.Headers.Authorization.Count > 0) + { + await headers.WaitAsync(context.RequestAborted); + } + await next(); + })); + await using var transport = CreateTransport(); + _authorization.Release.SetResult(); + var connecting = McpClient.CreateAsync(transport, Options(pinned: true), LoggerFactory, TestContext.Current.CancellationToken); + await headers.WaitUntilEnteredAsync(connecting); + _timeProvider.Advance(ProbeBudget); + await headers.Canceled.Task.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + await Assert.ThrowsAsync(() => connecting.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken)); + Assert.Equal(1, _callbackCount); + } + + [Fact] + public async Task AuthenticatedDiscoveryBodyTimeout_AbortsModernHandlerWithoutCancellationRpc() + { + var handler = new AsyncGate(); + ConfigureModernServer(); + Builder.Services.AddHttpContextAccessor(); + Builder.Services.AddMcpServer().WithMessageFilters(filters => filters.AddIncomingFilter(next => async (context, cancellationToken) => + { + if (context.JsonRpcMessage is JsonRpcRequest { Method: RequestMethods.ServerDiscover }) + { + var httpContext = context.Services!.GetRequiredService().HttpContext!; + httpContext.Response.ContentType = "text/event-stream"; + await httpContext.Response.WriteAsync(": waiting\n\n", cancellationToken); + await httpContext.Response.Body.FlushAsync(cancellationToken); + await handler.WaitAsync(cancellationToken); + } + await next(context, cancellationToken); + })); + await using var app = await StartMcpServerAsync(); + await using var transport = CreateTransport(); + _authorization.Release.SetResult(); + var connecting = McpClient.CreateAsync(transport, Options(pinned: true), LoggerFactory, TestContext.Current.CancellationToken); + await handler.WaitUntilEnteredAsync(connecting); + _timeProvider.Advance(ProbeBudget); + await handler.Canceled.Task.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + await Assert.ThrowsAsync(() => connecting.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken)); + Assert.Equal(1, _callbackCount); + Assert.Equal([RequestMethods.ServerDiscover], _methods); + } + + private void ConfigureModernServer() + { + Builder.Services.AddMcpServer().WithHttpTransport(options => options.Stateless = true) + .WithMessageFilters(filters => filters.AddIncomingFilter(next => async (context, cancellationToken) => + { + if (context.JsonRpcMessage is JsonRpcRequest request) + { + _methods.Enqueue(request.Method); + } + else if (context.JsonRpcMessage is JsonRpcNotification notification) + { + _methods.Enqueue(notification.Method); + } + await next(context, cancellationToken); + })); + } + + private HttpClientTransport CreateTransport(ITokenCache? cache = null) => new(new() + { + Endpoint = new(McpServerUrl), + TransportMode = HttpTransportMode.StreamableHttp, + OAuth = new() + { + ClientId = "demo-client", + ClientSecret = "demo-secret", + RedirectUri = new("http://localhost:1179/callback"), + TokenCache = cache, + AuthorizationCallbackHandler = async (context, cancellationToken) => + { + Interlocked.Increment(ref _callbackCount); + await _authorization.WaitAsync(cancellationToken); + return await HandleAuthorizationUrlAsync(context, cancellationToken); + }, + }, + }, HttpClient, LoggerFactory); + + private McpClientOptions Options(bool pinned = false) => new() + { + TimeProvider = _timeProvider, + DiscoverProbeTimeout = ProbeBudget, + ProtocolVersion = pinned ? McpProtocolVersions.July2026ProtocolVersion : null, + }; + + private sealed class GatedCache : ITokenCache + { + private TokenContainer? _tokens; + public AsyncGate Gate { get; } = new(); + public async ValueTask GetTokensAsync(CancellationToken cancellationToken) + { + await Gate.WaitAsync(cancellationToken); + return _tokens; + } + public ValueTask StoreTokensAsync(TokenContainer tokens, CancellationToken cancellationToken) + { + _tokens = tokens; + return default; + } + } +} diff --git a/tests/ModelContextProtocol.AspNetCore.Tests/OAuth/SseDiscoveryTests.cs b/tests/ModelContextProtocol.AspNetCore.Tests/OAuth/SseDiscoveryTests.cs new file mode 100644 index 000000000..96f2aa1ec --- /dev/null +++ b/tests/ModelContextProtocol.AspNetCore.Tests/OAuth/SseDiscoveryTests.cs @@ -0,0 +1,210 @@ +using Microsoft.AspNetCore.Authentication.JwtBearer; +using Microsoft.AspNetCore.Builder; +using Microsoft.AspNetCore.Http; +using Microsoft.Extensions.DependencyInjection; +using Microsoft.Extensions.Time.Testing; +using ModelContextProtocol.AspNetCore.Authentication; +using ModelContextProtocol.AspNetCore.Tests.Utils; +using ModelContextProtocol.Client; +using ModelContextProtocol.Protocol; +using ModelContextProtocol.Tests.Utils; +using System.Collections.Concurrent; +using System.Text.Json; + +namespace ModelContextProtocol.AspNetCore.Tests.OAuth; + +public class SseDiscoveryTests(ITestOutputHelper outputHelper) : OAuthTestBase(outputHelper) +{ + private static readonly TimeSpan ProbeBudget = TimeSpan.FromSeconds(5); + + [Theory] + [InlineData(HttpTransportMode.AutoDetect, null)] + [InlineData(HttpTransportMode.Sse, null)] + [InlineData(HttpTransportMode.AutoDetect, "2025-11-25")] + [InlineData(HttpTransportMode.Sse, "2025-11-25")] + [InlineData(HttpTransportMode.AutoDetect, "2026-07-28")] + [InlineData(HttpTransportMode.Sse, "2026-07-28")] + public async Task Sse_DefaultsToInitialize_AndHonorsExplicitTransportAndVersion(HttpTransportMode mode, string? version) + { + var methods = new ConcurrentQueue(); + var initialized = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + ConfigureSse(methods, initialized); + var timeProvider = new FakeTimeProvider(); + var authorization = new AsyncGate(); + var initialEndpointMethods = new ConcurrentQueue(); + await using var app = await StartMcpServerAsync(configureMiddleware: app => app.Use(async (context, next) => + { + if (context.Request.Method == HttpMethods.Post && context.Request.Path == "/sse") + { + context.Request.EnableBuffering(); + var message = await JsonSerializer.DeserializeAsync(context.Request.Body, McpJsonUtilities.DefaultOptions, context.RequestAborted); + initialEndpointMethods.Enqueue(Assert.IsType(message).Method); + context.Request.Body.Position = 0; + } + await next(); + })); + await using var transport = CreateTransport(mode, authorization); + var connecting = McpClient.CreateAsync(transport, new() + { + TimeProvider = timeProvider, + DiscoverProbeTimeout = ProbeBudget, + ProtocolVersion = version, + InitializationTimeout = mode == HttpTransportMode.Sse && version is null ? ProbeBudget : TestConstants.DefaultTimeout, + }, LoggerFactory, TestContext.Current.CancellationToken); + if (version is null) + { + // AutoDetect excludes GET establishment from the probe; explicit SSE precedes initialization. + await authorization.WaitUntilEnteredAsync(connecting); + timeProvider.Advance(ProbeBudget * 2); + Assert.False(authorization.Token.IsCancellationRequested); + } + authorization.Release.SetResult(); + bool modern = version == McpProtocolVersions.July2026ProtocolVersion; + if (modern && mode == HttpTransportMode.AutoDetect) + { + await Assert.ThrowsAsync(() => connecting.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken)); + Assert.Empty(methods); + } + else + { + await using var client = await connecting.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + Assert.Equal(version ?? McpProtocolVersions.November2025ProtocolVersion, client.NegotiatedProtocolVersion); + Assert.Empty(await client.ListToolsAsync(cancellationToken: TestContext.Current.CancellationToken)); + if (!modern) + { + await initialized.Task.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + } + Assert.Equal(modern + ? [RequestMethods.ServerDiscover, RequestMethods.ToolsList] + : new[] { RequestMethods.Initialize, RequestMethods.ToolsList }, methods); + } + Assert.Equal(1, TestOAuthServer.AuthorizationCodeTokenRequestCount); + Assert.Equal(mode == HttpTransportMode.Sse ? [] : + new[] { version == McpProtocolVersions.November2025ProtocolVersion ? RequestMethods.Initialize : RequestMethods.ServerDiscover }, + initialEndpointMethods); + } + + [Fact] + public async Task ExplicitModernSse_SilentDiscoveryTimesOutWithoutInitialize() + { + var methods = new ConcurrentQueue(); + var timeProvider = new FakeTimeProvider(); + var discoveryReceived = new AsyncGate(); + ConfigureSse(methods); + Builder.Services.AddMcpServer().WithMessageFilters(filters => filters.AddIncomingFilter(next => async (context, cancellationToken) => + { + if (context.JsonRpcMessage is JsonRpcRequest { Method: RequestMethods.ServerDiscover }) + { + await discoveryReceived.WaitAsync(cancellationToken); + } + else + { + await next(context, cancellationToken); + } + })); + var authorization = new AsyncGate(); + authorization.Release.SetResult(); + await using var app = await StartMcpServerAsync(); + await using var transport = CreateTransport(HttpTransportMode.Sse, authorization); + var connecting = McpClient.CreateAsync(transport, new() + { + TimeProvider = timeProvider, + ProtocolVersion = McpProtocolVersions.July2026ProtocolVersion, + DiscoverProbeTimeout = ProbeBudget, + }, LoggerFactory, TestContext.Current.CancellationToken); + await discoveryReceived.WaitUntilEnteredAsync(connecting); + timeProvider.Advance(ProbeBudget); + await Assert.ThrowsAsync(() => connecting.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken)); + Assert.Contains(RequestMethods.ServerDiscover, methods); + Assert.DoesNotContain(RequestMethods.Initialize, methods); + } + + [Theory] + [InlineData("caller")] + [InlineData("initialization")] + [InlineData("connection")] + public async Task SseGetAuthorization_PreservesExistingDeadlines(string deadline) + { + ConfigureSse(new()); + var timeProvider = new FakeTimeProvider(); + var authorization = new AsyncGate(); + await using var app = await StartMcpServerAsync(); + await using var transport = CreateTransport(HttpTransportMode.AutoDetect, authorization, + deadline == "connection" ? TimeSpan.FromSeconds(2) : null); + using var caller = CancellationTokenSource.CreateLinkedTokenSource(TestContext.Current.CancellationToken); + var connecting = McpClient.CreateAsync(transport, new() + { + TimeProvider = timeProvider, + DiscoverProbeTimeout = ProbeBudget, + InitializationTimeout = deadline == "initialization" ? ProbeBudget * 4 : TestConstants.DefaultTimeout, + }, LoggerFactory, caller.Token); + if (deadline != "connection") + { + await authorization.WaitUntilEnteredAsync(connecting); + } + if (deadline == "caller") + { + caller.Cancel(); + await Assert.ThrowsAnyAsync(() => connecting.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken)); + } + else if (deadline == "initialization") + { + timeProvider.Advance(ProbeBudget * 4); + var error = await Assert.ThrowsAsync(() => connecting.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken)); + Assert.Equal("Initialization timed out", error.Message); + } + else + { + // ConnectionTimeout uses real time and may expire before authorization starts. + var error = await Assert.ThrowsAsync(() => connecting.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken)); + Assert.IsType(error.InnerException); + } + if (authorization.Entered.Task.IsCompleted) + { + await authorization.Canceled.Task.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + } + Assert.Equal(0, TestOAuthServer.AuthorizationCodeTokenRequestCount); + } + + private void ConfigureSse(ConcurrentQueue methods, TaskCompletionSource? initialized = null) + { + TestOAuthServer.ValidResources = [.. TestOAuthServer.ValidResources, $"{McpServerUrl}/sse"]; + Builder.Services.Configure(JwtBearerDefaults.AuthenticationScheme, + options => options.TokenValidationParameters.ValidAudiences = [$"{McpServerUrl}/sse"]); + Builder.Services.Configure(McpAuthenticationDefaults.AuthenticationScheme, + options => options.ResourceMetadata!.Resource = $"{McpServerUrl}/sse"); + Builder.Services.AddMcpServer().WithHttpTransport(options => options.EnableLegacySse = true) + .WithListToolsHandler((_, _) => ValueTask.FromResult(new ListToolsResult { Tools = [] })) + .WithMessageFilters(filters => filters.AddIncomingFilter(next => async (context, cancellationToken) => + { + if (context.JsonRpcMessage is JsonRpcRequest request) + { + methods.Enqueue(request.Method); + } + else if (context.JsonRpcMessage is JsonRpcNotification { Method: NotificationMethods.InitializedNotification }) + { + initialized?.TrySetResult(); + } + await next(context, cancellationToken); + })); + } + + private HttpClientTransport CreateTransport(HttpTransportMode mode, AsyncGate authorization, TimeSpan? connectionTimeout = null) + => new(new() + { + Endpoint = new($"{McpServerUrl}/sse"), + TransportMode = mode, + ConnectionTimeout = connectionTimeout ?? TestConstants.DefaultTimeout, + OAuth = new() + { + ClientId = "demo-client", + ClientSecret = "demo-secret", + RedirectUri = new("http://localhost:1179/callback"), + AuthorizationCallbackHandler = async (context, cancellationToken) => + { + await authorization.WaitAsync(cancellationToken); + return await HandleAuthorizationUrlAsync(context, cancellationToken); + }, + }, + }, HttpClient, LoggerFactory); +} diff --git a/tests/ModelContextProtocol.AspNetCore.Tests/Utils/AsyncGate.cs b/tests/ModelContextProtocol.AspNetCore.Tests/Utils/AsyncGate.cs new file mode 100644 index 000000000..dafdbfd06 --- /dev/null +++ b/tests/ModelContextProtocol.AspNetCore.Tests/Utils/AsyncGate.cs @@ -0,0 +1,37 @@ +using ModelContextProtocol.Client; +using ModelContextProtocol.Tests.Utils; + +namespace ModelContextProtocol.AspNetCore.Tests.Utils; + +internal sealed class AsyncGate +{ + public TaskCompletionSource Entered { get; } = new(TaskCreationOptions.RunContinuationsAsynchronously); + public TaskCompletionSource Release { get; } = new(TaskCreationOptions.RunContinuationsAsynchronously); + public TaskCompletionSource Canceled { get; } = new(TaskCreationOptions.RunContinuationsAsynchronously); + public CancellationToken Token { get; private set; } + + public async Task WaitAsync(CancellationToken cancellationToken) + { + Token = cancellationToken; + Entered.TrySetResult(); + try + { + await Release.Task.WaitAsync(TestConstants.DefaultTimeout, cancellationToken); + } + catch (OperationCanceledException) + { + Canceled.TrySetResult(); + throw; + } + } + + public async Task WaitUntilEnteredAsync(Task connecting) + { + await Task.WhenAny(Entered.Task, connecting).WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + if (!Entered.Task.IsCompleted) + { + await using var client = await connecting; + Assert.Fail("The client connected without entering the expected phase."); + } + } +} diff --git a/tests/ModelContextProtocol.Tests/Client/July2026ProtocolFallbackTests.cs b/tests/ModelContextProtocol.Tests/Client/July2026ProtocolFallbackTests.cs index 557dc5655..c0cd0b40c 100644 --- a/tests/ModelContextProtocol.Tests/Client/July2026ProtocolFallbackTests.cs +++ b/tests/ModelContextProtocol.Tests/Client/July2026ProtocolFallbackTests.cs @@ -1,7 +1,7 @@ +using Microsoft.Extensions.Time.Testing; using ModelContextProtocol.Client; using ModelContextProtocol.Protocol; using ModelContextProtocol.Tests.Utils; -using System.Diagnostics; using System.Net; using System.Text; using System.Text.Json; @@ -153,35 +153,116 @@ public async Task Client_OnUnsupportedProtocolVersion_WithPerRequestMetadataVers Assert.Equal(McpProtocolVersions.July2026ProtocolVersion, client.NegotiatedProtocolVersion); } - [Fact] - public async Task Client_OnSilentProbe_FallsBackTo_Initialize_AfterConfiguredProbeTimeout() + [Theory] + [InlineData(false)] + [InlineData(true)] + public async Task Client_OnSilentProbe_FallsBackTo_Initialize_AfterConfiguredProbeTimeout(bool infiniteInitialization) { // Simulate an initialize-handshake server that silently drops the unknown server/discover method (it never // responds to the probe). The client must fall back to initialize once the configured // DiscoverProbeTimeout elapses, well before the much larger InitializationTimeout. - var ct = TestContext.Current.CancellationToken; + using var deadline = CancellationTokenSource.CreateLinkedTokenSource(TestContext.Current.CancellationToken); + deadline.CancelAfter(TestConstants.DefaultTimeout); + var ct = deadline.Token; await using var transport = new InitializeHandshakeServerTestTransport( serverNegotiatedVersion: McpProtocolVersions.November2025ProtocolVersion, silentDiscoverProbe: true); - var stopwatch = Stopwatch.StartNew(); + var timeProvider = new FakeTimeProvider(); + var probeBudget = TimeSpan.FromSeconds(5); // Default options (ProtocolVersion = null) prefer 2026-07-28 but allow automatic fallback. - await using var client = await McpClient.CreateAsync(transport, new McpClientOptions + var connecting = McpClient.CreateAsync(transport, new McpClientOptions { - DiscoverProbeTimeout = TimeSpan.FromMilliseconds(250), - InitializationTimeout = TestConstants.DefaultTimeout, + TimeProvider = timeProvider, + DiscoverProbeTimeout = probeBudget, + InitializationTimeout = infiniteInitialization ? Timeout.InfiniteTimeSpan : TestConstants.DefaultTimeout, }, loggerFactory: LoggerFactory, cancellationToken: ct); - stopwatch.Stop(); + await transport.DiscoverReceived.Task.WaitAsync(ct); + timeProvider.Advance(probeBudget - TimeSpan.FromMilliseconds(1)); + Assert.False(connecting.IsCompleted); + Assert.False(transport.InitializeReceived); + timeProvider.Advance(TimeSpan.FromMilliseconds(1)); + await using var client = await connecting.WaitAsync(ct); Assert.True(transport.ServerDiscoverProbed); Assert.True(transport.InitializeReceived); Assert.Equal(McpProtocolVersions.November2025ProtocolVersion, transport.InitializeProtocolVersion); Assert.Equal(McpProtocolVersions.November2025ProtocolVersion, client.NegotiatedProtocolVersion); - // The fallback was driven by the short probe timeout, not the 60s InitializationTimeout. - Assert.True( - stopwatch.Elapsed < TimeSpan.FromSeconds(30), - $"Fallback should have happened shortly after the {nameof(McpClientOptions.DiscoverProbeTimeout)}, but took {stopwatch.Elapsed}."); + Assert.False(ct.IsCancellationRequested); + } + + [Theory] + [InlineData(-1, 250)] + [InlineData(1000, 250)] + [InlineData(250, 250)] + public async Task Client_InitializationDeadlineWins_NoFallback(int probeMilliseconds, int initializationMilliseconds) + { + using var deadline = CancellationTokenSource.CreateLinkedTokenSource(TestContext.Current.CancellationToken); + deadline.CancelAfter(TestConstants.DefaultTimeout); + await using var transport = new InitializeHandshakeServerTestTransport( + McpProtocolVersions.November2025ProtocolVersion, silentDiscoverProbe: true); + var timeProvider = new FakeTimeProvider(); + var connecting = McpClient.CreateAsync(transport, new McpClientOptions + { + TimeProvider = timeProvider, + DiscoverProbeTimeout = TimeSpan.FromMilliseconds(probeMilliseconds), + InitializationTimeout = TimeSpan.FromMilliseconds(initializationMilliseconds), + }, LoggerFactory, deadline.Token); + await transport.DiscoverReceived.Task.WaitAsync(deadline.Token); + timeProvider.Advance(TimeSpan.FromMilliseconds(initializationMilliseconds)); + var exception = await Assert.ThrowsAsync(() => connecting.WaitAsync(deadline.Token)); + + Assert.Equal("Initialization timed out", exception.Message); + Assert.True(transport.ServerDiscoverProbed); + Assert.False(transport.InitializeReceived); + Assert.False(deadline.IsCancellationRequested); + } + + [Fact] + public async Task Client_InfiniteProbeAndInitialization_ObserveCallerCancellation() + { + using var deadline = CancellationTokenSource.CreateLinkedTokenSource(TestContext.Current.CancellationToken); + deadline.CancelAfter(TestConstants.DefaultTimeout); + using var caller = CancellationTokenSource.CreateLinkedTokenSource(deadline.Token); + await using var transport = new InitializeHandshakeServerTestTransport( + McpProtocolVersions.November2025ProtocolVersion, silentDiscoverProbe: true); + var timeProvider = new FakeTimeProvider(); + var connecting = McpClient.CreateAsync(transport, new McpClientOptions + { + TimeProvider = timeProvider, + DiscoverProbeTimeout = Timeout.InfiniteTimeSpan, + InitializationTimeout = Timeout.InfiniteTimeSpan, + }, LoggerFactory, caller.Token); + await transport.DiscoverReceived.Task.WaitAsync(deadline.Token); + timeProvider.Advance(TimeSpan.FromDays(1)); + Assert.False(connecting.IsCompleted); + caller.Cancel(); + await Assert.ThrowsAnyAsync(() => connecting); + Assert.False(transport.InitializeReceived); + } + + [Fact] + public async Task Client_PinnedModernVersion_ProbeExpiryDoesNotInitialize() + { + using var deadline = CancellationTokenSource.CreateLinkedTokenSource(TestContext.Current.CancellationToken); + deadline.CancelAfter(TestConstants.DefaultTimeout); + await using var transport = new InitializeHandshakeServerTestTransport( + McpProtocolVersions.November2025ProtocolVersion, silentDiscoverProbe: true); + var timeProvider = new FakeTimeProvider(); + var probeBudget = TimeSpan.FromSeconds(5); + var connecting = McpClient.CreateAsync(transport, new McpClientOptions + { + TimeProvider = timeProvider, + ProtocolVersion = McpProtocolVersions.July2026ProtocolVersion, + DiscoverProbeTimeout = probeBudget, + InitializationTimeout = Timeout.InfiniteTimeSpan, + }, LoggerFactory, deadline.Token); + await transport.DiscoverReceived.Task.WaitAsync(deadline.Token); + timeProvider.Advance(probeBudget); + await Assert.ThrowsAsync(() => connecting.WaitAsync(deadline.Token)); + Assert.True(transport.ServerDiscoverProbed); + Assert.False(transport.InitializeReceived); } [Theory] @@ -193,6 +274,20 @@ public void DiscoverProbeTimeout_Setter_Rejects_NonPositiveValues(int millisecon Assert.Throws(() => options.DiscoverProbeTimeout = TimeSpan.FromMilliseconds(milliseconds)); } + [Fact] + public async Task Client_RejectsProbeTimeoutBeyondRuntimeTimerRange() + { + await using var transport = new InitializeHandshakeServerTestTransport(McpProtocolVersions.November2025ProtocolVersion); + using var caller = new CancellationTokenSource(); + await Assert.ThrowsAsync(() => McpClient.CreateAsync(transport, new() + { + DiscoverProbeTimeout = TimeSpan.MaxValue, + InitializationTimeout = Timeout.InfiniteTimeSpan, + }, LoggerFactory, caller.Token)); + caller.Cancel(); + Assert.False(transport.ServerDiscoverProbed); + } + [Fact] public void DiscoverProbeTimeout_Setter_Accepts_PositiveAndInfiniteValues() { @@ -209,6 +304,54 @@ public void DiscoverProbeTimeout_Setter_Accepts_PositiveAndInfiniteValues() Assert.Equal(Timeout.InfiniteTimeSpan, options.DiscoverProbeTimeout); } + [Fact] + public void TimeProvider_DefaultsToSystem_AndRejectsNull() + { + var options = new McpClientOptions(); + Assert.Same(TimeProvider.System, options.TimeProvider); + Assert.Throws(() => options.TimeProvider = null!); + } + + [Fact] + public async Task Client_LegacyInitialization_UsesTimeProvider() + { + var timeProvider = new FakeTimeProvider(); + var timeout = TimeSpan.FromSeconds(10); + await using var transport = new InitializeHandshakeServerTestTransport( + McpProtocolVersions.November2025ProtocolVersion, silentInitialize: true); + var connecting = McpClient.CreateAsync(transport, new() + { + TimeProvider = timeProvider, + ProtocolVersion = McpProtocolVersions.November2025ProtocolVersion, + InitializationTimeout = timeout, + }, LoggerFactory, TestContext.Current.CancellationToken); + await transport.InitializeRequestReceived.Task.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + timeProvider.Advance(timeout - TimeSpan.FromMilliseconds(1)); + Assert.False(connecting.IsCompleted); + timeProvider.Advance(TimeSpan.FromMilliseconds(1)); + + var error = await Assert.ThrowsAsync(() => + connecting.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken)); + Assert.Equal("Initialization timed out", error.Message); + Assert.False(transport.ServerDiscoverProbed); + } + + [Fact] + public async Task Client_CompletedInitialization_IsNotCanceledByAdvancingTime() + { + var timeProvider = new FakeTimeProvider(); + await using var transport = new InitializeHandshakeServerTestTransport(McpProtocolVersions.November2025ProtocolVersion); + await using var client = await McpClient.CreateAsync(transport, new() + { + TimeProvider = timeProvider, + }, LoggerFactory, TestContext.Current.CancellationToken); + + timeProvider.Advance(TimeSpan.FromDays(1)); + await client.PingAsync(cancellationToken: TestContext.Current.CancellationToken).AsTask() + .WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + Assert.False(client.Completion.IsCompleted); + } + [Theory] [InlineData(HttpStatusCode.NotFound, HttpTransportMode.StreamableHttp)] [InlineData(HttpStatusCode.NotFound, HttpTransportMode.AutoDetect)] @@ -392,7 +535,8 @@ private static HttpResponseMessage EmptyResponse(HttpStatusCode status) private sealed class InitializeHandshakeServerTestTransport( string serverNegotiatedVersion, int probeErrorCode = (int)McpErrorCode.MethodNotFound, - bool silentDiscoverProbe = false) : IClientTransport + bool silentDiscoverProbe = false, + bool silentInitialize = false) : IClientTransport { private readonly Channel _incomingToClient = Channel.CreateUnbounded(); @@ -400,8 +544,12 @@ private sealed class InitializeHandshakeServerTestTransport( public bool ServerDiscoverProbed { get; private set; } + public TaskCompletionSource DiscoverReceived { get; } = new(TaskCreationOptions.RunContinuationsAsynchronously); + public bool InitializeReceived { get; private set; } + public TaskCompletionSource InitializeRequestReceived { get; } = new(TaskCreationOptions.RunContinuationsAsynchronously); + public string? InitializeProtocolVersion { get; private set; } public Task ConnectAsync(CancellationToken cancellationToken = default) @@ -418,6 +566,7 @@ private void HandleOutgoingMessage(JsonRpcMessage message) { case JsonRpcRequest { Method: RequestMethods.ServerDiscover } discoverReq: ServerDiscoverProbed = true; + DiscoverReceived.TrySetResult(true); if (silentDiscoverProbe) { // Model an initialize-handshake server that drops the unknown method without replying. @@ -439,6 +588,11 @@ private void HandleOutgoingMessage(JsonRpcMessage message) case JsonRpcRequest { Method: RequestMethods.Initialize } initReq: InitializeReceived = true; + InitializeRequestReceived.TrySetResult(true); + if (silentInitialize) + { + break; + } var initializeRequest = JsonSerializer.Deserialize(initReq.Params, McpJsonUtilities.DefaultOptions); InitializeProtocolVersion = initializeRequest?.ProtocolVersion; _ = WriteAsync(new JsonRpcResponse @@ -452,6 +606,10 @@ private void HandleOutgoingMessage(JsonRpcMessage message) }, McpJsonUtilities.DefaultOptions), }); break; + + case JsonRpcRequest { Method: RequestMethods.Ping } ping: + _ = WriteAsync(new JsonRpcResponse { Id = ping.Id, Result = new JsonObject() }); + break; } }