diff --git a/litellm/proxy/guardrails/guardrail_hooks/headroom/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/headroom/__init__.py index 9b7934b7705..d2b8a979261 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/headroom/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/headroom/__init__.py @@ -34,6 +34,7 @@ def initialize_guardrail(litellm_params: LitellmParams, guardrail: Guardrail) -> guardrail_name=guardrail["guardrail_name"], event_hook=_coerce_event_hook(litellm_params.mode), default_on=litellm_params.default_on or False, + unreachable_fallback=litellm_params.unreachable_fallback, ) litellm.logging_callback_manager.add_litellm_callback( # pyright: ignore[reportUnknownMemberType] _callback diff --git a/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py b/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py index 4badb48e2eb..575b4ca55dd 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py +++ b/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py @@ -214,6 +214,7 @@ class HeadroomGuardrail(CustomGuardrail): guardrail_name: str | None = None, event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None = None, default_on: bool = False, + unreachable_fallback: str | None = None, ): self.headroom_api_base = (api_base or get_secret_str("HEADROOM_API_BASE") or "").rstrip("/") if not self.headroom_api_base: @@ -223,6 +224,9 @@ class HeadroomGuardrail(CustomGuardrail): ) self.headroom_api_key = api_key or get_secret_str("HEADROOM_API_KEY") self.headroom_model = model + self.unreachable_fallback: Literal["fail_closed", "fail_open"] = ( + "fail_open" if unreachable_fallback == "fail_open" else "fail_closed" + ) self.async_handler = get_async_httpx_client( llm_provider=httpxSpecialProvider.GuardrailCallback, ) @@ -257,6 +261,21 @@ class HeadroomGuardrail(CustomGuardrail): if expiry > now } + def _handle_compress_failure( + self, + messages: list[dict[str, object]], + error: str, + detail: dict[str, object], + ) -> list[dict[str, object]]: + if self.unreachable_fallback == "fail_open": + verbose_proxy_logger.critical( + "Headroom: %s; fail_open configured, forwarding request uncompressed. detail=%s", + error, + detail, + ) + return messages + raise HTTPException(status_code=502, detail={"error": error, **detail}) + async def _call_compress( self, messages: list[dict[str, object]], @@ -273,67 +292,55 @@ class HeadroomGuardrail(CustomGuardrail): headers=self._request_headers(), ) except (httpx.ConnectError, httpx.TimeoutException, httpx.TransportError) as e: - raise HTTPException( - status_code=502, - detail={ - "error": "Headroom compression service unreachable", - "detail": str(e), - }, - ) from e + return self._handle_compress_failure( + messages, + "Headroom compression service unreachable", + {"detail": str(e)}, + ) if raw_response is None: - raise HTTPException( - status_code=502, - detail={"error": "Headroom compression service returned no response"}, + return self._handle_compress_failure( + messages, + "Headroom compression service returned no response", + {}, ) response: HttpxResponse = raw_response if response.status_code != 200: - raise HTTPException( - status_code=502, - detail={ - "error": "Headroom compression service returned an error", - "status_code": response.status_code, - "body": response.text, - }, + return self._handle_compress_failure( + messages, + "Headroom compression service returned an error", + {"status_code": response.status_code, "body": response.text}, ) try: body: object = response.json() except ValueError: - raise HTTPException( - status_code=502, - detail={ - "error": "Headroom compression service returned non-JSON response", - "body": response.text[:500], - }, + return self._handle_compress_failure( + messages, + "Headroom compression service returned non-JSON response", + {"body": response.text[:500]}, ) if not _is_str_object_dict(body): - raise HTTPException( - status_code=502, - detail={ - "error": "Headroom compression service returned unexpected response shape", - "body": response.text[:500], - }, + return self._handle_compress_failure( + messages, + "Headroom compression service returned unexpected response shape", + {"body": response.text[:500]}, ) compressed_messages = body.get("messages") if not _is_object_list(compressed_messages): - raise HTTPException( - status_code=502, - detail={ - "error": "Headroom compression service response missing 'messages'", - "body": response.text, - }, + return self._handle_compress_failure( + messages, + "Headroom compression service response missing 'messages'", + {"body": response.text}, ) filtered = [item for item in compressed_messages if _is_str_object_dict(item)] if not filtered: - raise HTTPException( - status_code=502, - detail={ - "error": "Headroom compression service returned empty message list", - "body": response.text, - }, + return self._handle_compress_failure( + messages, + "Headroom compression service returned empty message list", + {"body": response.text}, ) verbose_proxy_logger.debug( diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index 889e029b902..8d7d7311fad 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -697,7 +697,7 @@ class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch up default="fail_closed", description=( "Behavior when a guardrail endpoint is unreachable due to network errors. " - "Implemented by guardrail='generic_guardrail_api', 'akto', 'vigil_guard', and 'repelloai'. " + "Implemented by guardrail='generic_guardrail_api', 'akto', 'vigil_guard', 'repelloai', and 'headroom'. " "'fail_closed' raises an error (default). 'fail_open' logs a critical error and allows the request to proceed." ), ) diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/headroom.py b/litellm/types/proxy/guardrails/guardrail_hooks/headroom.py index 3186c9fc612..fd962b8ffa1 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/headroom.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/headroom.py @@ -1,4 +1,4 @@ -from typing import Optional +from typing import Literal, Optional from pydantic import BaseModel, Field @@ -18,6 +18,14 @@ class HeadroomGuardrailConfigModel(GuardrailConfigModel[BaseModel]): default=None, description="Model name forwarded to the headroom /v1/compress endpoint.", ) + unreachable_fallback: Optional[Literal["fail_closed", "fail_open"]] = Field( + default="fail_closed", + description=( + "Behavior when the headroom compression service is unreachable or errors. " + "'fail_closed' raises an error (default). 'fail_open' logs a critical error and " + "forwards the request uncompressed instead of blocking it." + ), + ) @staticmethod def ui_friendly_name() -> str: diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_headroom.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_headroom.py index f5ce6cedf64..954485d6b90 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_headroom.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_headroom.py @@ -6,8 +6,9 @@ Tests cover: - x-headroom-bypass: true header causes guardrail to skip compression - missing or empty messages are passed through unchanged - response-type input is passed through unchanged -- /v1/compress HTTP error raises HTTPException +- /v1/compress HTTP error raises HTTPException (fail_closed, the default) - /v1/compress returning malformed JSON raises HTTPException +- unreachable_fallback="fail_open" forwards the request uncompressed instead of raising - CCR: headroom_retrieve tool injected when compressed messages contain hashes - CCR: async_should_run_agentic_loop returns True when response has headroom_retrieve tool calls - CCR: async_build_agentic_loop_plan calls retrieve endpoint and builds follow-up messages @@ -806,6 +807,56 @@ async def test_apply_guardrail_transport_error_raises(): assert "unreachable" in str(exc_info.value.detail) +@pytest.mark.asyncio +async def test_apply_guardrail_transport_error_fail_open_forwards_uncompressed(): + guardrail = _make_guardrail(unreachable_fallback="fail_open") + + inputs = GenericGuardrailAPIInputs( + texts=["hello"], + structured_messages=ORIGINAL_MESSAGES, + ) + + with patch.object( + guardrail.async_handler, + "post", + new_callable=AsyncMock, + side_effect=httpx.ConnectError("Connection refused"), + ): + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data={}, + input_type="request", + ) + + assert result["structured_messages"] == ORIGINAL_MESSAGES + + +@pytest.mark.asyncio +async def test_apply_guardrail_http_error_fail_open_forwards_uncompressed(): + guardrail = _make_guardrail(unreachable_fallback="fail_open") + mock_response = _make_compress_response([], status=500) + mock_response.text = "Internal Server Error" + + inputs = GenericGuardrailAPIInputs( + texts=["hello"], + structured_messages=ORIGINAL_MESSAGES, + ) + + with patch.object( + guardrail.async_handler, + "post", + new_callable=AsyncMock, + return_value=mock_response, + ): + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data={}, + input_type="request", + ) + + assert result["structured_messages"] == ORIGINAL_MESSAGES + + @pytest.mark.asyncio async def test_apply_guardrail_missing_messages_key_raises(): guardrail = _make_guardrail() @@ -875,6 +926,16 @@ def test_init_raises_without_api_base(): HeadroomGuardrail(api_base=None) +def test_init_defaults_to_fail_closed(): + guardrail = _make_guardrail() + assert guardrail.unreachable_fallback == "fail_closed" + + +def test_init_rejects_invalid_unreachable_fallback_value(): + guardrail = _make_guardrail(unreachable_fallback="not-a-real-mode") + assert guardrail.unreachable_fallback == "fail_closed" + + def test_bypass_header_case_insensitive(): guardrail = _make_guardrail()