Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/cato_networks/cato_networks.py: 18%
380 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
1# +-------------------------------------------------------------+
2#
3# Use Cato Networks Guardrails for your LLM calls
4# https://www.catonetworks.com/
5#
6# +-------------------------------------------------------------+
7import asyncio
8import contextlib
9import json
10import os
11import ssl
12from collections.abc import AsyncGenerator, AsyncIterable, Mapping, Sequence
13from ssl import SSLContext
14from typing import TYPE_CHECKING, Any, Final
16from fastapi import HTTPException
17from pydantic import BaseModel
18from typing_extensions import NotRequired, TypedDict
19from websockets.asyncio.client import ClientConnection, connect
20from websockets.exceptions import ConnectionClosed
22from litellm import DualCache
23from litellm._logging import verbose_proxy_logger
24from litellm._version import version as litellm_version
25from litellm.integrations.custom_guardrail import CustomGuardrail
26from litellm.llms.custom_httpx.http_handler import (
27 get_async_httpx_client,
28 get_ssl_configuration,
29 httpxSpecialProvider,
30)
31from litellm.proxy._types import UserAPIKeyAuth
32from litellm.proxy.guardrails._content_utils import (
33 apply_redacted_messages_back,
34 build_inspection_messages,
35 is_non_conversational_call_type,
36 is_string_batch_input,
37)
38from litellm.types.guardrails import GuardrailEventHooks
39from litellm.types.utils import (
40 CallTypesLiteral,
41 Choices,
42 LLMResponseTypes,
43 Message,
44 ModelResponse,
45 ModelResponseStream,
46 ResponsesAPIResponse,
47)
49if TYPE_CHECKING: 49 ↛ 50line 49 didn't jump to line 50 because the condition on line 49 was never true
50 from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
53class CatoNetworksGuardrailMissingSecrets(Exception):
54 pass
57class _WsSslKwargs(TypedDict, total=False):
58 ssl: bool | str | SSLContext
61class _CatoRequiredAction(TypedDict, total=False):
62 action_type: str
63 detection_message: str
66class _CatoRedactedMessage(TypedDict):
67 role: NotRequired[str]
68 content: str | None
71class _CatoRedactedChat(TypedDict, total=False):
72 all_redacted_messages: Sequence[_CatoRedactedMessage]
75class _CatoAnalysisResult(TypedDict, total=False):
76 policy_drill_down: Mapping[str, object]
79class _CatoAnalyzeResponse(TypedDict):
80 required_action: NotRequired[_CatoRequiredAction | None]
81 analysis_result: NotRequired[_CatoAnalysisResult]
82 redacted_chat: NotRequired[_CatoRedactedChat]
85class _CatoOutputRedaction(TypedDict):
86 redacted_output: str
89class _CatoStreamMessage(TypedDict, total=False):
90 verified_chunk: Mapping[str, object]
91 done: bool
92 blocking_message: str
95class CatoNetworksGuardrail(CustomGuardrail):
96 @classmethod
97 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
98 return [
99 GuardrailEventHooks.pre_call,
100 GuardrailEventHooks.during_call,
101 GuardrailEventHooks.post_call,
102 ]
104 def __init__(
105 self,
106 api_key: str | None = None,
107 api_base: str | None = None,
108 inspect_embeddings: bool | None = None,
109 **kwargs,
110 ):
111 kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
112 self.inspect_embeddings: Final = inspect_embeddings is True
113 ssl_verify: Final = kwargs.pop("ssl_verify", None)
114 self.async_handler = get_async_httpx_client(
115 llm_provider=httpxSpecialProvider.GuardrailCallback,
116 params={"ssl_verify": ssl_verify} if ssl_verify is not None else None,
117 )
118 self.api_key = api_key or os.environ.get("CATO_API_KEY")
119 if not self.api_key:
120 msg: Final = (
121 "Couldn't get Cato Networks api key, either set the `CATO_API_KEY` in the environment or "
122 "pass it as a parameter to the guardrail in the config file"
123 )
124 raise CatoNetworksGuardrailMissingSecrets(msg)
125 self.api_base = api_base or os.environ.get("CATO_API_BASE") or "https://api.aisec.catonetworks.com"
126 self.api_base = self.api_base.rstrip("/")
127 self.ws_api_base = self.api_base.replace("http://", "ws://").replace("https://", "wss://")
128 self._ws_connect_ssl_kwargs = self._build_ws_ssl_kwargs(ssl_verify, self.ws_api_base)
129 super().__init__(**kwargs)
131 @staticmethod
132 def _build_ws_ssl_kwargs(ssl_verify: bool | str | None, ws_api_base: str) -> _WsSslKwargs:
133 """Resolve the ``ssl`` argument for ``websockets.connect``. Mirrors the
134 ``ssl_verify`` handling applied to the HTTP handler so a custom Cato instance
135 behind TLS honours the same verification settings for streaming."""
136 if ssl_verify is None or not ws_api_base.startswith("wss://"):
137 return {}
138 ssl_config = get_ssl_configuration(ssl_verify)
139 if ssl_config is False:
140 ssl_config = ssl.create_default_context()
141 ssl_config.check_hostname = False
142 ssl_config.verify_mode = ssl.CERT_NONE
143 return {"ssl": ssl_config}
145 @staticmethod
146 def _resolve_cato_user_email(user_api_key_dict: UserAPIKeyAuth) -> str | None:
147 """Only the key/JWT-bound user email is trusted. ``end_user_id`` is derived from
148 caller-supplied request fields (OpenAI ``user``, headers, metadata) and is spoofable,
149 so it must never be forwarded as the Cato user identity."""
150 return user_api_key_dict.user_email
152 @staticmethod
153 async def _cancel_background_task(task: asyncio.Task) -> None:
154 task.cancel()
155 with contextlib.suppress(asyncio.CancelledError, Exception):
156 await task
158 async def async_pre_call_hook(
159 self,
160 user_api_key_dict: UserAPIKeyAuth,
161 cache: DualCache,
162 data: dict,
163 call_type: CallTypesLiteral,
164 ) -> Exception | str | dict | None:
165 verbose_proxy_logger.debug("Inside Cato Pre-Call Hook")
166 # /embeddings carries documents being indexed, not a conversation to inspect.
167 if is_non_conversational_call_type(call_type) and not self.inspect_embeddings:
168 verbose_proxy_logger.debug("Cato: skipping non-conversational call type %s", call_type)
169 return data
170 return await self.call_cato_guardrail(
171 data,
172 hook="pre_call",
173 key_alias=user_api_key_dict.key_alias,
174 user_email=self._resolve_cato_user_email(user_api_key_dict),
175 )
177 async def async_moderation_hook(
178 self,
179 data: dict,
180 user_api_key_dict: UserAPIKeyAuth,
181 call_type: CallTypesLiteral,
182 ) -> Exception | str | dict | None:
183 verbose_proxy_logger.debug("Inside Cato Moderation Hook")
184 if is_non_conversational_call_type(call_type) and not self.inspect_embeddings:
185 verbose_proxy_logger.debug("Cato: skipping non-conversational call type %s", call_type)
186 return data
187 return await self.call_cato_guardrail(
188 data,
189 hook="moderation",
190 key_alias=user_api_key_dict.key_alias,
191 user_email=self._resolve_cato_user_email(user_api_key_dict),
192 )
194 @classmethod
195 def _inspection_messages(cls, data: dict) -> list:
196 """Flatten multimodal list ``content`` into plain text so Cato inspects
197 every text fragment. Chat ``messages`` stay 1:1 with the request so
198 redacted results map back by index, and every other field the proxy
199 forwards to the model (Responses-API ``input``/``instructions``, legacy
200 completion ``prompt`` and tool/function/``response_format`` schema strings)
201 is appended as synthetic messages so blocked text cannot bypass inspection
202 by hiding in one of them."""
203 flattened: Final = []
204 for message in data.get("messages") or []:
205 if isinstance(message, dict) and isinstance(message.get("content"), list):
206 parts = build_inspection_messages({"messages": [message]})
207 flattened.append({**message, "content": parts[0]["content"] if parts else ""})
208 else:
209 flattened.append(message)
210 for _field, messages in cls._extra_inspection_sources(data):
211 flattened.extend(messages)
212 return flattened
214 @staticmethod
215 def _prompt_inspection_messages(prompt: object) -> Sequence[Mapping[str, str]]:
216 """Synthetic user messages for a legacy completion ``prompt`` (a string
217 or a list of string prompts)."""
218 if isinstance(prompt, str):
219 return [{"role": "user", "content": prompt}] if prompt else []
220 if isinstance(prompt, list):
221 return [{"role": "user", "content": part} for part in prompt if isinstance(part, str) and part]
222 return []
224 @staticmethod
225 def _iter_schema_string_refs(data: Mapping[str, Any]):
226 """Yield ``(container, key)`` for every non-empty schema string the proxy
227 forwards to the model inside tool/function and structured-output schemas:
228 each ``tools[].function`` and legacy ``functions[]`` entry plus the
229 ``response_format`` JSON schema, walked recursively for the free-text and
230 value strings a caller could hide blocked text in (``description``,
231 ``title``, ``const``, ``default`` and every ``enum``/``examples`` item).
232 Blocked text in any of them must be inspected and redacted like any other
233 prompt."""
234 scalar_keys: Final = ("description", "title", "const", "default")
235 list_keys: Final = ("enum", "examples")
237 stack: Final[list] = []
238 for tool in data.get("tools") or []:
239 if isinstance(tool, dict) and isinstance(tool.get("function"), dict):
240 stack.append(tool["function"])
241 for function in data.get("functions") or []:
242 if isinstance(function, dict):
243 stack.append(function)
244 response_format: Final = data.get("response_format")
245 if isinstance(response_format, dict):
246 stack.append(response_format)
247 stack.reverse()
249 while stack:
250 node = stack.pop()
251 if isinstance(node, dict):
252 for key in scalar_keys:
253 value = node.get(key)
254 if isinstance(value, str) and value:
255 yield node, key
256 for key in list_keys:
257 items = node.get(key)
258 if isinstance(items, list):
259 for idx, item in enumerate(items):
260 if isinstance(item, str) and item:
261 yield items, idx
262 stack.extend(reversed(list(node.values())))
263 elif isinstance(node, list):
264 stack.extend(reversed(node))
266 @classmethod
267 def _extra_inspection_sources(cls, data: Mapping[str, object]) -> Sequence[tuple[str, Sequence[Mapping[str, str]]]]:
268 """Text the proxy forwards to the model outside chat ``messages``:
269 Responses-API ``input`` and ``instructions``, legacy completion
270 ``prompt`` and tool/function/``response_format`` schema strings. Returned
271 as ``(field, messages)`` in a fixed order so the anonymize path can slice
272 redactions back to the field they came from."""
273 sources: Final[list] = []
274 input_messages: Final = build_inspection_messages({"input": data.get("input")})
275 if input_messages:
276 sources.append(("input", input_messages))
277 instructions: Final = data.get("instructions")
278 if isinstance(instructions, str) and instructions:
279 sources.append(("instructions", [{"role": "system", "content": instructions}]))
280 prompt_messages: Final = cls._prompt_inspection_messages(data.get("prompt"))
281 if prompt_messages:
282 sources.append(("prompt", prompt_messages))
283 schema_strings: Final = [
284 {"role": "system", "content": container[key]} for container, key in cls._iter_schema_string_refs(data)
285 ]
286 if schema_strings:
287 sources.append(("schema_strings", schema_strings))
288 return sources
290 async def call_cato_guardrail(
291 self,
292 data: dict,
293 hook: str,
294 key_alias: str | None,
295 user_email: str | None = None,
296 ) -> dict:
297 call_id: Final = data.get("litellm_call_id")
298 headers: Final = self._build_cato_headers(
299 hook=hook,
300 key_alias=key_alias,
301 user_email=user_email,
302 litellm_call_id=call_id,
303 )
304 response: Final = await self.async_handler.post(
305 f"{self.api_base}/fw/v1/analyze",
306 headers=headers,
307 json={"messages": self._inspection_messages(data)},
308 )
309 response.raise_for_status()
310 res: Final[_CatoAnalyzeResponse] = response.json()
311 required_action: Final = res.get("required_action")
312 action_type: Final = required_action and required_action.get("action_type", None)
313 if action_type is None:
314 verbose_proxy_logger.debug("Cato: No required action specified")
315 return data
316 if action_type == "monitor_action":
317 verbose_proxy_logger.info("Cato: monitor action")
318 elif action_type == "block_action" and required_action is not None:
319 self._handle_block_action(res.get("analysis_result", {}), required_action)
320 elif action_type == "anonymize_action":
321 return self._anonymize_request(res, data)
322 else:
323 verbose_proxy_logger.error("Cato: %s action", action_type)
324 return data
326 def _handle_block_action(
327 self,
328 analysis_result: _CatoAnalysisResult,
329 required_action: _CatoRequiredAction,
330 ) -> None:
331 detection_message: Final = required_action.get("detection_message", None)
332 verbose_proxy_logger.info(
333 "Cato: Violation detected enabled policies: {policies}".format(
334 policies=list(analysis_result.get("policy_drill_down", {}).keys()),
335 ),
336 )
337 raise HTTPException(status_code=400, detail=detection_message)
339 def _anonymize_request(self, res: _CatoAnalyzeResponse, data: dict) -> dict:
340 verbose_proxy_logger.info("Cato: anonymize action")
341 redacted_chat: Final = res.get("redacted_chat")
342 if not redacted_chat:
343 return data
344 redacted_messages: Final = redacted_chat.get("all_redacted_messages") or []
345 original_messages: Final = data.get("messages")
346 sources: Final = self._extra_inspection_sources(data)
347 if is_string_batch_input(data) and len(redacted_messages) != sum(len(messages) for _, messages in sources):
348 raise HTTPException(
349 status_code=400,
350 detail=(
351 "Cato: anonymize action returned a redacted batch of a different "
352 "size than the inspected input, so the request cannot be rewritten "
353 "without forwarding unredacted text."
354 ),
355 )
356 offset = 0
357 if original_messages:
358 data["messages"] = [
359 (
360 {**original, "content": redacted_messages[idx]["content"]}
361 if idx < len(redacted_messages) and redacted_messages[idx].get("content") is not None
362 else original
363 )
364 for idx, original in enumerate(original_messages)
365 ]
366 offset = len(original_messages)
367 for field, messages in sources:
368 redacted_slice = redacted_messages[offset : offset + len(messages)]
369 offset += len(messages)
370 if not self._apply_extra_redaction(data, field, redacted_slice):
371 raise HTTPException(
372 status_code=400,
373 detail=(
374 "Cato: anonymize action returned a redacted batch of a different "
375 "size than the inspected input, so the request cannot be rewritten "
376 "without forwarding unredacted text."
377 ),
378 )
379 return data
381 @classmethod
382 def _apply_extra_redaction(cls, data: dict, field: str, redacted: Sequence[Mapping[str, object]]) -> bool:
383 if field == "input":
384 input_only: Final = {"input": data["input"]}
385 if not redacted:
386 return not is_string_batch_input(input_only)
387 if not apply_redacted_messages_back(input_only, redacted):
388 return False
389 data["input"] = input_only["input"]
390 return True
391 if not redacted:
392 return True
393 if field == "instructions":
394 if redacted[0].get("content") is not None:
395 data["instructions"] = redacted[0]["content"]
396 elif field == "prompt":
397 cls._apply_prompt_redaction(data, redacted)
398 elif field == "schema_strings":
399 cls._apply_schema_string_redaction(data, redacted)
400 return True
402 @classmethod
403 def _apply_schema_string_redaction(cls, data: dict, redacted: Sequence[Mapping[str, object]]) -> None:
404 redactions: Final = iter(redacted)
405 for container, key in cls._iter_schema_string_refs(data):
406 replacement = next(redactions, None)
407 if replacement is not None and replacement.get("content") is not None:
408 container[key] = replacement["content"]
410 @staticmethod
411 def _apply_prompt_redaction(data: dict, redacted: Sequence[Mapping[str, object]]) -> None:
412 contents: Final = [m.get("content") for m in redacted if isinstance(m, dict)]
413 prompt: Final = data.get("prompt")
414 if isinstance(prompt, str):
415 if contents and contents[0] is not None:
416 data["prompt"] = contents[0]
417 return
418 if isinstance(prompt, list):
419 new_prompt: Final = list(prompt)
420 redactions: Final = iter(contents)
421 for idx, part in enumerate(new_prompt):
422 if isinstance(part, str) and part:
423 replacement = next(redactions, None)
424 if replacement is not None:
425 new_prompt[idx] = replacement
426 data["prompt"] = new_prompt
428 async def call_cato_guardrail_on_output(
429 self,
430 request_data: dict,
431 output: str,
432 hook: str,
433 key_alias: str | None,
434 user_email: str | None = None,
435 ) -> _CatoOutputRedaction | None:
436 call_id: Final = request_data.get("litellm_call_id")
437 inspection_messages: Final = self._inspection_messages(request_data)
438 assistant_index: Final = len(inspection_messages)
439 response: Final = await self.async_handler.post(
440 f"{self.api_base}/fw/v1/analyze",
441 headers=self._build_cato_headers(
442 hook=hook,
443 key_alias=key_alias,
444 user_email=user_email,
445 litellm_call_id=call_id,
446 ),
447 json={"messages": inspection_messages + [{"role": "assistant", "content": output}]},
448 )
449 response.raise_for_status()
450 res: Final[_CatoAnalyzeResponse] = response.json()
451 required_action: Final = res.get("required_action")
452 action_type: Final = required_action and required_action.get("action_type", None)
453 if action_type == "block_action" and required_action is not None:
454 self._handle_block_action_on_output(res.get("analysis_result", {}), required_action)
455 redacted_chat: Final = res.get("redacted_chat", None)
457 if action_type and action_type == "anonymize_action" and redacted_chat:
458 all_redacted: Final = redacted_chat.get("all_redacted_messages") or []
459 if assistant_index < len(all_redacted):
460 redacted_output: Final = all_redacted[assistant_index].get("content")
461 if redacted_output is not None:
462 return {"redacted_output": redacted_output}
463 return None
465 def _handle_block_action_on_output(
466 self,
467 analysis_result: _CatoAnalysisResult,
468 required_action: _CatoRequiredAction,
469 ) -> None:
470 detection_message: Final = required_action.get("detection_message", None)
471 verbose_proxy_logger.info(
472 "Cato: detected: {detected}, enabled policies: {policies}".format(
473 detected=True,
474 policies=list(analysis_result.get("policy_drill_down", {}).keys()),
475 ),
476 )
477 raise HTTPException(status_code=400, detail=detection_message)
479 def _build_cato_headers(
480 self,
481 *,
482 hook: str,
483 key_alias: str | None,
484 user_email: str | None,
485 litellm_call_id: str | None,
486 ):
487 """
488 A helper function to build the http headers that are required by Cato guardrails.
489 """
490 return (
491 {
492 "Authorization": f"Bearer {self.api_key}",
493 # Used by Cato Networks to apply only the guardrails that should be applied in a specific request phase.
494 "x-cato-litellm-hook": hook,
495 # Used by Cato Networks to track LiteLLM version and provide backward compatibility.
496 "x-cato-litellm-version": litellm_version,
497 }
498 # Used by Cato Networks to track together single call input and output
499 | ({"x-cato-call-id": litellm_call_id} if litellm_call_id else {})
500 # Used by Cato Networks to track guardrails violations by user.
501 | ({"x-cato-user-email": user_email} if user_email else {})
502 | (
503 {
504 # Used by Cato Networks apply only the guardrails that are associated with the key alias.
505 "x-cato-gateway-key-alias": key_alias,
506 }
507 if key_alias
508 else {}
509 )
510 )
512 @staticmethod
513 def _output_fragments(message: Message) -> Sequence[tuple[tuple[str, int | None], str]]:
514 """Assistant text the proxy returns to the caller: ``content`` plus every
515 ``tool_calls[].function.arguments`` string, each tagged with where a
516 redaction must be written back. ``content`` is only included when present
517 so a tool-call-only choice keeps its ``None`` content (the text-vs-tool-call
518 signal downstream consumers rely on) while its arguments are still inspected."""
519 fragments: Final[list] = []
520 if message.content is not None:
521 fragments.append((("content", None), message.content))
522 for idx, tool_call in enumerate(message.tool_calls or []):
523 function = getattr(tool_call, "function", None)
524 arguments = getattr(function, "arguments", None)
525 if isinstance(arguments, str) and arguments:
526 fragments.append((("tool_call", idx), arguments))
527 return fragments
529 @staticmethod
530 def _apply_output_fragment(message: Any, target: tuple[str, int | None], redacted: str) -> None:
531 kind, idx = target
532 if kind == "content":
533 message.content = redacted
534 else:
535 message.tool_calls[idx].function.arguments = redacted
537 @staticmethod
538 def _responses_output_field(item: object, key: str) -> str | Sequence[object] | None:
539 return item.get(key) if isinstance(item, dict) else getattr(item, key, None)
541 @classmethod
542 def _responses_output_fragments(cls, response: ResponsesAPIResponse) -> Sequence[tuple[object, str, str]]:
543 """Assistant text the Responses API returns to the caller: every
544 ``output_text`` content block plus every function-call ``arguments``
545 string, each paired with the ``(container, key)`` a Cato redaction is
546 written back to. Output items and their content may be pydantic objects
547 or plain dicts, so both access patterns are handled."""
548 fragments: Final[list] = []
549 for item in response.output or []:
550 item_type = cls._responses_output_field(item, "type")
551 if item_type == "function_call":
552 arguments = cls._responses_output_field(item, "arguments")
553 if isinstance(arguments, str) and arguments:
554 fragments.append((item, "arguments", arguments))
555 elif item_type == "message":
556 for content in cls._responses_output_field(item, "content") or []:
557 if cls._responses_output_field(content, "type") != "output_text":
558 continue
559 text = cls._responses_output_field(content, "text")
560 if isinstance(text, str) and text:
561 fragments.append((content, "text", text))
562 return fragments
564 @staticmethod
565 def _apply_responses_output_fragment(container: object, key: str, redacted: str) -> None:
566 if isinstance(container, dict):
567 container[key] = redacted
568 else:
569 setattr(container, key, redacted)
571 async def _inspect_output_text(
572 self,
573 data: dict,
574 text: str,
575 user_api_key_dict: UserAPIKeyAuth,
576 user_email: str | None,
577 ) -> str | None:
578 """Run the Cato output guardrail on a single assistant text fragment.
579 Raises on a block action and returns the redacted replacement, or
580 ``None`` when the fragment must be left unchanged."""
581 cato_output_guardrail_result: Final = await self.call_cato_guardrail_on_output(
582 data,
583 text,
584 hook="output",
585 key_alias=user_api_key_dict.key_alias,
586 user_email=user_email,
587 )
588 if cato_output_guardrail_result:
589 return cato_output_guardrail_result.get("redacted_output")
590 return None
592 async def async_post_call_success_hook(
593 self,
594 data: dict,
595 user_api_key_dict: UserAPIKeyAuth,
596 response: LLMResponseTypes,
597 ) -> LLMResponseTypes:
598 user_email: Final = self._resolve_cato_user_email(user_api_key_dict)
599 if isinstance(response, ModelResponse) and response.choices:
600 for choice in response.choices:
601 if not isinstance(choice, Choices):
602 continue
603 for target, text in self._output_fragments(choice.message):
604 redacted_output = await self._inspect_output_text(data, text, user_api_key_dict, user_email)
605 if redacted_output is not None:
606 self._apply_output_fragment(choice.message, target, redacted_output)
607 elif isinstance(response, ResponsesAPIResponse):
608 for container, key, text in self._responses_output_fragments(response):
609 redacted_output = await self._inspect_output_text(data, text, user_api_key_dict, user_email)
610 if redacted_output is not None:
611 self._apply_responses_output_fragment(container, key, redacted_output)
612 return response
614 async def async_post_call_streaming_iterator_hook(
615 self,
616 user_api_key_dict: UserAPIKeyAuth,
617 response: AsyncIterable[object],
618 request_data: dict,
619 ) -> AsyncGenerator[ModelResponseStream, None]:
620 from litellm.proxy.proxy_server import StreamingCallbackError
622 user_email: Final = self._resolve_cato_user_email(user_api_key_dict)
623 call_id: Final = request_data.get("litellm_call_id")
624 async with connect(
625 f"{self.ws_api_base}/fw/v1/analyze/stream",
626 additional_headers=self._build_cato_headers(
627 hook="output",
628 key_alias=user_api_key_dict.key_alias,
629 user_email=user_email,
630 litellm_call_id=call_id,
631 ),
632 **self._ws_connect_ssl_kwargs,
633 ) as websocket:
634 sender: Final = asyncio.create_task(self.forward_the_stream_to_cato(websocket, response))
635 try:
636 while True:
637 raw_message = await self._await_cato_message(websocket, sender)
638 result: _CatoStreamMessage = json.loads(raw_message)
639 if verified_chunk := result.get("verified_chunk"):
640 yield ModelResponseStream.model_validate(verified_chunk)
641 continue
642 if result.get("done"):
643 return
644 if blocking_message := result.get("blocking_message"):
645 raise StreamingCallbackError(blocking_message)
646 verbose_proxy_logger.error("Unknown message received from Cato: %s", result)
647 return
648 finally:
649 await self._cancel_background_task(sender)
651 async def _await_cato_message(self, websocket: ClientConnection, sender: asyncio.Task[None]) -> str | bytes:
652 """Wait for the next Cato message, surfacing a dead forwarding task instead of blocking."""
653 from litellm.proxy.proxy_server import StreamingCallbackError
655 recv_task: Final = asyncio.ensure_future(websocket.recv())
656 pending: Final = {recv_task, sender} if not sender.done() else {recv_task}
657 await asyncio.wait(pending, return_when=asyncio.FIRST_COMPLETED)
658 if sender.done() and (sender_exc := sender.exception()) is not None:
659 await self._cancel_background_task(recv_task)
660 raise StreamingCallbackError("Cato guardrail upstream stream failed") from sender_exc
661 try:
662 return await recv_task
663 except ConnectionClosed as exc:
664 raise StreamingCallbackError("Cato guardrail connection closed unexpectedly") from exc
666 async def forward_the_stream_to_cato(
667 self,
668 websocket: ClientConnection,
669 response_iter: AsyncIterable[object],
670 ) -> None:
671 async for chunk in response_iter:
672 if isinstance(chunk, BaseModel):
673 chunk = chunk.model_dump_json()
674 elif not isinstance(chunk, (str, bytes)):
675 chunk = json.dumps(chunk)
676 await websocket.send(chunk)
677 await websocket.send(json.dumps({"done": True}))
679 @staticmethod
680 def get_config_model() -> type["GuardrailConfigModel"] | None:
681 from litellm.types.proxy.guardrails.guardrail_hooks.cato_networks import (
682 CatoNetworksGuardrailConfigModel,
683 )
685 return CatoNetworksGuardrailConfigModel