diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 2eb9cfb5042..99bb832e26c 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -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 diff --git a/litellm/integrations/rubrik.py b/litellm/integrations/rubrik.py index c9e511905a6..fe7264553df 100644 --- a/litellm/integrations/rubrik.py +++ b/litellm/integrations/rubrik.py @@ -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() diff --git a/litellm/proxy/guardrails/guardrail_hooks/aim/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/aim/__init__.py index e45c08c2256..0c791174b2d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/aim/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/aim/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/aim/aim.py b/litellm/proxy/guardrails/guardrail_hooks/aim/aim.py index 54c9d5760a7..61117fbc55e 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/aim/aim.py +++ b/litellm/proxy/guardrails/guardrail_hooks/aim/aim.py @@ -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() diff --git a/litellm/proxy/guardrails/guardrail_hooks/alice/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/alice/__init__.py index 75ea16f7a88..1ed62b0389f 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/alice/alice.py b/litellm/proxy/guardrails/guardrail_hooks/alice/alice.py index 287031c3528..5388f61277f 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/alice/alice.py +++ b/litellm/proxy/guardrails/guardrail_hooks/alice/alice.py @@ -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() diff --git a/litellm/proxy/guardrails/guardrail_hooks/aporia_ai/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/aporia_ai/__init__.py index 68141606a63..5d8cb45965d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/aporia_ai/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/aporia_ai/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/aporia_ai/aporia_ai.py b/litellm/proxy/guardrails/guardrail_hooks/aporia_ai/aporia_ai.py index dafa6e06652..593f8b797a5 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/aporia_ai/aporia_ai.py +++ b/litellm/proxy/guardrails/guardrail_hooks/aporia_ai/aporia_ai.py @@ -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: diff --git a/litellm/proxy/guardrails/guardrail_hooks/azure/base.py b/litellm/proxy/guardrails/guardrail_hooks/azure/base.py index d2aa11da7c9..4afc004808a 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/azure/base.py +++ b/litellm/proxy/guardrails/guardrail_hooks/azure/base.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 228b31604a3..6488fddd51e 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/cato_networks/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/cato_networks/__init__.py index f20b4ef9a59..6e98d11737a 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/cato_networks/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/cato_networks/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/cato_networks/cato_networks.py b/litellm/proxy/guardrails/guardrail_hooks/cato_networks/cato_networks.py index 2d203c31974..936f862b10b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/cato_networks/cato_networks.py +++ b/litellm/proxy/guardrails/guardrail_hooks/cato_networks/cato_networks.py @@ -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() diff --git a/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense.py b/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense.py index 017ef6e09f6..1f63851b216 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense.py +++ b/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense.py @@ -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")) diff --git a/litellm/proxy/guardrails/guardrail_hooks/compresr/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/compresr/__init__.py index d1806b76469..498f9bf4099 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/compresr/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/compresr/__init__.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/compresr/compresr.py b/litellm/proxy/guardrails/guardrail_hooks/compresr/compresr.py index 1ecdb1b0f63..bf3ca71f45c 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/compresr/compresr.py +++ b/litellm/proxy/guardrails/guardrail_hooks/compresr/compresr.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/__init__.py index 59f02817e5f..436bbe01314 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py b/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py index 3d4aba4ac02..739e6b1d865 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py +++ b/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py @@ -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() diff --git a/litellm/proxy/guardrails/guardrail_hooks/deepkeep/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/deepkeep/__init__.py index 3b73883d290..4278b4066e2 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/deepkeep/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/deepkeep/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py b/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py index 539dc1ea1e9..23803b636f2 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py +++ b/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py @@ -393,6 +393,7 @@ class DeepKeepGuardrail(CustomGuardrail): url=self.api_base, json=guardrail_request, headers=headers, + timeout=self.timeout, ) response.raise_for_status() diff --git a/litellm/proxy/guardrails/guardrail_hooks/dynamoai/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/dynamoai/__init__.py index 511dec7bae8..875335d7f54 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/dynamoai/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/dynamoai/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/dynamoai/dynamoai.py b/litellm/proxy/guardrails/guardrail_hooks/dynamoai/dynamoai.py index bc419b359c1..3a8bd54c587 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/dynamoai/dynamoai.py +++ b/litellm/proxy/guardrails/guardrail_hooks/dynamoai/dynamoai.py @@ -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() diff --git a/litellm/proxy/guardrails/guardrail_hooks/enkryptai/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/enkryptai/__init__.py index 18a26d3fde4..1747e3bc6c0 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/enkryptai/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/enkryptai/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/enkryptai/enkryptai.py b/litellm/proxy/guardrails/guardrail_hooks/enkryptai/enkryptai.py index efe959bd186..98db3822092 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/enkryptai/enkryptai.py +++ b/litellm/proxy/guardrails/guardrail_hooks/enkryptai/enkryptai.py @@ -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() diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py index e3511d46544..de389d8a945 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py index 3d1a173635e..786b65b1cc3 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py @@ -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() diff --git a/litellm/proxy/guardrails/guardrail_hooks/guardrails_ai/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/guardrails_ai/__init__.py index e0b884ef3b3..07678d549b4 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/guardrails_ai/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/guardrails_ai/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/guardrails_ai/guardrails_ai.py b/litellm/proxy/guardrails/guardrail_hooks/guardrails_ai/guardrails_ai.py index 18451df574f..cf6592a3e58 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/guardrails_ai/guardrails_ai.py +++ b/litellm/proxy/guardrails/guardrail_hooks/guardrails_ai/guardrails_ai.py @@ -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: diff --git a/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py b/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py index eb62b896784..46272af98ba 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py +++ b/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py @@ -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") diff --git a/litellm/proxy/guardrails/guardrail_hooks/hiddenlayer/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/hiddenlayer/__init__.py index 9408402ef7e..db487804dc5 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/hiddenlayer/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/hiddenlayer/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/hiddenlayer/hiddenlayer.py b/litellm/proxy/guardrails/guardrail_hooks/hiddenlayer/hiddenlayer.py index 68914a1989e..95e6b999825 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/hiddenlayer/hiddenlayer.py +++ b/litellm/proxy/guardrails/guardrail_hooks/hiddenlayer/hiddenlayer.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/ibm_guardrails/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/ibm_guardrails/__init__.py index 7dc85e51873..ad64f025b2c 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/ibm_guardrails/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/ibm_guardrails/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/ibm_guardrails/ibm_detector.py b/litellm/proxy/guardrails/guardrail_hooks/ibm_guardrails/ibm_detector.py index f4d9cbdec48..5da64de329b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/ibm_guardrails/ibm_detector.py +++ b/litellm/proxy/guardrails/guardrail_hooks/ibm_guardrails/ibm_detector.py @@ -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() diff --git a/litellm/proxy/guardrails/guardrail_hooks/javelin/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/javelin/__init__.py index 80d5f9e1b08..c85bfd0c7e8 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/javelin/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/javelin/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/javelin/javelin.py b/litellm/proxy/guardrails/guardrail_hooks/javelin/javelin.py index e54e07b6a1b..d5edcc19a02 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/javelin/javelin.py +++ b/litellm/proxy/guardrails/guardrail_hooks/javelin/javelin.py @@ -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() diff --git a/litellm/proxy/guardrails/guardrail_hooks/lakera_ai.py b/litellm/proxy/guardrails/guardrail_hooks/lakera_ai.py index c69f90282c3..cb1b223fb47 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/lakera_ai.py +++ b/litellm/proxy/guardrails/guardrail_hooks/lakera_ai.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/lakera_ai_v2.py b/litellm/proxy/guardrails/guardrail_hooks/lakera_ai_v2.py index 2f98a9afbd8..b9fb8c62969 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/lakera_ai_v2.py +++ b/litellm/proxy/guardrails/guardrail_hooks/lakera_ai_v2.py @@ -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()) diff --git a/litellm/proxy/guardrails/guardrail_hooks/lasso/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/lasso/__init__.py index f1a6870c5c3..af4b6810031 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/lasso/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/lasso/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py b/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py index 63821428c62..985812ca980 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py +++ b/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py @@ -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() diff --git a/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/__init__.py index 76bced17c9f..7c2d0dbc2fc 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/__init__.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py b/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py index 2c772c723e3..221c4b3752b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py +++ b/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py @@ -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() diff --git a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/__init__.py index 75f18336d7f..ed955ac829d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py index 3f666178970..f83314af548 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py +++ b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/model_armor/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/model_armor/__init__.py index eda505e2453..06875400f40 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/model_armor/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/model_armor/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py index 75e875c2384..77fc085d4bc 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py +++ b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/noma/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/noma/__init__.py index f82aaab4c0d..9391cf60cc3 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/noma/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/noma/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/noma/noma.py b/litellm/proxy/guardrails/guardrail_hooks/noma/noma.py index edd78e0bbc6..85f476ecd62 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/noma/noma.py +++ b/litellm/proxy/guardrails/guardrail_hooks/noma/noma.py @@ -751,6 +751,7 @@ class NomaGuardrail(CustomGuardrail): "requestId": llm_request_id, }, }, + timeout=self.timeout, ) response.raise_for_status() diff --git a/litellm/proxy/guardrails/guardrail_hooks/noma/noma_v2.py b/litellm/proxy/guardrails/guardrail_hooks/noma/noma_v2.py index 8b1fcda7f47..37a33023d54 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/noma/noma_v2.py +++ b/litellm/proxy/guardrails/guardrail_hooks/noma/noma_v2.py @@ -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", diff --git a/litellm/proxy/guardrails/guardrail_hooks/onyx/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/onyx/__init__.py index f2738050f6f..1054c6e8999 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/onyx/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/onyx/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py b/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py index c22d35509c1..b246b125909 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py +++ b/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py @@ -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()) diff --git a/litellm/proxy/guardrails/guardrail_hooks/ovalix/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/ovalix/__init__.py index 362ce6a4d44..0c651864bbb 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/ovalix/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/ovalix/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py b/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py index c69b24c0553..8409a801c0f 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py +++ b/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py @@ -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() diff --git a/litellm/proxy/guardrails/guardrail_hooks/pangea/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/pangea/__init__.py index fb60b9574ac..71f32f0b448 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/pangea/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/pangea/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/pangea/pangea.py b/litellm/proxy/guardrails/guardrail_hooks/pangea/pangea.py index aa61d98e76f..2194238cea0 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/pangea/pangea.py +++ b/litellm/proxy/guardrails/guardrail_hooks/pangea/pangea.py @@ -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() diff --git a/litellm/proxy/guardrails/guardrail_hooks/pillar/pillar.py b/litellm/proxy/guardrails/guardrail_hooks/pillar/pillar.py index 7021d41475b..3f025cfc8b3 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/pillar/pillar.py +++ b/litellm/proxy/guardrails/guardrail_hooks/pillar/pillar.py @@ -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) # ========================================================================= diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index 94750f08a9e..2c6b33838c2 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -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() diff --git a/litellm/proxy/guardrails/guardrail_hooks/prompt_security/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/prompt_security/__init__.py index be3cf4c82a4..3ff3a9bbf20 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/prompt_security/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/prompt_security/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/prompt_security/prompt_security.py b/litellm/proxy/guardrails/guardrail_hooks/prompt_security/prompt_security.py index e97b9229b83..2cb8110ab08 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/prompt_security/prompt_security.py +++ b/litellm/proxy/guardrails/guardrail_hooks/prompt_security/prompt_security.py @@ -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() diff --git a/litellm/proxy/guardrails/guardrail_hooks/promptguard/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/promptguard/__init__.py index 9b249fcb3ff..0f60470632d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/promptguard/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/promptguard/__init__.py @@ -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, diff --git a/litellm/proxy/guardrails/guardrail_hooks/promptguard/promptguard.py b/litellm/proxy/guardrails/guardrail_hooks/promptguard/promptguard.py index 7d3ae2ac521..7b509a25d35 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/promptguard/promptguard.py +++ b/litellm/proxy/guardrails/guardrail_hooks/promptguard/promptguard.py @@ -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()} diff --git a/litellm/proxy/guardrails/guardrail_hooks/qohash/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/qohash/__init__.py index 7d683211570..6a77d414733 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/qohash/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/qohash/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/qualifire/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/qualifire/__init__.py index c5cb066f281..a8785831c37 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/qualifire/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/qualifire/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/qualifire/qualifire.py b/litellm/proxy/guardrails/guardrail_hooks/qualifire/qualifire.py index eceb54681f6..c68d7e94717 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/qualifire/qualifire.py +++ b/litellm/proxy/guardrails/guardrail_hooks/qualifire/qualifire.py @@ -378,6 +378,7 @@ class QualifireGuardrail(CustomGuardrail): url=url, headers=headers, json=payload, + timeout=self.timeout, ) response.raise_for_status() result: Final = response.json() diff --git a/litellm/proxy/guardrails/guardrail_hooks/repelloai/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/repelloai/__init__.py index 37788b35ec7..7e58660ea4d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/repelloai/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/repelloai/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/repelloai/repelloai.py b/litellm/proxy/guardrails/guardrail_hooks/repelloai/repelloai.py index 8925cc5b3a6..b1f0f588ade 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/repelloai/repelloai.py +++ b/litellm/proxy/guardrails/guardrail_hooks/repelloai/repelloai.py @@ -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() diff --git a/litellm/proxy/guardrails/guardrail_hooks/rubrik/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/rubrik/__init__.py index c051368aab7..cb5592e7fb6 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/rubrik/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/rubrik/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/singulr/singulr.py b/litellm/proxy/guardrails/guardrail_hooks/singulr/singulr.py index bd5b18e368d..242280de3b9 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/singulr/singulr.py +++ b/litellm/proxy/guardrails/guardrail_hooks/singulr/singulr.py @@ -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: diff --git a/litellm/proxy/guardrails/guardrail_hooks/straiker/straiker.py b/litellm/proxy/guardrails/guardrail_hooks/straiker/straiker.py index e46458dfe5b..18cc229852c 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/straiker/straiker.py +++ b/litellm/proxy/guardrails/guardrail_hooks/straiker/straiker.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/typesafe/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/typesafe/__init__.py index dcea75d3a98..2e89c6b1566 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/typesafe/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/typesafe/__init__.py @@ -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 ), diff --git a/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py b/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py index 9df5c204a77..45cfbb2c4a1 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py +++ b/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/vigil_guard/vigil_guard.py b/litellm/proxy/guardrails/guardrail_hooks/vigil_guard/vigil_guard.py index e807da7079e..611738ede8a 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/vigil_guard/vigil_guard.py +++ b/litellm/proxy/guardrails/guardrail_hooks/vigil_guard/vigil_guard.py @@ -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 ( diff --git a/litellm/proxy/guardrails/guardrail_hooks/xecguard/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/xecguard/__init__.py index a3825cca7bc..a7ac0a2b305 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/xecguard/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/xecguard/__init__.py @@ -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, diff --git a/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py b/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py index ddf9cace8b9..b6d75b1f204 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py +++ b/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py @@ -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() diff --git a/litellm/proxy/guardrails/guardrail_hooks/zscaler_ai_guard/zscaler_ai_guard.py b/litellm/proxy/guardrails/guardrail_hooks/zscaler_ai_guard/zscaler_ai_guard.py index 1aefa38ecf8..9380a539ecd 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/zscaler_ai_guard/zscaler_ai_guard.py +++ b/litellm/proxy/guardrails/guardrail_hooks/zscaler_ai_guard/zscaler_ai_guard.py @@ -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 ...") diff --git a/litellm/proxy/guardrails/guardrail_initializers.py b/litellm/proxy/guardrails/guardrail_initializers.py index 31688b2e903..b7f3726d017 100644 --- a/litellm/proxy/guardrails/guardrail_initializers.py +++ b/litellm/proxy/guardrails/guardrail_initializers.py @@ -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) diff --git a/tests/integration/observability/test_guardrail_timeout_all_providers.py b/tests/integration/observability/test_guardrail_timeout_all_providers.py new file mode 100644 index 00000000000..df59d6ab9c4 --- /dev/null +++ b/tests/integration/observability/test_guardrail_timeout_all_providers.py @@ -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 diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py index a5e79f84ef1..e97de4686bf 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py @@ -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", diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_hiddenlayer.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_hiddenlayer.py index f5d51a601d7..954b9b99622 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_hiddenlayer.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_hiddenlayer.py @@ -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] diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py index a5625e45d75..08acec0d7ac 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py @@ -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): diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_repelloai.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_repelloai.py index 1ef25b6e7ab..77883e9af0e 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_repelloai.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_repelloai.py @@ -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) diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_coverage.py b/tests/test_litellm/proxy/guardrails/test_guardrail_coverage.py index 548677c70bc..49c64403313 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_coverage.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_coverage.py @@ -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, diff --git a/tests/unit/integrations/test_custom_guardrail.py b/tests/unit/integrations/test_custom_guardrail.py index 4649bddd281..7bfdfb00faf 100644 --- a/tests/unit/integrations/test_custom_guardrail.py +++ b/tests/unit/integrations/test_custom_guardrail.py @@ -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 diff --git a/tests/unit/integrations/test_rubrik.py b/tests/unit/integrations/test_rubrik.py index f3fea292bde..f8aec70a2f7 100644 --- a/tests/unit/integrations/test_rubrik.py +++ b/tests/unit/integrations/test_rubrik.py @@ -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"}] diff --git a/tests/unit/llms/openai/responses/test_openai_responses_guardrail_handler.py b/tests/unit/llms/openai/responses/test_openai_responses_guardrail_handler.py index 88d8169e196..87980a47f87 100644 --- a/tests/unit/llms/openai/responses/test_openai_responses_guardrail_handler.py +++ b/tests/unit/llms/openai/responses/test_openai_responses_guardrail_handler.py @@ -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",