From 8b14805ce6d26a3e74facba192e98122ca5287a8 Mon Sep 17 00:00:00 2001 From: Stephen Halter Date: Fri, 18 Sep 2026 21:23:16 -0700 Subject: [PATCH 1/2] 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. From 0748f44fa3bad805cad5ecf645b5528acb7fcf4d Mon Sep 17 00:00:00 2001 From: Stephen Halter Date: Mon, 21 Sep 2026 09:41:14 -0700 Subject: [PATCH 2/2] Use TimeProvider for deterministic client timeout tests Expose a per-client clock for initialization and discovery deadlines, defaulting to TimeProvider.System and supporting .NET Standard through Microsoft.Bcl.TimeProvider. Keep caller cancellation and transport connection deadlines unchanged. Drive timeout regressions with FakeTimeProvider and phase gates, surface connection failures while waiting for a phase, and observe initialized notifications without assuming server handler ordering. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- Directory.Packages.props | 1 + docs/concepts/transports/transports.md | 15 +++ .../Client/McpClientImpl.cs | 17 +-- .../Client/McpClientOptions.cs | 19 +++ .../ModelContextProtocol.Core.csproj | 1 + .../RequestTimeout.cs | 35 +++++- .../July2026ProtocolHttpFallbackTests.cs | 11 +- .../OAuth/AuthTests.cs | 10 +- .../OAuth/DiscoveryTimeoutTests.cs | 50 +++++--- .../OAuth/SseDiscoveryTests.cs | 59 +++++++--- .../Utils/AsyncGate.cs | 13 ++- .../Client/July2026ProtocolFallbackTests.cs | 110 +++++++++++++++--- 12 files changed, 264 insertions(+), 77 deletions(-) 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/Client/McpClientImpl.cs b/src/ModelContextProtocol.Core/Client/McpClientImpl.cs index 503c32e4e..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 { @@ -315,9 +316,9 @@ public async Task ConnectAsync(CancellationToken cancellationToken = default) var probeTimeout = _options.DiscoverProbeTimeout; using var probeTimeoutController = !fallbackToInitialize && probeTimeout != Timeout.InfiniteTimeSpan && (_options.InitializationTimeout == Timeout.InfiniteTimeSpan || probeTimeout < _options.InitializationTimeout) - ? new RequestTimeout(probeTimeout, initializationCts.Token) + ? new RequestTimeout(probeTimeout, timeProvider, initializationToken) : null; - var probeToken = probeTimeoutController?.Token ?? initializationCts.Token; + var probeToken = probeTimeoutController?.Token ?? initializationToken; try { @@ -399,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 (probeToken.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. @@ -441,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 { @@ -483,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 e2e200386..7d9c701cc 100644 --- a/src/ModelContextProtocol.Core/Client/McpClientOptions.cs +++ b/src/ModelContextProtocol.Core/Client/McpClientOptions.cs @@ -101,6 +101,25 @@ public sealed class McpClientOptions /// 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. 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/RequestTimeout.cs b/src/ModelContextProtocol.Core/RequestTimeout.cs index 15697aaf6..9fcde4e88 100644 --- a/src/ModelContextProtocol.Core/RequestTimeout.cs +++ b/src/ModelContextProtocol.Core/RequestTimeout.cs @@ -2,26 +2,45 @@ 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. +/// 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, CancellationToken cancellationToken) + public RequestTimeout(TimeSpan timeout, TimeProvider timeProvider, CancellationToken cancellationToken) { _timeout = timeout; _source = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken); Token = _source.Token; - _source.CancelAfter(timeout); + 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() => _source.CancelAfter(Timeout.InfiniteTimeSpan); + public void Stop() => _timer.Change(Timeout.InfiniteTimeSpan, Timeout.InfiniteTimeSpan); public Suspension Suspend() { @@ -29,7 +48,11 @@ public Suspension Suspend() return new Suspension(this); } - public void Dispose() => _source.Dispose(); + public void Dispose() + { + _timer.Dispose(); + _source.Dispose(); + } public readonly struct Suspension(RequestTimeout owner) : IDisposable { @@ -37,7 +60,7 @@ public void Dispose() { if (!owner.Token.IsCancellationRequested) { - owner._source.CancelAfter(owner._timeout); + owner._timer.Change(owner._timeout, Timeout.InfiniteTimeSpan); } } } diff --git a/tests/ModelContextProtocol.AspNetCore.Tests/July2026ProtocolHttpFallbackTests.cs b/tests/ModelContextProtocol.AspNetCore.Tests/July2026ProtocolHttpFallbackTests.cs index c97c4c8f7..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; @@ -85,7 +86,8 @@ private async Task StartServerAsync(RequestDelegate handler, bool acceptGet = fa [InlineData("application/json", 400)] public async Task SilentDiscoverHeadersOrBody_UseProbeBudget(string? contentType, int statusCode) { - var probeBudget = TimeSpan.FromMilliseconds(500); + var timeProvider = new FakeTimeProvider(); + var probeBudget = TimeSpan.FromSeconds(5); var stalled = new AsyncGate(); var methods = new List(); await StartServerAsync(async context => @@ -123,9 +125,10 @@ await StartServerAsync(async context => 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); + 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); } diff --git a/tests/ModelContextProtocol.AspNetCore.Tests/OAuth/AuthTests.cs b/tests/ModelContextProtocol.AspNetCore.Tests/OAuth/AuthTests.cs index 73779f200..476a0ccf3 100644 --- a/tests/ModelContextProtocol.AspNetCore.Tests/OAuth/AuthTests.cs +++ b/tests/ModelContextProtocol.AspNetCore.Tests/OAuth/AuthTests.cs @@ -5,6 +5,7 @@ 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; @@ -1508,7 +1509,8 @@ public async Task CanAuthenticate_WithResourceMetadataPathFallbacks() const string resourcePath = "/mcp"; List wellKnownRequests = []; var metadataGate = new AsyncGate(); - var probeBudget = TimeSpan.FromMilliseconds(500); + var timeProvider = new FakeTimeProvider(); + var probeBudget = TimeSpan.FromSeconds(5); Builder.Services.Configure(options => options.DefaultChallengeScheme = JwtBearerDefaults.AuthenticationScheme); Builder.Services.Configure(options => options.Stateless = true); @@ -1558,8 +1560,10 @@ public async Task CanAuthenticate_WithResourceMetadataPathFallbacks() }, HttpClient, LoggerFactory); var connecting = McpClient.CreateAsync( - transport, new() { DiscoverProbeTimeout = probeBudget }, loggerFactory: LoggerFactory, cancellationToken: TestContext.Current.CancellationToken); - await metadataGate.AssertStillWaitingAsync(probeBudget * 2); + 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); diff --git a/tests/ModelContextProtocol.AspNetCore.Tests/OAuth/DiscoveryTimeoutTests.cs b/tests/ModelContextProtocol.AspNetCore.Tests/OAuth/DiscoveryTimeoutTests.cs index 20e191f4c..287084628 100644 --- a/tests/ModelContextProtocol.AspNetCore.Tests/OAuth/DiscoveryTimeoutTests.cs +++ b/tests/ModelContextProtocol.AspNetCore.Tests/OAuth/DiscoveryTimeoutTests.cs @@ -1,6 +1,7 @@ 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; @@ -12,7 +13,8 @@ namespace ModelContextProtocol.AspNetCore.Tests.OAuth; public class DiscoveryTimeoutTests(ITestOutputHelper outputHelper) : OAuthTestBase(outputHelper) { - private static readonly TimeSpan ProbeBudget = TimeSpan.FromMilliseconds(500); + 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; @@ -26,7 +28,9 @@ public async Task SlowSilentAcquisition_IsExcludedBeforeTheInitialPost() 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); + await cache.Gate.WaitUntilEnteredAsync(connecting); + _timeProvider.Advance(ProbeBudget * 2); + Assert.False(cache.Gate.Token.IsCancellationRequested); Assert.Empty(_methods); cache.Gate.Release.SetResult(); @@ -48,16 +52,17 @@ public async Task Authorization_ObservesCallerAndInitializationCancellation(bool 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); + 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); + await Assert.ThrowsAnyAsync(() => connecting.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken)); } await _authorization.Canceled.Task.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); Assert.Equal(1, _callbackCount); @@ -85,14 +90,15 @@ public async Task SlowAuthorization_PreservesModernProtocol_WithoutSuspendingAno }); 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); + 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.Entered.Task.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + await secondPost.WaitUntilEnteredAsync(second); // 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); + _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); @@ -121,13 +127,16 @@ public async Task Authentication_RestartsProbeBudgetBeforeRetryHeaders() }); 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); + await initialHeaders.WaitUntilEnteredAsync(connecting); + _timeProvider.Advance(ProbeBudget * 0.6); initialHeaders.Release.SetResult(); - await _authorization.Entered.Task.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + await _authorization.WaitUntilEnteredAsync(connecting); + _timeProvider.Advance(ProbeBudget * 2); + Assert.False(_authorization.Token.IsCancellationRequested); _authorization.Release.SetResult(); - await retryHeaders.AssertStillWaitingAsync(options.DiscoverProbeTimeout * 0.6); + await retryHeaders.WaitUntilEnteredAsync(connecting); + _timeProvider.Advance(ProbeBudget * 0.6); retryHeaders.Release.SetResult(); await using var client = await connecting.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); @@ -152,9 +161,10 @@ public async Task AuthenticatedRetryHeaders_RemainProbeBounded() 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); + 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); } @@ -180,9 +190,10 @@ public async Task AuthenticatedDiscoveryBodyTimeout_AbortsModernHandlerWithoutCa 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); + 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); } @@ -223,8 +234,9 @@ private void ConfigureModernServer() }, }, HttpClient, LoggerFactory); - private static McpClientOptions Options(bool pinned = false) => new() + private McpClientOptions Options(bool pinned = false) => new() { + TimeProvider = _timeProvider, DiscoverProbeTimeout = ProbeBudget, ProtocolVersion = pinned ? McpProtocolVersions.July2026ProtocolVersion : null, }; diff --git a/tests/ModelContextProtocol.AspNetCore.Tests/OAuth/SseDiscoveryTests.cs b/tests/ModelContextProtocol.AspNetCore.Tests/OAuth/SseDiscoveryTests.cs index f024caf33..96f2aa1ec 100644 --- a/tests/ModelContextProtocol.AspNetCore.Tests/OAuth/SseDiscoveryTests.cs +++ b/tests/ModelContextProtocol.AspNetCore.Tests/OAuth/SseDiscoveryTests.cs @@ -2,6 +2,7 @@ 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; @@ -14,7 +15,7 @@ namespace ModelContextProtocol.AspNetCore.Tests.OAuth; public class SseDiscoveryTests(ITestOutputHelper outputHelper) : OAuthTestBase(outputHelper) { - private static readonly TimeSpan ProbeBudget = TimeSpan.FromMilliseconds(500); + private static readonly TimeSpan ProbeBudget = TimeSpan.FromSeconds(5); [Theory] [InlineData(HttpTransportMode.AutoDetect, null)] @@ -26,7 +27,9 @@ public class SseDiscoveryTests(ITestOutputHelper outputHelper) : OAuthTestBase(o public async Task Sse_DefaultsToInitialize_AndHonorsExplicitTransportAndVersion(HttpTransportMode mode, string? version) { var methods = new ConcurrentQueue(); - ConfigureSse(methods); + 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) => @@ -43,6 +46,7 @@ public async Task Sse_DefaultsToInitialize_AndHonorsExplicitTransportAndVersion( 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, @@ -50,13 +54,15 @@ public async Task Sse_DefaultsToInitialize_AndHonorsExplicitTransportAndVersion( if (version is null) { // AutoDetect excludes GET establishment from the probe; explicit SSE precedes initialization. - await authorization.AssertStillWaitingAsync(ProbeBudget * 2); + 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); + await Assert.ThrowsAsync(() => connecting.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken)); Assert.Empty(methods); } else @@ -64,9 +70,13 @@ public async Task Sse_DefaultsToInitialize_AndHonorsExplicitTransportAndVersion( 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, NotificationMethods.InitializedNotification, RequestMethods.ToolsList }, methods); + : new[] { RequestMethods.Initialize, RequestMethods.ToolsList }, methods); } Assert.Equal(1, TestOAuthServer.AuthorizationCodeTokenRequestCount); Assert.Equal(mode == HttpTransportMode.Sse ? [] : @@ -78,13 +88,14 @@ public async Task Sse_DefaultsToInitialize_AndHonorsExplicitTransportAndVersion( public async Task ExplicitModernSse_SilentDiscoveryTimesOutWithoutInitialize() { var methods = new ConcurrentQueue(); - var discoveryReceived = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + 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 }) { - discoveryReceived.TrySetResult(); + await discoveryReceived.WaitAsync(cancellationToken); } else { @@ -97,11 +108,13 @@ public async Task ExplicitModernSse_SilentDiscoveryTimesOutWithoutInitialize() 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.Task.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); - await Assert.ThrowsAsync(() => connecting.WaitAsync(ProbeBudget * 8, 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); } @@ -113,37 +126,47 @@ public async Task ExplicitModernSse_SilentDiscoveryTimesOutWithoutInitialize() 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" ? ProbeBudget * 4 : null); + 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); - await authorization.Entered.Task.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + if (deadline != "connection") + { + await authorization.WaitUntilEnteredAsync(connecting); + } if (deadline == "caller") { caller.Cancel(); - await Assert.ThrowsAnyAsync(() => connecting); + await Assert.ThrowsAnyAsync(() => connecting.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken)); } else if (deadline == "initialization") { - var error = await Assert.ThrowsAsync(() => connecting); + timeProvider.Advance(ProbeBudget * 4); + var error = await Assert.ThrowsAsync(() => connecting.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken)); Assert.Equal("Initialization timed out", error.Message); } else { - var error = await Assert.ThrowsAsync(() => connecting); + // 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); } - await authorization.Canceled.Task.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + 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) + private void ConfigureSse(ConcurrentQueue methods, TaskCompletionSource? initialized = null) { TestOAuthServer.ValidResources = [.. TestOAuthServer.ValidResources, $"{McpServerUrl}/sse"]; Builder.Services.Configure(JwtBearerDefaults.AuthenticationScheme, @@ -158,9 +181,9 @@ private void ConfigureSse(ConcurrentQueue methods) { methods.Enqueue(request.Method); } - else if (context.JsonRpcMessage is JsonRpcNotification notification) + else if (context.JsonRpcMessage is JsonRpcNotification { Method: NotificationMethods.InitializedNotification }) { - methods.Enqueue(notification.Method); + initialized?.TrySetResult(); } await next(context, cancellationToken); })); diff --git a/tests/ModelContextProtocol.AspNetCore.Tests/Utils/AsyncGate.cs b/tests/ModelContextProtocol.AspNetCore.Tests/Utils/AsyncGate.cs index 549ebc94b..dafdbfd06 100644 --- a/tests/ModelContextProtocol.AspNetCore.Tests/Utils/AsyncGate.cs +++ b/tests/ModelContextProtocol.AspNetCore.Tests/Utils/AsyncGate.cs @@ -1,3 +1,4 @@ +using ModelContextProtocol.Client; using ModelContextProtocol.Tests.Utils; namespace ModelContextProtocol.AspNetCore.Tests.Utils; @@ -7,9 +8,11 @@ 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 { @@ -22,9 +25,13 @@ public async Task WaitAsync(CancellationToken cancellationToken) } } - public async Task AssertStillWaitingAsync(TimeSpan duration) + public async Task WaitUntilEnteredAsync(Task connecting) { - await Entered.Task.WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); - await Assert.ThrowsAsync(() => Canceled.Task.WaitAsync(duration, TestContext.Current.CancellationToken)); + 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 1dd5c186b..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; @@ -168,24 +168,28 @@ public async Task Client_OnSilentProbe_FallsBackTo_Initialize_AfterConfiguredPro 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), + 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] @@ -198,12 +202,16 @@ public async Task Client_InitializationDeadlineWins_NoFallback(int probeMillisec deadline.CancelAfter(TestConstants.DefaultTimeout); await using var transport = new InitializeHandshakeServerTestTransport( McpProtocolVersions.November2025ProtocolVersion, silentDiscoverProbe: true); - - var exception = await Assert.ThrowsAsync(() => McpClient.CreateAsync(transport, new McpClientOptions + var timeProvider = new FakeTimeProvider(); + var connecting = McpClient.CreateAsync(transport, new McpClientOptions { + TimeProvider = timeProvider, DiscoverProbeTimeout = TimeSpan.FromMilliseconds(probeMilliseconds), InitializationTimeout = TimeSpan.FromMilliseconds(initializationMilliseconds), - }, LoggerFactory, deadline.Token)); + }, 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); @@ -219,12 +227,16 @@ public async Task Client_InfiniteProbeAndInitialization_ObserveCallerCancellatio 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); @@ -237,12 +249,18 @@ public async Task Client_PinnedModernVersion_ProbeExpiryDoesNotInitialize() deadline.CancelAfter(TestConstants.DefaultTimeout); await using var transport = new InitializeHandshakeServerTestTransport( McpProtocolVersions.November2025ProtocolVersion, silentDiscoverProbe: true); - await Assert.ThrowsAsync(() => McpClient.CreateAsync(transport, new McpClientOptions + var timeProvider = new FakeTimeProvider(); + var probeBudget = TimeSpan.FromSeconds(5); + var connecting = McpClient.CreateAsync(transport, new McpClientOptions { + TimeProvider = timeProvider, ProtocolVersion = McpProtocolVersions.July2026ProtocolVersion, - DiscoverProbeTimeout = TimeSpan.FromMilliseconds(250), + DiscoverProbeTimeout = probeBudget, InitializationTimeout = Timeout.InfiniteTimeSpan, - }, LoggerFactory, deadline.Token)); + }, 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); } @@ -286,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)] @@ -469,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(); @@ -481,6 +548,8 @@ private sealed class InitializeHandshakeServerTestTransport( 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) @@ -519,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 @@ -532,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; } }