mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
ea7db82049
commit
5835342cb3
2 changed files with 105 additions and 0 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue