fix(proxy): enforce HEADROOM_PROXY_TOKEN on WebSocket handshakes (#3305)

## Description

Closes #3281.

The security gate is registered with `@app.middleware("http")`
(`proxy/server.py`), which is a Starlette `BaseHTTPMiddleware` — and
that class hands any scope whose type is not `http` straight to the
wrapped app. WebSocket connections therefore never reached it, so
**every `app.websocket(...)` route accepted unauthenticated callers even
with `HEADROOM_PROXY_TOKEN` configured.**

Those routes are not incidental:

- `/v1/responses`, `/v1/codex/responses`, `/backend-api/responses`,
`/backend-api/codex/responses`
- `/v1/live`, `/v1/codex/live`, `/backend-api/live`,
`/backend-api/codex/live`

Both families are registered **unconditionally**
(`providers/proxy_routes.py:207` and `:230`), and they relay to the
upstream provider using the operator's own credentials. `/v1/responses`
is served on both transports, which makes the shape of the bug concrete:
the POST was authenticated, the upgrade on the very same path was not.

The existing WebSocket origin check (`_is_allowed_websocket_origin`) is
not a substitute — it defends against browser-driven cross-site
connections, and a non-browser client simply omits `Origin`.

## Type of Change

- [x] Bug fix (non-breaking change that fixes an issue)

## Changes Made

- Added `WebSocketAuthMiddleware`, a raw ASGI middleware, beside the
existing `WebSocketProjectPrefixMiddleware` (the codebase already uses
that idiom for the WebSocket scope). Registered after the HTTP gate so
it runs outermost — an unauthenticated handshake is refused before any
project-prefix or routing work.
- It applies exactly the HTTP gate's rule: loopback exempt
(`is_loopback_host`, including the `None` → loopback case for
UDS/TestClient), credential from `Authorization: Bearer` or
`X-Headroom-Proxy-Token`, `hmac.compare_digest` against a pre-encoded
token.
- Extracted the credential-reading rule into one shared
`read_proxy_token` used by both transports, so they cannot drift.
- Rejection sends `websocket.close` with **1008** *before* accept, after
receiving `websocket.connect` — that is what refuses the upgrade on the
wire rather than accepting and dropping it.

Deliberately **not** done: no query-string credential. Browsers cannot
set headers on a WebSocket, but these routes serve programmatic clients
that can, and a token in a URL lands in access logs and history.

## Testing

- [x] Unit tests pass (`pytest`)
- [x] Linting passes (CI-pinned `ruff` 0.16.3)
- [x] Type checking passes (`mypy headroom`)
- [x] New tests added for new functionality

### Test Output

```text
$ pytest tests/test_proxy_hardening.py
27 passed in 5.17s

$ pytest tests/test_proxy/
267 passed in 86.03s

$ pytest tests/ -k "hardening or websocket or ws or loopback or auth or security"
852 passed, 18 skipped, 11420 deselected in 83.95s

$ uvx ruff@0.16.3 check headroom/proxy/server.py tests/test_proxy_hardening.py
All checks passed!
$ mypy headroom/proxy/server.py
Success: no issues found in 1 source file
```

Tests are in two layers, deliberately:

**Unit (8)** — the middleware driven directly over ASGI. Asserted at
this layer because a pre-accept close surfaces through `TestClient` as a
bare `AttributeError`, indistinguishable from any other handshake
failure, so an exception-shape assertion would pass for the wrong
reason. These assert the downstream app is never invoked and that a
`websocket.close` with code 1008 was sent.

**Integration (4)** — that the middleware is actually wired into
`create_app`, asserted via the security property itself: the route
handler must never run for an unauthenticated handshake. Verified by
removing only the registration line — both `/v1/responses` and
`/v1/live` then fail:

```text
FAILED ...test_unauthenticated_handshake_never_reaches_the_handler[/v1/responses]
FAILED ...test_unauthenticated_handshake_never_reaches_the_handler[/v1/live]
2 failed, 2 passed
```

The 2 that still pass are the authenticated-path invariants, which must
hold either way.

## Real Behavior Proof

- **Environment:** macOS arm64, Python 3.12.13, `main` @ 0.36.5.
- **Exact command / steps:** build the real app with `proxy_token` set,
spy on both WebSocket route handlers, then attempt a handshake from a
non-loopback client (`203.0.113.5`, TEST-NET-3) with and without a
credential.
- **Observed result:** before — the handler ran for an unauthenticated
handshake on both route families. After — the handler is never reached
without a credential, and is reached with either accepted header form.
Loopback and no-token-configured both stay open, unchanged.
- **Not tested:** no live upstream WebSocket session end to end; the
upstream relay itself is unchanged by this PR. Not exercised against a
real browser client, which cannot send the header — see the query-string
note above.

## Runtime Rollout Safety

- **Rollout-managed feature(s):** none.
- **Minimum rollout channel:** n/a.
- **Stable/default behavior changed:** **no** for the default
deployment. With no `HEADROOM_PROXY_TOKEN` the middleware is a
passthrough, so nothing gains a new challenge. Behaviour changes only
where a token is already configured — where the WebSocket routes were
meant to be gated and silently were not.
- **Kill switch / disable path:** unset `HEADROOM_PROXY_TOKEN` (restores
the previous, open behaviour on both transports).
- **Unsafe override required:** none.
- **Qualification impact:** none.
- **Rollback path:** revert this commit; it is one middleware class plus
its registration.

## Review Readiness

- [x] I have performed a self-review
- [x] This PR is ready for human review

---------

Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
Tejas Chopra 2026-08-27 15:19:15 +05:30 committed by GitHub
parent 7c0b886004
commit 27b4e2d147
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 312 additions and 10 deletions

View file

@ -37,7 +37,7 @@ import sys
import threading import threading
import time import time
from collections import OrderedDict from collections import OrderedDict
from collections.abc import Callable from collections.abc import Callable, Mapping
from dataclasses import fields, is_dataclass, replace from dataclasses import fields, is_dataclass, replace
from datetime import datetime, timezone from datetime import datetime, timezone
from pathlib import Path from pathlib import Path
@ -2605,6 +2605,92 @@ _is_known_websocket_callback_failure = is_known_websocket_callback_failure
_tool_schema_saved_from_tags = tool_schema_saved_from_tags _tool_schema_saved_from_tags = tool_schema_saved_from_tags
def read_proxy_token(headers: Mapping[str, str]) -> str | None:
"""Return the caller-supplied proxy token from request headers, or ``None``.
Shared by the HTTP security gate and :class:`WebSocketAuthMiddleware` so the
two transports cannot drift on what counts as a credential. Header names are
expected to be lowercase (Starlette's ``Headers`` is case-insensitive; the
WebSocket middleware lowercases the raw ASGI pairs itself).
"""
auth = str(headers.get("authorization") or "")
if auth.lower().startswith("bearer "):
return auth[7:].strip() or None
raw = headers.get("x-headroom-proxy-token")
return str(raw) if raw else None
class WebSocketAuthMiddleware:
"""Enforce ``HEADROOM_PROXY_TOKEN`` on WebSocket handshakes.
The HTTP security gate is registered with ``@app.middleware("http")``, which
is a Starlette ``BaseHTTPMiddleware`` and that class hands any scope whose
type is not ``http`` straight to the wrapped app. WebSocket connections
therefore never reached the gate, so every ``app.websocket(...)`` route
accepted unauthenticated callers even with a token configured. Those routes
are not incidental: ``/v1/responses`` and ``/v1/live`` relay to the upstream
provider using the operator's own credentials, and they are registered
unconditionally. ``/v1/responses`` exists on both transports, so the POST was
authenticated while the upgrade on the very same path was not.
Written as a raw ASGI middleware rather than folded into the gate because
that is the only layer that sees the ``websocket`` scope at all.
Loopback callers are exempt, matching the HTTP gate exactly (same trust
boundary as the admin/debug routes). Credentials are read from headers only:
the handshake carries them fine for the programmatic clients these routes
serve, and accepting a token from the query string would put it in access
logs and browser history.
"""
def __init__(self, app: Any, *, proxy_token: str | None = None) -> None:
self.app = app
self.proxy_token = proxy_token
# Pre-encoded for constant-time comparison, mirroring the HTTP gate:
# compare_digest on str raises TypeError for non-ASCII input, which
# would turn a rejected handshake into a 500.
self.token_bytes = proxy_token.encode("utf-8") if proxy_token else b""
async def __call__(self, scope: Any, receive: Any, send: Any) -> None:
if scope["type"] != "websocket" or not self.proxy_token:
await self.app(scope, receive, send)
return
client = scope.get("client")
client_host = client[0] if client else None
if is_loopback_host(client_host):
await self.app(scope, receive, send)
return
# Starlette's own Headers rather than a hand-built dict: on a repeated
# header it returns the FIRST occurrence, which is what the HTTP gate
# sees. Building a dict here instead took the LAST one, so the two
# transports disagreed about which `Authorization` counted — exactly the
# drift the shared reader below exists to prevent.
from starlette.datastructures import Headers
provided = read_proxy_token(Headers(scope=scope))
if provided is not None and hmac.compare_digest(
provided.encode("utf-8", "replace"), self.token_bytes
):
await self.app(scope, receive, send)
return
logger.warning(
"event=proxy_auth_rejected transport=websocket path=%s client=%s reason=%s",
scope.get("path"),
client_host,
"missing_token" if provided is None else "bad_token",
)
# Receive the handshake before refusing it: ASGI servers send
# ``websocket.connect`` and wait for the application to answer, and
# answering with ``websocket.close`` *before* an accept is what refuses
# the upgrade on the wire instead of accepting and then dropping it.
message = await receive()
if message["type"] == "websocket.connect":
await send({"type": "websocket.close", "code": 1008})
class WebSocketProjectPrefixMiddleware: class WebSocketProjectPrefixMiddleware:
"""Normalize project-prefixed WebSocket paths before route matching.""" """Normalize project-prefixed WebSocket paths before route matching."""
@ -3450,10 +3536,12 @@ def create_app(config: ProxyConfig | None = None) -> FastAPI:
) )
return response return response
# ── Security gate (registered last → runs outermost) ────────────────── # ── Security gate (outermost of the HTTP middlewares) ─────────────────
# Three concerns, kept together because they all wrap every inbound # Three concerns, kept together because they all wrap every inbound HTTP
# request: optional inbound auth on the data plane, response security # request: optional inbound auth on the data plane, response security
# headers, and an audit trail for state-mutating admin endpoints. # headers, and an audit trail for state-mutating admin endpoints.
# WebSocket handshakes are covered separately — see the
# WebSocketAuthMiddleware registration just below this block.
_proxy_token = config.proxy_token or os.environ.get("HEADROOM_PROXY_TOKEN") or None _proxy_token = config.proxy_token or os.environ.get("HEADROOM_PROXY_TOKEN") or None
# Pre-encode once for constant-time comparison (compare_digest on str raises # Pre-encode once for constant-time comparison (compare_digest on str raises
# TypeError for non-ASCII input, which would turn a 401 into a 500). # TypeError for non-ASCII input, which would turn a 401 into a 500).
@ -3483,12 +3571,9 @@ def create_app(config: ProxyConfig | None = None) -> FastAPI:
"Strict-Transport-Security", "max-age=31536000; includeSubDomains" "Strict-Transport-Security", "max-age=31536000; includeSubDomains"
) )
def _extract_proxy_token(headers) -> str | None: # Delegates so the HTTP gate and WebSocketAuthMiddleware read a credential
auth = str(headers.get("authorization") or "") # by exactly one rule; they guard the same token on two transports.
if auth.lower().startswith("bearer "): _extract_proxy_token = read_proxy_token
return auth[7:].strip() or None
raw = headers.get("x-headroom-proxy-token")
return str(raw) if raw else None
@app.middleware("http") @app.middleware("http")
async def _security_gate(request, call_next): async def _security_gate(request, call_next):
@ -3530,6 +3615,12 @@ def create_app(config: ProxyConfig | None = None) -> FastAPI:
logger.debug("admin audit emission failed", exc_info=True) logger.debug("admin audit emission failed", exc_info=True)
return response return response
# The gate above is http-only (BaseHTTPMiddleware ignores every other
# scope), so the same token rule is applied to the `websocket` scope here.
# Added after it, which makes it the outermost layer — an unauthenticated
# handshake is refused before any project-prefix or routing work happens.
app.add_middleware(WebSocketAuthMiddleware, proxy_token=_proxy_token)
# Third-party proxy extensions (Enterprise, custom plugins). Discovered via # Third-party proxy extensions (Enterprise, custom plugins). Discovered via
# the `headroom.proxy_extension` entry-point group, but **opt-in only**: # the `headroom.proxy_extension` entry-point group, but **opt-in only**:
# only names listed in config.proxy_extensions (CLI: --proxy-extension, # only names listed in config.proxy_extensions (CLI: --proxy-extension,

View file

@ -19,7 +19,7 @@ from fastapi.testclient import TestClient
from headroom.cache.compression_store import reset_compression_store from headroom.cache.compression_store import reset_compression_store
from headroom.offline import apply_offline_env, is_offline from headroom.offline import apply_offline_env, is_offline
from headroom.proxy.audit import is_auditable_path from headroom.proxy.audit import is_auditable_path
from headroom.proxy.server import ProxyConfig, create_app from headroom.proxy.server import ProxyConfig, WebSocketAuthMiddleware, create_app
NONLOOPBACK = ("203.0.113.5", 44444) # TEST-NET-3, never loopback NONLOOPBACK = ("203.0.113.5", 44444) # TEST-NET-3, never loopback
LOOPBACK = ("127.0.0.1", 12345) LOOPBACK = ("127.0.0.1", 12345)
@ -85,6 +85,217 @@ class TestInboundAuthToken:
assert c.get("/readyz").status_code in (200, 503) # ready/not-ready, never 401 assert c.get("/readyz").status_code in (200, 503) # ready/not-ready, never 401
# ──────────────────── 2.1b inbound auth token over WebSocket ──────────────
WS_PATHS = ("/v1/responses", "/v1/live")
class _SpyApp:
"""Downstream ASGI app that records whether it was ever reached."""
def __init__(self) -> None:
self.called = False
async def __call__(self, scope, receive, send) -> None:
self.called = True
def _ws_scope(*, client=NONLOOPBACK, headers=(), path="/v1/responses"):
return {
"type": "websocket",
"path": path,
"client": client,
"headers": [(k.lower().encode("latin-1"), v.encode("latin-1")) for k, v in headers],
}
async def _drive(middleware, scope):
"""Run one connection through the middleware, returning (sent, downstream)."""
inbox = [{"type": "websocket.connect"}]
sent: list[dict] = []
async def receive():
return inbox.pop(0) if inbox else {"type": "websocket.disconnect"}
async def send(message):
sent.append(message)
await middleware(scope, receive, send)
return sent
def _closed_with_policy_violation(sent) -> bool:
return any(m.get("type") == "websocket.close" and m.get("code") == 1008 for m in sent)
class TestWebSocketAuthMiddleware:
"""The middleware itself, driven directly over ASGI.
Asserted at this layer because a pre-accept close surfaces through
``TestClient`` as a bare ``AttributeError`` indistinguishable from any
other handshake failure so an exception-shape assertion would pass for
the wrong reason.
"""
async def test_rejects_missing_credential(self):
downstream = _SpyApp()
mw = WebSocketAuthMiddleware(downstream, proxy_token="s3cr3t-token")
sent = await _drive(mw, _ws_scope())
assert downstream.called is False
assert _closed_with_policy_violation(sent)
async def test_rejects_wrong_credential(self):
downstream = _SpyApp()
mw = WebSocketAuthMiddleware(downstream, proxy_token="s3cr3t-token")
sent = await _drive(mw, _ws_scope(headers=[("authorization", "Bearer wrong")]))
assert downstream.called is False
assert _closed_with_policy_violation(sent)
async def test_accepts_correct_bearer(self):
downstream = _SpyApp()
mw = WebSocketAuthMiddleware(downstream, proxy_token="s3cr3t-token")
sent = await _drive(mw, _ws_scope(headers=[("authorization", "Bearer s3cr3t-token")]))
assert downstream.called is True
assert not _closed_with_policy_violation(sent)
async def test_accepts_custom_header(self):
downstream = _SpyApp()
mw = WebSocketAuthMiddleware(downstream, proxy_token="s3cr3t-token")
sent = await _drive(mw, _ws_scope(headers=[("x-headroom-proxy-token", "s3cr3t-token")]))
assert downstream.called is True
assert not _closed_with_policy_violation(sent)
async def test_loopback_is_exempt(self):
"""Same trust boundary the HTTP gate already grants loopback."""
downstream = _SpyApp()
mw = WebSocketAuthMiddleware(downstream, proxy_token="s3cr3t-token")
sent = await _drive(mw, _ws_scope(client=LOOPBACK))
assert downstream.called is True
assert not _closed_with_policy_violation(sent)
async def test_unknown_client_is_treated_as_loopback(self):
"""Mirrors is_loopback_host(None) -> True, as the HTTP gate does."""
downstream = _SpyApp()
mw = WebSocketAuthMiddleware(downstream, proxy_token="s3cr3t-token")
sent = await _drive(mw, _ws_scope(client=None))
assert downstream.called is True
assert not _closed_with_policy_violation(sent)
async def test_repeated_header_resolves_like_the_http_gate(self):
"""A duplicated Authorization must mean the same thing on both transports.
Starlette's Headers (what the HTTP gate reads) returns the FIRST
occurrence. A hand-built dict returns the last, which would let the two
paths disagree about which credential counted.
"""
downstream = _SpyApp()
mw = WebSocketAuthMiddleware(downstream, proxy_token="s3cr3t-token")
sent = await _drive(
mw,
_ws_scope(
headers=[
("authorization", "Bearer s3cr3t-token"),
("authorization", "Bearer wrong"),
]
),
)
# First header wins → authenticated, same as the HTTP gate.
assert downstream.called is True
assert not _closed_with_policy_violation(sent)
async def test_no_token_configured_is_a_passthrough(self):
"""Default deployment must gain no new challenge."""
downstream = _SpyApp()
mw = WebSocketAuthMiddleware(downstream, proxy_token=None)
sent = await _drive(mw, _ws_scope())
assert downstream.called is True
assert not _closed_with_policy_violation(sent)
async def test_http_scope_is_left_to_the_http_gate(self):
downstream = _SpyApp()
mw = WebSocketAuthMiddleware(downstream, proxy_token="s3cr3t-token")
sent = await _drive(mw, {**_ws_scope(), "type": "http"})
assert downstream.called is True
assert not _closed_with_policy_violation(sent)
class TestWebSocketRoutesAreGatedInTheApp:
"""The middleware is actually wired into ``create_app``.
Asserts the security property directly the route handler must never run
for an unauthenticated handshake rather than inspecting the exception the
client happens to see.
"""
@pytest.mark.parametrize("path", WS_PATHS)
def test_unauthenticated_handshake_never_reaches_the_handler(self, path, monkeypatch):
app = _make_app(proxy_token="s3cr3t-token")
reached = _record_ws_handler_reached(app, monkeypatch)
with TestClient(app, base_url="http://testserver", client=NONLOOPBACK) as c:
try:
with c.websocket_connect(path):
pass
except Exception: # noqa: BLE001 - the refusal shape is asserted above
pass
assert reached() is False
@pytest.mark.parametrize("path", WS_PATHS)
def test_authenticated_handshake_reaches_the_handler(self, path, monkeypatch):
app = _make_app(proxy_token="s3cr3t-token")
reached = _record_ws_handler_reached(app, monkeypatch)
with TestClient(app, base_url="http://testserver", client=NONLOOPBACK) as c:
try:
with c.websocket_connect(path, headers={"X-Headroom-Proxy-Token": "s3cr3t-token"}):
pass
except Exception: # noqa: BLE001 - route may fail with no upstream
pass
assert reached() is True
def _record_ws_handler_reached(app, monkeypatch):
"""Spy both WebSocket route families; returns a callable reporting arrival."""
from headroom.providers import proxy_routes
seen: list[str] = []
# Each spy must terminate the handshake itself: a handler that returns
# without accepting or closing leaves the client waiting forever.
async def _responses_spy(websocket):
seen.append("responses")
await websocket.close(code=1000)
async def _live_spy(websocket, *args, **kwargs):
seen.append("live")
await websocket.close(code=1000)
monkeypatch.setattr(app.state.proxy, "handle_openai_responses_ws", _responses_spy)
monkeypatch.setattr(proxy_routes, "handle_codex_live_websocket", _live_spy)
return lambda: bool(seen)
# ───────────────────────────── 3.1 security headers ─────────────────────── # ───────────────────────────── 3.1 security headers ───────────────────────