diff --git a/src/ReverseProxy/Transforms/Builder/StructuredTransformer.cs b/src/ReverseProxy/Transforms/Builder/StructuredTransformer.cs
index 12e2270c2b..a10022df64 100644
--- a/src/ReverseProxy/Transforms/Builder/StructuredTransformer.cs
+++ b/src/ReverseProxy/Transforms/Builder/StructuredTransformer.cs
@@ -18,6 +18,8 @@ namespace Yarp.ReverseProxy.Transforms.Builder;
///
internal sealed class StructuredTransformer : HttpTransformer
{
+ private readonly bool _canUseRequestTransformFastPath;
+
///
/// Creates a new instance.
///
@@ -36,6 +38,7 @@ internal StructuredTransformer(bool? copyRequestHeaders, bool? copyResponseHeade
RequestTransforms = requestTransforms.ToArray();
ResponseTransforms = responseTransforms.ToArray();
ResponseTrailerTransforms = responseTrailerTransforms.ToArray();
+ _canUseRequestTransformFastPath = RequestTransforms.All(CanUseRequestTransformFastPath);
}
///
@@ -68,6 +71,22 @@ internal StructuredTransformer(bool? copyRequestHeaders, bool? copyResponseHeade
///
internal ResponseTrailersTransform[] ResponseTrailerTransforms { get; }
+ private static bool CanUseRequestTransformFastPath(RequestTransform transform)
+ {
+ // Built-in transforms are inheritable, so exact type checks preserve derived ApplyAsync overrides.
+ var transformType = transform.GetType();
+ return transformType == typeof(RequestHeaderOriginalHostTransform)
+ || transformType == typeof(RequestHeaderXForwardedForTransform)
+ || transformType == typeof(RequestHeaderXForwardedHostTransform)
+ || transformType == typeof(RequestHeaderXForwardedProtoTransform)
+ || transformType == typeof(RequestHeaderXForwardedPrefixTransform)
+ || transformType == typeof(RequestHeaderForwardedTransform)
+ || transformType == typeof(RequestHeaderRemoveTransform)
+ || transformType == typeof(RequestHeaderValueTransform)
+ || transformType == typeof(RequestHeaderRouteValueTransform)
+ || transformType == typeof(RequestHeadersAllowedTransform);
+ }
+
#pragma warning disable CS0672 // We're overriding the obsolete overloads to preserve backwards compatibility.
public override ValueTask TransformRequestAsync(HttpContext httpContext, HttpRequestMessage proxyRequest, string destinationPrefix) =>
TransformRequestAsync(httpContext, proxyRequest, destinationPrefix, CancellationToken.None);
@@ -98,6 +117,19 @@ public override async ValueTask TransformRequestAsync(HttpContext httpContext, H
return;
}
+ if (_canUseRequestTransformFastPath)
+ {
+ var headersCopied = ShouldCopyRequestHeaders.GetValueOrDefault(true);
+ foreach (var requestTransform in RequestTransforms)
+ {
+ requestTransform.ApplyFast(httpContext, proxyRequest, ref headersCopied);
+ }
+
+ proxyRequest.RequestUri ??= RequestUtilities.MakeDestinationAddress(
+ destinationPrefix, httpContext.Request.Path, httpContext.Request.QueryString);
+ return;
+ }
+
var transformContext = new RequestTransformContext()
{
DestinationPrefix = destinationPrefix,
diff --git a/src/ReverseProxy/Transforms/RequestHeaderForwardedTransform.cs b/src/ReverseProxy/Transforms/RequestHeaderForwardedTransform.cs
index 6a87e6f5b2..ae60ff41a5 100644
--- a/src/ReverseProxy/Transforms/RequestHeaderForwardedTransform.cs
+++ b/src/ReverseProxy/Transforms/RequestHeaderForwardedTransform.cs
@@ -4,6 +4,7 @@
using System;
using System.Diagnostics;
using System.Net;
+using System.Net.Http;
using System.Threading.Tasks;
using Microsoft.AspNetCore.Http;
using Microsoft.Extensions.Primitives;
@@ -51,28 +52,35 @@ public RequestHeaderForwardedTransform(IRandomFactory randomFactory, NodeFormat
public override ValueTask ApplyAsync(RequestTransformContext context)
{
ArgumentNullException.ThrowIfNull(context);
+ Apply(context.HttpContext, context.ProxyRequest, context.HeadersCopied);
+ return default;
+ }
- var httpContext = context.HttpContext;
+ internal override void ApplyFast(HttpContext httpContext, HttpRequestMessage proxyRequest, ref bool headersCopied)
+ {
+ Apply(httpContext, proxyRequest, headersCopied);
+ }
+ private void Apply(HttpContext httpContext, HttpRequestMessage proxyRequest, bool headersCopied)
+ {
switch (TransformAction)
{
case ForwardedTransformActions.Set:
- RemoveHeader(context, ForwardedHeaderName);
- AddHeader(context, ForwardedHeaderName, GetHeaderValue(httpContext));
+ RemoveHeader(proxyRequest, ForwardedHeaderName);
+ AddHeader(proxyRequest, ForwardedHeaderName, GetHeaderValue(httpContext));
break;
case ForwardedTransformActions.Append:
- var existingValues = TakeHeader(context, ForwardedHeaderName);
+ var existingValues = TakeHeader(httpContext, proxyRequest, headersCopied, ForwardedHeaderName);
var values = StringValues.Concat(existingValues, GetHeaderValue(httpContext));
- AddHeader(context, ForwardedHeaderName, values);
+ AddHeader(proxyRequest, ForwardedHeaderName, values);
break;
case ForwardedTransformActions.Remove:
- RemoveHeader(context, ForwardedHeaderName);
+ RemoveHeader(proxyRequest, ForwardedHeaderName);
break;
default:
throw new NotImplementedException(TransformAction.ToString());
}
- return default;
}
private string GetHeaderValue(HttpContext httpContext)
diff --git a/src/ReverseProxy/Transforms/RequestHeaderOriginalHostTransform.cs b/src/ReverseProxy/Transforms/RequestHeaderOriginalHostTransform.cs
index a940d41ac4..f433e4d611 100644
--- a/src/ReverseProxy/Transforms/RequestHeaderOriginalHostTransform.cs
+++ b/src/ReverseProxy/Transforms/RequestHeaderOriginalHostTransform.cs
@@ -2,7 +2,9 @@
// The .NET Foundation licenses this file to you under the MIT license.
using System;
+using System.Net.Http;
using System.Threading.Tasks;
+using Microsoft.AspNetCore.Http;
using Microsoft.Net.Http.Headers;
using Yarp.ReverseProxy.Forwarder;
using Yarp.ReverseProxy.Model;
@@ -32,17 +34,28 @@ private RequestHeaderOriginalHostTransform(bool useOriginalHost)
public override ValueTask ApplyAsync(RequestTransformContext context)
{
- var destinationConfigHost = context.HttpContext.Features.Get()?.ProxiedDestination?.Model.Config?.Host;
- var originalHost = context.HttpContext.Request.Host.Value is { Length: > 0 } host ? host : null;
- var existingHost = RequestUtilities.TryGetValues(context.ProxyRequest.Headers, HeaderNames.Host, out var currentHost) ? currentHost.ToString() : null;
+ Apply(context.HttpContext, context.ProxyRequest, context.HeadersCopied);
+ return default;
+ }
+
+ internal override void ApplyFast(HttpContext httpContext, HttpRequestMessage proxyRequest, ref bool headersCopied)
+ {
+ Apply(httpContext, proxyRequest, headersCopied);
+ }
+
+ private void Apply(HttpContext httpContext, HttpRequestMessage proxyRequest, bool headersCopied)
+ {
+ var destinationConfigHost = httpContext.Features.Get()?.ProxiedDestination?.Model.Config?.Host;
+ var originalHost = httpContext.Request.Host.Value is { Length: > 0 } host ? host : null;
+ var existingHost = RequestUtilities.TryGetValues(proxyRequest.Headers, HeaderNames.Host, out var currentHost) ? currentHost.ToString() : null;
if (UseOriginalHost)
{
- if (!context.HeadersCopied && existingHost is null)
+ if (!headersCopied && existingHost is null)
{
// Propagate the host if the transform pipeline didn't already override it.
// If there was no original host specified, allow the destination config host to flow through.
- context.ProxyRequest.Headers.TryAddWithoutValidation(HeaderNames.Host, originalHost ?? destinationConfigHost);
+ proxyRequest.Headers.TryAddWithoutValidation(HeaderNames.Host, originalHost ?? destinationConfigHost);
}
}
else if (existingHost is null || string.Equals(originalHost, existingHost, StringComparison.Ordinal))
@@ -50,9 +63,7 @@ public override ValueTask ApplyAsync(RequestTransformContext context)
// Use the host from destination configuration (which may be null) if either:
// * there is no host header set, or
// * the original host header is being suppressed and has not been modified by the transform pipeline
- context.ProxyRequest.Headers.Host = destinationConfigHost;
+ proxyRequest.Headers.Host = destinationConfigHost;
}
-
- return default;
}
}
diff --git a/src/ReverseProxy/Transforms/RequestHeaderRemoveTransform.cs b/src/ReverseProxy/Transforms/RequestHeaderRemoveTransform.cs
index 444e2dfca6..c3fc1d8270 100644
--- a/src/ReverseProxy/Transforms/RequestHeaderRemoveTransform.cs
+++ b/src/ReverseProxy/Transforms/RequestHeaderRemoveTransform.cs
@@ -2,7 +2,9 @@
// The .NET Foundation licenses this file to you under the MIT license.
using System;
+using System.Net.Http;
using System.Threading.Tasks;
+using Microsoft.AspNetCore.Http;
namespace Yarp.ReverseProxy.Transforms;
@@ -27,9 +29,12 @@ public RequestHeaderRemoveTransform(string headerName)
public override ValueTask ApplyAsync(RequestTransformContext context)
{
ArgumentNullException.ThrowIfNull(context);
-
- RemoveHeader(context, HeaderName);
-
+ RemoveHeader(context.ProxyRequest, HeaderName);
return default;
}
+
+ internal override void ApplyFast(HttpContext httpContext, HttpRequestMessage proxyRequest, ref bool headersCopied)
+ {
+ RemoveHeader(proxyRequest, HeaderName);
+ }
}
diff --git a/src/ReverseProxy/Transforms/RequestHeaderRouteValueTransform.cs b/src/ReverseProxy/Transforms/RequestHeaderRouteValueTransform.cs
index a5bdf3f98e..b2b0c29328 100644
--- a/src/ReverseProxy/Transforms/RequestHeaderRouteValueTransform.cs
+++ b/src/ReverseProxy/Transforms/RequestHeaderRouteValueTransform.cs
@@ -2,6 +2,10 @@
// The .NET Foundation licenses this file to you under the MIT license.
using System;
+using System.Net.Http;
+using System.Threading.Tasks;
+using Microsoft.AspNetCore.Http;
+using Microsoft.Extensions.Primitives;
namespace Yarp.ReverseProxy.Transforms;
@@ -35,5 +39,24 @@ public RequestHeaderRouteValueTransform(string headerName, string routeValueKey,
return value?.ToString();
}
-}
+ internal override void ApplyFast(HttpContext httpContext, HttpRequestMessage proxyRequest, ref bool headersCopied)
+ {
+ if (!httpContext.Request.RouteValues.TryGetValue(RouteValueKey, out var routeValue)
+ || routeValue?.ToString() is not { } value)
+ {
+ return;
+ }
+
+ if (Append)
+ {
+ var existingValues = TakeHeader(httpContext, proxyRequest, headersCopied, HeaderName);
+ AddHeader(proxyRequest, HeaderName, StringValues.Concat(existingValues, value));
+ }
+ else
+ {
+ RemoveHeader(proxyRequest, HeaderName);
+ AddHeader(proxyRequest, HeaderName, value);
+ }
+ }
+}
diff --git a/src/ReverseProxy/Transforms/RequestHeaderValueTransform.cs b/src/ReverseProxy/Transforms/RequestHeaderValueTransform.cs
index 42a40aefed..730655857a 100644
--- a/src/ReverseProxy/Transforms/RequestHeaderValueTransform.cs
+++ b/src/ReverseProxy/Transforms/RequestHeaderValueTransform.cs
@@ -2,7 +2,10 @@
// The .NET Foundation licenses this file to you under the MIT license.
using System;
+using System.Net.Http;
using System.Threading.Tasks;
+using Microsoft.AspNetCore.Http;
+using Microsoft.Extensions.Primitives;
namespace Yarp.ReverseProxy.Transforms;
@@ -31,4 +34,18 @@ protected override string GetValue(RequestTransformContext context)
{
return Value;
}
+
+ internal override void ApplyFast(HttpContext httpContext, HttpRequestMessage proxyRequest, ref bool headersCopied)
+ {
+ if (Append)
+ {
+ var existingValues = TakeHeader(httpContext, proxyRequest, headersCopied, HeaderName);
+ AddHeader(proxyRequest, HeaderName, StringValues.Concat(existingValues, Value));
+ }
+ else
+ {
+ RemoveHeader(proxyRequest, HeaderName);
+ AddHeader(proxyRequest, HeaderName, Value);
+ }
+ }
}
diff --git a/src/ReverseProxy/Transforms/RequestHeaderXForwardedForTransform.cs b/src/ReverseProxy/Transforms/RequestHeaderXForwardedForTransform.cs
index 50b702043b..26cb418a22 100644
--- a/src/ReverseProxy/Transforms/RequestHeaderXForwardedForTransform.cs
+++ b/src/ReverseProxy/Transforms/RequestHeaderXForwardedForTransform.cs
@@ -3,7 +3,9 @@
using System;
using System.Diagnostics;
+using System.Net.Http;
using System.Threading.Tasks;
+using Microsoft.AspNetCore.Http;
using Microsoft.Extensions.Primitives;
namespace Yarp.ReverseProxy.Transforms;
@@ -38,9 +40,19 @@ public RequestHeaderXForwardedForTransform(string headerName, ForwardedTransform
public override ValueTask ApplyAsync(RequestTransformContext context)
{
ArgumentNullException.ThrowIfNull(context);
+ Apply(context.HttpContext, context.ProxyRequest, context.HeadersCopied);
+ return default;
+ }
+ internal override void ApplyFast(HttpContext httpContext, HttpRequestMessage proxyRequest, ref bool headersCopied)
+ {
+ Apply(httpContext, proxyRequest, headersCopied);
+ }
+
+ private void Apply(HttpContext httpContext, HttpRequestMessage proxyRequest, bool headersCopied)
+ {
string? remoteIp = null;
- var remoteIpAddress = context.HttpContext.Connection.RemoteIpAddress;
+ var remoteIpAddress = httpContext.Connection.RemoteIpAddress;
if (remoteIpAddress is not null)
{
remoteIp = remoteIpAddress.IsIPv4MappedToIPv6 ?
@@ -51,39 +63,38 @@ public override ValueTask ApplyAsync(RequestTransformContext context)
switch (TransformAction)
{
case ForwardedTransformActions.Set:
- RemoveHeader(context, HeaderName);
+ RemoveHeader(proxyRequest, HeaderName);
if (remoteIp is not null)
{
- AddHeader(context, HeaderName, remoteIp);
+ AddHeader(proxyRequest, HeaderName, remoteIp);
}
break;
case ForwardedTransformActions.Append:
- Append(context, remoteIp);
+ Append(httpContext, proxyRequest, headersCopied, remoteIp);
break;
case ForwardedTransformActions.Remove:
- RemoveHeader(context, HeaderName);
+ RemoveHeader(proxyRequest, HeaderName);
break;
default:
throw new NotImplementedException(TransformAction.ToString());
}
- return default;
}
- private void Append(RequestTransformContext context, string? remoteIp)
+ private void Append(HttpContext httpContext, HttpRequestMessage proxyRequest, bool headersCopied, string? remoteIp)
{
- var existingValues = TakeHeader(context, HeaderName);
+ var existingValues = TakeHeader(httpContext, proxyRequest, headersCopied, HeaderName);
if (remoteIp is null)
{
if (!string.IsNullOrEmpty(existingValues))
{
- AddHeader(context, HeaderName, existingValues);
+ AddHeader(proxyRequest, HeaderName, existingValues);
}
}
else
{
var values = StringValues.Concat(existingValues, remoteIp);
- AddHeader(context, HeaderName, values);
+ AddHeader(proxyRequest, HeaderName, values);
}
}
}
diff --git a/src/ReverseProxy/Transforms/RequestHeaderXForwardedHostTransform.cs b/src/ReverseProxy/Transforms/RequestHeaderXForwardedHostTransform.cs
index 3309b2a3e0..f63cb5b556 100644
--- a/src/ReverseProxy/Transforms/RequestHeaderXForwardedHostTransform.cs
+++ b/src/ReverseProxy/Transforms/RequestHeaderXForwardedHostTransform.cs
@@ -3,7 +3,9 @@
using System;
using System.Diagnostics;
+using System.Net.Http;
using System.Threading.Tasks;
+using Microsoft.AspNetCore.Http;
using Microsoft.Extensions.Primitives;
namespace Yarp.ReverseProxy.Transforms;
@@ -34,45 +36,54 @@ public RequestHeaderXForwardedHostTransform(string headerName, ForwardedTransfor
public override ValueTask ApplyAsync(RequestTransformContext context)
{
ArgumentNullException.ThrowIfNull(context);
+ Apply(context.HttpContext, context.ProxyRequest, context.HeadersCopied);
+ return default;
+ }
- var host = context.HttpContext.Request.Host;
+ internal override void ApplyFast(HttpContext httpContext, HttpRequestMessage proxyRequest, ref bool headersCopied)
+ {
+ Apply(httpContext, proxyRequest, headersCopied);
+ }
+
+ private void Apply(HttpContext httpContext, HttpRequestMessage proxyRequest, bool headersCopied)
+ {
+ var host = httpContext.Request.Host;
switch (TransformAction)
{
case ForwardedTransformActions.Set:
- RemoveHeader(context, HeaderName);
+ RemoveHeader(proxyRequest, HeaderName);
if (host.HasValue)
{
- AddHeader(context, HeaderName, host.ToUriComponent());
+ AddHeader(proxyRequest, HeaderName, host.ToUriComponent());
}
break;
case ForwardedTransformActions.Append:
- Append(context, host);
+ Append(httpContext, proxyRequest, headersCopied, host);
break;
case ForwardedTransformActions.Remove:
- RemoveHeader(context, HeaderName);
+ RemoveHeader(proxyRequest, HeaderName);
break;
default:
throw new NotImplementedException(TransformAction.ToString());
}
- return default;
}
- private void Append(RequestTransformContext context, Microsoft.AspNetCore.Http.HostString host)
+ private void Append(HttpContext httpContext, HttpRequestMessage proxyRequest, bool headersCopied, HostString host)
{
- var existingValues = TakeHeader(context, HeaderName);
+ var existingValues = TakeHeader(httpContext, proxyRequest, headersCopied, HeaderName);
if (!host.HasValue)
{
if (!string.IsNullOrEmpty(existingValues))
{
- AddHeader(context, HeaderName, existingValues);
+ AddHeader(proxyRequest, HeaderName, existingValues);
}
}
else
{
var values = StringValues.Concat(existingValues, host.ToUriComponent());
- AddHeader(context, HeaderName, values);
+ AddHeader(proxyRequest, HeaderName, values);
}
}
}
diff --git a/src/ReverseProxy/Transforms/RequestHeaderXForwardedPrefixTransform.cs b/src/ReverseProxy/Transforms/RequestHeaderXForwardedPrefixTransform.cs
index c493e7d7ef..353a78e15b 100644
--- a/src/ReverseProxy/Transforms/RequestHeaderXForwardedPrefixTransform.cs
+++ b/src/ReverseProxy/Transforms/RequestHeaderXForwardedPrefixTransform.cs
@@ -3,7 +3,9 @@
using System;
using System.Diagnostics;
+using System.Net.Http;
using System.Threading.Tasks;
+using Microsoft.AspNetCore.Http;
using Microsoft.Extensions.Primitives;
namespace Yarp.ReverseProxy.Transforms;
@@ -32,45 +34,54 @@ public RequestHeaderXForwardedPrefixTransform(string headerName, ForwardedTransf
public override ValueTask ApplyAsync(RequestTransformContext context)
{
ArgumentNullException.ThrowIfNull(context);
+ Apply(context.HttpContext, context.ProxyRequest, context.HeadersCopied);
+ return default;
+ }
- var pathBase = context.HttpContext.Request.PathBase;
+ internal override void ApplyFast(HttpContext httpContext, HttpRequestMessage proxyRequest, ref bool headersCopied)
+ {
+ Apply(httpContext, proxyRequest, headersCopied);
+ }
+
+ private void Apply(HttpContext httpContext, HttpRequestMessage proxyRequest, bool headersCopied)
+ {
+ var pathBase = httpContext.Request.PathBase;
switch (TransformAction)
{
case ForwardedTransformActions.Set:
- RemoveHeader(context, HeaderName);
+ RemoveHeader(proxyRequest, HeaderName);
if (pathBase.HasValue)
{
- AddHeader(context, HeaderName, pathBase.ToUriComponent());
+ AddHeader(proxyRequest, HeaderName, pathBase.ToUriComponent());
}
break;
case ForwardedTransformActions.Append:
- Append(context, pathBase);
+ Append(httpContext, proxyRequest, headersCopied, pathBase);
break;
case ForwardedTransformActions.Remove:
- RemoveHeader(context, HeaderName);
+ RemoveHeader(proxyRequest, HeaderName);
break;
default:
throw new NotImplementedException(TransformAction.ToString());
}
- return default;
}
- private void Append(RequestTransformContext context, Microsoft.AspNetCore.Http.PathString pathBase)
+ private void Append(HttpContext httpContext, HttpRequestMessage proxyRequest, bool headersCopied, PathString pathBase)
{
- var existingValues = TakeHeader(context, HeaderName);
+ var existingValues = TakeHeader(httpContext, proxyRequest, headersCopied, HeaderName);
if (!pathBase.HasValue)
{
if (!string.IsNullOrEmpty(existingValues))
{
- AddHeader(context, HeaderName, existingValues);
+ AddHeader(proxyRequest, HeaderName, existingValues);
}
}
else
{
var values = StringValues.Concat(existingValues, pathBase.ToUriComponent());
- AddHeader(context, HeaderName, values);
+ AddHeader(proxyRequest, HeaderName, values);
}
}
}
diff --git a/src/ReverseProxy/Transforms/RequestHeaderXForwardedProtoTransform.cs b/src/ReverseProxy/Transforms/RequestHeaderXForwardedProtoTransform.cs
index 2acbbc2f25..c71af28776 100644
--- a/src/ReverseProxy/Transforms/RequestHeaderXForwardedProtoTransform.cs
+++ b/src/ReverseProxy/Transforms/RequestHeaderXForwardedProtoTransform.cs
@@ -3,7 +3,9 @@
using System;
using System.Diagnostics;
+using System.Net.Http;
using System.Threading.Tasks;
+using Microsoft.AspNetCore.Http;
using Microsoft.Extensions.Primitives;
namespace Yarp.ReverseProxy.Transforms;
@@ -37,27 +39,36 @@ public RequestHeaderXForwardedProtoTransform(string headerName, ForwardedTransfo
public override ValueTask ApplyAsync(RequestTransformContext context)
{
ArgumentNullException.ThrowIfNull(context);
+ Apply(context.HttpContext, context.ProxyRequest, context.HeadersCopied);
+ return default;
+ }
- var scheme = context.HttpContext.Request.Scheme;
+ internal override void ApplyFast(HttpContext httpContext, HttpRequestMessage proxyRequest, ref bool headersCopied)
+ {
+ Apply(httpContext, proxyRequest, headersCopied);
+ }
+
+ private void Apply(HttpContext httpContext, HttpRequestMessage proxyRequest, bool headersCopied)
+ {
+ var scheme = httpContext.Request.Scheme;
switch (TransformAction)
{
case ForwardedTransformActions.Set:
- RemoveHeader(context, HeaderName);
- AddHeader(context, HeaderName, scheme);
+ RemoveHeader(proxyRequest, HeaderName);
+ AddHeader(proxyRequest, HeaderName, scheme);
break;
case ForwardedTransformActions.Append:
- var existingValues = TakeHeader(context, HeaderName);
+ var existingValues = TakeHeader(httpContext, proxyRequest, headersCopied, HeaderName);
var values = StringValues.Concat(existingValues, scheme);
- AddHeader(context, HeaderName, values);
+ AddHeader(proxyRequest, HeaderName, values);
break;
case ForwardedTransformActions.Remove:
- RemoveHeader(context, HeaderName);
+ RemoveHeader(proxyRequest, HeaderName);
break;
default:
throw new NotImplementedException(TransformAction.ToString());
}
- return default;
}
}
diff --git a/src/ReverseProxy/Transforms/RequestHeadersAllowedTransform.cs b/src/ReverseProxy/Transforms/RequestHeadersAllowedTransform.cs
index 93757f6b39..4507cc435f 100644
--- a/src/ReverseProxy/Transforms/RequestHeadersAllowedTransform.cs
+++ b/src/ReverseProxy/Transforms/RequestHeadersAllowedTransform.cs
@@ -5,7 +5,9 @@
using System.Collections.Frozen;
using System.Collections.Generic;
using System.Diagnostics;
+using System.Net.Http;
using System.Threading.Tasks;
+using Microsoft.AspNetCore.Http;
using Microsoft.Extensions.Primitives;
namespace Yarp.ReverseProxy.Transforms;
@@ -33,20 +35,29 @@ public override ValueTask ApplyAsync(RequestTransformContext context)
ArgumentNullException.ThrowIfNull(context);
Debug.Assert(!context.HeadersCopied);
+ Apply(context.HttpContext, context.ProxyRequest);
+ context.HeadersCopied = true;
+ return default;
+ }
+
+ internal override void ApplyFast(HttpContext httpContext, HttpRequestMessage proxyRequest, ref bool headersCopied)
+ {
+ Debug.Assert(!headersCopied);
+ Apply(httpContext, proxyRequest);
+ headersCopied = true;
+ }
- foreach (var header in context.HttpContext.Request.Headers)
+ private void Apply(HttpContext httpContext, HttpRequestMessage proxyRequest)
+ {
+ foreach (var header in httpContext.Request.Headers)
{
var headerName = header.Key;
var headerValue = header.Value;
if (!StringValues.IsNullOrEmpty(headerValue)
&& AllowedHeadersSet.Contains(headerName))
{
- AddHeader(context, headerName, headerValue);
+ AddHeader(proxyRequest, headerName, headerValue);
}
}
-
- context.HeadersCopied = true;
-
- return default;
}
}
diff --git a/src/ReverseProxy/Transforms/RequestTransform.cs b/src/ReverseProxy/Transforms/RequestTransform.cs
index 3d199e61a5..6d0926dc78 100644
--- a/src/ReverseProxy/Transforms/RequestTransform.cs
+++ b/src/ReverseProxy/Transforms/RequestTransform.cs
@@ -2,7 +2,9 @@
// The .NET Foundation licenses this file to you under the MIT license.
using System;
+using System.Net.Http;
using System.Threading.Tasks;
+using Microsoft.AspNetCore.Http;
using Microsoft.Extensions.Primitives;
using Yarp.ReverseProxy.Forwarder;
@@ -18,6 +20,11 @@ public abstract class RequestTransform
///
public abstract ValueTask ApplyAsync(RequestTransformContext context);
+ internal virtual void ApplyFast(HttpContext httpContext, HttpRequestMessage proxyRequest, ref bool headersCopied)
+ {
+ throw new NotSupportedException();
+ }
+
///
/// Removes and returns the current header value by first checking the HttpRequestMessage,
/// then the HttpContent, and falling back to the HttpContext only if
@@ -34,8 +41,11 @@ public static StringValues TakeHeader(RequestTransformContext context, string he
throw new ArgumentException($"'{nameof(headerName)}' cannot be null or empty.", nameof(headerName));
}
- var proxyRequest = context.ProxyRequest;
+ return TakeHeader(context.HttpContext, context.ProxyRequest, context.HeadersCopied, headerName);
+ }
+ internal static StringValues TakeHeader(HttpContext httpContext, HttpRequestMessage proxyRequest, bool headersCopied, string headerName)
+ {
if (RequestUtilities.TryGetValues(proxyRequest.Headers, headerName, out var existingValues))
{
proxyRequest.Headers.Remove(headerName);
@@ -44,9 +54,9 @@ public static StringValues TakeHeader(RequestTransformContext context, string he
{
content.Headers.Remove(headerName);
}
- else if (!context.HeadersCopied)
+ else if (!headersCopied)
{
- existingValues = context.HttpContext.Request.Headers[headerName];
+ existingValues = httpContext.Request.Headers[headerName];
}
return existingValues;
@@ -60,7 +70,12 @@ public static void AddHeader(RequestTransformContext context, string headerName,
ArgumentNullException.ThrowIfNull(context);
ArgumentException.ThrowIfNullOrEmpty(headerName);
- RequestUtilities.AddHeader(context.ProxyRequest, headerName, values);
+ AddHeader(context.ProxyRequest, headerName, values);
+ }
+
+ internal static void AddHeader(HttpRequestMessage proxyRequest, string headerName, StringValues values)
+ {
+ RequestUtilities.AddHeader(proxyRequest, headerName, values);
}
///
@@ -71,6 +86,11 @@ public static void RemoveHeader(RequestTransformContext context, string headerNa
ArgumentNullException.ThrowIfNull(context);
ArgumentException.ThrowIfNullOrEmpty(headerName);
- RequestUtilities.RemoveHeader(context.ProxyRequest, headerName);
+ RemoveHeader(context.ProxyRequest, headerName);
+ }
+
+ internal static void RemoveHeader(HttpRequestMessage proxyRequest, string headerName)
+ {
+ RequestUtilities.RemoveHeader(proxyRequest, headerName);
}
}
diff --git a/test/ReverseProxy.Tests/Transforms/Builder/StructuredTransformerFastPathTests.cs b/test/ReverseProxy.Tests/Transforms/Builder/StructuredTransformerFastPathTests.cs
new file mode 100644
index 0000000000..6e5a95eeac
--- /dev/null
+++ b/test/ReverseProxy.Tests/Transforms/Builder/StructuredTransformerFastPathTests.cs
@@ -0,0 +1,437 @@
+// Licensed to the .NET Foundation under one or more agreements.
+// The .NET Foundation licenses this file to you under the MIT license.
+
+#nullable enable
+
+using System;
+using System.Collections.Generic;
+using System.Linq;
+using System.Net;
+using System.Net.Http;
+using System.Reflection;
+using System.Threading;
+using System.Threading.Tasks;
+using Microsoft.AspNetCore.Http;
+using Microsoft.AspNetCore.Http.Features;
+using Microsoft.Extensions.Primitives;
+using Microsoft.Net.Http.Headers;
+using Xunit;
+using Yarp.ReverseProxy.Forwarder;
+using Yarp.ReverseProxy.Utilities;
+using Yarp.Tests.Common;
+
+namespace Yarp.ReverseProxy.Transforms.Builder.Tests;
+
+public class StructuredTransformerFastPathTests
+{
+ [Theory]
+ [InlineData("HTTP/1.1", "127.0.0.1")]
+ [InlineData("HTTP/2", "::ffff:127.0.0.1")]
+ [InlineData("HTTP/2", "2001:db8::1")]
+ public async Task BuiltInXForwardedAndHeaderTransforms_MatchContextFallback(string protocol, string remoteIp)
+ {
+ static void Configure(TransformBuilderContext context)
+ {
+ context.AddXForwarded(ForwardedTransformActions.Append);
+ context.AddRequestHeader("x-MiXeD-Set", "set", append: false);
+ context.AddRequestHeader("X-Duplicate", "tail", append: true);
+ context.AddRequestHeaderRemove("x-REMOVE");
+ }
+
+ using var cancellation = new CancellationTokenSource();
+ cancellation.Cancel();
+
+ var optimized = CreateTransformer(Configure);
+ var fallback = CreateTransformerWithFallback(Configure);
+ Assert.True(UsesFastPath(optimized));
+ Assert.False(UsesFastPath(fallback));
+
+ var optimizedResult = await TransformRequestAsync(optimized, protocol, IPAddress.Parse(remoteIp), cancellation.Token);
+ var fallbackResult = await TransformRequestAsync(fallback, protocol, IPAddress.Parse(remoteIp), cancellation.Token);
+
+ Assert.Equal(fallbackResult, optimizedResult);
+ Assert.Contains("X-Duplicate:one\u001ftwo\u001ftail", optimizedResult.Headers);
+ Assert.Contains("x-MiXeD-Set:set", optimizedResult.Headers);
+ Assert.DoesNotContain(optimizedResult.Headers, header => header.StartsWith("x-REMOVE:", StringComparison.OrdinalIgnoreCase));
+ }
+
+ [Theory]
+ [InlineData("HTTP/1.1", "127.0.0.1", 1234)]
+ [InlineData("HTTP/2", "2001:db8::1", 4321)]
+ public async Task ForwardedTransform_MatchesContextFallback(string protocol, string remoteIp, int remotePort)
+ {
+ static void Configure(TransformBuilderContext context)
+ {
+ context.UseDefaultForwarders = false;
+ context.RequestTransforms.Add(new RequestHeaderForwardedTransform(
+ new TestRandomFactory(),
+ forFormat: NodeFormat.IpAndPort,
+ byFormat: NodeFormat.None,
+ host: true,
+ proto: true,
+ action: ForwardedTransformActions.Append));
+ context.AddRequestHeader("X-Set", "replacement", append: false);
+ }
+
+ var optimized = CreateTransformer(Configure);
+ var fallback = CreateTransformerWithFallback(Configure);
+ Assert.True(UsesFastPath(optimized));
+ Assert.False(UsesFastPath(fallback));
+
+ var optimizedResult = await TransformRequestAsync(optimized, protocol, IPAddress.Parse(remoteIp), default, remotePort);
+ var fallbackResult = await TransformRequestAsync(fallback, protocol, IPAddress.Parse(remoteIp), default, remotePort);
+
+ Assert.Equal(fallbackResult, optimizedResult);
+ Assert.Contains(optimizedResult.Headers, header =>
+ header.StartsWith("Forwarded:for=\"unterminated\u001ffor=", StringComparison.Ordinal));
+ }
+
+ [Theory]
+ [InlineData("HTTP/1.1")]
+ [InlineData("HTTP/2")]
+ public async Task PathQueryEncodingAndOrdering_ArePreserved(string protocol)
+ {
+ static void Configure(TransformBuilderContext context)
+ {
+ context.AddPathRemovePrefix("/api");
+ context.AddPathPrefix("/v2/%2F");
+ context.AddQueryValue("key", "replacement", append: false);
+ context.AddQueryValue("added", "a/b", append: true);
+ context.AddQueryRemoveKey("remove");
+ }
+
+ var transformer = CreateTransformer(Configure);
+ Assert.False(UsesFastPath(transformer));
+ var result = await TransformRequestAsync(transformer, protocol, IPAddress.IPv6Loopback, default);
+
+ Assert.Equal("http://destination/base/v2/%252F/items/%252Fvalue?key=replacement&escaped=a%2Fb&added=a%2Fb", result.Uri);
+ Assert.Equal(protocol == "HTTP/2", result.Headers.Contains("TE:traiLers"));
+ }
+
+ [Fact]
+ public async Task CustomSyncAndAsyncTransforms_PreserveReassignmentAndCancellation()
+ {
+ using var cancellation = new CancellationTokenSource();
+ cancellation.Cancel();
+ var observedToken = default(CancellationToken);
+ var transformer = CreateTransformer(context =>
+ {
+ context.AddRequestTransform(transformContext =>
+ {
+ observedToken = transformContext.CancellationToken;
+ transformContext.DestinationPrefix = "https://other.example/base";
+ transformContext.Path = "/reassigned/%2F";
+ transformContext.Query = new QueryTransformContext(transformContext.HttpContext.Request);
+ transformContext.Query.Collection["custom"] = "one/two";
+ return default;
+ });
+ context.AddRequestTransform(async transformContext =>
+ {
+ await Task.Yield();
+ transformContext.ProxyRequest.Headers.TryAddWithoutValidation("X-Async", "set");
+ });
+ });
+ Assert.False(UsesFastPath(transformer));
+
+ var result = await TransformRequestAsync(transformer, "HTTP/2", IPAddress.Loopback, cancellation.Token);
+
+ Assert.Equal(cancellation.Token, observedToken);
+ Assert.Equal("https://other.example/base/reassigned/%252F?key=one&key=two&remove=yes&escaped=a%2Fb&custom=one%2Ftwo", result.Uri);
+ Assert.Contains("X-Async:set", result.Headers);
+ }
+
+ [Fact]
+ public async Task DerivedBuiltInTransform_UsesOverriddenApplyAsync()
+ {
+ var transformer = CreateTransformer(context =>
+ {
+ context.UseDefaultForwarders = false;
+ context.RequestTransforms.Add(new DerivedRequestHeaderValueTransform());
+ });
+ Assert.False(UsesFastPath(transformer));
+
+ var result = await TransformRequestAsync(transformer, "HTTP/2", IPAddress.Loopback, default);
+
+ Assert.Contains("X-Derived:derived", result.Headers);
+ }
+
+ [Fact]
+ public async Task HeadersAllowed_OrderingAndHeadersCopied_MatchContextFallback()
+ {
+ static void Configure(TransformBuilderContext context)
+ {
+ context.UseDefaultForwarders = false;
+ context.AddRequestHeader("X-Before", "before", append: true);
+ context.AddRequestHeadersAllowed("X-Before", "X-After");
+ context.AddRequestHeader("X-After", "after", append: true);
+ }
+
+ var optimized = CreateTransformer(Configure);
+ var fallback = CreateTransformerWithFallback(Configure);
+ Assert.True(UsesFastPath(optimized));
+ Assert.False(UsesFastPath(fallback));
+
+ var optimizedResult = await TransformRequestAsync(optimized, "HTTP/2", IPAddress.Loopback, default);
+ var fallbackResult = await TransformRequestAsync(fallback, "HTTP/2", IPAddress.Loopback, default);
+
+ Assert.Equal(fallbackResult, optimizedResult);
+ Assert.Contains("X-Before:original-before\u001fbefore\u001foriginal-before", optimizedResult.Headers);
+ Assert.Contains("X-After:original-after\u001fafter", optimizedResult.Headers);
+ }
+
+ [Fact]
+ public async Task OriginalHostAndRouteValue_MatchContextFallback()
+ {
+ static void Configure(TransformBuilderContext context)
+ {
+ context.UseDefaultForwarders = false;
+ context.AddOriginalHost(useOriginal: true);
+ context.AddRequestHeaderRouteValue("X-Route", "id", append: false);
+ }
+
+ var optimized = CreateTransformer(Configure);
+ var fallback = CreateTransformerWithFallback(Configure);
+ Assert.True(UsesFastPath(optimized));
+ Assert.False(UsesFastPath(fallback));
+
+ var optimizedResult = await TransformRequestAsync(optimized, "HTTP/2", IPAddress.Loopback, default);
+ var fallbackResult = await TransformRequestAsync(fallback, "HTTP/2", IPAddress.Loopback, default);
+
+ Assert.Equal(fallbackResult, optimizedResult);
+ Assert.Contains("Host:xn--host-6j1i.example:8443", optimizedResult.Headers);
+ Assert.Contains("X-Route:route-value", optimizedResult.Headers);
+ }
+
+ [Fact]
+ public void EligibleBuiltIns_OverrideFastPath()
+ {
+ RequestTransform[] transforms =
+ [
+ RequestHeaderOriginalHostTransform.OriginalHost,
+ new RequestHeaderXForwardedForTransform("X-Forwarded-For", ForwardedTransformActions.Set),
+ new RequestHeaderXForwardedHostTransform("X-Forwarded-Host", ForwardedTransformActions.Set),
+ new RequestHeaderXForwardedProtoTransform("X-Forwarded-Proto", ForwardedTransformActions.Set),
+ new RequestHeaderXForwardedPrefixTransform("X-Forwarded-Prefix", ForwardedTransformActions.Set),
+ new RequestHeaderForwardedTransform(
+ new TestRandomFactory(),
+ NodeFormat.Ip,
+ NodeFormat.None,
+ host: true,
+ proto: true,
+ ForwardedTransformActions.Set),
+ new RequestHeaderRemoveTransform("X-Remove"),
+ new RequestHeaderValueTransform("X-Value", "value", append: false),
+ new RequestHeaderRouteValueTransform("X-Route", "id", append: false),
+ new RequestHeadersAllowedTransform(["X-Allowed"]),
+ ];
+
+ foreach (var transform in transforms)
+ {
+ var method = transform.GetType().GetMethod(
+ "ApplyFast",
+ BindingFlags.Instance | BindingFlags.NonPublic);
+
+ Assert.NotNull(method);
+ Assert.Equal(transform.GetType(), method.DeclaringType);
+ }
+ }
+
+ [Theory]
+ [InlineData("HTTP/1.1")]
+ [InlineData("HTTP/2")]
+ public async Task ResponseHeadersAndTrailers_MatchContextFallback(string protocol)
+ {
+ static void Configure(TransformBuilderContext context)
+ {
+ context.AddXForwarded();
+ context.AddResponseHeader("X-Append", "tail", append: true, ResponseCondition.Always);
+ context.AddResponseHeaderRemove("X-Remove", ResponseCondition.Always);
+ context.AddResponseTrailer("X-Trailer", "tail", append: true, ResponseCondition.Always);
+ }
+
+ var optimized = CreateTransformer(Configure);
+ var fallback = CreateTransformerWithFallback(Configure);
+ Assert.True(UsesFastPath(optimized));
+ Assert.False(UsesFastPath(fallback));
+
+ var optimizedResult = await TransformResponseAsync(optimized, protocol);
+ var fallbackResult = await TransformResponseAsync(fallback, protocol);
+
+ Assert.Equal(fallbackResult, optimizedResult);
+ Assert.Contains("X-Append:one\u001ftwo\u001ftail", optimizedResult.Headers);
+ Assert.Contains("X-Trailer:one\u001ftwo\u001ftail", optimizedResult.Trailers);
+ }
+
+ [Fact]
+ public async Task BuiltInTransforms_AreConcurrencySafe()
+ {
+ static void Configure(TransformBuilderContext context)
+ {
+ context.AddXForwarded(ForwardedTransformActions.Append);
+ context.AddRequestHeader("X-Duplicate", "tail", append: true);
+ }
+
+ var optimized = CreateTransformer(Configure);
+ var fallback = CreateTransformerWithFallback(Configure);
+ Assert.True(UsesFastPath(optimized));
+ Assert.False(UsesFastPath(fallback));
+
+ await Parallel.ForEachAsync(Enumerable.Range(0, 200), async (index, _) =>
+ {
+ var protocol = index % 2 == 0 ? "HTTP/1.1" : "HTTP/2";
+ var ipAddress = index % 3 == 0 ? IPAddress.Loopback : IPAddress.IPv6Loopback;
+ var optimizedResult = await TransformRequestAsync(optimized, protocol, ipAddress, default, index);
+ var fallbackResult = await TransformRequestAsync(fallback, protocol, ipAddress, default, index);
+ Assert.Equal(fallbackResult, optimizedResult);
+ });
+ }
+
+ private static StructuredTransformer CreateTransformer(Action configure)
+ {
+ return TransformBuilderTests.CreateTransformBuilder().CreateInternal(configure);
+ }
+
+ private static StructuredTransformer CreateTransformerWithFallback(Action configure)
+ {
+ return TransformBuilderTests.CreateTransformBuilder().CreateInternal(context =>
+ {
+ configure(context);
+ context.AddRequestTransform(static _ => default);
+ });
+ }
+
+ private static async Task TransformRequestAsync(
+ StructuredTransformer transformer,
+ string protocol,
+ IPAddress remoteIp,
+ CancellationToken cancellationToken,
+ int remotePort = 4321)
+ {
+ var context = CreateHttpContext(protocol, remoteIp, remotePort);
+ using var request = new HttpRequestMessage
+ {
+ Content = new ByteArrayContent(Array.Empty()),
+ };
+
+ await transformer.TransformRequestAsync(context, request, "http://destination/base", cancellationToken);
+
+ return new RequestSnapshot(request.RequestUri!.AbsoluteUri, GetHeaders(request));
+ }
+
+ private static async Task TransformResponseAsync(StructuredTransformer transformer, string protocol)
+ {
+ var context = CreateHttpContext(protocol, IPAddress.IPv6Loopback, 4321);
+ var trailersFeature = new TestTrailersFeature();
+ context.Features.Set(trailersFeature);
+ using var response = new HttpResponseMessage(HttpStatusCode.OK)
+ {
+ Content = new ByteArrayContent(Array.Empty()),
+ };
+ response.Headers.TryAddWithoutValidation("X-Append", new[] { "one", "two" });
+ response.Headers.TryAddWithoutValidation("X-Remove", "remove");
+ response.TrailingHeaders.TryAddWithoutValidation("X-Trailer", new[] { "one", "two" });
+
+ await transformer.TransformResponseAsync(context, response, default);
+ await transformer.TransformResponseTrailersAsync(context, response, default);
+
+ return new ResponseSnapshot(GetHeaders(context.Response.Headers), GetHeaders(trailersFeature.Trailers));
+ }
+
+ private static DefaultHttpContext CreateHttpContext(string protocol, IPAddress remoteIp, int remotePort)
+ {
+ var context = new DefaultHttpContext();
+ context.Request.Protocol = protocol;
+ context.Request.Scheme = "https";
+ context.Request.Host = new HostString("ho本st.example", 8443);
+ context.Request.PathBase = "/base";
+ context.Request.Path = "/api/items/%2Fvalue";
+ context.Request.QueryString = new QueryString("?key=one&key=two&remove=yes&escaped=a%2Fb");
+ context.Connection.RemoteIpAddress = remoteIp;
+ context.Connection.RemotePort = remotePort;
+ context.Connection.LocalIpAddress = IPAddress.Loopback;
+ context.Connection.LocalPort = 443;
+
+ context.Request.Headers["x-forwarded-for"] = new StringValues(new[] { "192.0.2.1", "2001:db8::2" });
+ context.Request.Headers["X-FORWARDED-HOST"] = new StringValues(new[] { "prior.example", "older.example" });
+ context.Request.Headers["x-forwarded-proto"] = new StringValues(new[] { "http", "https" });
+ context.Request.Headers["X-Forwarded-Prefix"] = new StringValues(new[] { "/prior", "/older" });
+ context.Request.Headers["Forwarded"] = new StringValues(new[] { "for=\"unterminated", "for=192.0.2.1;proto=http" });
+ context.Request.Headers["X-Duplicate"] = new StringValues(new[] { "one", "two" });
+ context.Request.Headers["X-Remove"] = "remove";
+ context.Request.Headers["X-Mixed-Set"] = "old";
+ context.Request.Headers["X-Before"] = "original-before";
+ context.Request.Headers["X-After"] = "original-after";
+ context.Request.Headers[HeaderNames.TE] = "gzip, traiLers";
+ context.Request.RouteValues["id"] = "route-value";
+ return context;
+ }
+
+ private static bool UsesFastPath(StructuredTransformer transformer)
+ {
+ var field = typeof(StructuredTransformer).GetField(
+ "_canUseRequestTransformFastPath",
+ BindingFlags.Instance | BindingFlags.NonPublic);
+ Assert.NotNull(field);
+ return Assert.IsType(field.GetValue(transformer));
+ }
+
+ private static string[] GetHeaders(HttpRequestMessage request)
+ {
+ return request.Headers.NonValidated
+ .Concat(request.Content!.Headers.NonValidated)
+ .Select(header => $"{header.Key}:{string.Join('\u001f', header.Value)}")
+ .OrderBy(header => header, StringComparer.Ordinal)
+ .ToArray();
+ }
+
+ private static string[] GetHeaders(IHeaderDictionary headers)
+ {
+ return headers
+ .Select(header => $"{header.Key}:{string.Join('\u001f', header.Value.ToArray())}")
+ .OrderBy(header => header, StringComparer.Ordinal)
+ .ToArray();
+ }
+
+ private sealed record RequestSnapshot(string Uri, string[] Headers)
+ {
+ public bool Equals(RequestSnapshot? other)
+ {
+ return other is not null
+ && string.Equals(Uri, other.Uri, StringComparison.Ordinal)
+ && Headers.SequenceEqual(other.Headers, StringComparer.Ordinal);
+ }
+
+ public override int GetHashCode() => HashCode.Combine(Uri, Headers.Length);
+ }
+
+ private sealed record ResponseSnapshot(string[] Headers, string[] Trailers)
+ {
+ public bool Equals(ResponseSnapshot? other)
+ {
+ return other is not null
+ && Headers.SequenceEqual(other.Headers, StringComparer.Ordinal)
+ && Trailers.SequenceEqual(other.Trailers, StringComparer.Ordinal);
+ }
+
+ public override int GetHashCode() => HashCode.Combine(Headers.Length, Trailers.Length);
+ }
+
+ private sealed class TestRandomFactory : IRandomFactory
+ {
+ public Random CreateRandomInstance() => Random.Shared;
+ }
+
+ private sealed class DerivedRequestHeaderValueTransform : RequestHeaderValueTransform
+ {
+ public DerivedRequestHeaderValueTransform()
+ : base("X-Derived", "base", append: false)
+ {
+ }
+
+ public override ValueTask ApplyAsync(RequestTransformContext context)
+ {
+ context.ProxyRequest.Headers.TryAddWithoutValidation("X-Derived", "derived");
+ return default;
+ }
+ }
+}