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