feat(guardrails): honor litellm_params.timeout in every HTTP guardrail (#43134)

* feat(guardrails): honor litellm_params.timeout in every HTTP guardrail

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(guardrails): accept timeout kwarg in presidio and responses-handler post stubs

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(guardrails): bound hiddenlayer startup jwt call by configured timeout, drop akto from timeout coverage

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* refactor(guardrails): narrow hiddenlayer startup auth timeout without cast

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(guardrails): bound hiddenlayer jwt refresh by configured timeout

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(guardrails): keep provider timeout defaults when unset and bound only rubrik moderation calls

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(guardrails): cover model_armor and run timeout probes concurrently

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(guardrails): match sink calls to the exact guardrail name

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: kerry <kerry@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-30 21:18:47 -07:00 • committed by GitHub
parent ae05f7d2c1
commit a308a8e579
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
83 changed files with 688 additions and 62 deletions

View file

@ -8,6 +8,8 @@ from datetime import datetime
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, Optional, get_args
import httpx
from litellm._logging import verbose_logger
from litellm.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger
@ -176,6 +178,8 @@ class CustomGuardrail(CustomLogger):
records_own_guardrail_information: ClassVar[bool] = False
timeout: float | httpx.Timeout | None = None
def __init_subclass__(cls, **kwargs: object) -> None: # kwargs-ok: forwarded to cooperative __init_subclass__ hooks
super().__init_subclass__(**kwargs)
own_apply_guardrail: Final[object] = cls.__dict__.get("apply_guardrail")
@ -201,6 +205,7 @@ class CustomGuardrail(CustomLogger):
run_in_parallel: bool = False,
scan_raw_request: bool = False,
only_scan_new_messages: bool = False,
timeout: float | None = None,
**kwargs,
):
"""
@ -229,6 +234,8 @@ class CustomGuardrail(CustomLogger):
guardrails: any data this guardrail returns is discarded, matching run_in_parallel's
contract, since applying its mutations on top of a stale snapshot would silently
undo whatever later guardrails already did to the live request.
timeout: Per-request timeout in seconds for the guardrail provider's API call. When
None, the guardrail keeps whatever default its HTTP handler or SDK already uses.
"""
self.guardrail_name = guardrail_name
self.supported_event_hooks = supported_event_hooks
@ -246,6 +253,8 @@ class CustomGuardrail(CustomLogger):
self.run_in_parallel: bool = run_in_parallel
self.scan_raw_request: bool = scan_raw_request
self.only_scan_new_messages: bool = only_scan_new_messages
if timeout is not None:
self.timeout = timeout
if supported_event_hooks:
## validate event_hook is in supported_event_hooks

View file

@ -1120,6 +1120,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
endpoint,
json=dict(payload),
headers=dict(self._headers),
timeout=self.timeout,
)
http_response.raise_for_status()
result: Final[_ModerationResponse | None] = http_response.json()

View file

@ -19,6 +19,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
inspect_embeddings=litellm_params.inspect_embeddings,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_aim_callback)

View file

@ -181,6 +181,7 @@ class AimGuardrail(CustomGuardrail):
f"{self.api_base}/fw/v1/analyze",
headers=headers,
json={"messages": self._build_aim_inspection_messages(data)},
timeout=self.timeout,
)
response.raise_for_status()
res: Final[AimAnalyzeResponse] = response.json()
@ -285,6 +286,7 @@ class AimGuardrail(CustomGuardrail):
"messages": self._build_aim_inspection_messages(request_data)
+ [{"role": "assistant", "content": output}]
},
timeout=self.timeout,
)
response.raise_for_status()
res: Final[AimAnalyzeResponse] = response.json()

View file

@ -18,6 +18,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_alice_guardrail_callback)

View file

@ -227,6 +227,7 @@ class AliceGuardrail(CustomGuardrail):
"Content-Type": "application/json",
"af-api-key": self.alice_api_key,
},
timeout=self.timeout,
)
response.raise_for_status()
body = response.json()

View file

@ -17,6 +17,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_aporia_callback)

View file

@ -123,6 +123,7 @@ class AporiaGuardrail(CustomGuardrail):
"X-APORIA-API-KEY": self.aporia_api_key,
"Content-Type": "application/json",
},
timeout=self.timeout,
)
verbose_proxy_logger.debug("Aporia AI response: %s", response.text)
if response.status_code == 200:

View file

@ -1,6 +1,8 @@
import re
from typing import TYPE_CHECKING, Any, Final
import httpx
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.prompt_templates.common_utils import (
get_last_user_message,
@ -49,6 +51,7 @@ class AzureGuardrailBase:
# (typically CustomGuardrail).
super().__init__(**kwargs)
self.timeout: float | httpx.Timeout | None
self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
self.api_key = api_key
self.api_base = api_base
@ -77,6 +80,7 @@ class AzureGuardrailBase:
url=url,
headers=headers,
json=request_body,
timeout=self.timeout,
)
response_json: Final[dict[str, Any]] = response.json()
verbose_proxy_logger.debug("Azure Content Safety response [%s]: %s", endpoint_path, response_json)

View file

@ -1787,6 +1787,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
url=prepared_request.url,
data=prepared_request.body,
headers=prepared_request.headers,
timeout=self.timeout,
)
except HTTPException:
# Propagate HTTPException (e.g. from non-200 path) as-is

View file

@ -22,6 +22,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
default_on=litellm_params.default_on,
inspect_embeddings=litellm_params.inspect_embeddings,
ssl_verify=getattr(litellm_params, "ssl_verify", None),
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_cato_callback)

View file

@ -305,6 +305,7 @@ class CatoNetworksGuardrail(CustomGuardrail):
f"{self.api_base}/fw/v1/analyze",
headers=headers,
json={"messages": self._inspection_messages(data)},
timeout=self.timeout,
)
response.raise_for_status()
res: Final[_CatoAnalyzeResponse] = response.json()
@ -445,6 +446,7 @@ class CatoNetworksGuardrail(CustomGuardrail):
litellm_call_id=call_id,
),
json={"messages": inspection_messages + [{"role": "assistant", "content": output}]},
timeout=self.timeout,
)
response.raise_for_status()
res: Final[_CatoAnalyzeResponse] = response.json()

View file

@ -214,8 +214,6 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail):
else:
env_timeout: Final = os.environ.get("CISCO_AI_DEFENSE_TIMEOUT")
resolved_timeout = self._coerce_timeout(env_timeout) if env_timeout is not None else None
self.timeout: float = resolved_timeout if resolved_timeout is not None else DEFAULT_TIMEOUT_SECONDS
self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
# Register broadly; runtime filtering happens in ``_surface_matches``.
@ -224,6 +222,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail):
supported_event_hooks=list(self.get_supported_event_hooks()),
**kwargs,
)
self.timeout = resolved_timeout if resolved_timeout is not None else DEFAULT_TIMEOUT_SECONDS
self._warn_if_mode_surface_mismatch(kwargs.get("event_hook"))

View file

@ -59,6 +59,7 @@ def initialize_guardrail(litellm_params: LitellmParams, guardrail: Guardrail) ->
event_hook=_coerce_event_hook(litellm_params.mode),
default_on=litellm_params.default_on or False,
unreachable_fallback=litellm_params.unreachable_fallback,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback( # pyright: ignore[reportUnknownMemberType] # callback manager is untyped
_callback

View file

@ -520,6 +520,7 @@ class CompresrGuardrail(CustomGuardrail):
dynamic_min_ratio: float | None = None,
dynamic_max_ratio: float | None = None,
compression_params: dict[str, object] | None = None,
timeout: float | None = None,
):
raw_api_base: Final = (api_base or get_secret_str("COMPRESR_API_BASE") or DEFAULT_API_BASE).rstrip("/")
self.compresr_api_base = _validate_api_base(raw_api_base)
@ -583,6 +584,7 @@ class CompresrGuardrail(CustomGuardrail):
guardrail_name=guardrail_name,
event_hook=event_hook,
default_on=default_on,
timeout=timeout,
)
def _should_bypass(self, request_data: dict) -> bool:
@ -755,7 +757,7 @@ class CompresrGuardrail(CustomGuardrail):
url=url,
json=payload,
headers=self._request_headers(),
timeout=_COMPRESS_TIMEOUT_SECONDS,
timeout=self.timeout if self.timeout is not None else _COMPRESS_TIMEOUT_SECONDS,
)
except asyncio.CancelledError:
raise

View file

@ -27,6 +27,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
streaming_buffer_release_on_scan=streaming_params.streaming_buffer_release_on_scan,
streaming_end_of_stream_only=streaming_params.streaming_end_of_stream_only,
streaming_sampling_rate=streaming_params.streaming_sampling_rate,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_crowdstrike_aidr_callback)

View file

@ -355,7 +355,9 @@ class CrowdStrikeAIDRHandler(CustomGuardrail):
"CrowdStrike AIDR Guardrail (%s): Calling endpoint %s with payload: %s", hook_name, endpoint, payload
)
response: Final = await self.async_handler.post(url=endpoint, json=payload, headers=headers)
response: Final = await self.async_handler.post(
url=endpoint, json=payload, headers=headers, timeout=self.timeout
)
assert response is not None
response.raise_for_status()

View file

@ -20,6 +20,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_deepkeep_guardrail_callback)

View file

@ -393,6 +393,7 @@ class DeepKeepGuardrail(CustomGuardrail):
url=self.api_base,
json=guardrail_request,
headers=headers,
timeout=self.timeout,
)
response.raise_for_status()

View file

@ -17,6 +17,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_dynamoai_callback)

View file

@ -130,6 +130,7 @@ class DynamoAIGuardrails(CustomGuardrail):
url=self.api_url,
json=dict(payload),
headers=headers,
timeout=self.timeout,
)
response.raise_for_status()
response_json: Final = response.json()

View file

@ -24,6 +24,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
block_on_violation=litellm_params.block_on_violation,
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_enkryptai_callback)

View file

@ -123,6 +123,7 @@ class EnkryptAIGuardrails(CustomGuardrail):
url=self.api_url,
json=payload,
headers=headers,
timeout=self.timeout,
)
response.raise_for_status()
response_json: Final = response.json()

View file

@ -39,6 +39,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
streaming_end_of_stream_only=_get_config_value(litellm_params, optional_params, "streaming_end_of_stream_only"),
streaming_sampling_rate=_get_config_value(litellm_params, optional_params, "streaming_sampling_rate"),
streaming_transform_mode=_get_config_value(litellm_params, optional_params, "streaming_transform_mode"),
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_generic_guardrail_api_callback)

View file

@ -477,6 +477,7 @@ class GenericGuardrailAPI(CustomGuardrail):
url=self.api_base,
json=guardrail_request.model_dump(mode="json"),
headers=headers,
timeout=self.timeout,
)
response.raise_for_status()

View file

@ -24,6 +24,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
default_on=litellm_params.default_on,
guard_name=litellm_params.guard_name,
guardrails_ai_api_input_format=getattr(litellm_params, "guardrails_ai_api_input_format", "llmOutput"),
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_guardrails_ai_callback)

View file

@ -80,6 +80,7 @@ class GuardrailsAI(CustomGuardrail):
headers={
"Content-Type": "application/json",
},
timeout=self.timeout,
)
verbose_proxy_logger.debug("guardrails_ai response: %s", response)
_json_response: Final = GuardrailsAIResponse(**response.json())
@ -117,6 +118,7 @@ class GuardrailsAI(CustomGuardrail):
headers={
"Content-Type": "application/json",
},
timeout=self.timeout,
)
verbose_proxy_logger.debug("guardrails_ai response: %s", response)
if response.status_code == 400:

View file

@ -508,7 +508,6 @@ class HeadroomGuardrail(CustomGuardrail):
self.unreachable_fallback: Literal["fail_closed", "fail_open"] = (
"fail_open" if unreachable_fallback == "fail_open" else "fail_closed"
)
self.timeout: httpx.Timeout = self._resolve_timeout(timeout)
self.ccr_retrieval = ccr_retrieval
self.async_handler = get_async_httpx_client(
llm_provider=httpxSpecialProvider.GuardrailCallback,
@ -520,6 +519,7 @@ class HeadroomGuardrail(CustomGuardrail):
default_on=default_on,
supported_event_hooks=list(self.get_supported_event_hooks()),
)
self.timeout = self._resolve_timeout(timeout)
def _should_bypass(self, request_data: dict) -> bool:
psr: Final = request_data.get("proxy_server_request")

View file

@ -25,6 +25,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
else:
_hiddenlayer_callback = HiddenlayerGuardrailV2(
@ -35,6 +36,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_hiddenlayer_callback)

View file

@ -243,15 +243,19 @@ class HiddenlayerGuardrail(CustomGuardrail):
if not self.hiddenlayer_client_secret:
raise RuntimeError("`api_key` cannot be None when using the SaaS version of HiddenLayer.")
ctor_timeout: Final = kwargs.get("timeout")
auth_timeout: Final = ctor_timeout if isinstance(ctor_timeout, (int, float)) else _AUTH_TIMEOUT_SECONDS
self.jwt_token = _get_jwt(
auth_url=auth_url,
api_id=self.hiddenlayer_client_id,
api_key=self.hiddenlayer_client_secret,
timeout=auth_timeout,
)
self.refresh_jwt_func = lambda: _get_jwt(
auth_url=auth_url,
api_id=self.hiddenlayer_client_id,
api_key=self.hiddenlayer_client_secret,
timeout=auth_timeout,
)
self._http_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
@ -382,6 +386,7 @@ class HiddenlayerGuardrail(CustomGuardrail):
f"{self.api_base}/detection/v1/interactions",
json=data,
headers=headers,
timeout=self.timeout,
)
response.raise_for_status()
result: _HiddenlayerResponse = _interaction_body(response)
@ -403,6 +408,7 @@ class HiddenlayerGuardrail(CustomGuardrail):
f"{self.api_base}/detection/v1/interactions",
json=data,
headers=headers,
timeout=self.timeout,
)
else:
raise e
@ -447,15 +453,19 @@ class HiddenlayerGuardrailV2(CustomGuardrail):
if not self.hiddenlayer_client_secret:
raise RuntimeError("`api_key` cannot be None when using the SaaS version of HiddenLayer.")
ctor_timeout: Final = kwargs.get("timeout")
auth_timeout: Final = ctor_timeout if isinstance(ctor_timeout, (int, float)) else _AUTH_TIMEOUT_SECONDS
self.jwt_token = _get_jwt(
auth_url=auth_url,
api_id=self.hiddenlayer_client_id,
api_key=self.hiddenlayer_client_secret,
timeout=auth_timeout,
)
self.refresh_jwt_func = lambda: _get_jwt(
auth_url=auth_url,
api_id=self.hiddenlayer_client_id,
api_key=self.hiddenlayer_client_secret,
timeout=auth_timeout,
)
self._http_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
@ -584,6 +594,7 @@ class HiddenlayerGuardrailV2(CustomGuardrail):
f"{self.api_base}/{path}",
json=payload,
headers=headers,
timeout=self.timeout,
)
response.raise_for_status()
@ -604,6 +615,7 @@ class HiddenlayerGuardrailV2(CustomGuardrail):
f"{self.api_base}/{path}",
json=payload,
headers=headers,
timeout=self.timeout,
)
else:
raise e

View file

@ -49,6 +49,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
verify_ssl=verify_ssl,
default_on=litellm_params.default_on,
event_hook=litellm_params.mode,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(ibm_guardrail)

View file

@ -140,6 +140,7 @@ class IBMGuardrailDetector(CustomGuardrail):
url=self.api_url,
json=payload,
headers=headers,
timeout=self.timeout,
)
response.raise_for_status()
response_json: Final[list[list[IBMDetectorDetection]]] = response.json()
@ -231,6 +232,7 @@ class IBMGuardrailDetector(CustomGuardrail):
url=self.api_url,
json=payload,
headers=headers,
timeout=self.timeout,
)
response.raise_for_status()
response_json: Final[IBMDetectorResponseOrchestrator] = response.json()

View file

@ -27,6 +27,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
config=litellm_params.config,
metadata=litellm_params.metadata,
application=litellm_params.application,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_javelin_callback)

View file

@ -111,6 +111,7 @@ class JavelinGuardrail(CustomGuardrail):
url=url,
headers=headers,
json=dict(request),
timeout=self.timeout,
)
verbose_proxy_logger.debug("Javelin Guardrail: Javelin guard API response: %s", response.json())
response_data: Final = response.json()

View file

@ -250,6 +250,7 @@ class lakeraAI_Moderation(CustomGuardrail):
"Authorization": "Bearer " + self.lakera_api_key,
"Content-Type": "application/json",
},
timeout=self.timeout,
)
except httpx.HTTPStatusError as e:
raise Exception(e.response.text)

View file

@ -402,6 +402,7 @@ class LakeraAIGuardrail(CustomGuardrail):
url=f"{self.api_base}/v2/guard",
headers={"Authorization": f"Bearer {self.lakera_api_key}"},
json=request,
timeout=self.timeout,
)
verbose_proxy_logger.debug("Lakera AI v2 guard response: %s", response.json())
lakera_response = LakeraAIResponse(**response.json())

View file

@ -19,6 +19,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
conversation_id=litellm_params.lasso_conversation_id,
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_lasso_callback)

View file

@ -814,7 +814,7 @@ class LassoGuardrail(CustomGuardrail):
url=url,
headers=headers,
json=payload,
timeout=10.0,
timeout=self.timeout if self.timeout is not None else 10.0,
)
response.raise_for_status()
return response.json()

View file

@ -62,6 +62,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
debug_headers=_get("debug_headers") or False,
# FR-10: configurable scopes
allowed_scopes=_get("allowed_scopes"),
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(signer)
return signer

View file

@ -76,6 +76,7 @@ import time
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, Optional
import httpx
import jwt
from cryptography.hazmat.primitives import serialization
from cryptography.hazmat.primitives.asymmetric import rsa
@ -173,7 +174,7 @@ def _compute_kid(public_key: RSAPublicKey) -> str:
return hashlib.sha256(der_bytes).hexdigest()[:16]
async def _fetch_jwks(jwks_uri: str) -> Sequence[Mapping[str, object]]:
async def _fetch_jwks(jwks_uri: str, timeout: float | httpx.Timeout | None = None) -> Sequence[Mapping[str, object]]:
"""
Fetch and cache a JWKS from the given URI.
@ -192,7 +193,7 @@ async def _fetch_jwks(jwks_uri: str) -> Sequence[Mapping[str, object]]:
)
client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
resp: Final = await client.get(jwks_uri, headers={"Accept": "application/json"})
resp: Final = await client.get(jwks_uri, headers={"Accept": "application/json"}, timeout=timeout)
resp.raise_for_status()
jwks_body: Final[Mapping[str, Sequence[Mapping[str, object]]]] = resp.json()
fetched_keys: Final = jwks_body.get("keys", [])
@ -200,7 +201,9 @@ async def _fetch_jwks(jwks_uri: str) -> Sequence[Mapping[str, object]]:
return fetched_keys
async def _fetch_oidc_discovery(discovery_uri: str) -> _OIDCDiscoveryDocument:
async def _fetch_oidc_discovery(
discovery_uri: str, timeout: float | httpx.Timeout | None = None
) -> _OIDCDiscoveryDocument:
"""Fetch an OIDC discovery document and return its parsed JSON."""
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
@ -208,7 +211,7 @@ async def _fetch_oidc_discovery(discovery_uri: str) -> _OIDCDiscoveryDocument:
)
client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
resp: Final = await client.get(discovery_uri, headers={"Accept": "application/json"})
resp: Final = await client.get(discovery_uri, headers={"Accept": "application/json"}, timeout=timeout)
resp.raise_for_status()
document: Final[_OIDCDiscoveryDocument] = resp.json()
return document
@ -417,7 +420,7 @@ class MCPJWTSigner(CustomGuardrail):
now: Final = time.time()
cache_expired: Final = (now - self._oidc_discovery_fetched_at) >= self._OIDC_DISCOVERY_TTL
if (self._oidc_discovery_doc is None or cache_expired) and self.access_token_discovery_uri:
doc: Final = await _fetch_oidc_discovery(self.access_token_discovery_uri)
doc: Final = await _fetch_oidc_discovery(self.access_token_discovery_uri, timeout=self.timeout)
if "jwks_uri" in doc:
self._oidc_discovery_doc = doc
self._oidc_discovery_fetched_at = now
@ -440,7 +443,7 @@ class MCPJWTSigner(CustomGuardrail):
f"at {self.access_token_discovery_uri!r} has no 'jwks_uri'."
)
jwks_keys: Final = await _fetch_jwks(jwks_uri)
jwks_keys: Final = await _fetch_jwks(jwks_uri, timeout=self.timeout)
# Only read `kid` from the unverified header — never `alg`.
# Reading `alg` from an attacker-controlled header enables algorithm
@ -511,6 +514,7 @@ class MCPJWTSigner(CustomGuardrail):
self.token_introspection_endpoint,
data={"token": token},
headers={"Accept": "application/json"},
timeout=self.timeout,
)
resp.raise_for_status()
result: Final[dict[str, object]] = resp.json()

View file

@ -38,6 +38,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
user_id_field=str(getattr(litellm_params, "user_id_field", None) or "user_id"),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(purview_guardrail)

View file

@ -5,6 +5,7 @@ from collections import OrderedDict
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final
import httpx
from typing_extensions import NotRequired, TypedDict
from litellm._logging import verbose_proxy_logger
@ -56,6 +57,7 @@ class PurviewGuardrailBase:
# (typically CustomGuardrail).
super().__init__(**kwargs)
self.timeout: float | httpx.Timeout | None
self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
self.tenant_id = tenant_id
self.client_id = client_id
@ -107,6 +109,7 @@ class PurviewGuardrailBase:
url=url,
data=data,
headers={"Content-Type": "application/x-www-form-urlencoded"},
timeout=self.timeout,
)
response.raise_for_status()
token_data: Final[GraphTokenResponse] = response.json()
@ -143,7 +146,7 @@ class PurviewGuardrailBase:
headers.update(extra_headers)
verbose_proxy_logger.debug("Purview Graph POST %s", url)
response: Final = await self.async_handler.post(url=url, headers=headers, json=json_body)
response: Final = await self.async_handler.post(url=url, headers=headers, json=json_body, timeout=self.timeout)
response.raise_for_status()
response_json: Final[dict[str, object]] = response.json()
response_headers: Final = dict(response.headers)

View file

@ -28,6 +28,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
fail_on_error=litellm_params.fail_on_error,
skip_unscannable_attachments=litellm_params.skip_unscannable_attachments,
sanitize_error_detail=litellm_params.sanitize_error_detail,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_model_armor_callback)

View file

@ -337,6 +337,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
url=url,
json=body,
headers=headers,
timeout=self.timeout,
)
except httpx.HTTPStatusError as e:
detail = self._build_api_error_detail(e.response.status_code, e.response.text)

View file

@ -28,6 +28,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
anonymize_input=litellm_params.anonymize_input,
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_noma_callback)
@ -47,6 +48,7 @@ def initialize_guardrail_v2(litellm_params: "LitellmParams", guardrail: "Guardra
block_failures=litellm_params.block_failures,
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_noma_v2_callback)

View file

@ -751,6 +751,7 @@ class NomaGuardrail(CustomGuardrail):
"requestId": llm_request_id,
},
},
timeout=self.timeout,
)
response.raise_for_status()

View file

@ -220,6 +220,7 @@ class NomaV2Guardrail(CustomGuardrail):
url=endpoint,
headers=headers,
json=sanitized_payload,
timeout=self.timeout,
)
verbose_proxy_logger.debug(
"Noma v2 AIDR response: status_code=%s body=%s",

View file

@ -16,6 +16,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_onyx_callback)

View file

@ -116,6 +116,7 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail):
"Content-Type": "application/json",
},
json=request_body,
timeout=self.timeout,
)
verbose_proxy_logger.debug("OpenAI Moderation guard response: %s", response.json())

View file

@ -29,6 +29,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
post_checkpoint_id=post_checkpoint_id,
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_ovalix_callback)

View file

@ -194,7 +194,7 @@ class OvalixGuardrail(CustomGuardrail):
"data_type": "TEXT",
"data": {"content": content},
}
response: Final = await self._async_handler.post(url, headers=headers, json=payload)
response: Final = await self._async_handler.post(url, headers=headers, json=payload, timeout=self.timeout)
response.raise_for_status()
return response.json()

View file

@ -23,6 +23,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
api_key=litellm_params.api_key,
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_pangea_callback)

View file

@ -131,7 +131,9 @@ class PangeaHandler(CustomGuardrail):
"Pangea Guardrail (%s): Calling endpoint %s with payload: %s", hook_name, endpoint, payload
)
response: Final = await self.async_handler.post(url=endpoint, json=payload, headers=headers)
response: Final = await self.async_handler.post(
url=endpoint, json=payload, headers=headers, timeout=self.timeout
)
response.raise_for_status()
result: Final = response.json()

View file

@ -11,6 +11,8 @@ import os
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol
from urllib.parse import quote
import httpx
# Third-party imports
from fastapi import HTTPException
from typing_extensions import NotRequired, ReadOnly, TypedDict
@ -66,7 +68,7 @@ class _PillarProtectHTTPClient(Protocol):
url: str,
headers: dict[str, str],
json: dict[str, object],
timeout: float,
timeout: float | httpx.Timeout | None,
) -> _PillarProtectHTTPResponse: ...
@ -284,7 +286,12 @@ class PillarGuardrail(CustomGuardrail):
verbose_proxy_logger.debug("Pillar Guardrail: Initialized with fallback_on_error: %s", self.fallback_on_error)
# Set timeout with graceful fallback on invalid configuration
super().__init__(
guardrail_name=guardrail_name,
supported_event_hooks=list(self.get_supported_event_hooks()),
**kwargs,
)
if timeout is not None:
self.timeout = timeout
else:
@ -298,12 +305,6 @@ class PillarGuardrail(CustomGuardrail):
)
self.timeout = self.DEFAULT_TIMEOUT
super().__init__(
guardrail_name=guardrail_name,
supported_event_hooks=list(self.get_supported_event_hooks()),
**kwargs,
)
# =========================================================================
# PUBLIC HOOK METHODS (Main Interface)
# =========================================================================

View file

@ -460,6 +460,11 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
analyze_url,
json=analyze_payload,
headers={"Accept": "application/json"},
timeout=(
aiohttp.ClientTimeout(total=self.timeout)
if isinstance(self.timeout, (int, float))
else aiohttp.client.DEFAULT_TIMEOUT
),
) as response:
# Validate HTTP status
if response.status >= 400:
@ -745,6 +750,11 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
anonymize_url,
json=anonymize_payload,
headers={"Accept": "application/json"},
timeout=(
aiohttp.ClientTimeout(total=self.timeout)
if isinstance(self.timeout, (int, float))
else aiohttp.client.DEFAULT_TIMEOUT
),
) as response:
if response.status >= 400:
error_body = await response.text()

View file

@ -23,6 +23,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
streaming_transform_mode=getattr(litellm_params, "streaming_transform_mode", None),
file_sanitization_fail_open=getattr(litellm_params, "file_sanitization_fail_open", None),
block_on_file_modify=getattr(litellm_params, "block_on_file_modify", None),
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_prompt_security_callback)

View file

@ -290,6 +290,7 @@ class PromptSecurityGuardrail(CustomGuardrail):
f"{self.api_base}/api/protect",
headers=headers,
json=payload,
timeout=self.timeout,
)
response.raise_for_status()
res: Final[_ProtectResponse] = response.json()
@ -407,6 +408,7 @@ class PromptSecurityGuardrail(CustomGuardrail):
f"{self.api_base}/api/protect",
headers=headers,
json=payload,
timeout=self.timeout,
)
response.raise_for_status()
res: Final[_ProtectResponse] = response.json()
@ -522,6 +524,7 @@ class PromptSecurityGuardrail(CustomGuardrail):
f"{self.api_base}/api/sanitizeFile",
headers=headers,
files=files,
timeout=self.timeout,
)
upload_response.raise_for_status()
upload_result: Final[_SanitizeUploadResponse] = upload_response.json()
@ -552,6 +555,7 @@ class PromptSecurityGuardrail(CustomGuardrail):
f"{self.api_base}/api/sanitizeFile",
headers=headers,
params={"jobId": job_id},
timeout=self.timeout,
)
poll_response.raise_for_status()
result: _SanitizeStatusResponse = poll_response.json()

View file

@ -24,6 +24,7 @@ def initialize_guardrail(
),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(
_cb,

View file

@ -168,7 +168,7 @@ class PromptGuardGuardrail(CustomGuardrail):
"Content-Type": "application/json",
},
json=payload,
timeout=10.0,
timeout=self.timeout if self.timeout is not None else 10.0,
)
response.raise_for_status()
view: Final[PromptGuardHTTPView] = {"guard_response": response.json()}

View file

@ -18,6 +18,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
default_on=litellm_params.default_on,
additional_provider_specific_params=litellm_params.additional_provider_specific_params,
extra_headers=getattr(litellm_params, "extra_headers", None),
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_instance)

View file

@ -26,6 +26,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_qualifire_callback)

View file

@ -378,6 +378,7 @@ class QualifireGuardrail(CustomGuardrail):
url=url,
headers=headers,
json=payload,
timeout=self.timeout,
)
response.raise_for_status()
result: Final = response.json()

View file

@ -33,6 +33,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
unreachable_fallback=litellm_params.unreachable_fallback,
event_hook=_event_hook_from_mode(litellm_params.mode),
default_on=litellm_params.default_on or False,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_repelloai_callback)

View file

@ -148,6 +148,7 @@ class RepelloAIGuardrail(CustomGuardrail):
guardrail_name: str | None = None,
event_hook: (GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None) = None,
default_on: bool = False,
timeout: float | None = None,
):
self.repelloai_api_key = api_key or get_secret_str("ARGUS_API_KEY") or get_secret_str("REPELLOAI_API_KEY") or ""
if not self.repelloai_api_key:
@ -176,6 +177,7 @@ class RepelloAIGuardrail(CustomGuardrail):
event_hook=event_hook,
default_on=default_on,
supported_event_hooks=list(self.get_supported_event_hooks()),
timeout=timeout,
)
async def _call_analyze(
@ -201,6 +203,7 @@ class RepelloAIGuardrail(CustomGuardrail):
url=endpoint,
headers={"X-API-Key": self.repelloai_api_key},
json=request,
timeout=self.timeout,
)
self._raise_for_config_error(response)
response.raise_for_status()

View file

@ -30,6 +30,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(rubrik_callback)

View file

@ -85,8 +85,6 @@ class SingulrGuardrail(CustomGuardrail):
else:
self.block_on_error = block_on_error
self.timeout = _DEFAULT_TIMEOUT if timeout is None else timeout
self.async_handler = get_async_httpx_client(
llm_provider=httpxSpecialProvider.GuardrailCallback,
)
@ -101,6 +99,7 @@ class SingulrGuardrail(CustomGuardrail):
]
super().__init__(**kwargs)
self.timeout = _DEFAULT_TIMEOUT if timeout is None else timeout
@staticmethod
def get_config_model() -> type["GuardrailConfigModel"] | None:

View file

@ -909,7 +909,6 @@ class StraikerGuardrail(CustomGuardrail):
max_size_in_memory=V3_BLOCKED_TURN_MEMORY, default_ttl=V3_BLOCKED_TURN_TTL_SECONDS
)
self.source = source
self.timeout = float(timeout)
self.max_retries = max(0, int(max_retries))
self.initial_backoff = max(0.0, float(initial_backoff))
self.max_backoff = max(self.initial_backoff, float(max_backoff))
@ -928,7 +927,8 @@ class StraikerGuardrail(CustomGuardrail):
)
kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
super().__init__(**kwargs)
super().__init__(**kwargs) # pyright: ignore[reportArgumentType] # kwargs splat carries object-typed values
self.timeout = float(timeout)
self.configured_modes = _configured_modes(self.event_hook)

View file

@ -55,6 +55,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,
timeout=litellm_params.timeout,
unreachable_fallback=(
litellm_params.unreachable_fallback if "unreachable_fallback" in litellm_params.model_fields_set else None
),

View file

@ -161,6 +161,7 @@ class TypeSafeGuardrail(CustomGuardrail):
event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None = None,
default_on: bool = False,
async_handler: AsyncHTTPHandler | None = None,
timeout: float | None = None,
) -> None:
raw_api_base: Final = (api_base or get_secret_str("TYPESAFE_API_BASE") or DEFAULT_API_BASE).rstrip("/")
self.typesafe_api_base = raw_api_base
@ -188,6 +189,7 @@ class TypeSafeGuardrail(CustomGuardrail):
guardrail_name=guardrail_name,
event_hook=event_hook,
default_on=default_on,
timeout=timeout,
)
def _handle_failure(self, error: str, log_detail: dict[str, object]) -> None:
@ -271,7 +273,7 @@ class TypeSafeGuardrail(CustomGuardrail):
"Authorization": f"Bearer {self.typesafe_api_key}",
"Content-Type": "application/json",
},
timeout=_JEV_TIMEOUT_SECONDS,
timeout=self.timeout if self.timeout is not None else _JEV_TIMEOUT_SECONDS,
)
except asyncio.CancelledError:
raise

View file

@ -85,7 +85,7 @@ class _AsyncPostHandler(Protocol):
url: str,
headers: dict[str, str],
json: _AnalyzePayload,
timeout: httpx.Timeout,
timeout: float | httpx.Timeout | None,
) -> Awaitable[httpx.Response]: ...
@ -122,10 +122,6 @@ class VigilGuardGuardrail(CustomGuardrail):
fallback: Final = (unreachable_fallback or "fail_closed").lower()
self.unreachable_fallback: _FallbackMode = "fail_open" if fallback == "fail_open" else "fail_closed"
self.timeout: httpx.Timeout = (
_DEFAULT_VIGIL_TIMEOUT if timeout is None else httpx.Timeout(timeout, connect=min(timeout, 5.0))
)
self.async_handler: _AsyncPostHandler = async_handler or get_async_httpx_client(
llm_provider=httpxSpecialProvider.GuardrailCallback,
)
@ -137,6 +133,8 @@ class VigilGuardGuardrail(CustomGuardrail):
super().__init__(**forwarded)
self.timeout = _DEFAULT_VIGIL_TIMEOUT if timeout is None else httpx.Timeout(timeout, connect=min(timeout, 5.0))
@staticmethod
def get_config_model() -> type["GuardrailConfigModel"] | None:
from litellm.types.proxy.guardrails.guardrail_hooks.vigil_guard import (

View file

@ -27,6 +27,7 @@ def initialize_guardrail(
),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(
_cb,

View file

@ -360,7 +360,7 @@ class XecGuardGuardrail(CustomGuardrail):
"Content-Type": "application/json",
},
json=payload,
timeout=10.0,
timeout=self.timeout if self.timeout is not None else 10.0,
)
response.raise_for_status()
return response.json()

View file

@ -70,8 +70,6 @@ class ZscalerAIGuard(CustomGuardrail):
if send_user_api_key_team_id is not None
else os.getenv("SEND_USER_API_KEY_TEAM_ID", "False").lower() in ("true", "1")
)
self.timeout = self._resolve_timeout(timeout)
verbose_proxy_logger.debug(
"send_user_api_key_alias: %s, \n send_user_api_key_user_id:%s, \n send_user_api_key_team_id:%s",
self.send_user_api_key_alias,
@ -80,6 +78,7 @@ class ZscalerAIGuard(CustomGuardrail):
)
super().__init__(**kwargs)
self.timeout = self._resolve_timeout(timeout)
verbose_proxy_logger.debug("ZscalerAIGuard Initializing ...")

View file

@ -45,6 +45,7 @@ def initialize_bedrock(litellm_params: LitellmParams, guardrail: Guardrail):
streaming_sampling_rate=streaming_params.streaming_sampling_rate,
streaming_end_of_stream_only=streaming_params.streaming_end_of_stream_only,
streaming_buffer_release_on_scan=streaming_params.streaming_buffer_release_on_scan,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_bedrock_callback)
return _bedrock_callback
@ -60,6 +61,7 @@ def initialize_lakera(litellm_params: LitellmParams, guardrail: Guardrail):
event_hook=litellm_params.mode,
category_thresholds=litellm_params.category_thresholds,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_lakera_callback)
return _lakera_callback
@ -83,6 +85,7 @@ def initialize_lakera_v2(litellm_params: LitellmParams, guardrail: Guardrail):
skip_system_message_in_guardrail=litellm_params.skip_system_message_in_guardrail,
skip_tool_message_in_guardrail=litellm_params.skip_tool_message_in_guardrail,
advisory_system_message=litellm_params.advisory_system_message,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_lakera_v2_callback)
return _lakera_v2_callback
@ -154,6 +157,7 @@ def initialize_presidio(litellm_params: LitellmParams, guardrail: Guardrail) ->
presidio_language=litellm_params.presidio_language,
presidio_entities_deny_list=litellm_params.presidio_entities_deny_list,
apply_to_output=False,
timeout=litellm_params.timeout,
_callback_role="scan",
)
params.update(overrides)
@ -251,6 +255,7 @@ def initialize_lasso(
mask=litellm_params.mask,
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_lasso_callback)

View file

@ -0,0 +1,446 @@
"""litellm_params.timeout bounds every HTTP guardrail's outbound call, through a real proxy.
Each guardrail is configured against an owned sink that records the request and then sleeps
~20s. With `timeout: 1` the outbound call must abort near the bound, so the chat round trip
completes in seconds instead of waiting on the sink. A control guardrail without `timeout`
points at a sink path that sleeps ~3s and must wait for the reply, proving unset keeps the
handler default. All probes are sent concurrently so their waits overlap.
"""
from __future__ import annotations
import json
import re
import socket
import threading
import time
from collections.abc import Iterator, Mapping
from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass, field
from functools import partial
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from types import MappingProxyType
from typing import Final, cast
import httpx
import pytest
import yaml
from cryptography.hazmat.primitives import serialization
from cryptography.hazmat.primitives.asymmetric import rsa
from integration._support.client import Gateway, gateway_from_environment
from integration._support.process import owned_proxy_process
from integration._support.wire import Reply, Request, wire_server
SLOW_SECONDS: Final = 20
FAST_SECONDS: Final = 3
BOUND_SECONDS: Final = 8
TOKEN_PATH: Final = "/token"
TOKEN_REPLY: Final = json.dumps(
{"access_token": "synthetic-google-token", "expires_in": 3600, "token_type": "Bearer"}
).encode()
EXCLUDED: Final = {
"microsoft_purview": "token endpoint is the fixed login.microsoftonline.com and cannot point at a sink",
"agent_365": "honors its own request_timeout param, not litellm_params.timeout",
"mcp_jwt_signer": "only runs for pre_mcp_call, which /v1/chat/completions cannot trigger",
"semantic_guard": "routes through litellm embeddings, not a guardrail provider HTTP client",
"llm_as_a_judge": "routes through litellm completions, not a guardrail provider HTTP client",
"litellm_content_filter": "local pattern matching with no outbound HTTP",
"tool_permission": "policy evaluation with no outbound HTTP",
"mcp_end_user_permission": "policy evaluation with no outbound HTTP",
"block_code_execution": "local code analysis with no outbound HTTP",
"custom_code": "runs user code with no provider HTTP client",
"hide-secrets": "in-process masking with no outbound HTTP",
"mcp_security": "MCP tool scanning with no provider HTTP client",
"unified_guardrail": "delegates to other guardrails, makes no HTTP call of its own",
"conduct": "requires the optional conduct-litellm-guard package, which is not installed",
"grayswan": "honors its own guardrail_timeout param, not litellm_params.timeout",
"akto": "honors its own guardrail_timeout param, not litellm_params.timeout",
}
PROVIDERS: Final = (
pytest.param("aim", "aim", {}, "pre_call", False, id="aim"),
pytest.param("aporia", "aporia", {}, "post_call", False, id="aporia"),
pytest.param("alice", "alice", {}, "pre_call", False, id="alice"),
pytest.param("azure-prompt-shield", "azure/prompt_shield", {}, "pre_call", False, id="azure-prompt-shield"),
pytest.param(
"azure-text-moderations", "azure/text_moderations", {}, "pre_call", False, id="azure-text-moderations"
),
pytest.param("cato", "cato_networks", {}, "pre_call", False, id="cato-networks"),
pytest.param("crowdstrike", "crowdstrike_aidr", {}, "pre_call", False, id="crowdstrike-aidr"),
pytest.param(
"deepkeep", "deepkeep", {"deepkeep_firewall_id": "synthetic-firewall"}, "pre_call", False, id="deepkeep"
),
pytest.param("dynamoai", "dynamoai", {}, "pre_call", False, id="dynamoai"),
pytest.param("enkryptai", "enkryptai", {}, "pre_call", False, id="enkryptai"),
pytest.param("generic", "generic_guardrail_api", {}, "pre_call", False, id="generic-guardrail-api"),
pytest.param(
"ibm",
"ibm_guardrails",
{"auth_token": "synthetic-ibm-token", "detector_id": "synthetic-detector"},
"pre_call",
False,
id="ibm-guardrails",
),
pytest.param("javelin", "javelin", {"guard_name": "synthetic-guard"}, "pre_call", False, id="javelin"),
pytest.param("lasso", "lasso", {}, "pre_call", False, id="lasso"),
pytest.param("qualifire", "qualifire", {}, "pre_call", False, id="qualifire"),
pytest.param("noma", "noma", {}, "pre_call", False, id="noma"),
pytest.param("noma-v2", "noma_v2", {}, "pre_call", False, id="noma-v2"),
pytest.param(
"ovalix",
"ovalix",
{
"tracker_api_key": "synthetic-tracker-key",
"application_id": "synthetic-app",
"pre_checkpoint_id": "synthetic-pre",
},
"pre_call",
False,
id="ovalix",
),
pytest.param("pangea", "pangea", {}, "pre_call", False, id="pangea"),
pytest.param("openai-moderation", "openai_moderation", {}, "pre_call", False, id="openai-moderation"),
pytest.param("lakera", "lakera", {}, "pre_call", False, id="lakera"),
pytest.param("lakera-v2", "lakera_v2", {}, "pre_call", False, id="lakera-v2"),
pytest.param("promptguard", "promptguard", {}, "pre_call", False, id="promptguard"),
pytest.param("xecguard", "xecguard", {"xecguard_model": "synthetic-model"}, "pre_call", False, id="xecguard"),
pytest.param("typesafe", "typesafe", {}, "pre_call", True, id="typesafe"),
pytest.param("compresr", "compresr", {}, "pre_call", True, id="compresr"),
pytest.param("repelloai", "repelloai", {"asset_id": "synthetic-asset"}, "pre_call", False, id="repelloai"),
pytest.param("prompt-security", "prompt_security", {}, "pre_call", False, id="prompt-security"),
pytest.param("hiddenlayer", "hiddenlayer", {}, "pre_call", False, id="hiddenlayer"),
pytest.param(
"guardrails-ai", "guardrails_ai", {"guard_name": "synthetic-guard"}, "pre_call", False, id="guardrails-ai"
),
pytest.param(
"presidio",
"presidio",
{"pii_entities_config": {"EMAIL_ADDRESS": "BLOCK"}},
"pre_call",
False,
id="presidio",
),
pytest.param(
"bedrock",
"bedrock",
{
"guardrailIdentifier": "synthetic-guardrail",
"guardrailVersion": "DRAFT",
"aws_region_name": "us-east-1",
},
"pre_call",
False,
id="bedrock",
),
pytest.param("rubrik", "rubrik", {}, "pre_call", False, id="rubrik"),
pytest.param("qostodian", "qostodian_nexus", {}, "pre_call", False, id="qostodian-nexus"),
pytest.param("straiker", "straiker", {"default_app": "synthetic-app"}, "pre_call", False, id="straiker"),
pytest.param("zscaler", "zscaler_ai_guard", {}, "pre_call", False, id="zscaler-ai-guard"),
pytest.param("pillar", "pillar", {}, "pre_call", False, id="pillar"),
pytest.param("cisco", "cisco_ai_defense", {}, "pre_call", False, id="cisco-ai-defense"),
pytest.param("vigil", "vigil_guard", {}, "pre_call", False, id="vigil-guard"),
pytest.param("singulr", "singulr", {}, "pre_call", False, id="singulr"),
pytest.param("headroom", "headroom", {}, "pre_call", True, id="headroom"),
pytest.param("onyx", "onyx", {}, "post_call", False, id="onyx"),
pytest.param("panw", "panw_prisma_airs", {}, "pre_call", False, id="panw-prisma-airs"),
pytest.param(
"model-armor",
"model_armor",
{"project_id": "synthetic-project", "location": "us-central1", "template_id": "synthetic-template"},
"pre_call",
False,
id="model-armor",
),
)
@dataclass(frozen=True, slots=True)
class Seen:
target: str
headers: dict[str, str]
body: str
@dataclass(slots=True)
class Sink:
port: int
seen: list[Seen] = field(default_factory=list)
lock: threading.Lock = field(default_factory=threading.Lock)
server: ThreadingHTTPServer | None = None
thread: threading.Thread | None = None
@property
def url(self) -> str:
return f"http://127.0.0.1:{self.port}"
def start(self) -> None:
sink: Final = self
class Handler(BaseHTTPRequestHandler):
protocol_version = "HTTP/1.1"
def _handle(self) -> None:
raw: Final = self.rfile.read(int(self.headers.get("content-length", "0")))
with sink.lock:
sink.seen.append(
Seen(self.path, {k.lower(): v for k, v in self.headers.items()}, raw.decode(errors="replace"))
)
is_token: Final = self.path.startswith(TOKEN_PATH)
if not is_token:
time.sleep(SLOW_SECONDS if self.path.startswith("/slow/") else FAST_SECONDS)
payload: Final = TOKEN_REPLY if is_token else b"{}"
self.send_response(200)
self.send_header("content-type", "application/json")
self.send_header("content-length", str(len(payload)))
self.send_header("connection", "close")
self.end_headers()
self.wfile.write(payload)
do_POST = _handle
do_GET = _handle
do_PUT = _handle
def log_message(self, format: str, *args: object) -> None:
pass
class Server(ThreadingHTTPServer):
allow_reuse_address = True
daemon_threads = True
self.server = Server(("127.0.0.1", self.port), Handler)
self.thread = threading.Thread(target=self.server.serve_forever, daemon=True)
self.thread.start()
def stop(self) -> None:
assert self.server is not None and self.thread is not None
self.server.shutdown()
self.server.server_close()
self.thread.join(timeout=5)
self.server = None
self.thread = None
def calls_for(self, name: str) -> tuple[Seen, ...]:
mention: Final = re.compile(rf"(?:/|key-){re.escape(name)}(?![\w-])")
with self.lock:
return tuple(
s
for s in self.seen
if mention.search(s.target)
or any(mention.search(v) for v in s.headers.values())
or mention.search(s.body)
)
def _provider(request: Request) -> Reply:
body: Final = json.dumps(
{
"id": "chatcmpl-timeout",
"object": "chat.completion",
"created": 1,
"model": "gpt-4o-mini",
"choices": [
{"index": 0, "message": {"role": "assistant", "content": "synthetic answer"}, "finish_reason": "stop"}
],
"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15},
}
).encode()
return Reply(body=body)
def _guardrail(
name: str,
provider: str,
sink: str,
timeout: object,
extra: dict[str, object],
mode: str,
) -> dict[str, object]:
base: Final = f"{sink}/slow/{name}/" if timeout is not None else f"{sink}/fast/{name}/"
return {
"guardrail_name": name,
"litellm_params": {
"guardrail": provider,
"mode": mode,
"default_on": False,
"api_key": f"key-{name}",
**extra,
**_bases(provider, base, sink),
**({"timeout": timeout} if timeout is not None else {}),
},
}
def _synthetic_private_key() -> str:
key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048)
return key.private_bytes(
serialization.Encoding.PEM, serialization.PrivateFormat.PKCS8, serialization.NoEncryption()
).decode()
def _bases(provider: str, base: str, sink: str) -> dict[str, object]:
if provider == "model_armor":
return {
"api_endpoint": base.rstrip("/"),
"credentials": json.dumps(
{
"type": "service_account",
"client_email": "synthetic@synthetic-project.iam.gserviceaccount.com",
"private_key": _synthetic_private_key(),
"token_uri": sink + TOKEN_PATH,
}
),
}
if provider == "ibm_guardrails":
return {"base_url": base}
if provider == "ovalix":
return {"tracker_api_base": base}
if provider == "akto":
return {"akto_base_url": base}
if provider == "singulr":
return {"singulr_api_base": base}
if provider == "presidio":
return {"presidio_analyzer_api_base": base + "/", "presidio_anonymizer_api_base": base + "/"}
if provider == "bedrock":
return {"aws_bedrock_runtime_endpoint": base}
return {"api_base": base}
def _provider_values() -> Iterator[tuple[str, str, dict[str, object], str, bool]]:
for param in PROVIDERS:
yield cast("tuple[str, str, dict[str, object], str, bool]", param.values)
def _rig_config(sink_url: str, root: Path) -> Path:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["litellm_settings"]["cache"] = False
config["guardrails"] = [
_guardrail(name, provider, sink_url, 1, dict(extra), mode)
for name, provider, extra, mode, _ in _provider_values()
] + [
_guardrail("control-generic", "generic_guardrail_api", sink_url, None, {}, "pre_call"),
]
path: Final = root / "guardrail-timeout.yaml"
path.write_text(yaml.safe_dump(config))
return path
@dataclass(frozen=True, slots=True)
class Rig:
proxy: Gateway
sink: Sink
chat_model: str
@pytest.fixture(scope="module")
def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Rig]:
root: Final = tmp_path_factory.mktemp("guardrail-timeout")
with socket.socket() as reserve:
reserve.bind(("127.0.0.1", 0))
port: Final = reserve.getsockname()[1]
sink: Final = Sink(port)
sink.start()
with gateway_from_environment() as gateway, wire_server(_provider) as provider:
config: Final = _rig_config(sink.url, root)
overrides: Final = {
"AWS_ACCESS_KEY_ID": "synthetic-aws-key",
"AWS_SECRET_ACCESS_KEY": "synthetic-aws-secret",
"AWS_REGION_NAME": "us-east-1",
}
with (
owned_proxy_process(gateway, root, overrides, config=config, workers=2) as owned,
owned.gateway.scenario() as scenario,
):
chat: Final = scenario.model(
model="openai/gpt-4o-mini", api_base=provider.url + "/v1", api_key="synthetic-openai-key"
)
yield Rig(owned.gateway, sink, chat)
if sink.server is not None:
sink.stop()
@dataclass(frozen=True, slots=True)
class Outcome:
response: httpx.Response | httpx.TimeoutException
elapsed: float
def _chat(rig: Rig, guardrail_name: str, exchange: bool = False) -> Outcome:
def tool_call(index: int) -> dict[str, object]:
return {
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": f"call_synthetic_{index}",
"type": "function",
"function": {"name": "lookup", "arguments": "{}"},
}
],
}
messages: Final = (
[
{"role": "user", "content": f"look up a fact for {guardrail_name}"},
tool_call(0),
{"role": "tool", "tool_call_id": "call_synthetic_0", "content": "synthetic tool output " * 200},
tool_call(1),
{"role": "tool", "tool_call_id": "call_synthetic_1", "content": "synthetic newer output " * 200},
{"role": "user", "content": f"guardrail timeout probe {guardrail_name}"},
]
if exchange
else [{"role": "user", "content": f"guardrail timeout probe {guardrail_name}"}]
)
start: Final = time.monotonic()
try:
response: Final = rig.proxy.client.post(
"/v1/chat/completions",
json={"model": rig.chat_model, "messages": messages, "guardrails": [guardrail_name]},
headers={"Authorization": f"Bearer {rig.proxy.key}"},
)
except httpx.TimeoutException as error:
return Outcome(error, time.monotonic() - start)
return Outcome(response, time.monotonic() - start)
@pytest.fixture(scope="module")
def outcomes(rig: Rig) -> Mapping[str, Outcome]:
values: Final = tuple(_provider_values())
names: Final = (*(value[0] for value in values), "control-generic")
exchanges: Final = (*(value[4] for value in values), False)
with ThreadPoolExecutor(max_workers=len(names)) as pool:
results: Final = tuple(pool.map(partial(_chat, rig), names, exchanges))
return MappingProxyType(dict(zip(names, results, strict=True)))
@pytest.mark.parametrize("name,provider,extra,mode,exchange", PROVIDERS)
def test_litellm_params_timeout_bounds_outbound_call(
rig: Rig,
outcomes: Mapping[str, Outcome],
name: str,
provider: str,
extra: dict[str, object],
mode: str,
exchange: bool,
) -> None:
outcome: Final = outcomes[name]
calls: Final = rig.sink.calls_for(name)
assert calls, f"{name}: sink saw no request for {provider}"
assert outcome.elapsed < BOUND_SECONDS, (
f"{name}: elapsed {outcome.elapsed:.2f}s, expected under {BOUND_SECONDS}s with timeout=1"
)
assert isinstance(outcome.response, httpx.Response), f"{name}: client gave up: {outcome.response!r}"
assert outcome.response.status_code != 504, outcome.response.text
def test_unset_timeout_waits_for_sink_response(rig: Rig, outcomes: Mapping[str, Outcome]) -> None:
outcome: Final = outcomes["control-generic"]
calls: Final = rig.sink.calls_for("control-generic")
assert calls, "control-generic: sink saw no request"
assert outcome.elapsed >= FAST_SECONDS - 0.5, (
f"control-generic: elapsed {outcome.elapsed:.2f}s, expected to wait for the {FAST_SECONDS}s sink response"
)
assert isinstance(outcome.response, httpx.Response), f"control-generic: client gave up: {outcome.response!r}"
assert outcome.response.status_code in (200, 400, 500), outcome.response.text

View file

@ -630,7 +630,7 @@ class TestStructuredMessagesInResponse:
{"role": "tool", "tool_call_id": "call_1", "content": '{"ssn": "123-45-6789"}'},
]
def echo_with_tool_output_redacted(url, json, headers):
def echo_with_tool_output_redacted(url, json, headers, **_kwargs):
shown_rows = json["structured_messages"]
assert "index" not in shown_rows[1]["tool_calls"][0]
assert "name" not in shown_rows[0]
@ -670,7 +670,7 @@ class TestStructuredMessagesInResponse:
{"role": "user", "content": "Look up 123-45-6789 for me."},
]
def echo_rows_and_rewrite_texts(url, json, headers):
def echo_rows_and_rewrite_texts(url, json, headers, **_kwargs):
answer = MagicMock()
answer.json.return_value = {
"action": "NONE",

View file

@ -428,6 +428,7 @@ class TestHiddenlayerGuardrail:
"hl-runtime-edge-provider": "litellm",
"hl-runtime-edge-provider-version": "1",
},
timeout=None,
)
@pytest.mark.asyncio
@ -1137,3 +1138,18 @@ def test_get_jwt_gives_up_at_the_timeout_instead_of_blocking_the_event_loop(hang
_get_jwt(auth_url=hanging_auth_server, api_id="id", api_key="secret", timeout=1)
assert time.monotonic() - started < 10
with patch(
"litellm.proxy.guardrails.guardrail_hooks.hiddenlayer.hiddenlayer._get_jwt",
return_value="tok",
) as get_jwt:
guardrail = HiddenlayerGuardrail(
guardrail_name="hiddenlayer",
api_id="id",
api_key="secret",
api_base="https://api.hiddenlayer.ai",
timeout=2,
)
guardrail.refresh_jwt_func()
assert [call.kwargs["timeout"] for call in get_jwt.call_args_list] == [2, 2]

View file

@ -3581,7 +3581,7 @@ def _make_marker_session_iterator(
return False
class MockSession:
def post(self, url, json=None, headers=None):
def post(self, url, json=None, headers=None, timeout=None):
payload = json
if url.endswith("analyze"):
recorded_analyze_payloads.append(payload)
@ -3940,7 +3940,7 @@ async def test_chunked_analyze_concurrency_is_bounded():
return False
class MockSession:
def post(self, url, json=None, headers=None):
def post(self, url, json=None, headers=None, timeout=None):
return MockResponse()
async def __aenter__(self):
@ -4010,7 +4010,7 @@ async def test_chunked_analyze_applies_score_threshold_before_merge():
return False
class MockSession:
def post(self, url, json=None, headers=None):
def post(self, url, json=None, headers=None, timeout=None):
text = json["text"]
idx = text.find(CHUNK_MARKER_ONE)
if idx == -1:
@ -4082,7 +4082,7 @@ async def test_chunk_fanout_bound_is_shared_across_concurrent_calls():
return False
class MockSession:
def post(self, url, json=None, headers=None):
def post(self, url, json=None, headers=None, timeout=None):
return MockResponse()
async def __aenter__(self):

View file

@ -233,7 +233,7 @@ class TestRepelloAIPreCall:
data = {"messages": [{"role": "user", "content": "check me"}]}
captured = {}
async def capture(url, headers, json):
async def capture(url, headers, json, **_kwargs):
captured["url"] = url
captured["headers"] = headers
captured["json"] = json
@ -282,7 +282,7 @@ class TestRepelloAIInputCoverage:
async def _scanned_prompt(guardrail, data, monkeypatch) -> str:
captured = {}
async def capture(url, headers, json):
async def capture(url, headers, json, **_kwargs):
captured["json"] = json
return _verdict_response("passed", url)
@ -609,7 +609,7 @@ class TestRepelloAIPostCall:
response = _model_response("the answer content")
captured = {}
async def capture(url, headers, json):
async def capture(url, headers, json, **_kwargs):
captured["url"] = url
captured["json"] = json
return _verdict_response("passed", url)
@ -630,7 +630,7 @@ class TestRepelloAIPostCall:
response = {"choices": [{"text": "text completion answer"}]}
captured = {}
async def capture(url, headers, json):
async def capture(url, headers, json, **_kwargs):
captured["url"] = url
captured["json"] = json
return _verdict_response("passed", url)
@ -662,7 +662,7 @@ class TestRepelloAIPostCall:
)
captured = {}
async def capture(url, headers, json):
async def capture(url, headers, json, **_kwargs):
captured["json"] = json
return _verdict_response("passed", url)
@ -689,7 +689,7 @@ class TestRepelloAIPostCall:
}
captured = {}
async def capture(url, headers, json):
async def capture(url, headers, json, **_kwargs):
captured["json"] = json
return _verdict_response("passed", url)
@ -720,7 +720,7 @@ class TestRepelloAIPostCall:
}
captured = {}
async def capture(url, headers, json):
async def capture(url, headers, json, **_kwargs):
captured["json"] = json
return _verdict_response("passed", url)
@ -745,7 +745,7 @@ class TestRepelloAIPostCall:
)
captured = {}
async def capture(url, headers, json):
async def capture(url, headers, json, **_kwargs):
captured["json"] = json
return _verdict_response("passed", url)
@ -805,7 +805,7 @@ class TestRepelloAIPostCall:
}
captured = {}
async def capture(url, headers, json):
async def capture(url, headers, json, **_kwargs):
captured["json"] = json
return _verdict_response("passed", url)
@ -839,7 +839,7 @@ class TestRepelloAIPostCall:
}
captured = {}
async def capture(url, headers, json):
async def capture(url, headers, json, **_kwargs):
captured["json"] = json
return _verdict_response("passed", url)
@ -1057,7 +1057,7 @@ class TestRepelloAIStreaming:
data = {"messages": [{"role": "user", "content": "q"}]}
captured = {}
async def capture(url, headers, json):
async def capture(url, headers, json, **_kwargs):
captured["json"] = json
return _verdict_response("blocked", url)

View file

@ -49,7 +49,7 @@ async def test_aim_inspects_multimodal_list_content(user_api_key, monkeypatch):
guard = AimGuardrail()
sent_payload: Dict[str, Any] = {}
async def capture(url, headers, json):
async def capture(url, headers, json, **_kwargs):
sent_payload.update(json)
return _aim_no_action_response()
@ -83,7 +83,7 @@ async def test_aim_inspects_responses_api_input(user_api_key, monkeypatch):
guard = AimGuardrail()
sent_payload: Dict[str, Any] = {}
async def capture(url, headers, json):
async def capture(url, headers, json, **_kwargs):
sent_payload.update(json)
return _aim_no_action_response()
@ -219,7 +219,7 @@ async def test_aim_responses_api_input_anonymize_writeback(user_api_key, monkeyp
},
}
async def capture(url, headers, json):
async def capture(url, headers, json, **_kwargs):
return Response(
status_code=200,
json=aim_response_body,

View file

@ -1963,7 +1963,7 @@ class TestOnlyScanNewMessages:
def _guardrail(self, **overrides):
params = dict(guardrail_name="test-guard", only_scan_new_messages=True)
params.update(overrides)
return CustomGuardrail(**params)
return CustomGuardrail(**params) # pyright: ignore[reportArgumentType] # params values mix str/bool
def _cache(self):
from litellm.caching import DualCache
@ -2939,9 +2939,7 @@ async def test_native_lifecycle_guardrail_logging_only_scans_assembled_response(
from litellm.types.utils import Choices, Message, ModelResponse
guardrail = _NativeLifecycleLoggingGuardrail()
assembled = ModelResponse(
choices=[Choices(message=Message(role="assistant", content="assembled stream text"))]
)
assembled = ModelResponse(choices=[Choices(message=Message(role="assistant", content="assembled stream text"))])
sentinel_result = object()
kwargs = {
"model": "gpt-5.4-mini",
@ -3166,3 +3164,37 @@ class TestPreCallHookResponseIsNotLoggedVerbatim:
)
assert self._logged_response(data) == "allow"
class TestCustomGuardrailTimeout:
def test_timeout_constructor_exposes_it(self):
guardrail = CustomGuardrail(guardrail_name="g1", timeout=2.5)
assert guardrail.timeout == 2.5
def test_timeout_unset_stays_none(self):
guardrail = CustomGuardrail(guardrail_name="g1")
assert guardrail.timeout is None
@pytest.mark.parametrize("configured, expected", [(None, 10.0), (3, 3)])
def test_unset_timeout_keeps_default_assigned_before_super_init(self, configured, expected):
class PresetTimeoutGuardrail(CustomGuardrail):
def __init__(self, **kwargs):
self.timeout = 10.0
super().__init__(guardrail_name="preset", **kwargs)
guardrail = PresetTimeoutGuardrail(timeout=configured)
assert guardrail.timeout == expected
def test_update_in_memory_litellm_params_refreshes_timeout(self):
from litellm.types.guardrails import LitellmParams
guardrail = CustomGuardrail(guardrail_name="g1", timeout=2.5)
guardrail.update_in_memory_litellm_params(
LitellmParams(guardrail="generic_guardrail_api", mode="pre_call", timeout=7)
)
assert guardrail.timeout == 7.0

View file

@ -302,6 +302,24 @@ class TestBatchLogging:
handler.async_httpx_client.post.assert_called_once()
assert len(handler.log_queue) == 0
async def test_flush_queue_does_not_inherit_guardrail_timeout(self, mock_env):
with patch("asyncio.create_task", Mock()):
handler = RubrikLogger(timeout=0.5)
handler.log_queue = [{"msg": "a"}]
sent: list[dict] = []
async def capture(**kwargs):
sent.append(kwargs)
return Mock()
handler.async_httpx_client = AsyncMock()
handler.async_httpx_client.post = capture
await handler.flush_queue()
assert handler.timeout == 0.5
assert [call.get("timeout") for call in sent] == [None], sent
async def test_flush_queue_preserves_events_added_during_send(self, handler):
handler.log_queue = [{"msg": "a"}, {"msg": "b"}]

View file

@ -2763,7 +2763,7 @@ def _per_message_guardrail_server(structured_messages_in_answer: bool) -> Callab
"""Answers one redacted text per chat row it was shown, the way a guardrail
that scans per message does, and optionally the rewritten rows themselves."""
def post(url: str, json: dict, headers: dict) -> MagicMock:
def post(url: str, json: dict, headers: dict, timeout=None) -> MagicMock:
rows = json["structured_messages"]
answer: dict = {
"action": "GUARDRAIL_INTERVENED",