mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
feat(guardrails): honor litellm_params.timeout in every HTTP guardrail (#43134)
* feat(guardrails): honor litellm_params.timeout in every HTTP guardrail Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(guardrails): accept timeout kwarg in presidio and responses-handler post stubs Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(guardrails): bound hiddenlayer startup jwt call by configured timeout, drop akto from timeout coverage Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(guardrails): narrow hiddenlayer startup auth timeout without cast Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(guardrails): bound hiddenlayer jwt refresh by configured timeout Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(guardrails): keep provider timeout defaults when unset and bound only rubrik moderation calls Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(guardrails): cover model_armor and run timeout probes concurrently Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(guardrails): match sink calls to the exact guardrail name Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: kerry <kerry@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
ae05f7d2c1
commit
a308a8e579
83 changed files with 688 additions and 62 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -393,6 +393,7 @@ class DeepKeepGuardrail(CustomGuardrail):
|
|||
url=self.api_base,
|
||||
json=guardrail_request,
|
||||
headers=headers,
|
||||
timeout=self.timeout,
|
||||
)
|
||||
|
||||
response.raise_for_status()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -751,6 +751,7 @@ class NomaGuardrail(CustomGuardrail):
|
|||
"requestId": llm_request_id,
|
||||
},
|
||||
},
|
||||
timeout=self.timeout,
|
||||
)
|
||||
response.raise_for_status()
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
# =========================================================================
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -378,6 +378,7 @@ class QualifireGuardrail(CustomGuardrail):
|
|||
url=url,
|
||||
headers=headers,
|
||||
json=payload,
|
||||
timeout=self.timeout,
|
||||
)
|
||||
response.raise_for_status()
|
||||
result: Final = response.json()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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 ...")
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"}]
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue