diff --git a/src/ModelContextProtocol.Core/Client/SseClientSessionTransport.cs b/src/ModelContextProtocol.Core/Client/SseClientSessionTransport.cs index 03969b4d9..99bdc1eb9 100644 --- a/src/ModelContextProtocol.Core/Client/SseClientSessionTransport.cs +++ b/src/ModelContextProtocol.Core/Client/SseClientSessionTransport.cs @@ -23,6 +23,7 @@ internal sealed partial class SseClientSessionTransport : TransportBase private Task? _receiveTask; private readonly ILogger _logger; private readonly TaskCompletionSource _connectionEstablished; + private volatile bool _sseAdopted; /// /// SSE transport for a single session. Unlike stdio it does not launch a process, but connects to an existing server. @@ -107,6 +108,8 @@ public override async Task SendMessageAsync( throw HttpResponseMessageExtensions.CreateHttpRequestException(response, responseBody); } + + _sseAdopted = true; } private async Task CloseAsync() @@ -129,7 +132,7 @@ private async Task CloseAsync() } finally { - SetDisconnected(new ClientTransportClosedException(new HttpClientCompletionDetails())); + SetSseDisconnected(new ClientTransportClosedException(new HttpClientCompletionDetails())); } } @@ -190,7 +193,7 @@ private async Task ReceiveMessagesAsync(CancellationToken cancellationToken) } else { - SetDisconnected(new ClientTransportClosedException(new HttpClientCompletionDetails + SetSseDisconnected(new ClientTransportClosedException(new HttpClientCompletionDetails { HttpStatusCode = failureStatusCode, Exception = ex, @@ -202,7 +205,18 @@ private async Task ReceiveMessagesAsync(CancellationToken cancellationToken) } finally { - SetDisconnected(new ClientTransportClosedException(new HttpClientCompletionDetails())); + SetSseDisconnected(new ClientTransportClosedException(new HttpClientCompletionDetails())); + } + } + + private void SetSseDisconnected(Exception error) + { + // If AutoDetect is still probing SSE, leave its shared message channel open so it can + // retry with another transport. A successful POST means SSE was selected and owns the + // channel from that point on, matching Streamable HTTP's adoption behavior. + if (_options.TransportMode is not HttpTransportMode.AutoDetect || _sseAdopted) + { + SetDisconnected(error); } } diff --git a/tests/ModelContextProtocol.Tests/Transport/HttpClientTransportAutoDetectTests.cs b/tests/ModelContextProtocol.Tests/Transport/HttpClientTransportAutoDetectTests.cs index e6386adc6..45919dda1 100644 --- a/tests/ModelContextProtocol.Tests/Transport/HttpClientTransportAutoDetectTests.cs +++ b/tests/ModelContextProtocol.Tests/Transport/HttpClientTransportAutoDetectTests.cs @@ -2,6 +2,7 @@ using ModelContextProtocol.Protocol; using ModelContextProtocol.Tests.Utils; using Microsoft.Extensions.Logging; +using System.IO.Pipelines; using System.Net; namespace ModelContextProtocol.Tests.Transport; @@ -163,6 +164,115 @@ public async Task AutoDetectMode_FallsBackToSse_WhenStreamableHttpFails() Assert.NotNull(session); } + [Fact] + public async Task AutoDetectMode_WhenProvisionalSseFails_LeavesSharedMessageChannelOpen() + { + var options = new HttpClientTransportOptions + { + Endpoint = new Uri("http://localhost"), + TransportMode = HttpTransportMode.AutoDetect, + Name = "AutoDetect shared channel test client" + }; + + using var mockHttpHandler = new MockHttpHandler(); + using var httpClient = new HttpClient(mockHttpHandler); + await using var transport = new HttpClientTransport(options, httpClient, LoggerFactory); + var streamableHttpPostCount = 0; + + mockHttpHandler.RequestHandler = request => + { + if (request.Method == HttpMethod.Get) + { + return Task.FromResult(new HttpResponseMessage(HttpStatusCode.MethodNotAllowed)); + } + + if (request.Method == HttpMethod.Post && ++streamableHttpPostCount == 1) + { + return Task.FromResult(new HttpResponseMessage(HttpStatusCode.NotFound) + { + Content = new StringContent("Invalid session ID"), + }); + } + + return Task.FromResult(new HttpResponseMessage(HttpStatusCode.OK) + { + Content = new StringContent( + "{\"jsonrpc\":\"2.0\",\"id\":2,\"result\":{\"protocolVersion\":\"2025-11-25\",\"capabilities\":{},\"serverInfo\":{\"name\":\"test\",\"version\":\"1.0\"}}}", + System.Text.Encoding.UTF8, + "application/json"), + }); + }; + + await using var session = await transport.ConnectAsync(TestContext.Current.CancellationToken); + + await Assert.ThrowsAsync(() => + session.SendMessageAsync( + new JsonRpcRequest { Method = RequestMethods.ServerDiscover, Id = new RequestId(1) }, + TestContext.Current.CancellationToken)); + + await session.SendMessageAsync( + new JsonRpcRequest { Method = RequestMethods.Initialize, Id = new RequestId(2) }, + TestContext.Current.CancellationToken); + + var response = await session.MessageReader.ReadAsync(TestContext.Current.CancellationToken); + Assert.Equal(new RequestId(2), Assert.IsType(response).Id); + } + + [Fact] + public async Task AutoDetectMode_WhenAdoptedSseDisconnects_CompletesSharedMessageChannel() + { + var options = new HttpClientTransportOptions + { + Endpoint = new Uri("http://localhost"), + TransportMode = HttpTransportMode.AutoDetect, + Name = "AutoDetect adopted SSE test client" + }; + + using var mockHttpHandler = new MockHttpHandler(); + using var httpClient = new HttpClient(mockHttpHandler); + await using var transport = new HttpClientTransport(options, httpClient, LoggerFactory); + var ssePipe = new Pipe(); + var postCount = 0; + + await ssePipe.Writer.WriteAsync( + System.Text.Encoding.UTF8.GetBytes("event: endpoint\r\ndata: /sse-endpoint\r\n\r\n"), + TestContext.Current.CancellationToken); + + mockHttpHandler.RequestHandler = request => + { + if (request.Method == HttpMethod.Get) + { + var content = new StreamContent(ssePipe.Reader.AsStream()); + content.Headers.ContentType = new("text/event-stream"); + return Task.FromResult(new HttpResponseMessage(HttpStatusCode.OK) { Content = content }); + } + + if (request.Method == HttpMethod.Post && ++postCount == 1) + { + return Task.FromResult(new HttpResponseMessage(HttpStatusCode.NotFound) + { + Content = new StringContent("Streamable HTTP not supported"), + }); + } + + return Task.FromResult(new HttpResponseMessage(HttpStatusCode.OK)); + }; + + await using var session = await transport.ConnectAsync(TestContext.Current.CancellationToken); + + await session.SendMessageAsync( + new JsonRpcRequest { Method = RequestMethods.Initialize, Id = new RequestId(1) }, + TestContext.Current.CancellationToken); + + await ssePipe.Writer.CompleteAsync(); + + var exception = await Assert.ThrowsAsync( + async () => await session.MessageReader.Completion.WaitAsync( + TestConstants.DefaultTimeout, + TestContext.Current.CancellationToken)); + Assert.IsType(exception.Details); + } + // Regression test for https://github.com/modelcontextprotocol/csharp-sdk/issues/1526 // When Streamable HTTP returns 415 (e.g. wrong Content-Type) and the SSE fallback also fails // (e.g. a Streamable-HTTP-only server returns 405 to the GET), the surfaced exception must