diff --git a/src/Common/HttpResponseMessageExtensions.cs b/src/Common/HttpResponseMessageExtensions.cs index 6acc6f7a3..d155e845a 100644 --- a/src/Common/HttpResponseMessageExtensions.cs +++ b/src/Common/HttpResponseMessageExtensions.cs @@ -69,10 +69,52 @@ public static HttpRequestException CreateHttpRequestException(HttpResponseMessag ? $"Response status code does not indicate success: {statusCodeInt} ({response.ReasonPhrase})." : $"Response status code does not indicate success: {statusCodeInt} ({response.ReasonPhrase}). Response body: {responseBody}"; + return HttpRequestExceptionExtensions.Create(message, innerException: null, response.StatusCode); + } +} + +/// +/// Helpers for preserving HTTP status codes on across all target frameworks. +/// +internal static class HttpRequestExceptionExtensions +{ + internal const string StatusCodeDataKey = "ModelContextProtocol.HttpStatusCode"; + + /// + /// Creates an and preserves its HTTP status code. + /// + public static HttpRequestException Create(string message, Exception? innerException, HttpStatusCode? statusCode) + { #if NET - return new HttpRequestException(message, inner: null, response.StatusCode); + var exception = new HttpRequestException(message, innerException, statusCode); #else - return new HttpRequestException(message); + var exception = new HttpRequestException(message, innerException); #endif + + if (statusCode is not null) + { + exception.Data[StatusCodeDataKey] = statusCode.Value; + } + + return exception; + } + + /// + /// Gets the preserved HTTP status code from an . + /// + public static HttpStatusCode? GetStatusCode(this HttpRequestException exception) + { +#if NET + if (exception.StatusCode is { } statusCode) + { + return statusCode; + } +#endif + + return exception.Data[StatusCodeDataKey] switch + { + HttpStatusCode storedStatusCode => storedStatusCode, + _ => null, + }; } } diff --git a/src/ModelContextProtocol.Core/Client/AutoDetectingClientSessionTransport.cs b/src/ModelContextProtocol.Core/Client/AutoDetectingClientSessionTransport.cs index b19e8a803..c9e821adf 100644 --- a/src/ModelContextProtocol.Core/Client/AutoDetectingClientSessionTransport.cs +++ b/src/ModelContextProtocol.Core/Client/AutoDetectingClientSessionTransport.cs @@ -146,14 +146,10 @@ private async Task InitializeSseTransportAsync(JsonRpcMessage message, HttpReque // keep HttpRequestException as the surfaced type so existing callers can still catch it and read StatusCode. await sseTransport.DisposeAsync().ConfigureAwait(false); LogSseFallbackFailedAfterStreamableHttp(_name, sseError); -#if NET - throw new HttpRequestException(streamableHttpError.Message, sseError, streamableHttpError.StatusCode); -#else - // net472 has no HttpRequestException overload that carries a status code, so this target - // preserves the status text in the message but not a programmatic StatusCode. Preserving the - // status code is intentionally net5+ only rather than an oversight. - throw new HttpRequestException(streamableHttpError.Message, sseError); -#endif + throw HttpRequestExceptionExtensions.Create( + streamableHttpError.Message, + sseError, + streamableHttpError.GetStatusCode()); } catch { diff --git a/tests/ModelContextProtocol.Tests/Transport/HttpClientTransportAutoDetectTests.cs b/tests/ModelContextProtocol.Tests/Transport/HttpClientTransportAutoDetectTests.cs index e6386adc6..c94e4f800 100644 --- a/tests/ModelContextProtocol.Tests/Transport/HttpClientTransportAutoDetectTests.cs +++ b/tests/ModelContextProtocol.Tests/Transport/HttpClientTransportAutoDetectTests.cs @@ -101,6 +101,10 @@ public async Task AutoDetectMode_WhenBothTransportsFail_PreservesStreamableHttpE Assert.Contains("403", ex.Message); Assert.IsType(ex.InnerException); Assert.Contains("405", ex.InnerException.Message); + Assert.Equal(HttpStatusCode.Forbidden, ex.Data["ModelContextProtocol.HttpStatusCode"]); +#if NET + Assert.Equal(HttpStatusCode.Forbidden, ex.StatusCode); +#endif } [Fact] @@ -300,6 +304,7 @@ public async Task AutoDetectMode_SurfacesStreamableHttpError_WithSseAsInner_When var httpEx = Assert.IsType(ex); Assert.Contains("415", httpEx.Message); Assert.Contains(streamableHttpBody, httpEx.Message); + Assert.Equal(HttpStatusCode.UnsupportedMediaType, httpEx.Data["ModelContextProtocol.HttpStatusCode"]); #if NET Assert.Equal(HttpStatusCode.UnsupportedMediaType, httpEx.StatusCode); #endif