mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge pull request #32026 from BerriAI/litellm_headroom_unreachable_fallback
feat(guardrails): add unreachable_fallback fail-open option to headroom guardrail
This commit is contained in:
commit
55e3296e0c
5 changed files with 393 additions and 55 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -8,6 +8,8 @@ from typing import TYPE_CHECKING, Any, Literal, Optional
|
|||
|
||||
import httpx
|
||||
from fastapi import HTTPException
|
||||
|
||||
import litellm
|
||||
from httpx import Response as HttpxResponse
|
||||
from typing_extensions import TypeGuard
|
||||
|
||||
|
|
@ -214,6 +216,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 +226,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,11 +263,26 @@ 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]],
|
||||
model: str | None,
|
||||
) -> list[dict[str, object]]:
|
||||
) -> tuple[list[dict[str, object]], bool]:
|
||||
payload: dict[str, object] = {"messages": messages}
|
||||
if model:
|
||||
payload["model"] = model
|
||||
|
|
@ -272,69 +293,63 @@ class HeadroomGuardrail(CustomGuardrail):
|
|||
json=payload,
|
||||
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
|
||||
except httpx.HTTPStatusError as e:
|
||||
return self._handle_compress_failure(
|
||||
messages,
|
||||
"Headroom compression service returned an error",
|
||||
{"status_code": e.response.status_code, "body": e.response.text},
|
||||
), False
|
||||
except (httpx.ConnectError, httpx.TimeoutException, httpx.TransportError, litellm.Timeout) as e:
|
||||
return self._handle_compress_failure(
|
||||
messages,
|
||||
"Headroom compression service unreachable",
|
||||
{"detail": str(e)},
|
||||
), False
|
||||
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",
|
||||
{},
|
||||
), False
|
||||
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},
|
||||
), False
|
||||
|
||||
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]},
|
||||
), False
|
||||
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]},
|
||||
), False
|
||||
|
||||
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},
|
||||
), False
|
||||
|
||||
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},
|
||||
), False
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
"Headroom: compressed %s tokens -> %s tokens (ratio %.2f)",
|
||||
|
|
@ -342,7 +357,7 @@ class HeadroomGuardrail(CustomGuardrail):
|
|||
body.get("tokens_after", "?"),
|
||||
body.get("compression_ratio", 0),
|
||||
)
|
||||
return filtered
|
||||
return filtered, True
|
||||
|
||||
async def _call_retrieve(self, hash_value: str, query: str | None = None) -> str:
|
||||
params: dict[str, str] = {}
|
||||
|
|
@ -355,7 +370,7 @@ class HeadroomGuardrail(CustomGuardrail):
|
|||
params=params,
|
||||
headers=self._request_headers(),
|
||||
)
|
||||
except (httpx.ConnectError, httpx.TimeoutException, httpx.TransportError) as e:
|
||||
except (httpx.ConnectError, httpx.TimeoutException, httpx.TransportError, litellm.Timeout) as e:
|
||||
verbose_proxy_logger.warning("Headroom: retrieve failed for hash=%s: %s", hash_value, e)
|
||||
return f"[Headroom: retrieval failed for hash={hash_value}]"
|
||||
|
||||
|
|
@ -406,11 +421,14 @@ class HeadroomGuardrail(CustomGuardrail):
|
|||
return inputs
|
||||
|
||||
model = self.headroom_model or request_data.get("model")
|
||||
compressed = await self._call_compress(
|
||||
compressed, compression_succeeded = await self._call_compress(
|
||||
messages=messages,
|
||||
model=model if isinstance(model, str) else None,
|
||||
)
|
||||
|
||||
if not compression_succeeded:
|
||||
return {**inputs, "structured_messages": compressed} # pyright: ignore[reportReturnType]
|
||||
|
||||
hashes = extract_hashes_from_messages(compressed)
|
||||
if not hashes:
|
||||
return {**inputs, "structured_messages": compressed} # pyright: ignore[reportReturnType]
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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: 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:
|
||||
|
|
|
|||
|
|
@ -6,8 +6,11 @@ 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
|
||||
- /v1/compress non-2xx surfaces as httpx.HTTPStatusError (raise_for_status),
|
||||
not a status_code check on the returned response -- both are handled
|
||||
- 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
|
||||
|
|
@ -21,6 +24,8 @@ import httpx
|
|||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
import litellm
|
||||
|
||||
from litellm.proxy.guardrails.guardrail_hooks.headroom.headroom import (
|
||||
HeadroomGuardrail,
|
||||
extract_hashes_from_messages,
|
||||
|
|
@ -806,6 +811,244 @@ async def test_apply_guardrail_transport_error_raises():
|
|||
assert "unreachable" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
def _make_http_status_error(status: int, body: str) -> httpx.HTTPStatusError:
|
||||
request = httpx.Request("POST", f"{FAKE_API_BASE}/v1/compress")
|
||||
response = httpx.Response(status, request=request, text=body)
|
||||
return httpx.HTTPStatusError(
|
||||
f"Server error '{status}' for url",
|
||||
request=request,
|
||||
response=response,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_http_status_error_raises():
|
||||
"""Regression test: litellm's async httpx client calls raise_for_status()
|
||||
internally, so a non-2xx /v1/compress response surfaces as
|
||||
httpx.HTTPStatusError, not as a returned MagicMock with status_code set.
|
||||
A prior version of _call_compress only checked response.status_code and
|
||||
never caught this exception, so it went unhandled instead of blocking
|
||||
the request per fail_closed policy."""
|
||||
guardrail = _make_guardrail()
|
||||
|
||||
inputs = GenericGuardrailAPIInputs(
|
||||
texts=["hello"],
|
||||
structured_messages=ORIGINAL_MESSAGES,
|
||||
)
|
||||
|
||||
with patch.object(
|
||||
guardrail.async_handler,
|
||||
"post",
|
||||
new_callable=AsyncMock,
|
||||
side_effect=_make_http_status_error(500, "headroom internal error"),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data={},
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 502
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_http_status_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=_make_http_status_error(500, "headroom internal error"),
|
||||
):
|
||||
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_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_non_json_response_fail_open_forwards_uncompressed():
|
||||
guardrail = _make_guardrail(unreachable_fallback="fail_open")
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.side_effect = ValueError("not JSON")
|
||||
mock_response.text = "<!DOCTYPE html><html>not json</html>"
|
||||
|
||||
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_fail_open_forwards_uncompressed():
|
||||
guardrail = _make_guardrail(unreachable_fallback="fail_open")
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {"tokens_before": 100, "tokens_after": 10}
|
||||
mock_response.text = "{}"
|
||||
|
||||
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_empty_compressed_messages_fail_open_forwards_uncompressed():
|
||||
guardrail = _make_guardrail(unreachable_fallback="fail_open")
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {
|
||||
"messages": ["not-a-dict", 42, None],
|
||||
"tokens_before": 1000,
|
||||
"tokens_after": 0,
|
||||
"compression_ratio": 0,
|
||||
}
|
||||
mock_response.text = "{}"
|
||||
|
||||
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_fail_open_does_not_register_hashes_from_original_messages():
|
||||
"""When compression fails with fail_open, user-supplied messages that
|
||||
happen to contain hash-shaped strings must NOT cause those hashes to be
|
||||
registered as valid for CCR retrieval. Otherwise an attacker can plant a
|
||||
hash= string in their prompt, trigger a compression failure, and have
|
||||
that hash honored by a later headroom_retrieve tool call."""
|
||||
messages_with_fake_hash = [
|
||||
{"role": "user", "content": "Please fetch hash=deadbeef000000000000dead for me"},
|
||||
]
|
||||
guardrail = _make_guardrail(unreachable_fallback="fail_open")
|
||||
|
||||
inputs = GenericGuardrailAPIInputs(
|
||||
texts=["hello"],
|
||||
structured_messages=messages_with_fake_hash,
|
||||
)
|
||||
|
||||
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"] == messages_with_fake_hash
|
||||
assert not has_headroom_retrieve_tool(result.get("tools") or [])
|
||||
assert not guardrail._issued_hashes_by_call_id
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_missing_messages_key_raises():
|
||||
guardrail = _make_guardrail()
|
||||
|
|
@ -875,6 +1118,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()
|
||||
|
||||
|
|
@ -1059,3 +1312,61 @@ async def test_async_should_run_agentic_loop_detects_responses_api_output_format
|
|||
assert should_run is True
|
||||
assert len(ctx["tool_calls"]) == 1
|
||||
assert ctx["tool_calls"][0]["arguments"]["hash"] == "b573993006976af767214fac"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_litellm_timeout_raises_when_fail_closed():
|
||||
guardrail = _make_guardrail()
|
||||
|
||||
inputs = GenericGuardrailAPIInputs(
|
||||
texts=["hello"],
|
||||
structured_messages=ORIGINAL_MESSAGES,
|
||||
)
|
||||
|
||||
with patch.object(
|
||||
guardrail.async_handler,
|
||||
"post",
|
||||
new_callable=AsyncMock,
|
||||
side_effect=litellm.Timeout(
|
||||
message="Connection timed out after 10 seconds.",
|
||||
model="default-model-name",
|
||||
llm_provider="litellm-httpx-handler",
|
||||
),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data={},
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 502
|
||||
assert "unreachable" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_litellm_timeout_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=litellm.Timeout(
|
||||
message="Connection timed out after 10 seconds.",
|
||||
model="default-model-name",
|
||||
llm_provider="litellm-httpx-handler",
|
||||
),
|
||||
):
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data={},
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
assert result["structured_messages"] == ORIGINAL_MESSAGES
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue