Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/repelloai/repelloai.py: 15%
347 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
1from __future__ import annotations
3from collections.abc import AsyncGenerator
4from datetime import datetime
5from typing import Final, Literal, TypeGuard
7from fastapi import HTTPException
8from httpx import HTTPError
9from httpx import Response as HttpxResponse
10from pydantic import BaseModel, TypeAdapter, ValidationError
12import litellm
13from litellm._logging import verbose_proxy_logger
14from litellm.integrations.custom_guardrail import CustomGuardrail
15from litellm.llms.custom_httpx.http_handler import (
16 get_async_httpx_client, # pyright: ignore[reportUnknownVariableType]
17 httpxSpecialProvider,
18)
19from litellm.proxy._types import UserAPIKeyAuth
20from litellm.proxy.common_utils.callback_utils import (
21 add_guardrail_to_applied_guardrails_header, # pyright: ignore[reportUnknownVariableType]
22)
23from litellm.proxy.guardrails._content_utils import build_inspection_messages
24from litellm.secret_managers.main import get_secret_str
25from litellm.types.guardrails import GuardrailEventHooks, Mode
26from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
27from litellm.types.proxy.guardrails.guardrail_hooks.repelloai import (
28 RepelloAIAnalyzeResponse,
29)
30from litellm.types.utils import (
31 CallTypesLiteral,
32 GuardrailStatus,
33 LLMResponseTypes,
34 ModelResponse,
35 ModelResponseStream,
36)
38DEFAULT_REPELLOAI_API_BASE: Final = "https://argusapi.repello.ai/sdk/v1"
39DEFAULT_REPELLOAI_TIMEOUT: Final = 30.0
40BLOCKED_VERDICT: Final = "blocked"
41FLAGGED_VERDICT: Final = "flagged"
42PASSED_VERDICT: Final = "passed"
44# Argus returns these for a permanently broken guardrail (bad key, unknown
45# asset_id, malformed payload), not a transient outage. They must always
46# block, never honour fail_open.
47CONFIG_ERROR_STATUS_CODES: Final = frozenset({400, 401, 403, 404, 422})
48_SCHEMA_SCALAR_KEYS: Final = frozenset(("name", "description", "title", "const", "default"))
49_SCHEMA_LIST_KEYS: Final = frozenset(("enum", "examples"))
50_SCHEMA_EXTRACTED_KEYS: Final = _SCHEMA_SCALAR_KEYS | _SCHEMA_LIST_KEYS
53class RepelloAIGuardrailMissingSecrets(Exception):
54 pass
57def _is_object_dict(value: object) -> TypeGuard[dict[str, object]]: # guard-ok: isinstance narrows correctly; predicate is trivially correct # fmt: skip
58 return isinstance(value, dict)
61def _is_object_list(value: object) -> TypeGuard[list[object]]: # guard-ok: isinstance narrows correctly; predicate is trivially correct # fmt: skip
62 return isinstance(value, list)
65class RepelloAIGuardrail(CustomGuardrail):
66 @classmethod
67 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
68 return [
69 GuardrailEventHooks.pre_call,
70 GuardrailEventHooks.post_call,
71 ]
73 @staticmethod
74 def _get_field(obj: object, key: str) -> object:
75 if _is_object_dict(obj):
76 return obj.get(key)
77 return getattr(obj, key, None)
79 @classmethod
80 def _extract_tool_call_args_from_message(cls, message: object) -> list[str]:
81 args: Final[list[str]] = []
83 tool_calls: Final = cls._get_field(message, "tool_calls")
84 if _is_object_list(tool_calls):
85 for tool_call in tool_calls:
86 function = cls._get_field(tool_call, "function")
87 arguments = cls._get_field(function, "arguments")
88 if isinstance(arguments, str) and arguments.strip():
89 args.append(arguments)
91 function_call: Final = cls._get_field(message, "function_call")
92 arguments = cls._get_field(function_call, "arguments")
93 if isinstance(arguments, str) and arguments.strip():
94 args.append(arguments)
96 return args
98 @staticmethod
99 def _iter_schema_text(node: object) -> list[str]:
100 texts: Final[list[str]] = []
101 stack: Final[list[object]] = [node]
103 while stack:
104 current = stack.pop()
105 if _is_object_dict(current):
106 for key in _SCHEMA_SCALAR_KEYS:
107 value = current.get(key)
108 if isinstance(value, str) and value:
109 texts.append(value)
110 for key in _SCHEMA_LIST_KEYS:
111 items = current.get(key)
112 if _is_object_list(items):
113 for item in items:
114 if isinstance(item, str) and item:
115 texts.append(item)
116 remaining: list[object] = [v for k, v in current.items() if k not in _SCHEMA_EXTRACTED_KEYS]
117 stack.extend(reversed(remaining))
118 elif _is_object_list(current):
119 stack.extend(reversed(current))
121 return texts
123 @classmethod
124 def _extract_tool_definition_text(cls, data: dict[str, object]) -> list[str]:
125 texts: Final[list[str]] = []
127 tools: Final = data.get("tools")
128 for tool in tools if _is_object_list(tools) else []:
129 if not _is_object_dict(tool):
130 continue
131 function = tool.get("function")
132 if _is_object_dict(function):
133 texts.extend(cls._iter_schema_text(function))
135 functions: Final = data.get("functions")
136 for function in functions if _is_object_list(functions) else []:
137 if _is_object_dict(function):
138 texts.extend(cls._iter_schema_text(function))
140 return texts
142 def __init__(
143 self,
144 api_key: str | None = None,
145 api_base: str | None = None,
146 asset_id: str | None = None,
147 unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed",
148 guardrail_name: str | None = None,
149 event_hook: (GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None) = None,
150 default_on: bool = False,
151 ):
152 self.repelloai_api_key = api_key or get_secret_str("ARGUS_API_KEY") or get_secret_str("REPELLOAI_API_KEY") or ""
153 if not self.repelloai_api_key:
154 raise RepelloAIGuardrailMissingSecrets(
155 "Couldn't get Repello API key. Set `ARGUS_API_KEY` in the environment "
156 "or pass `api_key` to the guardrail in the config file."
157 )
159 self.asset_id = asset_id
160 if not self.asset_id:
161 raise ValueError(
162 "Repello guardrail requires an `asset_id`. Create an asset in the Repello "
163 "dashboard and set `asset_id` on the guardrail in the config file."
164 )
166 self.api_base = api_base or get_secret_str("REPELLOAI_API_BASE") or DEFAULT_REPELLOAI_API_BASE
167 self.unreachable_fallback: Literal["fail_closed", "fail_open"] = (
168 "fail_open" if unreachable_fallback == "fail_open" else "fail_closed"
169 )
170 self.async_handler = get_async_httpx_client(
171 llm_provider=httpxSpecialProvider.GuardrailCallback,
172 params={"timeout": DEFAULT_REPELLOAI_TIMEOUT},
173 )
174 super().__init__( # pyright: ignore[reportUnknownMemberType]
175 guardrail_name=guardrail_name,
176 event_hook=event_hook,
177 default_on=default_on,
178 supported_event_hooks=list(self.get_supported_event_hooks()),
179 )
181 async def _call_analyze(
182 self,
183 text: str,
184 stage: Literal["prompt", "response"],
185 request_data: dict[str, object],
186 event_type: GuardrailEventHooks,
187 ) -> RepelloAIAnalyzeResponse | None:
188 endpoint: Final = f"{self.api_base}/analyze/{stage}"
189 request: Final[dict[str, object]] = {
190 "asset_id": self.asset_id or "",
191 "scan_data": {stage: text},
192 }
194 status: GuardrailStatus = "success"
195 guardrail_json_response: str | dict[str, object] | list[dict[str, object]] = ""
196 start_time: Final[datetime] = datetime.now()
197 repelloai_response: RepelloAIAnalyzeResponse | None = None
198 try:
199 verbose_proxy_logger.debug("RepelloAI Argus request: %s", request)
200 response: Final[HttpxResponse] = await self.async_handler.post( # pyright: ignore[reportUnknownMemberType] # AsyncHTTPHandler.post is untyped
201 url=endpoint,
202 headers={"X-API-Key": self.repelloai_api_key},
203 json=request,
204 )
205 self._raise_for_config_error(response)
206 response.raise_for_status()
207 try:
208 repelloai_response = TypeAdapter(RepelloAIAnalyzeResponse).validate_json(response.text)
209 except ValidationError as e:
210 raise HTTPException(
211 status_code=500,
212 detail={
213 "error": "RepelloAI Argus guardrail returned invalid JSON",
214 "status_code": response.status_code,
215 },
216 ) from e
217 verbose_proxy_logger.debug("RepelloAI Argus response: %s", repelloai_response)
218 if self._verdict_blocks(repelloai_response):
219 status = "guardrail_intervened"
220 return repelloai_response
221 except HTTPException as e:
222 status = "guardrail_failed_to_respond"
223 guardrail_json_response = str(e.detail) if not isinstance(e.detail, (dict, list)) else e.detail
224 raise
225 except HTTPError as e:
226 status = "guardrail_failed_to_respond"
227 guardrail_json_response = str(e)
228 return self._handle_unreachable(e)
229 except Exception as e:
230 status = "guardrail_failed_to_respond"
231 guardrail_json_response = str(e)
232 raise HTTPException(status_code=500, detail={"error": "RepelloAI Argus guardrail failed"}) from e
233 finally:
234 end_time: Final = datetime.now()
235 if repelloai_response is not None:
236 guardrail_json_response = dict(repelloai_response)
237 self.add_standard_logging_guardrail_information_to_request_data( # pyright: ignore[reportUnknownMemberType]
238 guardrail_json_response=guardrail_json_response,
239 guardrail_status=status,
240 request_data=request_data,
241 start_time=start_time.timestamp(),
242 end_time=end_time.timestamp(),
243 duration=(end_time - start_time).total_seconds(),
244 masked_entity_count={},
245 event_type=event_type,
246 )
248 @staticmethod
249 def _raise_for_config_error(response: HttpxResponse) -> None:
250 if response.status_code in CONFIG_ERROR_STATUS_CODES:
251 raise HTTPException(
252 status_code=500,
253 detail={
254 "error": "RepelloAI Argus guardrail is misconfigured",
255 "status_code": response.status_code,
256 },
257 )
259 def _verdict_blocks(self, repelloai_response: RepelloAIAnalyzeResponse | None) -> bool:
260 if repelloai_response is None:
261 return False
262 verdict: Final = repelloai_response.get("verdict")
263 if verdict == BLOCKED_VERDICT:
264 return True
265 if verdict in (PASSED_VERDICT, FLAGGED_VERDICT):
266 return False
267 verbose_proxy_logger.warning(
268 "RepelloAI Argus returned an unrecognized verdict (%s) - blocking.",
269 verdict,
270 )
271 return True
273 def _handle_unreachable(self, error: Exception) -> RepelloAIAnalyzeResponse | None:
274 verbose_proxy_logger.warning("RepelloAI Argus unreachable: %s", str(error))
275 if self.unreachable_fallback == "fail_closed":
276 raise HTTPException(
277 status_code=500,
278 detail={"error": "RepelloAI Argus guardrail unreachable"},
279 )
280 return None
282 def _raise_if_blocked(self, repelloai_response: RepelloAIAnalyzeResponse | None) -> None:
283 if repelloai_response is None:
284 return
285 if self._verdict_blocks(repelloai_response):
286 raise HTTPException(
287 status_code=400,
288 detail=self._format_blocked_detail(repelloai_response),
289 )
290 self._log_flagged_verdict(repelloai_response)
292 @classmethod
293 def _format_blocked_detail(cls, repelloai_response: RepelloAIAnalyzeResponse) -> str:
294 policies: Final = repelloai_response.get("policies_violated")
295 if not isinstance(policies, list) or not policies:
296 return "Blocked by RepelloAI Argus guardrail."
298 formatted_policies: Final[list[str]] = []
299 for policy in policies:
300 policy_name = policy.get("policy_name") or "unknown_policy"
301 details: list[str] = []
302 action_taken = policy.get("action_taken")
303 if action_taken:
304 details.append(f"action: {action_taken}")
305 policy_details = policy.get("details")
306 if isinstance(policy_details, dict):
307 score = policy_details.get("score")
308 if score is not None:
309 details.append(f"score: {score}")
310 suffix = f" ({', '.join(details)})" if details else ""
311 formatted_policies.append(f"{policy_name}{suffix}")
313 if not formatted_policies:
314 return "Blocked by RepelloAI Argus guardrail."
315 return f"Blocked by RepelloAI Argus guardrail. Policies violated: {'; '.join(formatted_policies)}."
317 @staticmethod
318 def _log_flagged_verdict(repelloai_response: RepelloAIAnalyzeResponse) -> None:
319 if repelloai_response.get("verdict") == FLAGGED_VERDICT:
320 verbose_proxy_logger.warning(
321 "RepelloAI Argus flagged content (allowed): %s",
322 repelloai_response.get("policies_violated"),
323 )
325 @staticmethod
326 def _extract_prompt_message_text(data: dict[str, object]) -> list[str]:
327 messages: Final = build_inspection_messages(data)
328 return [content for message in messages if isinstance(content := message.get("content"), str) and content]
330 @staticmethod
331 def _extract_input_text_parts(content: object) -> list[str]:
332 if not _is_object_list(content):
333 return []
334 return [
335 text
336 for part in content
337 if _is_object_dict(part) and part.get("type") == "input_text"
338 if isinstance(text := part.get("text"), str) and text
339 ]
341 @staticmethod
342 def _extract_prompt_field_text(data: dict[str, object]) -> list[str]:
343 prompt: Final = data.get("prompt")
344 if isinstance(prompt, str) and prompt:
345 return [prompt]
346 if _is_object_list(prompt):
347 return [item for item in prompt if isinstance(item, str) and item]
348 return []
350 @classmethod
351 def _extract_prompt_text(cls, data: dict[str, object]) -> str | None:
352 texts: Final = cls._extract_prompt_message_text(data)
353 texts.extend(cls._extract_prompt_field_text(data))
355 instructions: Final = data.get("instructions")
356 if isinstance(instructions, str) and instructions:
357 texts.append(instructions)
359 raw_messages: Final = data.get("messages")
360 if _is_object_list(raw_messages):
361 for message in raw_messages:
362 texts.extend(cls._extract_tool_call_args_from_message(message))
364 raw_input: Final = data.get("input")
365 if _is_object_list(raw_input):
366 for item in raw_input:
367 if _is_object_dict(item):
368 if "role" not in item:
369 continue
370 texts.extend(cls._extract_tool_call_args_from_message(item))
371 texts.extend(cls._extract_input_text_parts(item.get("content")))
373 texts.extend(cls._extract_tool_definition_text(data))
374 return "\n".join(text for text in texts if text) if texts else None
376 async def async_pre_call_hook(
377 self,
378 user_api_key_dict: UserAPIKeyAuth,
379 cache: litellm.DualCache,
380 data: dict[str, object],
381 call_type: CallTypesLiteral,
382 ) -> Exception | str | dict[str, object] | None:
383 verbose_proxy_logger.debug("RepelloAI Argus: pre_call_hook")
385 event_type: Final = GuardrailEventHooks.pre_call
386 if (
387 self.should_run_guardrail( # pyright: ignore[reportUnknownMemberType]
388 data=data, event_type=event_type
389 )
390 is not True
391 ):
392 return data
394 text: Final = self._extract_prompt_text(data)
395 if not text:
396 verbose_proxy_logger.warning("RepelloAI Argus: no inspectable prompt text in data - skipping.")
397 return data
399 repelloai_response: Final = await self._call_analyze(
400 text=text,
401 stage="prompt",
402 request_data=data,
403 event_type=event_type,
404 )
405 self._raise_if_blocked(repelloai_response)
407 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name)
408 return data
410 async def async_post_call_success_hook(
411 self,
412 data: dict[str, object],
413 user_api_key_dict: UserAPIKeyAuth,
414 response: LLMResponseTypes,
415 ):
416 verbose_proxy_logger.debug("RepelloAI Argus: post_call_success_hook")
418 event_type: Final = GuardrailEventHooks.post_call
419 if (
420 self.should_run_guardrail( # pyright: ignore[reportUnknownMemberType]
421 data=data, event_type=event_type
422 )
423 is not True
424 ):
425 return response
427 text: Final = self._extract_response_text(response)
428 if not text:
429 verbose_proxy_logger.warning("RepelloAI Argus: no inspectable response text - skipping.")
430 return response
432 repelloai_response: Final = await self._call_analyze(
433 text=text,
434 stage="response",
435 request_data=data,
436 event_type=event_type,
437 )
438 self._raise_if_blocked(repelloai_response)
440 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name)
441 return response
443 async def async_post_call_streaming_iterator_hook(
444 self,
445 user_api_key_dict: UserAPIKeyAuth,
446 response: AsyncGenerator[ModelResponseStream, None],
447 request_data: dict[str, object],
448 ) -> AsyncGenerator[ModelResponseStream, None]:
449 from litellm import main as litellm_main
451 event_type: Final = GuardrailEventHooks.post_call
452 if (
453 self.should_run_guardrail( # pyright: ignore[reportUnknownMemberType]
454 data=request_data, event_type=event_type
455 )
456 is not True
457 ):
458 async for chunk in response:
459 yield chunk
460 return
462 chunks: Final[list[ModelResponseStream]] = []
463 async for chunk in response:
464 chunks.append(chunk)
466 assembled = litellm_main.stream_chunk_builder( # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType]
467 chunks=chunks
468 )
469 text: Final = self._extract_response_text(assembled) if isinstance(assembled, ModelResponse) else None
470 if text:
471 repelloai_response: Final = await self._call_analyze(
472 text=text,
473 stage="response",
474 request_data=request_data,
475 event_type=event_type,
476 )
477 if repelloai_response is not None:
478 self._log_flagged_verdict(repelloai_response)
479 if self._verdict_blocks(repelloai_response):
480 from litellm.proxy.proxy_server import StreamingCallbackError
482 raise StreamingCallbackError("Blocked by RepelloAI Argus guardrail")
483 add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name)
484 else:
485 verbose_proxy_logger.warning(
486 "RepelloAI Argus: no inspectable text in streamed response; skipping scan. "
487 "guardrail=%s assembled_type=%s",
488 self.guardrail_name,
489 type(assembled).__name__,
490 )
492 for chunk in chunks:
493 yield chunk
495 @staticmethod
496 def _extract_response_text(response: object) -> str | None:
497 if _is_object_dict(response):
498 response_dict = response
499 elif isinstance(response, ModelResponse):
500 response_dict = (
501 response.model_dump() # pyright: ignore[reportUnknownMemberType]
502 )
503 else:
504 output_text: Final = getattr(response, "output_text", None)
505 if isinstance(output_text, str) and output_text:
506 return output_text
507 response_dict = {}
509 text: Final = RepelloAIGuardrail._extract_chat_completion_text(response_dict)
510 if text:
511 return text
512 return RepelloAIGuardrail._extract_responses_api_text(response_dict)
514 @classmethod
515 def _extract_chat_completion_text(cls, response_dict: dict[str, object]) -> str | None:
516 choices: Final = response_dict.get("choices")
517 if not _is_object_list(choices):
518 return None
519 parts: Final[list[str]] = []
520 for choice in choices:
521 if not _is_object_dict(choice):
522 continue
523 message = choice.get("message")
524 if _is_object_dict(message):
525 content = message.get("content")
526 if isinstance(content, str) and content:
527 parts.append(content)
528 parts.extend(cls._extract_tool_call_args_from_message(message))
529 text = choice.get("text")
530 if isinstance(text, str) and text:
531 parts.append(text)
532 return "\n".join(parts) if parts else None
534 @staticmethod
535 def _extract_responses_api_text(response_dict: dict[str, object]) -> str | None:
536 output: Final = response_dict.get("output")
537 if not _is_object_list(output):
538 return None
539 texts: Final[list[str]] = []
540 for output_item in output:
541 if not _is_object_dict(output_item):
542 continue
543 item_type = output_item.get("type")
544 if item_type == "function_call":
545 arguments = output_item.get("arguments")
546 if isinstance(arguments, str) and arguments:
547 texts.append(arguments)
548 continue
549 if item_type != "message":
550 continue
551 content = output_item.get("content")
552 if not _is_object_list(content):
553 continue
554 for content_item in content:
555 if not _is_object_dict(content_item):
556 continue
557 if content_item.get("type") not in ("output_text", "text"):
558 continue
559 text = content_item.get("text")
560 if isinstance(text, str) and text:
561 texts.append(text)
562 return "".join(texts) if texts else None
564 @staticmethod
565 def get_config_model() -> type[GuardrailConfigModel[BaseModel]] | None:
566 from litellm.types.proxy.guardrails.guardrail_hooks.repelloai import (
567 RepelloAIGuardrailConfigModel,
568 )
570 return RepelloAIGuardrailConfigModel