Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/singulr/singulr.py: 23%
196 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1import json
2import os
3from collections.abc import Mapping, Sequence
4from types import MappingProxyType
5from typing import Any, Final
6from urllib.parse import urlparse
8import httpx
9import pydantic
10from typing_extensions import TypedDict, Unpack
12from litellm._logging import verbose_proxy_logger
13from litellm.exceptions import GuardrailRaisedException
14from litellm.integrations.custom_guardrail import (
15 CustomGuardrail,
16 log_guardrail_information,
17)
18from litellm.litellm_core_utils.litellm_logging import (
19 Logging as LiteLLMLoggingObj,
20)
21from litellm.llms.custom_httpx.http_handler import (
22 get_async_httpx_client,
23 httpxSpecialProvider,
24)
25from litellm.proxy._types import UserAPIKeyAuth
26from litellm.types.guardrails import GuardrailEventHooks
27from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolCallChunk
28from litellm.types.proxy.guardrails.guardrail_hooks.base import (
29 GuardrailConfigModel,
30)
31from litellm.types.proxy.guardrails.guardrail_hooks.singulr import (
32 AssistantMessage,
33 SingulrGuardrailPayload,
34 SingulrGuardrailResponse,
35 SingulrMcpGuardrailPayload,
36 ToolCall,
37 ToolCallFunction,
38)
39from litellm.types.utils import CallTypes, ChatCompletionMessageToolCall, GenericGuardrailAPIInputs
41_DEFAULT_API_BASE: Final = "http://localhost:8003"
42_GUARD_ENDPOINT: Final = "/api/v1/ai-gateway/litellm-v2"
43_DEFAULT_TIMEOUT: Final = 30.0
44_EMPTY_MAPPING: Final[Mapping[str, Mapping[str, object]]] = MappingProxyType({})
45_MCP_MODEL_PREFIX: Final = "MCP:"
48class _CustomGuardrailOptions(TypedDict, total=False, extra_items=object):
49 """Base-class constructor options this guardrail forwards untouched to CustomGuardrail."""
52class SingulrGuardrail(CustomGuardrail):
53 def __init__(
54 self,
55 singulr_api_key: str | None = None,
56 singulr_api_base: str | None = None,
57 singulr_application_id: str | None = None,
58 singulr_guardrail_id: str | None = None,
59 block_on_error: bool | None = None,
60 timeout: float | None = None,
61 **kwargs: Unpack[_CustomGuardrailOptions],
62 ) -> None:
63 self.singulr_api_key = singulr_api_key or os.environ.get("SINGULR_API_KEY")
64 self.singulr_api_base = (
65 (singulr_api_base or os.environ.get("SINGULR_API_BASE") or _DEFAULT_API_BASE).strip().rstrip("/")
66 )
67 parsed: Final = urlparse(self.singulr_api_base)
68 if parsed.scheme == "http" and parsed.hostname not in (
69 "localhost",
70 "127.0.0.1",
71 ):
72 raise ValueError(
73 f"Singulr: api_base {self.singulr_api_base} uses plain HTTP for a "
74 "non-local endpoint. Guardrail payloads contain the API token, full "
75 "conversation content, and the guardrail decision, so this endpoint "
76 "must use HTTPS."
77 )
79 self.singulr_application_id = singulr_application_id or os.environ.get("SINGULR_ENFORCEMENT_ENTITY_ID")
80 self.singulr_guardrail_id = singulr_guardrail_id or os.environ.get("SINGULR_GUARDRAIL_ID")
82 if block_on_error is None:
83 env: Final = os.environ.get("SINGULR_BLOCK_ON_ERROR", "true")
84 self.block_on_error = env.lower() in ("true", "1", "yes")
85 else:
86 self.block_on_error = block_on_error
88 self.timeout = _DEFAULT_TIMEOUT if timeout is None else timeout
90 self.async_handler = get_async_httpx_client(
91 llm_provider=httpxSpecialProvider.GuardrailCallback,
92 )
94 if "supported_event_hooks" not in kwargs:
95 kwargs["supported_event_hooks"] = [
96 GuardrailEventHooks.pre_call,
97 GuardrailEventHooks.post_call,
98 GuardrailEventHooks.logging_only,
99 GuardrailEventHooks.pre_mcp_call,
100 GuardrailEventHooks.post_mcp_call,
101 ]
103 super().__init__(**kwargs)
105 @staticmethod
106 def get_config_model() -> type["GuardrailConfigModel"] | None:
107 from litellm.types.proxy.guardrails.guardrail_hooks.singulr import (
108 SingulrGuardrailConfigModel,
109 )
111 return SingulrGuardrailConfigModel
113 @staticmethod
114 def _metadata_containers(request_data: Mapping[str, Any]) -> tuple[Mapping[str, Any], ...]:
115 litellm_params: Final = request_data.get("litellm_params") or _EMPTY_MAPPING
116 return tuple(
117 container
118 for container in (
119 request_data.get("litellm_metadata"),
120 request_data.get("metadata"),
121 litellm_params.get("litellm_metadata") if litellm_params else None,
122 litellm_params.get("metadata") if litellm_params else None,
123 )
124 if container
125 )
127 @classmethod
128 def _resolve_metadata_value(cls, request_data: Mapping[str, Any], key: str) -> str | None:
129 for container in cls._metadata_containers(request_data=request_data):
130 value = container.get(key)
131 if value:
132 return value
133 return None
135 @classmethod
136 def _resolve_user_role_from_request_data(cls, request_data: Mapping[str, Any]) -> str | None:
137 for container in cls._metadata_containers(request_data=request_data):
138 auth = container.get("user_api_key_auth")
139 if isinstance(auth, UserAPIKeyAuth) and auth.user_role:
140 return auth.user_role.value
141 return None
143 @classmethod
144 def _build_metadata(cls, request_data: Mapping[str, Any]) -> Mapping[str, str] | None:
145 fields: Final = (
146 "user_api_key_alias",
147 "user_api_key_user_id",
148 "user_api_key_user_email",
149 "user_api_key_org_id",
150 "user_api_key_org_alias",
151 "user_api_key_team_id",
152 "user_api_key_team_alias",
153 )
154 resolved: Final = (
155 *((field, cls._resolve_metadata_value(request_data=request_data, key=field)) for field in fields),
156 ("user_api_key_user_role", cls._resolve_user_role_from_request_data(request_data=request_data)),
157 )
158 if not any(value for _, value in resolved):
159 return None
160 return {key: value for key, value in resolved if value} # mutable-ok: short-lived JSON payload dict
162 @staticmethod
163 def _build_user_message(text: str) -> Mapping[str, str]:
164 return {"role": "user", "content": text} # mutable-ok: short-lived JSON payload dict
166 def _build_headers(self) -> Mapping[str, str]:
167 all_headers: Final = MappingProxyType(
168 {
169 "Content-Type": "application/json",
170 "X-Singulr-Gateway-Token": self.singulr_api_key,
171 "X-Singulr-Enforcement-Entity-Id": self.singulr_application_id,
172 "X-Singulr-Guardrail-Id": self.singulr_guardrail_id,
173 }
174 )
175 return MappingProxyType({header: value for header, value in all_headers.items() if value})
177 async def _call_api(self, payload: dict[str, object]) -> SingulrGuardrailResponse | None:
178 endpoint: Final = f"{self.singulr_api_base}{_GUARD_ENDPOINT}"
179 verbose_proxy_logger.debug("Singulr: %s", endpoint)
181 try:
182 response: Final = await self.async_handler.post(
183 url=endpoint,
184 headers=self._build_headers(),
185 json=payload,
186 timeout=self.timeout,
187 )
188 response.raise_for_status()
189 result: Final = SingulrGuardrailResponse.model_validate(response.json())
190 verbose_proxy_logger.debug("Singulr: result=%s", result)
191 return result
193 except httpx.HTTPStatusError as exc:
194 verbose_proxy_logger.error(
195 "Singulr API returned HTTP %s: %s",
196 exc.response.status_code,
197 str(exc),
198 )
199 if self.block_on_error:
200 raise GuardrailRaisedException(
201 guardrail_name=self.guardrail_name,
202 message=f"Singulr API returned HTTP {exc.response.status_code}: {exc.response.text}",
203 ) from exc
204 return None
206 except httpx.TransportError as exc:
207 verbose_proxy_logger.error("Singulr API unreachable: %s", str(exc))
208 if self.block_on_error:
209 raise GuardrailRaisedException(
210 guardrail_name=self.guardrail_name,
211 message=f"Singulr API unreachable (block_on_error=True): {exc}",
212 ) from exc
213 return None
215 except (ValueError, pydantic.ValidationError) as exc:
216 verbose_proxy_logger.error("Singulr API returned an invalid response: %s", str(exc))
217 if self.block_on_error:
218 raise GuardrailRaisedException(
219 guardrail_name=self.guardrail_name,
220 message=f"Singulr API returned an invalid response: {exc}",
221 ) from exc
222 return None
224 async def _apply_guardrail_on_request(
225 self,
226 inputs: GenericGuardrailAPIInputs,
227 texts: Sequence[str],
228 structured_messages: Sequence[AllMessageValues],
229 request_data: Mapping[str, Any],
230 ) -> GenericGuardrailAPIInputs:
231 messages: Final = (
232 tuple(structured_messages)
233 if structured_messages
234 else tuple(self._build_user_message(text) for text in texts)
235 )
237 images: Final = inputs.get("images")
238 tools: Final = inputs.get("tools")
240 if not messages and not images and not tools:
241 verbose_proxy_logger.debug("Singulr: No messages, images, or tools to check after filtering")
242 return inputs
244 metadata: Final = self._build_metadata(request_data=request_data)
246 singulr_req_obj = SingulrGuardrailPayload(
247 correlation_id=request_data.get("litellm_call_id"),
248 model_name=inputs.get("model"),
249 guardrail_scope="request",
250 messages=messages,
251 images=images,
252 tools=tools,
253 metadata=metadata,
254 )
255 payload = singulr_req_obj.model_dump(mode="json")
256 guardrail_resp = await self._call_api(payload)
258 if guardrail_resp is None:
259 return inputs
261 if guardrail_resp.should_block:
262 raise GuardrailRaisedException(
263 guardrail_name=self.guardrail_name,
264 status_code=400,
265 message=f"Blocked by Singulr, Blocking due to {guardrail_resp.blocking_due_to or 'unknown'}",
266 blocked_content=True,
267 )
268 return inputs
270 @staticmethod
271 def _mcp_tool_name(request_data: Mapping[str, Any]) -> str | None:
272 return request_data.get("mcp_tool_name") or request_data.get("name")
274 @staticmethod
275 def _mcp_arguments(request_data: Mapping[str, object]) -> object:
276 arguments: Final = request_data.get("mcp_arguments")
277 return arguments if arguments is not None else request_data.get("arguments")
279 @staticmethod
280 def _is_mcp_call(request_data: Mapping[str, object], logging_obj: LiteLLMLoggingObj | None) -> bool:
281 call_type: Final = logging_obj.call_type if logging_obj is not None else request_data.get("call_type")
282 if call_type is not None:
283 return call_type == CallTypes.call_mcp_tool.value
284 model: Final = request_data.get("model")
285 return "mcp_tool_name" in request_data or (isinstance(model, str) and model.startswith(_MCP_MODEL_PREFIX))
287 async def _apply_guardrail_on_mcp_request(self, request_data: Mapping[str, Any]) -> None:
288 metadata: Final = self._build_metadata(request_data=request_data)
290 singulr_mcp_obj = SingulrMcpGuardrailPayload(
291 guardrail_scope="mcp_request",
292 tool_name=self._mcp_tool_name(request_data),
293 tool_arguments=self._mcp_arguments(request_data),
294 mcp_server_name=request_data.get("mcp_server_name"),
295 metadata=metadata,
296 )
297 payload = singulr_mcp_obj.model_dump(mode="json")
298 guardrail_resp = await self._call_api(payload)
300 if guardrail_resp is None:
301 return
303 if guardrail_resp.should_block:
304 raise GuardrailRaisedException(
305 guardrail_name=self.guardrail_name,
306 status_code=400,
307 message=f"Blocked by Singulr, Blocking due to {guardrail_resp.blocking_due_to or 'unknown'}",
308 blocked_content=True,
309 )
311 async def _apply_guardrail_on_mcp_response(
312 self, inputs: GenericGuardrailAPIInputs, texts: Sequence[str], request_data: Mapping[str, Any]
313 ) -> GenericGuardrailAPIInputs:
314 if not texts:
315 return inputs
317 metadata: Final = self._build_metadata(request_data=request_data)
319 singulr_mcp_obj = SingulrMcpGuardrailPayload(
320 model_name=request_data.get("model"),
321 guardrail_scope="mcp_response",
322 tool_result=texts,
323 metadata=metadata,
324 )
325 payload = singulr_mcp_obj.model_dump(mode="json")
326 guardrail_resp = await self._call_api(payload)
328 if guardrail_resp is None:
329 return inputs
331 if guardrail_resp.should_block:
332 raise GuardrailRaisedException(
333 guardrail_name=self.guardrail_name,
334 status_code=400,
335 message=f"Blocked by Singulr, Blocking due to {guardrail_resp.blocking_due_to or 'unknown'}",
336 blocked_content=True,
337 )
339 return inputs
341 @staticmethod
342 def _build_tool_call(tool_call: ChatCompletionToolCallChunk | ChatCompletionMessageToolCall) -> "ToolCall | None":
343 tool_call_id: Final = tool_call.get("id")
344 fun: Final = tool_call.get("function")
345 if not tool_call_id or not fun:
346 return None
347 func_name: Final = fun.get("name")
348 args: Final = fun.get("arguments")
349 if not func_name or args is None:
350 return None
351 call_type: Final = tool_call.get("type")
352 return ToolCall(
353 id=tool_call_id,
354 type=call_type if isinstance(call_type, str) and call_type else "function",
355 function=ToolCallFunction(
356 name=func_name,
357 arguments=args if isinstance(args, str) else json.dumps(args, default=str),
358 ),
359 )
361 async def _apply_guardrail_on_response(
362 self, inputs: GenericGuardrailAPIInputs, texts: Sequence[str], request_data: Mapping[str, Any]
363 ) -> GenericGuardrailAPIInputs:
364 combined_texts: Final = "\n".join(texts) if texts else None
366 tool_calls: Final = inputs.get("tool_calls", ())
367 tool_calls_res: Final = tuple(
368 tool_call_res
369 for tool_call_res in (self._build_tool_call(tool_call) for tool_call in tool_calls)
370 if tool_call_res is not None
371 )
373 assistant_message: Final = AssistantMessage(
374 role="assistant",
375 content=combined_texts,
376 tool_calls=tool_calls_res,
377 )
379 metadata: Final = self._build_metadata(request_data=request_data)
381 singulr_resp_obj = SingulrGuardrailPayload(
382 correlation_id=request_data.get("litellm_call_id"),
383 guardrail_scope="response",
384 model_name=request_data.get("model"),
385 messages=request_data.get("messages"),
386 images=inputs.get("images"),
387 response=assistant_message,
388 metadata=metadata,
389 )
391 payload = singulr_resp_obj.model_dump(mode="json")
392 guardrail_resp = await self._call_api(payload)
394 if guardrail_resp is None:
395 return inputs
397 if guardrail_resp.should_block:
398 raise GuardrailRaisedException(
399 guardrail_name=self.guardrail_name,
400 status_code=400,
401 message=f"Blocked by Singulr, Blocking due to {guardrail_resp.blocking_due_to or 'unknown'}",
402 blocked_content=True,
403 )
404 return inputs
406 @log_guardrail_information
407 async def apply_guardrail(
408 self,
409 inputs: GenericGuardrailAPIInputs,
410 request_data: dict, # mutable-ok: required by CustomGuardrail.apply_guardrail override signature
411 input_type: str,
412 logging_obj: "LiteLLMLoggingObj | None" = None,
413 ) -> GenericGuardrailAPIInputs:
414 texts: Final = inputs.get("texts", ())
415 structured_messages: Final = inputs.get("structured_messages", ())
417 verbose_proxy_logger.debug(
418 "Singulr Guardrail: apply_guardrail called with input_type=%s, texts=%d, structured_messages=%d",
419 input_type,
420 len(texts),
421 len(structured_messages),
422 )
424 is_mcp_call: Final = self._is_mcp_call(request_data, logging_obj)
425 if input_type == "request":
426 if is_mcp_call:
427 await self._apply_guardrail_on_mcp_request(request_data=request_data)
428 return inputs
429 return await self._apply_guardrail_on_request(
430 inputs=inputs, texts=texts, structured_messages=structured_messages, request_data=request_data
431 )
432 elif input_type == "response":
433 if is_mcp_call:
434 return await self._apply_guardrail_on_mcp_response(
435 inputs=inputs, texts=texts, request_data=request_data
436 )
437 return await self._apply_guardrail_on_response(inputs=inputs, texts=texts, request_data=request_data)
438 return inputs