diff --git a/starlette/middleware/trustedhost.py b/starlette/middleware/trustedhost.py index 98451e29f..0d50d0199 100644 --- a/starlette/middleware/trustedhost.py +++ b/starlette/middleware/trustedhost.py @@ -37,7 +37,16 @@ async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: return headers = Headers(scope=scope) - host = headers.get("host", "").split(":")[0] + host = headers.get("host", "") + if host.startswith("["): + closing_bracket = host.find("]") + suffix = host[closing_bracket + 1 :] if closing_bracket != -1 else "" + if closing_bracket != -1 and ( + not suffix or (suffix.startswith(":") and suffix[1:].isascii() and suffix[1:].isdigit()) + ): + host = host[: closing_bracket + 1] + else: + host = host.split(":", 1)[0] is_valid_host = False found_www_redirect = False for pattern in self.allowed_hosts: diff --git a/tests/middleware/test_trusted_host.py b/tests/middleware/test_trusted_host.py index 5b8b217c3..5ac86cebc 100644 --- a/tests/middleware/test_trusted_host.py +++ b/tests/middleware/test_trusted_host.py @@ -29,6 +29,29 @@ def homepage(request: Request) -> PlainTextResponse: assert response.status_code == 400 +def test_trusted_host_middleware_with_ipv6(test_client_factory: TestClientFactory) -> None: + def homepage(request: Request) -> PlainTextResponse: + return PlainTextResponse("OK", status_code=200) + + app = Starlette( + routes=[Route("/", endpoint=homepage)], + middleware=[Middleware(TrustedHostMiddleware, allowed_hosts=["[::1]"])], + ) + + client = test_client_factory(app) + response = client.get("/", headers={"host": "[::1]:8000"}) + assert response.status_code == 200 + + response = client.get("/", headers={"host": "[::1]:attacker"}) + assert response.status_code == 400 + + response = client.get("/", headers={"host": "[::1]:"}) + assert response.status_code == 400 + + response = client.get("/", headers={b"host": b"[::1]:\xb2"}) + assert response.status_code == 400 + + def test_default_allowed_hosts() -> None: app = Starlette() middleware = TrustedHostMiddleware(app)