diff --git a/headroom/proxy/helpers.py b/headroom/proxy/helpers.py index e03634c7d..b931b9339 100644 --- a/headroom/proxy/helpers.py +++ b/headroom/proxy/helpers.py @@ -39,6 +39,14 @@ from headroom.proxy.body_forwarding import ( ) from headroom.proxy.body_forwarding import serialize_body_canonical from headroom.proxy.ccr_session_tracker import SessionCcrTracker as _SessionCcrTracker +from headroom.proxy.internal_header_policy import ( + INTERNAL_HEADER_PREFIX, + STRIP_INTERNAL_HEADERS_DEFAULT, + STRIP_INTERNAL_HEADERS_ENV, + StripInternalHeadersMode, + resolve_strip_internal_headers_mode, + strip_internal_headers, +) from headroom.proxy.tool_injection_config import ( ToolInjectionStickyMode, ) @@ -1474,16 +1482,9 @@ def is_anthropic_auth(headers: dict[str, str]) -> bool: # tell its client about its own work. This helper only filters # request-side headers. -_INTERNAL_HEADER_PREFIX = "x-headroom-" - -# Operator opt-in env var. ``enabled`` (default) strips internal -# ``x-headroom-*`` headers from every upstream-bound forwarder. -# ``disabled`` is an explicit operator opt-in for diagnostic shadow -# tracing — NOT a fallback. Per realignment build constraint #4 the -# behaviour is loud, configurable, and never silent. -_STRIP_INTERNAL_HEADERS_ENV = "HEADROOM_STRIP_INTERNAL_HEADERS" -StripInternalHeadersMode = Literal["enabled", "disabled"] -_STRIP_INTERNAL_HEADERS_DEFAULT: StripInternalHeadersMode = "enabled" +_INTERNAL_HEADER_PREFIX = INTERNAL_HEADER_PREFIX +_STRIP_INTERNAL_HEADERS_ENV = STRIP_INTERNAL_HEADERS_ENV +_STRIP_INTERNAL_HEADERS_DEFAULT = STRIP_INTERNAL_HEADERS_DEFAULT def get_strip_internal_headers_mode() -> StripInternalHeadersMode: @@ -1493,14 +1494,7 @@ def get_strip_internal_headers_mode() -> StripInternalHeadersMode: restart. Unknown values raise loudly per the no-silent-fallback build constraint. """ - raw = os.environ.get(_STRIP_INTERNAL_HEADERS_ENV, "").strip().lower() - if not raw: - return _STRIP_INTERNAL_HEADERS_DEFAULT - if raw in ("enabled", "disabled"): - return cast(StripInternalHeadersMode, raw) - raise ValueError( - f"Invalid {_STRIP_INTERNAL_HEADERS_ENV}={raw!r}; expected 'enabled' or 'disabled'" - ) + return resolve_strip_internal_headers_mode(os.environ.get(_STRIP_INTERNAL_HEADERS_ENV)) def _strip_internal_headers(headers: dict[str, str]) -> dict[str, str]: @@ -1516,11 +1510,7 @@ def _strip_internal_headers(headers: dict[str, str]) -> dict[str, str]: is set, returns a shallow copy unchanged. That mode is for diagnostic shadow tracing only and is documented as a per-deploy choice. """ - mode = get_strip_internal_headers_mode() - if mode == "disabled": - # Always return a copy so callers can mutate without surprise. - return dict(headers) - return {k: v for k, v in headers.items() if not k.lower().startswith(_INTERNAL_HEADER_PREFIX)} + return strip_internal_headers(headers, mode=get_strip_internal_headers_mode()) def log_outbound_headers( diff --git a/headroom/proxy/internal_header_policy.py b/headroom/proxy/internal_header_policy.py new file mode 100644 index 000000000..422a7c280 --- /dev/null +++ b/headroom/proxy/internal_header_policy.py @@ -0,0 +1,40 @@ +"""Policy for stripping proxy-internal request headers before upstream calls.""" + +from __future__ import annotations + +from collections.abc import Mapping +from typing import Literal, cast + +INTERNAL_HEADER_PREFIX = "x-headroom-" +STRIP_INTERNAL_HEADERS_ENV = "HEADROOM_STRIP_INTERNAL_HEADERS" +StripInternalHeadersMode = Literal["enabled", "disabled"] +STRIP_INTERNAL_HEADERS_DEFAULT: StripInternalHeadersMode = "enabled" + + +def resolve_strip_internal_headers_mode(raw: str | None) -> StripInternalHeadersMode: + """Resolve the configured internal-header strip mode.""" + + normalized = (raw or "").strip().lower() + if not normalized: + return STRIP_INTERNAL_HEADERS_DEFAULT + if normalized in ("enabled", "disabled"): + return cast(StripInternalHeadersMode, normalized) + raise ValueError( + f"Invalid {STRIP_INTERNAL_HEADERS_ENV}={normalized!r}; expected 'enabled' or 'disabled'" + ) + + +def strip_internal_headers( + headers: Mapping[str, str], + *, + mode: StripInternalHeadersMode, +) -> dict[str, str]: + """Return a copy of headers with internal x-headroom-* request headers removed.""" + + if mode == "disabled": + return dict(headers) + return { + key: value + for key, value in headers.items() + if not key.lower().startswith(INTERNAL_HEADER_PREFIX) + } diff --git a/tests/test_internal_header_policy.py b/tests/test_internal_header_policy.py new file mode 100644 index 000000000..8d28d0344 --- /dev/null +++ b/tests/test_internal_header_policy.py @@ -0,0 +1,50 @@ +from __future__ import annotations + +import pytest + +from headroom.proxy.internal_header_policy import ( + STRIP_INTERNAL_HEADERS_ENV, + resolve_strip_internal_headers_mode, + strip_internal_headers, +) + + +def test_resolve_strip_internal_headers_mode_defaults_to_enabled() -> None: + assert resolve_strip_internal_headers_mode(None) == "enabled" + assert resolve_strip_internal_headers_mode(" ") == "enabled" + + +def test_resolve_strip_internal_headers_mode_accepts_known_values() -> None: + assert resolve_strip_internal_headers_mode("ENABLED") == "enabled" + assert resolve_strip_internal_headers_mode(" disabled ") == "disabled" + + +def test_resolve_strip_internal_headers_mode_rejects_unknown_values() -> None: + with pytest.raises(ValueError, match=STRIP_INTERNAL_HEADERS_ENV): + resolve_strip_internal_headers_mode("maybe") + + +def test_strip_internal_headers_removes_headroom_headers_case_insensitively() -> None: + headers = { + "Authorization": "Bearer token", + "x-headroom-bypass": "true", + "X-Headroom-User-Id": "user-1", + "content-type": "application/json", + } + + stripped = strip_internal_headers(headers, mode="enabled") + + assert stripped == { + "Authorization": "Bearer token", + "content-type": "application/json", + } + assert "x-headroom-bypass" in headers + + +def test_strip_internal_headers_disabled_returns_copy_unchanged() -> None: + headers = {"x-headroom-mode": "passthrough", "content-type": "application/json"} + + copied = strip_internal_headers(headers, mode="disabled") + + assert copied == headers + assert copied is not headers