diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/DefaultHttpRequestHandler.cs b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/DefaultHttpRequestHandler.cs index 606a716c207..72ad1d57409 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/DefaultHttpRequestHandler.cs +++ b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/DefaultHttpRequestHandler.cs @@ -2,7 +2,9 @@ using System; using System.Collections.Generic; +using System.Diagnostics.CodeAnalysis; using System.Linq; +using System.Net; using System.Net.Http; using System.Text; using System.Threading; @@ -24,9 +26,17 @@ namespace Microsoft.Agents.AI.Workflows.Declarative; /// The handler applies the per-request using a linked /// so it does not mutate on shared instances. /// +/// +/// Redirects are handled by this handler only when using its internally owned client, which disables automatic +/// redirects so per-request headers are not forwarded to redirect destinations. If a supplied client returns a +/// redirect response, this handler does not follow it because the client's redirect behavior is opaque. Supplied +/// clients should disable automatic redirects and handle redirect responses before returning them to this handler. +/// /// public sealed class DefaultHttpRequestHandler : IHttpRequestHandler, IAsyncDisposable { + private const int MaxAutomaticRedirections = 50; + private readonly Func>? _httpClientProvider; private readonly Lazy _ownedHttpClient; @@ -80,7 +90,7 @@ public DefaultHttpRequestHandler(HttpClient httpClient) public DefaultHttpRequestHandler(Func>? httpClientProvider) { this._httpClientProvider = httpClientProvider; - this._ownedHttpClient = new Lazy(() => new HttpClient(), LazyThreadSafetyMode.ExecutionAndPublication); + this._ownedHttpClient = new Lazy(CreateOwnedHttpClient, LazyThreadSafetyMode.ExecutionAndPublication); } private static Func> CreateSingleClientProvider(HttpClient httpClient) @@ -111,15 +121,8 @@ public async Task SendAsync(HttpRequestInfo request, Cancella throw new ArgumentException("Request method must be provided.", nameof(request)); } - HttpClient? providedClient = null; - if (this._httpClientProvider is not null) - { - providedClient = await this._httpClientProvider(request, cancellationToken).ConfigureAwait(false); - } - - HttpClient client = providedClient ?? this._ownedHttpClient.Value; - - using HttpRequestMessage httpRequest = BuildHttpRequestMessage(request); + HttpRequestInfo currentRequest = request; + Uri currentUri = CreateAbsoluteUri(ResolveRequestUri(request)); using CancellationTokenSource? timeoutCts = request.Timeout is { } timeout && timeout > TimeSpan.Zero ? CancellationTokenSource.CreateLinkedTokenSource(cancellationToken) @@ -129,32 +132,76 @@ public async Task SendAsync(HttpRequestInfo request, Cancella CancellationToken effectiveToken = timeoutCts?.Token ?? cancellationToken; - using HttpResponseMessage httpResponse = await client - .SendAsync(httpRequest, HttpCompletionOption.ResponseContentRead, effectiveToken) - .ConfigureAwait(false); + for (int redirectCount = 0; redirectCount <= MaxAutomaticRedirections; redirectCount++) + { + HttpClient? providedClient = null; + if (this._httpClientProvider is not null) + { + providedClient = await this._httpClientProvider(currentRequest, effectiveToken).ConfigureAwait(false); + } + + HttpClient client = providedClient ?? this._ownedHttpClient.Value; + + using HttpRequestMessage httpRequest = BuildHttpRequestMessage(currentRequest); + + using HttpResponseMessage httpResponse = await client + .SendAsync(httpRequest, HttpCompletionOption.ResponseHeadersRead, effectiveToken) + .ConfigureAwait(false); + + if (providedClient is null && + TryCreateRedirectRequest(httpResponse, currentRequest, currentUri, out HttpRequestInfo? redirectRequest, out Uri? redirectUri)) + { + currentRequest = redirectRequest; + currentUri = redirectUri; + continue; + } + + string? body = httpResponse.Content is null + ? null + : await ReadResponseBodyAsStringAsync(httpResponse.Content, effectiveToken).ConfigureAwait(false); - string? body = httpResponse.Content is null - ? null + Dictionary> headers = new(StringComparer.OrdinalIgnoreCase); + AppendHeaders(headers, httpResponse.Headers); + if (httpResponse.Content is not null) + { + AppendHeaders(headers, httpResponse.Content.Headers); + } + + return new HttpRequestResult + { + StatusCode = (int)httpResponse.StatusCode, + IsSuccessStatusCode = httpResponse.IsSuccessStatusCode, + Body = body, + Headers = headers, + }; + } + + throw new HttpRequestException($"The maximum number of HTTP redirects ({MaxAutomaticRedirections}) was exceeded."); + } + + private static async Task ReadResponseBodyAsStringAsync(HttpContent content, CancellationToken cancellationToken) + { #if NET - : await httpResponse.Content.ReadAsStringAsync(effectiveToken).ConfigureAwait(false); + return await content.ReadAsStringAsync(cancellationToken).ConfigureAwait(false); #else - : await httpResponse.Content.ReadAsStringAsync().ConfigureAwait(false); -#endif - - Dictionary> headers = new(StringComparer.OrdinalIgnoreCase); - AppendHeaders(headers, httpResponse.Headers); - if (httpResponse.Content is not null) + Task readTask = content.ReadAsStringAsync(); + Task cancellationTask = Task.Delay(Timeout.Infinite, cancellationToken); + Task completedTask = await Task.WhenAny(readTask, cancellationTask).ConfigureAwait(false); + if (completedTask == readTask) { - AppendHeaders(headers, httpResponse.Content.Headers); + return await readTask.ConfigureAwait(false); } - return new HttpRequestResult - { - StatusCode = (int)httpResponse.StatusCode, - IsSuccessStatusCode = httpResponse.IsSuccessStatusCode, - Body = body, - Headers = headers, - }; + content.Dispose(); + _ = readTask.ContinueWith( + static task => _ = task.Exception, + CancellationToken.None, + TaskContinuationOptions.OnlyOnFaulted | TaskContinuationOptions.ExecuteSynchronously, + TaskScheduler.Default); + + cancellationToken.ThrowIfCancellationRequested(); + throw new OperationCanceledException(cancellationToken); +#endif } /// @@ -213,6 +260,96 @@ private static HttpRequestMessage BuildHttpRequestMessage(HttpRequestInfo reques return httpRequest; } + private static HttpClient CreateOwnedHttpClient() + { + HttpClientHandler handler = new() + { + AllowAutoRedirect = false, + UseCookies = false, + CheckCertificateRevocationList = true + }; + + return new HttpClient(handler); + } + + private static Uri CreateAbsoluteUri(string requestUri) + { + if (!Uri.TryCreate(requestUri, UriKind.Absolute, out Uri? uri)) + { + throw new ArgumentException("Request URL must be an absolute URL.", nameof(requestUri)); + } + + return uri; + } + + private static bool TryCreateRedirectRequest( + HttpResponseMessage response, + HttpRequestInfo currentRequest, + Uri currentUri, + [NotNullWhen(true)] out HttpRequestInfo? redirectRequest, + [NotNullWhen(true)] out Uri? redirectUri) + { + redirectRequest = null; + redirectUri = null; + + if (!IsRedirectStatusCode(response.StatusCode) || response.Headers.Location is null) + { + return false; + } + + redirectUri = response.Headers.Location.IsAbsoluteUri + ? response.Headers.Location + : new Uri(currentUri, response.Headers.Location); + + if (IsHttpsToHttpRedirect(currentUri, redirectUri)) + { + throw new HttpRequestException("Redirects from HTTPS to HTTP are not allowed."); + } + + bool rewriteToGet = ShouldRewriteRedirectMethodToGet(response.StatusCode, currentRequest.Method); + if (!rewriteToGet && currentRequest.Body is not null && !HaveSameOrigin(currentUri, redirectUri)) + { + throw new HttpRequestException("Redirects that preserve the request body to a different origin are not allowed."); + } + + redirectRequest = new HttpRequestInfo + { + Method = rewriteToGet ? "GET" : currentRequest.Method, + Url = redirectUri.ToString(), + Body = rewriteToGet ? null : currentRequest.Body, + BodyContentType = rewriteToGet ? null : currentRequest.BodyContentType, + Timeout = currentRequest.Timeout, + ConnectionName = currentRequest.ConnectionName, + }; + + return true; + } + + private static bool IsRedirectStatusCode(HttpStatusCode statusCode) + { + int code = (int)statusCode; + return code is 301 or 302 or 303 or 307 or 308; + } + + private static bool ShouldRewriteRedirectMethodToGet(HttpStatusCode statusCode, string method) + { + string normalized = method.Trim().ToUpperInvariant(); + int code = (int)statusCode; + return code == 303 || ((code == 301 || code == 302) && string.Equals(normalized, "POST", StringComparison.Ordinal)); + } + + private static bool IsHttpsToHttpRedirect(Uri currentUri, Uri redirectUri) => + string.Equals(currentUri.Scheme, Uri.UriSchemeHttps, StringComparison.OrdinalIgnoreCase) && + string.Equals(redirectUri.Scheme, Uri.UriSchemeHttp, StringComparison.OrdinalIgnoreCase); + + private static bool HaveSameOrigin(Uri currentUri, Uri redirectUri) => + Uri.Compare( + currentUri, + redirectUri, + UriComponents.SchemeAndServer, + UriFormat.Unescaped, + StringComparison.OrdinalIgnoreCase) == 0; + private static HttpMethod ResolveMethod(string method) { string normalized = method.Trim().ToUpperInvariant(); diff --git a/dotnet/tests/Microsoft.Agents.AI.Workflows.Declarative.UnitTests/DefaultHttpRequestHandlerTests.cs b/dotnet/tests/Microsoft.Agents.AI.Workflows.Declarative.UnitTests/DefaultHttpRequestHandlerTests.cs index 6970bc336c2..33dbfd17cef 100644 --- a/dotnet/tests/Microsoft.Agents.AI.Workflows.Declarative.UnitTests/DefaultHttpRequestHandlerTests.cs +++ b/dotnet/tests/Microsoft.Agents.AI.Workflows.Declarative.UnitTests/DefaultHttpRequestHandlerTests.cs @@ -2,8 +2,10 @@ using System; using System.Collections.Generic; +using System.IO; using System.Net; using System.Net.Http; +using System.Reflection; using System.Text; using System.Threading; using System.Threading.Tasks; @@ -364,6 +366,254 @@ public async Task SendAsyncTimeoutCancelsRequestAsync() await Assert.ThrowsAnyAsync(actAsync); } + [Fact] + public async Task SendAsyncTimeoutCancelsResponseBodyReadAsync() + { + // Arrange + CancellationToken cancellationToken = TestContext.Current.CancellationToken; + int requestCount = 0; + using var response = new HttpResponseMessage(HttpStatusCode.OK) + { + Content = new StallingContent(), + }; + TestHttpMessageHandler messageHandler = new((_, _) => + { + requestCount++; +#pragma warning disable CA2025 // Do not pass 'IDisposable' instances into unawaited tasks + return Task.FromResult(response); +#pragma warning restore CA2025 // Do not pass 'IDisposable' instances into unawaited tasks + }); + + using HttpClient httpClient = new(messageHandler); +#pragma warning disable CA2025 // Do not pass 'IDisposable' instances into unawaited tasks + await using DefaultHttpRequestHandler handler = new((_, _) => Task.FromResult(httpClient)); +#pragma warning restore CA2025 // Do not pass 'IDisposable' instances into unawaited tasks + HttpRequestInfo request = new() + { + Method = "GET", + Url = TestUrl, + Timeout = TimeSpan.FromMilliseconds(50), + }; + + // Act + async Task actAsync() => await handler.SendAsync(request, cancellationToken); + + // Assert + await Assert.ThrowsAnyAsync(actAsync); + Assert.Equal(1, requestCount); + } + + [Fact] + public async Task SendAsyncTimeoutAppliesAcrossRedirectsAsync() + { + // Arrange + CancellationToken cancellationToken = TestContext.Current.CancellationToken; + int requestCount = 0; + using HttpResponseMessage redirectResponse = new(HttpStatusCode.TemporaryRedirect) + { + Headers = { Location = new Uri("https://api.example.test/next") }, + }; + using HttpResponseMessage okResponse = new(HttpStatusCode.OK) + { + Content = new StringContent("ok", Encoding.UTF8, "text/plain"), + }; + TestHttpMessageHandler messageHandler = new(async (req, ct) => + { + requestCount++; + TimeSpan delay = requestCount == 1 ? TimeSpan.FromMilliseconds(1) : TimeSpan.FromSeconds(5); + await Task.Delay(delay, ct).ConfigureAwait(false); + + if (requestCount == 1) + { + return redirectResponse; + } + + return okResponse; + }); + + await using DefaultHttpRequestHandler handler = CreateHandlerWithOwnedMessageHandler(messageHandler); + HttpRequestInfo request = new() + { + Method = "GET", + Url = TestUrl, + Timeout = TimeSpan.FromMilliseconds(300), + }; + + // Act + async Task actAsync() => await handler.SendAsync(request, cancellationToken); + + // Assert + await Assert.ThrowsAnyAsync(actAsync); + Assert.Equal(2, requestCount); + } + + [Fact] + public async Task SendAsyncPostFoundRedirectRewritesToGetAsync() + { + // Arrange + CancellationToken cancellationToken = TestContext.Current.CancellationToken; + using HttpResponseMessage redirectResponse = new(HttpStatusCode.Found) + { + Headers = { Location = new Uri("https://api.example.test/next") }, + }; + using HttpResponseMessage okResponse = new(HttpStatusCode.OK) + { + Content = new StringContent("redirected", Encoding.UTF8, "text/plain"), + }; + int requestCount = 0; +#pragma warning disable CA2025 + TestHttpMessageHandler messageHandler = new((_, _) => + { + requestCount++; + return Task.FromResult(requestCount == 1 ? redirectResponse : okResponse); + }); +#pragma warning restore CA2025 + + await using DefaultHttpRequestHandler handler = CreateHandlerWithOwnedMessageHandler(messageHandler); + HttpRequestInfo request = new() + { + Method = "POST", + Url = TestUrl, + Body = "request-body", + BodyContentType = "text/plain", + }; + + // Act + HttpRequestResult result = await handler.SendAsync(request, cancellationToken); + + // Assert + Assert.Equal("redirected", result.Body); + Assert.Equal(["POST", "GET"], messageHandler.RequestMethods); + Assert.Equal(["request-body", null], messageHandler.RequestBodies); + } + + [Theory] + [InlineData(307)] + [InlineData(308)] + public async Task SendAsyncPostPreserveMethodRedirectPreservesBodyAsync(int redirectStatusCode) + { + // Arrange + CancellationToken cancellationToken = TestContext.Current.CancellationToken; + using HttpResponseMessage redirectResponse = new((HttpStatusCode)redirectStatusCode) + { + Headers = { Location = new Uri("https://api.example.test/next") }, + }; + using HttpResponseMessage okResponse = new(HttpStatusCode.OK) + { + Content = new StringContent("redirected", Encoding.UTF8, "text/plain"), + }; + int requestCount = 0; +#pragma warning disable CA2025 + TestHttpMessageHandler messageHandler = new((req, _) => + { + requestCount++; + return Task.FromResult(requestCount == 1 ? redirectResponse : okResponse); + }); +#pragma warning restore CA2025 + + await using DefaultHttpRequestHandler handler = CreateHandlerWithOwnedMessageHandler(messageHandler); + HttpRequestInfo request = new() + { + Method = "POST", + Url = TestUrl, + Body = "request-body", + BodyContentType = "text/plain", + }; + + // Act + HttpRequestResult result = await handler.SendAsync(request, cancellationToken); + + // Assert + Assert.Equal("redirected", result.Body); + Assert.Equal(["POST", "POST"], messageHandler.RequestMethods); + Assert.Equal(["request-body", "request-body"], messageHandler.RequestBodies); + } + + [Theory] + [InlineData(307)] + [InlineData(308)] + public async Task SendAsyncPostPreserveMethodRedirectToDifferentOriginWithBodyThrowsAsync(int redirectStatusCode) + { + // Arrange + CancellationToken cancellationToken = TestContext.Current.CancellationToken; + int requestCount = 0; + using HttpResponseMessage redirectResponse = new((HttpStatusCode)redirectStatusCode) + { + Headers = { Location = new Uri("https://secondary.example.test/next") }, + }; +#pragma warning disable CA2025 + TestHttpMessageHandler messageHandler = new((_, _) => + { + requestCount++; + return Task.FromResult(redirectResponse); + }); +#pragma warning restore CA2025 + + await using DefaultHttpRequestHandler handler = CreateHandlerWithOwnedMessageHandler(messageHandler); + HttpRequestInfo request = new() + { + Method = "POST", + Url = TestUrl, + Body = "request-body", + BodyContentType = "text/plain", + }; + + // Act + async Task actAsync() => await handler.SendAsync(request, cancellationToken); + + // Assert + HttpRequestException exception = await Assert.ThrowsAsync(actAsync); + Assert.Contains("preserve the request body to a different origin", exception.Message, StringComparison.Ordinal); + Assert.Equal(1, requestCount); + } + + [Fact] + public async Task SendAsyncTooManyRedirectsThrowsAsync() + { + // Arrange + CancellationToken cancellationToken = TestContext.Current.CancellationToken; + int requestCount = 0; + List createdResponses = []; + TestHttpMessageHandler messageHandler = new((_, _) => + { + requestCount++; + HttpResponseMessage response = new(HttpStatusCode.TemporaryRedirect) + { + Headers = { Location = new Uri($"https://api.example.test/redirect/{requestCount}") }, + }; + + createdResponses.Add(response); +#pragma warning disable CA2025 + return Task.FromResult(response); +#pragma warning restore CA2025 + }); + + try + { + await using DefaultHttpRequestHandler handler = CreateHandlerWithOwnedMessageHandler(messageHandler); + HttpRequestInfo request = new() + { + Method = "GET", + Url = TestUrl, + }; + + // Act + async Task actAsync() => await handler.SendAsync(request, cancellationToken); + + // Assert + HttpRequestException exception = await Assert.ThrowsAsync(actAsync); + Assert.Contains("maximum number of HTTP redirects", exception.Message, StringComparison.Ordinal); + Assert.Equal(51, requestCount); + } + finally + { + foreach (HttpResponseMessage response in createdResponses) + { + response.Dispose(); + } + } + } + [Fact] public async Task SendAsyncFallsBackToOwnedClientWhenProviderReturnsNullAsync() { @@ -385,6 +635,302 @@ public async Task SendAsyncFallsBackToOwnedClientWhenProviderReturnsNullAsync() Assert.Equal(1, providerCallCount); } + [Fact] + public async Task SendAsyncDoesNotApplyRequestHeadersToRedirectedEndpointAsync() + { + // Arrange + CancellationToken cancellationToken = TestContext.Current.CancellationToken; + List requestsWithHeader = []; + using HttpResponseMessage redirectResponse = new(HttpStatusCode.TemporaryRedirect) + { + Headers = { Location = new Uri("https://api.example.test/next") }, + }; + using HttpResponseMessage okResponse = new(HttpStatusCode.OK) + { + Content = new StringContent("redirected", Encoding.UTF8, "text/plain"), + }; +#pragma warning disable CA2025 + TestHttpMessageHandler messageHandler = new((req, _) => + { + requestsWithHeader.Add(req.Headers.Contains("X-Trace-Id")); + if (requestsWithHeader.Count == 1) + { + return Task.FromResult(redirectResponse); + } + + return Task.FromResult(okResponse); + }); +#pragma warning restore CA2025 + + await using DefaultHttpRequestHandler handler = CreateHandlerWithOwnedMessageHandler(messageHandler); + HttpRequestInfo request = new() + { + Method = "GET", + Url = TestUrl, + Headers = new Dictionary + { + ["X-Trace-Id"] = "trace-1", + }, + }; + + // Act + HttpRequestResult result = await handler.SendAsync(request, cancellationToken); + + // Assert + Assert.Equal("redirected", result.Body); + Assert.Equal([true, false], requestsWithHeader); + } + + [Fact] + public async Task SendAsyncDoesNotApplyRequestHeadersToDifferentRedirectedEndpointAsync() + { + // Arrange + CancellationToken cancellationToken = TestContext.Current.CancellationToken; + List requestsWithHeader = []; + List requestUrls = []; + using HttpResponseMessage redirectResponse = new(HttpStatusCode.TemporaryRedirect) + { + Headers = { Location = new Uri("https://secondary.example.test/next") }, + }; + using HttpResponseMessage okResponse = new(HttpStatusCode.OK) + { + Content = new StringContent("redirected", Encoding.UTF8, "text/plain"), + }; +#pragma warning disable CA2025 + TestHttpMessageHandler messageHandler = new((req, _) => + { + requestUrls.Add(req.RequestUri!.ToString()); + requestsWithHeader.Add(req.Headers.Contains("X-Trace-Id")); + return Task.FromResult(requestUrls.Count == 1 ? redirectResponse : okResponse); + }); +#pragma warning restore CA2025 + + await using DefaultHttpRequestHandler handler = CreateHandlerWithOwnedMessageHandler(messageHandler); + + HttpRequestInfo request = new() + { + Method = "GET", + Url = TestUrl, + Headers = new Dictionary + { + ["X-Trace-Id"] = "trace-1", + }, + }; + + // Act + HttpRequestResult result = await handler.SendAsync(request, cancellationToken); + + // Assert + Assert.Equal("redirected", result.Body); + Assert.Equal([true, false], requestsWithHeader); + Assert.Equal([TestUrl, "https://secondary.example.test/next"], requestUrls); + } + + [Fact] + public async Task SendAsyncDoesNotReadRedirectedResponseBodyAsync() + { + // Arrange + CancellationToken cancellationToken = TestContext.Current.CancellationToken; + TrackingContent redirectedContent = new("not returned"); + using HttpResponseMessage redirectResponse = new(HttpStatusCode.TemporaryRedirect) + { + Content = redirectedContent, + Headers = { Location = new Uri("https://api.example.test/next") }, + }; + using HttpResponseMessage okResponse = new(HttpStatusCode.OK) + { + Content = new StringContent("redirected", Encoding.UTF8, "text/plain"), + }; +#pragma warning disable CA2025 + TestHttpMessageHandler messageHandler = new((req, _) => + { + if (req.RequestUri!.AbsolutePath == "/resource") + { + return Task.FromResult(redirectResponse); + } + + return Task.FromResult(okResponse); + }); +#pragma warning restore CA2025 + + await using DefaultHttpRequestHandler handler = CreateHandlerWithOwnedMessageHandler(messageHandler); + HttpRequestInfo request = new() + { + Method = "GET", + Url = TestUrl, + }; + + // Act + HttpRequestResult result = await handler.SendAsync(request, cancellationToken); + + // Assert + Assert.Equal("redirected", result.Body); + Assert.False(redirectedContent.WasRead); + } + + [Fact] + public async Task SendAsyncRejectsHttpsToHttpRedirectAsync() + { + // Arrange + CancellationToken cancellationToken = TestContext.Current.CancellationToken; + int requestCount = 0; + using HttpResponseMessage redirectResponse = new(HttpStatusCode.TemporaryRedirect) + { + Headers = { Location = new Uri("http://api.example.test/next") }, + }; +#pragma warning disable CA2025 + TestHttpMessageHandler messageHandler = new((req, _) => + { + requestCount++; + return Task.FromResult(redirectResponse); + }); +#pragma warning restore CA2025 + + await using DefaultHttpRequestHandler handler = CreateHandlerWithOwnedMessageHandler(messageHandler); + HttpRequestInfo request = new() + { + Method = "POST", + Url = TestUrl, + Body = "request-body", + BodyContentType = "text/plain", + }; + + // Act + async Task actAsync() => await handler.SendAsync(request, cancellationToken); + + // Assert + HttpRequestException exception = await Assert.ThrowsAsync(actAsync); + Assert.Contains("HTTPS to HTTP", exception.Message, StringComparison.Ordinal); + Assert.Equal(1, requestCount); + } + + [Fact] + public async Task SendAsyncProviderClientAllowsScopedDefaultHeadersOnInitialRequestAsync() + { + // Arrange + CancellationToken cancellationToken = TestContext.Current.CancellationToken; + using HttpResponseMessage okResponse = new(HttpStatusCode.OK) + { + Content = new StringContent("ok", Encoding.UTF8, "text/plain"), + }; +#pragma warning disable CA2025 + TestHttpMessageHandler messageHandler = new((req, _) => + Task.FromResult(okResponse)); +#pragma warning restore CA2025 + using HttpClient providerClient = new(messageHandler); + providerClient.DefaultRequestHeaders.TryAddWithoutValidation("X-Client-Token", "provider-header-value"); + + int providerCallCount = 0; +#pragma warning disable CA2025 + await using DefaultHttpRequestHandler handler = new((_, _) => + { + providerCallCount++; + return Task.FromResult(providerClient); + }); +#pragma warning restore CA2025 + + HttpRequestInfo request = new() + { + Method = "GET", + Url = TestUrl, + }; + + // Act + HttpRequestResult result = await handler.SendAsync(request, cancellationToken); + + // Assert + Assert.Equal("ok", result.Body); + Assert.Equal(1, providerCallCount); + } + + [Fact] + public async Task SendAsyncSuppliedClientReturnsRedirectResponseAsync() + { + // Arrange + CancellationToken cancellationToken = TestContext.Current.CancellationToken; + using HttpResponseMessage redirectResponse = new(HttpStatusCode.TemporaryRedirect) + { + Headers = { Location = new Uri("https://secondary.example.test/next") }, + }; +#pragma warning disable CA2025 + TestHttpMessageHandler primaryMessageHandler = new((req, _) => + Task.FromResult(redirectResponse)); +#pragma warning restore CA2025 + using HttpClient primaryClient = new(primaryMessageHandler); + primaryClient.DefaultRequestHeaders.TryAddWithoutValidation("Ocp-Apim-Subscription-Key", "provider-header-value"); + + int providerCallCount = 0; +#pragma warning disable CA2025 + await using DefaultHttpRequestHandler handler = new((_, _) => + { + providerCallCount++; + return Task.FromResult(primaryClient); + }); +#pragma warning restore CA2025 + + HttpRequestInfo request = new() + { + Method = "GET", + Url = TestUrl, + }; + + // Act + HttpRequestResult result = await handler.SendAsync(request, cancellationToken); + + // Assert + Assert.Equal(307, result.StatusCode); + Assert.False(result.IsSuccessStatusCode); + Assert.NotNull(result.Headers); + Assert.Equal("https://secondary.example.test/next", Assert.Single(result.Headers!["Location"])); + Assert.Equal(1, providerCallCount); + } + + [Fact] + public async Task SendAsyncInvokesProviderForRedirectDestinationAsync() + { + // Arrange + CancellationToken cancellationToken = TestContext.Current.CancellationToken; + List providerUrls = []; + List requestUrls = []; + using HttpResponseMessage redirectResponse = new(HttpStatusCode.TemporaryRedirect) + { + Headers = { Location = new Uri("https://secondary.example.test/next") }, + }; + using HttpResponseMessage okResponse = new(HttpStatusCode.OK) + { + Content = new StringContent("redirected", Encoding.UTF8, "text/plain"), + }; +#pragma warning disable CA2025 + TestHttpMessageHandler messageHandler = new((req, _) => + { + requestUrls.Add(req.RequestUri!.ToString()); + return Task.FromResult(requestUrls.Count == 1 ? redirectResponse : okResponse); + }); +#pragma warning restore CA2025 + + await using DefaultHttpRequestHandler handler = CreateHandlerWithOwnedMessageHandler(messageHandler, (info, _) => + { + providerUrls.Add(info.Url); + return Task.FromResult(null); + }); + + HttpRequestInfo request = new() + { + Method = "GET", + Url = TestUrl, + }; + + // Act + HttpRequestResult result = await handler.SendAsync(request, cancellationToken); + + // Assert + Assert.Equal("redirected", result.Body); + Assert.Equal(2, providerUrls.Count); + Assert.Equal(TestUrl, providerUrls[0]); + Assert.Equal("https://secondary.example.test/next", providerUrls[1]); + Assert.Equal(providerUrls, requestUrls); + } + #endregion #region DisposeAsync @@ -476,6 +1022,17 @@ public async Task QueryParametersPreserveExistingQueryStringAsync() #endregion + private static DefaultHttpRequestHandler CreateHandlerWithOwnedMessageHandler( + HttpMessageHandler ownedHttpMessageHandler, + Func>? httpClientProvider = null) + { + DefaultHttpRequestHandler handler = new(httpClientProvider); + FieldInfo? ownedHttpClientField = typeof(DefaultHttpRequestHandler).GetField("_ownedHttpClient", BindingFlags.Instance | BindingFlags.NonPublic); + Assert.NotNull(ownedHttpClientField); + ownedHttpClientField.SetValue(handler, new Lazy(() => new HttpClient(ownedHttpMessageHandler), LazyThreadSafetyMode.ExecutionAndPublication)); + return handler; + } + private sealed class TestHttpMessageHandler : HttpMessageHandler { private readonly Func> _responseFactory; @@ -491,9 +1048,14 @@ public TestHttpMessageHandler(Func RequestMethods { get; } = []; + + public List RequestBodies { get; } = []; + protected override async Task SendAsync(HttpRequestMessage request, CancellationToken cancellationToken) { this.LastRequest = request; + this.RequestMethods.Add(request.Method.Method); if (request.Content is not null) { #if NET @@ -501,9 +1063,68 @@ protected override async Task SendAsync(HttpRequestMessage #else this.LastRequestBody = await request.Content.ReadAsStringAsync().ConfigureAwait(false); #endif + this.RequestBodies.Add(this.LastRequestBody); this.LastRequestContentType = request.Content.Headers.ContentType?.MediaType; } + else + { + this.RequestBodies.Add(null); + } return await this._responseFactory(request, cancellationToken).ConfigureAwait(false); } } + + private sealed class TrackingContent : HttpContent + { + private readonly string _content; + + public TrackingContent(string content) + { + this._content = content; + } + + public bool WasRead { get; private set; } + + protected override Task SerializeToStreamAsync(Stream stream, TransportContext? context) + { + this.WasRead = true; + byte[] bytes = Encoding.UTF8.GetBytes(this._content); + return stream.WriteAsync(bytes, 0, bytes.Length); + } + + protected override bool TryComputeLength(out long length) + { + length = Encoding.UTF8.GetByteCount(this._content); + return true; + } + } + + private sealed class StallingContent : HttpContent + { + private readonly TaskCompletionSource _stall = new(TaskCreationOptions.RunContinuationsAsynchronously); + + protected override Task SerializeToStreamAsync(Stream stream, TransportContext? context) => + this._stall.Task; + +#if NET + protected override Task SerializeToStreamAsync(Stream stream, TransportContext? context, CancellationToken cancellationToken) => + Task.Delay(Timeout.Infinite, cancellationToken); +#endif + + protected override bool TryComputeLength(out long length) + { + length = 0; + return false; + } + + protected override void Dispose(bool disposing) + { + if (disposing) + { + this._stall.TrySetCanceled(); + } + + base.Dispose(disposing); + } + } }