feat(guardrails): add unreachable_fallback fail-open option to headroom guardrail

Reuses the existing unreachable_fallback flag (already implemented by
generic_guardrail_api, akto, vigil_guard, repelloai) so headroom compression
failures can forward the request uncompressed instead of blocking it with a
502.
This commit is contained in:
Krrish Dholakia 2026-07-02 22:00:24 -07:00
parent dfbbda4f19
commit 659127bd0d
5 changed files with 121 additions and 44 deletions

View file

@ -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

View file

@ -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(

View file

@ -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."
),
)

View file

@ -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:

View file

@ -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()