From 8b14805ce6d26a3e74facba192e98122ca5287a8 Mon Sep 17 00:00:00 2001 From: Stephen Halter Date: Fri, 18 Sep 2026 21:23:16 -0700 Subject: [PATCH] Keep OAuth outside the discovery probe timeout Pause the request-local discovery timer during SDK token acquisition while keeping response waits, initialization, and caller cancellation bounded. Stop the abandoned probe when AutoDetect selects SSE and initialize by default, while honoring explicit modern SSE configuration. Preserve the existing transport lifecycle and add focused authentication, deadline, and protocol fallback regression coverage on top of #1855. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../Authentication/ClientOAuthProvider.cs | 13 +- .../AutoDetectingClientSessionTransport.cs | 15 +- .../Client/McpClientImpl.cs | 33 ++- .../Client/McpClientOptions.cs | 15 ++ .../ServerDiscoverSkippedForSseException.cs | 5 + .../McpSession.Methods.cs | 5 +- .../Protocol/JsonRpcMessageContext.cs | 5 + .../RequestTimeout.cs | 44 ++++ .../July2026ProtocolHttpFallbackTests.cs | 52 ++++ .../MapMcpTests.Mrtr.cs | 8 +- .../MapMcpTests.cs | 10 +- .../OAuth/AuthTests.cs | 14 +- .../OAuth/DiscoveryTimeoutTests.cs | 247 ++++++++++++++++++ .../OAuth/SseDiscoveryTests.cs | 187 +++++++++++++ .../Utils/AsyncGate.cs | 30 +++ .../Client/July2026ProtocolFallbackTests.cs | 88 ++++++- 16 files changed, 738 insertions(+), 33 deletions(-) create mode 100644 src/ModelContextProtocol.Core/Client/ServerDiscoverSkippedForSseException.cs create mode 100644 src/ModelContextProtocol.Core/RequestTimeout.cs create mode 100644 tests/ModelContextProtocol.AspNetCore.Tests/OAuth/DiscoveryTimeoutTests.cs create mode 100644 tests/ModelContextProtocol.AspNetCore.Tests/OAuth/SseDiscoveryTests.cs create mode 100644 tests/ModelContextProtocol.AspNetCore.Tests/Utils/AsyncGate.cs 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..503c32e4e 100644 --- a/src/ModelContextProtocol.Core/Client/McpClientImpl.cs +++ b/src/ModelContextProtocol.Core/Client/McpClientImpl.cs @@ -296,31 +296,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, initializationCts.Token) + : null; + var probeToken = probeTimeoutController?.Token ?? initializationCts.Token; 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 +354,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 +399,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 && !initializationCts.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. @@ -465,6 +473,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); } } diff --git a/src/ModelContextProtocol.Core/Client/McpClientOptions.cs b/src/ModelContextProtocol.Core/Client/McpClientOptions.cs index 61a0613df..e2e200386 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,6 +90,11 @@ 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. /// @@ -121,6 +130,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/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..15697aaf6 --- /dev/null +++ b/src/ModelContextProtocol.Core/RequestTimeout.cs @@ -0,0 +1,44 @@ +namespace ModelContextProtocol; + +/// A request-local timer that can be suspended without suspending linked cancellation. +/// +/// Owned by one awaited discovery request, linked to the enclosing initialization scope. +/// 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 TimeSpan _timeout; + + public RequestTimeout(TimeSpan timeout, CancellationToken cancellationToken) + { + _timeout = timeout; + _source = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken); + Token = _source.Token; + _source.CancelAfter(timeout); + } + + public CancellationToken Token { get; } + + public void Stop() => _source.CancelAfter(Timeout.InfiniteTimeSpan); + + public Suspension Suspend() + { + Stop(); + return new Suspension(this); + } + + public void Dispose() => _source.Dispose(); + + public readonly struct Suspension(RequestTimeout owner) : IDisposable + { + public void Dispose() + { + if (!owner.Token.IsCancellationRequested) + { + owner._source.CancelAfter(owner._timeout); + } + } + } +} diff --git a/tests/ModelContextProtocol.AspNetCore.Tests/July2026ProtocolHttpFallbackTests.cs b/tests/ModelContextProtocol.AspNetCore.Tests/July2026ProtocolHttpFallbackTests.cs index 9126331de..c97c4c8f7 100644 --- a/tests/ModelContextProtocol.AspNetCore.Tests/July2026ProtocolHttpFallbackTests.cs +++ b/tests/ModelContextProtocol.AspNetCore.Tests/July2026ProtocolHttpFallbackTests.cs @@ -78,6 +78,58 @@ 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 probeBudget = TimeSpan.FromMilliseconds(500); + 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() { DiscoverProbeTimeout = probeBudget }, LoggerFactory, TestContext.Current.CancellationToken); + await stalled.Entered.Task.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + await using var client = await connecting.WaitAsync(probeBudget * 8, 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..73779f200 100644 --- a/tests/ModelContextProtocol.AspNetCore.Tests/OAuth/AuthTests.cs +++ b/tests/ModelContextProtocol.AspNetCore.Tests/OAuth/AuthTests.cs @@ -7,6 +7,7 @@ using Microsoft.Extensions.DependencyInjection; using ModelContextProtocol; using ModelContextProtocol.AspNetCore.Authentication; +using ModelContextProtocol.AspNetCore.Tests.Utils; using ModelContextProtocol.Authentication; using ModelContextProtocol.Client; using ModelContextProtocol.Protocol; @@ -1506,8 +1507,11 @@ public async Task CanAuthenticate_WithResourceMetadataPathFallbacks() { const string resourcePath = "/mcp"; List wellKnownRequests = []; + var metadataGate = new AsyncGate(); + var probeBudget = TimeSpan.FromMilliseconds(500); 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 +1527,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 +1557,14 @@ 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() { DiscoverProbeTimeout = probeBudget }, loggerFactory: LoggerFactory, cancellationToken: TestContext.Current.CancellationToken); + await metadataGate.AssertStillWaitingAsync(probeBudget * 2); + 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..20e191f4c --- /dev/null +++ b/tests/ModelContextProtocol.AspNetCore.Tests/OAuth/DiscoveryTimeoutTests.cs @@ -0,0 +1,247 @@ +using Microsoft.AspNetCore.Builder; +using Microsoft.AspNetCore.Http; +using Microsoft.Extensions.DependencyInjection; +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.FromMilliseconds(500); + 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.AssertStillWaitingAsync(ProbeBudget * 2); + 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.Entered.Task.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + if (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); + } + 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.Entered.Task.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + var secondOptions = Options(pinned: true); + secondOptions.DiscoverProbeTimeout = ProbeBudget * 2; + var second = McpClient.CreateAsync(transport, secondOptions, LoggerFactory, TestContext.Current.CancellationToken); + await secondPost.Entered.Task.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + // The second probe expiring proves the first authorization survived more than its own budget. + await Assert.ThrowsAsync(() => second.WaitAsync(ProbeBudget * 8, TestContext.Current.CancellationToken)); + Assert.False(_authorization.Canceled.Task.IsCompleted); + _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); + options.DiscoverProbeTimeout = TimeSpan.FromSeconds(2); + var connecting = McpClient.CreateAsync(transport, options, LoggerFactory, TestContext.Current.CancellationToken); + await initialHeaders.AssertStillWaitingAsync(options.DiscoverProbeTimeout * 0.6); + initialHeaders.Release.SetResult(); + await _authorization.Entered.Task.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + _authorization.Release.SetResult(); + await retryHeaders.AssertStillWaitingAsync(options.DiscoverProbeTimeout * 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.Entered.Task.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + await headers.Canceled.Task.WaitAsync(ProbeBudget * 8, TestContext.Current.CancellationToken); + await Assert.ThrowsAsync(() => connecting); + 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.Entered.Task.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + await handler.Canceled.Task.WaitAsync(ProbeBudget * 4, TestContext.Current.CancellationToken); + await Assert.ThrowsAsync(() => connecting); + 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 static McpClientOptions Options(bool pinned = false) => new() + { + 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..f024caf33 --- /dev/null +++ b/tests/ModelContextProtocol.AspNetCore.Tests/OAuth/SseDiscoveryTests.cs @@ -0,0 +1,187 @@ +using Microsoft.AspNetCore.Authentication.JwtBearer; +using Microsoft.AspNetCore.Builder; +using Microsoft.AspNetCore.Http; +using Microsoft.Extensions.DependencyInjection; +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.FromMilliseconds(500); + + [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(); + ConfigureSse(methods); + 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() + { + 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.AssertStillWaitingAsync(ProbeBudget * 2); + } + authorization.Release.SetResult(); + bool modern = version == McpProtocolVersions.July2026ProtocolVersion; + if (modern && mode == HttpTransportMode.AutoDetect) + { + await Assert.ThrowsAsync(() => connecting); + 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)); + Assert.Equal(modern + ? [RequestMethods.ServerDiscover, RequestMethods.ToolsList] + : new[] { RequestMethods.Initialize, NotificationMethods.InitializedNotification, 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 discoveryReceived = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + ConfigureSse(methods); + Builder.Services.AddMcpServer().WithMessageFilters(filters => filters.AddIncomingFilter(next => async (context, cancellationToken) => + { + if (context.JsonRpcMessage is JsonRpcRequest { Method: RequestMethods.ServerDiscover }) + { + discoveryReceived.TrySetResult(); + } + 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() + { + ProtocolVersion = McpProtocolVersions.July2026ProtocolVersion, + DiscoverProbeTimeout = ProbeBudget, + }, LoggerFactory, TestContext.Current.CancellationToken); + await discoveryReceived.Task.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + await Assert.ThrowsAsync(() => connecting.WaitAsync(ProbeBudget * 8, 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 authorization = new AsyncGate(); + await using var app = await StartMcpServerAsync(); + await using var transport = CreateTransport(HttpTransportMode.AutoDetect, authorization, + deadline == "connection" ? ProbeBudget * 4 : null); + using var caller = CancellationTokenSource.CreateLinkedTokenSource(TestContext.Current.CancellationToken); + var connecting = McpClient.CreateAsync(transport, new() + { + DiscoverProbeTimeout = ProbeBudget, + InitializationTimeout = deadline == "initialization" ? ProbeBudget * 4 : TestConstants.DefaultTimeout, + }, LoggerFactory, caller.Token); + await authorization.Entered.Task.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + if (deadline == "caller") + { + caller.Cancel(); + await Assert.ThrowsAnyAsync(() => connecting); + } + else if (deadline == "initialization") + { + var error = await Assert.ThrowsAsync(() => connecting); + Assert.Equal("Initialization timed out", error.Message); + } + else + { + var error = await Assert.ThrowsAsync(() => connecting); + Assert.IsType(error.InnerException); + } + await authorization.Canceled.Task.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + Assert.Equal(0, TestOAuthServer.AuthorizationCodeTokenRequestCount); + } + + private void ConfigureSse(ConcurrentQueue methods) + { + 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 notification) + { + methods.Enqueue(notification.Method); + } + 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..549ebc94b --- /dev/null +++ b/tests/ModelContextProtocol.AspNetCore.Tests/Utils/AsyncGate.cs @@ -0,0 +1,30 @@ +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 async Task WaitAsync(CancellationToken cancellationToken) + { + Entered.TrySetResult(); + try + { + await Release.Task.WaitAsync(TestConstants.DefaultTimeout, cancellationToken); + } + catch (OperationCanceledException) + { + Canceled.TrySetResult(); + throw; + } + } + + public async Task AssertStillWaitingAsync(TimeSpan duration) + { + await Entered.Task.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + await Assert.ThrowsAsync(() => Canceled.Task.WaitAsync(duration, TestContext.Current.CancellationToken)); + } +} diff --git a/tests/ModelContextProtocol.Tests/Client/July2026ProtocolFallbackTests.cs b/tests/ModelContextProtocol.Tests/Client/July2026ProtocolFallbackTests.cs index 557dc5655..1dd5c186b 100644 --- a/tests/ModelContextProtocol.Tests/Client/July2026ProtocolFallbackTests.cs +++ b/tests/ModelContextProtocol.Tests/Client/July2026ProtocolFallbackTests.cs @@ -153,13 +153,17 @@ 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); @@ -169,7 +173,7 @@ public async Task Client_OnSilentProbe_FallsBackTo_Initialize_AfterConfiguredPro await using var client = await McpClient.CreateAsync(transport, new McpClientOptions { DiscoverProbeTimeout = TimeSpan.FromMilliseconds(250), - InitializationTimeout = TestConstants.DefaultTimeout, + InitializationTimeout = infiniteInitialization ? Timeout.InfiniteTimeSpan : TestConstants.DefaultTimeout, }, loggerFactory: LoggerFactory, cancellationToken: ct); stopwatch.Stop(); @@ -184,6 +188,65 @@ public async Task Client_OnSilentProbe_FallsBackTo_Initialize_AfterConfiguredPro $"Fallback should have happened shortly after the {nameof(McpClientOptions.DiscoverProbeTimeout)}, but took {stopwatch.Elapsed}."); } + [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 exception = await Assert.ThrowsAsync(() => McpClient.CreateAsync(transport, new McpClientOptions + { + DiscoverProbeTimeout = TimeSpan.FromMilliseconds(probeMilliseconds), + InitializationTimeout = TimeSpan.FromMilliseconds(initializationMilliseconds), + }, LoggerFactory, 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 connecting = McpClient.CreateAsync(transport, new McpClientOptions + { + DiscoverProbeTimeout = Timeout.InfiniteTimeSpan, + InitializationTimeout = Timeout.InfiniteTimeSpan, + }, LoggerFactory, caller.Token); + await transport.DiscoverReceived.Task.WaitAsync(deadline.Token); + 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); + await Assert.ThrowsAsync(() => McpClient.CreateAsync(transport, new McpClientOptions + { + ProtocolVersion = McpProtocolVersions.July2026ProtocolVersion, + DiscoverProbeTimeout = TimeSpan.FromMilliseconds(250), + InitializationTimeout = Timeout.InfiniteTimeSpan, + }, LoggerFactory, deadline.Token)); + Assert.True(transport.ServerDiscoverProbed); + Assert.False(transport.InitializeReceived); + } + [Theory] [InlineData(0)] [InlineData(-1000)] @@ -193,6 +256,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() { @@ -400,6 +477,8 @@ private sealed class InitializeHandshakeServerTestTransport( public bool ServerDiscoverProbed { get; private set; } + public TaskCompletionSource DiscoverReceived { get; } = new(TaskCreationOptions.RunContinuationsAsynchronously); + public bool InitializeReceived { get; private set; } public string? InitializeProtocolVersion { get; private set; } @@ -418,6 +497,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.