Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 9 additions & 0 deletions litellm/integrations/custom_guardrail.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,8 @@
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
Expand Down Expand Up @@ -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")
Expand All @@ -201,6 +205,7 @@ def __init__(
run_in_parallel: bool = False,
scan_raw_request: bool = False,
only_scan_new_messages: bool = False,
timeout: float | None = None,
**kwargs,
):
"""
Expand Down Expand Up @@ -229,6 +234,8 @@ def __init__(
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
Expand All @@ -246,6 +253,8 @@ def __init__(
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
Expand Down
1 change: 1 addition & 0 deletions litellm/integrations/rubrik.py
Original file line number Diff line number Diff line change
Expand Up @@ -1120,6 +1120,7 @@ async def _post_json(self, endpoint: str, payload: Mapping[str, object], service
endpoint,
json=dict(payload),
headers=dict(self._headers),
timeout=self.timeout,
)
http_response.raise_for_status()
result: Final[_ModerationResponse | None] = http_response.json()
Expand Down
1 change: 1 addition & 0 deletions litellm/proxy/guardrails/guardrail_hooks/aim/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down
2 changes: 2 additions & 0 deletions litellm/proxy/guardrails/guardrail_hooks/aim/aim.py
Original file line number Diff line number Diff line change
Expand Up @@ -181,6 +181,7 @@ async def call_aim_guardrail(self, data: dict, hook: str, key_alias: str | None)
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()
Expand Down Expand Up @@ -285,6 +286,7 @@ async def call_aim_guardrail_on_output(
"messages": self._build_aim_inspection_messages(request_data)
+ [{"role": "assistant", "content": output}]
},
timeout=self.timeout,
)
response.raise_for_status()
res: Final[AimAnalyzeResponse] = response.json()
Expand Down
1 change: 1 addition & 0 deletions litellm/proxy/guardrails/guardrail_hooks/alice/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
1 change: 1 addition & 0 deletions litellm/proxy/guardrails/guardrail_hooks/alice/alice.py
Original file line number Diff line number Diff line change
Expand Up @@ -227,6 +227,7 @@ async def _evaluate(
"Content-Type": "application/json",
"af-api-key": self.alice_api_key,
},
timeout=self.timeout,
)
response.raise_for_status()
body = response.json()
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -123,6 +123,7 @@ async def make_aporia_api_request(
"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:
Expand Down
4 changes: 4 additions & 0 deletions litellm/proxy/guardrails/guardrail_hooks/azure/base.py
Original file line number Diff line number Diff line change
@@ -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,
Expand Down Expand Up @@ -49,6 +51,7 @@ def __init__(
# (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
Expand Down Expand Up @@ -77,6 +80,7 @@ async def _post_to_content_safety(self, endpoint_path: str, request_body: dict[s
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)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1787,6 +1787,7 @@ async def _sign_and_post(
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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -305,6 +305,7 @@ async def call_cato_guardrail(
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()
Expand Down Expand Up @@ -445,6 +446,7 @@ async def call_cato_guardrail_on_output(
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()
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -214,8 +214,6 @@ def __init__(
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``.
Expand All @@ -224,6 +222,7 @@ def __init__(
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"))

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -520,6 +520,7 @@ def __init__(
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)
Expand Down Expand Up @@ -583,6 +584,7 @@ def __init__(
guardrail_name=guardrail_name,
event_hook=event_hook,
default_on=default_on,
timeout=timeout,
)

def _should_bypass(self, request_data: dict) -> bool:
Expand Down Expand Up @@ -755,7 +757,7 @@ async def _call_compress(
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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -355,7 +355,9 @@ async def _call_crowdstrike_aidr_guard(
"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()

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -393,6 +393,7 @@ async def apply_guardrail(
url=self.api_base,
json=guardrail_request,
headers=headers,
timeout=self.timeout,
)

response.raise_for_status()
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -130,6 +130,7 @@ async def _call_dynamoai_guardrails(
url=self.api_url,
json=dict(payload),
headers=headers,
timeout=self.timeout,
)
response.raise_for_status()
response_json: Final = response.json()
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -123,6 +123,7 @@ async def _call_enkryptai_guardrails(
url=self.api_url,
json=payload,
headers=headers,
timeout=self.timeout,
)
response.raise_for_status()
response_json: Final = response.json()
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -477,6 +477,7 @@ async def apply_guardrail(
url=self.api_base,
json=guardrail_request.model_dump(mode="json"),
headers=headers,
timeout=self.timeout,
)

response.raise_for_status()
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -80,6 +80,7 @@ async def make_guardrails_ai_api_request(self, llm_output: str, request_data: di
headers={
"Content-Type": "application/json",
},
timeout=self.timeout,
)
verbose_proxy_logger.debug("guardrails_ai response: %s", response)
_json_response: Final = GuardrailsAIResponse(**response.json())
Expand Down Expand Up @@ -117,6 +118,7 @@ async def make_guardrails_ai_api_request_pre_call_request(self, text_input: str,
headers={
"Content-Type": "application/json",
},
timeout=self.timeout,
)
verbose_proxy_logger.debug("guardrails_ai response: %s", response)
if response.status_code == 400:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -508,7 +508,6 @@ def __init__(
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,
Expand All @@ -520,6 +519,7 @@ def __init__(
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")
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Comment thread
greptile-apps[bot] marked this conversation as resolved.
)
else:
_hiddenlayer_callback = HiddenlayerGuardrailV2(
Expand All @@ -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)
Expand Down
Loading
Loading