fix: add ForwardedHostMiddleware to rewrite ASGI scope from X-Forwarded-Host

The 307 trailing-slash redirects generated by Starlette's Mount and
StaticFiles classes construct their Location header from the raw host
header in the ASGI scope. Uvicorn's ProxyHeadersMiddleware only handles
X-Forwarded-Proto and X-Forwarded-For — it does not rewrite the host.

This adds a lightweight ASGI middleware that rewrites scope["headers"]
host and scope["server"] from X-Forwarded-Host when present, so all
redirect URLs use the external hostname instead of internal pod IPs.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
yuneng-jiang 2026-03-06 16:36:30 -08:00
parent ea7db82049
commit 5835342cb3
2 changed files with 105 additions and 0 deletions

View file

@ -1424,6 +1424,57 @@ app.add_middleware(
app.add_middleware(PrometheusAuthMiddleware)
class ForwardedHostMiddleware:
"""
ASGI middleware that rewrites the ASGI scope when X-Forwarded-Host
and/or X-Forwarded-Proto headers are present.
This ensures that all redirect URLs generated by Starlette
(StaticFiles trailing-slash redirects, Mount redirects, etc.)
use the external hostname instead of the internal pod IP.
Uvicorn's ProxyHeadersMiddleware handles X-Forwarded-Proto and
X-Forwarded-For but does NOT handle X-Forwarded-Host.
"""
def __init__(self, app: Any) -> None:
self.app = app
async def __call__(self, scope: Any, receive: Any, send: Any) -> None:
if scope["type"] in ("http", "websocket"):
headers = dict(scope.get("headers", []))
forwarded_host = headers.get(b"x-forwarded-host")
if forwarded_host is not None:
# Rewrite the host header so URL(scope=scope) uses the
# external hostname for all redirect Location headers.
new_headers = []
for key, value in scope["headers"]:
if key == b"host":
new_headers.append((b"host", forwarded_host))
else:
new_headers.append((key, value))
scope = dict(scope)
scope["headers"] = new_headers
# Also update scope["server"] for consistency
decoded = forwarded_host.decode("latin-1")
if ":" in decoded and not decoded.startswith("["):
host, port_str = decoded.rsplit(":", 1)
try:
scope["server"] = (host, int(port_str))
except ValueError:
scope["server"] = (decoded, None)
else:
scope["server"] = (decoded, None)
await self.app(scope, receive, send)
# ForwardedHostMiddleware must be added last so it runs first in the
# request chain, before any other middleware constructs redirect URLs.
app.add_middleware(ForwardedHostMiddleware)
def mount_swagger_ui():
swagger_directory = os.path.join(current_dir, "swagger")
swagger_path = "/" if server_root_path is None else server_root_path

View file

@ -190,6 +190,60 @@ def test_get_custom_url_proxy_base_url_takes_priority(monkeypatch):
assert url == "https://configured.company.com/ui/"
@pytest.mark.asyncio
async def test_forwarded_host_middleware_rewrites_host():
"""Test that ForwardedHostMiddleware rewrites scope host from X-Forwarded-Host"""
from litellm.proxy.proxy_server import ForwardedHostMiddleware
captured_scope = {}
async def mock_app(scope, receive, send):
captured_scope.update(scope)
middleware = ForwardedHostMiddleware(mock_app)
scope = {
"type": "http",
"headers": [
(b"host", b"100.64.1.22:4000"),
(b"x-forwarded-host", b"external.company.com"),
(b"x-forwarded-proto", b"https"),
],
"server": ("100.64.1.22", 4000),
}
await middleware(scope, None, None)
# Verify host header was rewritten
host_values = [v for k, v in captured_scope["headers"] if k == b"host"]
assert host_values == [b"external.company.com"]
assert captured_scope["server"] == ("external.company.com", None)
@pytest.mark.asyncio
async def test_forwarded_host_middleware_no_forwarded_host():
"""Test that middleware is a no-op without X-Forwarded-Host"""
from litellm.proxy.proxy_server import ForwardedHostMiddleware
captured_scope = {}
async def mock_app(scope, receive, send):
captured_scope.update(scope)
middleware = ForwardedHostMiddleware(mock_app)
scope = {
"type": "http",
"headers": [
(b"host", b"100.64.1.22:4000"),
],
"server": ("100.64.1.22", 4000),
}
await middleware(scope, None, None)
# Verify host header was NOT rewritten
host_values = [v for k, v in captured_scope["headers"] if k == b"host"]
assert host_values == [b"100.64.1.22:4000"]
assert captured_scope["server"] == ("100.64.1.22", 4000)
def _patch_today(monkeypatch, year, month, day):
class PatchedDate(real_datetime.date):
@classmethod