Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -14,17 +14,19 @@ namespace ModelContextProtocol.Client;
internal sealed partial class AutoDetectingClientSessionTransport : ITransport
{
private readonly HttpClientTransportOptions _options;
private readonly Uri _endpoint;
private readonly McpHttpClient _httpClient;
private readonly ILoggerFactory? _loggerFactory;
private readonly ILogger _logger;
private readonly string _name;
private readonly Channel<JsonRpcMessage> _messageChannel;

public AutoDetectingClientSessionTransport(string endpointName, HttpClientTransportOptions transportOptions, McpHttpClient httpClient, ILoggerFactory? loggerFactory)
public AutoDetectingClientSessionTransport(string endpointName, Uri endpoint, HttpClientTransportOptions transportOptions, McpHttpClient httpClient, ILoggerFactory? loggerFactory)
{
Throw.IfNull(transportOptions);
Throw.IfNull(httpClient);

_endpoint = endpoint;
_options = transportOptions;
_httpClient = httpClient;
_loggerFactory = loggerFactory;
Expand Down Expand Up @@ -62,7 +64,7 @@ public Task SendMessageAsync(JsonRpcMessage message, CancellationToken cancellat
private async Task InitializeAsync(JsonRpcMessage message, CancellationToken cancellationToken)
{
// Try StreamableHttp first
var streamableHttpTransport = new StreamableHttpClientSessionTransport(_name, _options, _httpClient, _messageChannel, _loggerFactory);
var streamableHttpTransport = new StreamableHttpClientSessionTransport(_name, _endpoint, _options, _httpClient, _messageChannel, _loggerFactory);

try
{
Expand Down Expand Up @@ -117,7 +119,7 @@ private async Task InitializeSseTransportAsync(JsonRpcMessage message, Cancellat
throw new InvalidOperationException("Streamable HTTP transport is required to resume an existing session.");
}

var sseTransport = new SseClientSessionTransport(_name, _options, _httpClient, _messageChannel, _loggerFactory);
var sseTransport = new SseClientSessionTransport(_name, _endpoint, _options, _httpClient, _messageChannel, _loggerFactory);

try
{
Expand Down Expand Up @@ -166,4 +168,4 @@ public async ValueTask DisposeAsync()

[LoggerMessage(Level = LogLevel.Information, Message = "{EndpointName} using SSE transport.")]
private partial void LogUsingSSE(string endpointName);
}
}
15 changes: 10 additions & 5 deletions src/ModelContextProtocol.Core/Client/HttpClientTransport.cs
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@ public sealed class HttpClientTransport : IClientTransport, IAsyncDisposable
private readonly HttpClientTransportOptions _options;
private readonly McpHttpClient _mcpHttpClient;
private readonly ILoggerFactory? _loggerFactory;
private readonly Uri _endpoint;

private readonly HttpClient? _ownedHttpClient;

Expand Down Expand Up @@ -49,11 +50,15 @@ public HttpClientTransport(HttpClientTransportOptions transportOptions, HttpClie

_options = transportOptions;
_loggerFactory = loggerFactory;
Name = transportOptions.Name ?? transportOptions.Endpoint.ToString();
_endpoint = transportOptions.Endpoint ?? httpClient.BaseAddress ??
throw new ArgumentException(
$"Either '{nameof(HttpClientTransportOptions)}.{nameof(HttpClientTransportOptions.Endpoint)}' or '{nameof(HttpClient)}.{nameof(HttpClient.BaseAddress)}' must be set.",
nameof(transportOptions));
Name = transportOptions.Name ?? _endpoint.ToString();

if (transportOptions.OAuth is { } clientOAuthOptions)
{
_mcpHttpClient = new ClientOAuthProvider(_options.Endpoint, clientOAuthOptions, httpClient, loggerFactory);
_mcpHttpClient = new ClientOAuthProvider(_endpoint, clientOAuthOptions, httpClient, loggerFactory);
}
else
{
Expand All @@ -79,16 +84,16 @@ public async Task<ITransport> ConnectAsync(CancellationToken cancellationToken =

return _options.TransportMode switch
{
HttpTransportMode.AutoDetect => new AutoDetectingClientSessionTransport(Name, _options, _mcpHttpClient, _loggerFactory),
HttpTransportMode.StreamableHttp => new StreamableHttpClientSessionTransport(Name, _options, _mcpHttpClient, messageChannel: null, _loggerFactory),
HttpTransportMode.AutoDetect => new AutoDetectingClientSessionTransport(Name, _endpoint, _options, _mcpHttpClient, _loggerFactory),
HttpTransportMode.StreamableHttp => new StreamableHttpClientSessionTransport(Name, _endpoint, _options, _mcpHttpClient, messageChannel: null, _loggerFactory),
HttpTransportMode.Sse => await ConnectSseTransportAsync(cancellationToken).ConfigureAwait(false),
_ => throw new InvalidOperationException($"Unsupported transport mode: {_options.TransportMode}"),
};
}

private async Task<ITransport> ConnectSseTransportAsync(CancellationToken cancellationToken)
{
var sessionTransport = new SseClientSessionTransport(Name, _options, _mcpHttpClient, messageChannel: null, _loggerFactory);
var sessionTransport = new SseClientSessionTransport(Name, _endpoint, _options, _mcpHttpClient, messageChannel: null, _loggerFactory);

try
{
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -8,24 +8,23 @@ namespace ModelContextProtocol.Client;
public sealed class HttpClientTransportOptions
{
/// <summary>
/// Gets or sets the base address of the server for SSE connections.
/// Gets or sets the base address of the server for HTTP connections.
/// </summary>
/// <exception cref="ArgumentNullException">The value is <see langword="null"/>.</exception>
/// <exception cref="ArgumentException">The value is not an absolute URI, or does not use the HTTP or HTTPS scheme.</exception>
public required Uri Endpoint
/// <remarks>
/// This can be omitted when constructing the transport with an <see cref="HttpClient"/> that has a
/// <see cref="HttpClient.BaseAddress"/>. An explicitly configured endpoint takes precedence over the client base address.
/// </remarks>
public Uri? Endpoint
{
get;
set
{
if (value is null)
{
throw new ArgumentNullException(nameof(value), "Endpoint cannot be null.");
}
if (!value.IsAbsoluteUri)
if (value is not null && !value.IsAbsoluteUri)
{
throw new ArgumentException("Endpoint must be an absolute URI.", nameof(value));
}
if (value.Scheme != Uri.UriSchemeHttp && value.Scheme != Uri.UriSchemeHttps)
if (value is not null && value.Scheme != Uri.UriSchemeHttp && value.Scheme != Uri.UriSchemeHttps)
{
throw new ArgumentException("Endpoint must use HTTP or HTTPS scheme.", nameof(value));
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,7 @@ internal sealed partial class SseClientSessionTransport : TransportBase
/// </summary>
public SseClientSessionTransport(
string endpointName,
Uri endpoint,
HttpClientTransportOptions transportOptions,
McpHttpClient httpClient,
Channel<JsonRpcMessage>? messageChannel,
Expand All @@ -40,7 +41,7 @@ public SseClientSessionTransport(
Throw.IfNull(httpClient);

_options = transportOptions;
_sseEndpoint = transportOptions.Endpoint;
_sseEndpoint = endpoint;
_httpClient = httpClient;
_connectionCts = new CancellationTokenSource();
_logger = (ILogger?)loggerFactory?.CreateLogger<HttpClientTransport>() ?? NullLogger.Instance;
Expand Down Expand Up @@ -265,4 +266,4 @@ private void HandleEndpointEvent(string data)

[LoggerMessage(Level = LogLevel.Trace, Message = "{EndpointName} rejected SSE transport POST for message ID '{MessageId}'. Server response: '{responseContent}'.")]
private partial void LogRejectedPostSensitive(string endpointName, string messageId, string responseContent);
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@ internal sealed partial class StreamableHttpClientSessionTransport : TransportBa
private static readonly MediaTypeWithQualityHeaderValue s_textEventStreamMediaType = new("text/event-stream");

private readonly McpHttpClient _httpClient;
private readonly Uri _endpoint;
private readonly HttpClientTransportOptions _options;
private readonly CancellationTokenSource _connectionCts = new();
private readonly ILogger _logger;
Expand All @@ -34,6 +35,7 @@ internal sealed partial class StreamableHttpClientSessionTransport : TransportBa

public StreamableHttpClientSessionTransport(
string endpointName,
Uri endpoint,
HttpClientTransportOptions transportOptions,
McpHttpClient httpClient,
Channel<JsonRpcMessage>? messageChannel,
Expand All @@ -43,6 +45,7 @@ public StreamableHttpClientSessionTransport(
Throw.IfNull(transportOptions);
Throw.IfNull(httpClient);

_endpoint = endpoint;
_options = transportOptions;
_httpClient = httpClient;
_logger = (ILogger?)loggerFactory?.CreateLogger<HttpClientTransport>() ?? NullLogger.Instance;
Expand Down Expand Up @@ -155,7 +158,7 @@ internal async Task<HttpResponseMessage> SendHttpRequestAsync(JsonRpcMessage mes
using var sendCts = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken, _connectionCts.Token);
cancellationToken = sendCts.Token;

using var httpRequestMessage = new HttpRequestMessage(HttpMethod.Post, _options.Endpoint)
using var httpRequestMessage = new HttpRequestMessage(HttpMethod.Post, _endpoint)
{
Headers =
{
Expand Down Expand Up @@ -377,7 +380,7 @@ await SendGetSseRequestWithRetriesAsync(
}
shouldDelay = true;

using var request = new HttpRequestMessage(HttpMethod.Get, _options.Endpoint);
using var request = new HttpRequestMessage(HttpMethod.Get, _endpoint);
request.Headers.Accept.Add(s_textEventStreamMediaType);
CopyAdditionalHeaders(request.Headers, _options.AdditionalHeaders, SessionId, _negotiatedProtocolVersion, state.LastEventId);

Expand Down Expand Up @@ -516,7 +519,7 @@ message is JsonRpcMessageWithId rpcResponseOrError &&

private async Task SendDeleteRequest()
{
using var deleteRequest = new HttpRequestMessage(HttpMethod.Delete, _options.Endpoint);
using var deleteRequest = new HttpRequestMessage(HttpMethod.Delete, _endpoint);
CopyAdditionalHeaders(deleteRequest.Headers, _options.AdditionalHeaders, SessionId, _negotiatedProtocolVersion);

// Do not validate we get a successful status code, because server support for the DELETE request is optional
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,75 @@ public void Constructor_Throws_For_Null_HttpClient()
Assert.Equal("httpClient", exception.ParamName);
}

[Fact]
public void Constructor_Throws_When_No_Endpoint_Is_Available()
{
var options = new HttpClientTransportOptions();
using var httpClient = new HttpClient();

var exception = Assert.Throws<ArgumentException>(() => new HttpClientTransport(options, httpClient, LoggerFactory));

Assert.Equal("transportOptions", exception.ParamName);
Assert.Contains(nameof(HttpClient.BaseAddress), exception.Message);
}

[Fact]
public async Task ConnectAsync_Uses_Injected_HttpClient_BaseAddress_When_Endpoint_Is_Omitted()
{
var options = new HttpClientTransportOptions
{
TransportMode = HttpTransportMode.Sse,
};
using var mockHttpHandler = new MockHttpHandler();
using var httpClient = new HttpClient(mockHttpHandler)
{
BaseAddress = new Uri("https+http://mcp-server/sse"),
};
await using var transport = new HttpClientTransport(options, httpClient, LoggerFactory);

mockHttpHandler.RequestHandler = request =>
{
Assert.Equal(httpClient.BaseAddress, request.RequestUri);
return Task.FromResult(new HttpResponseMessage
{
StatusCode = HttpStatusCode.OK,
Content = new StringContent("event: endpoint\r\ndata: /messages\r\n\r\n"),
});
};

await using var session = await transport.ConnectAsync(TestContext.Current.CancellationToken);
Assert.NotNull(session);
}

[Fact]
public async Task ConnectAsync_Prefers_Explicit_Endpoint_Over_HttpClient_BaseAddress()
{
var options = new HttpClientTransportOptions
{
Endpoint = new Uri("https://explicit.example/sse"),
TransportMode = HttpTransportMode.Sse,
};
using var mockHttpHandler = new MockHttpHandler();
using var httpClient = new HttpClient(mockHttpHandler)
{
BaseAddress = new Uri("https://base-address.example/sse"),
};
await using var transport = new HttpClientTransport(options, httpClient, LoggerFactory);

mockHttpHandler.RequestHandler = request =>
{
Assert.Equal(options.Endpoint, request.RequestUri);
return Task.FromResult(new HttpResponseMessage
{
StatusCode = HttpStatusCode.OK,
Content = new StringContent("event: endpoint\r\ndata: /messages\r\n\r\n"),
});
};

await using var session = await transport.ConnectAsync(TestContext.Current.CancellationToken);
Assert.NotNull(session);
}

[Fact]
public async Task ConnectAsync_Should_Connect_Successfully()
{
Expand Down