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

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 

15 

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 

21 

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) 

48 

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 

51 

52 

53class CatoNetworksGuardrailMissingSecrets(Exception): 

54 pass 

55 

56 

57class _WsSslKwargs(TypedDict, total=False): 

58 ssl: bool | str | SSLContext 

59 

60 

61class _CatoRequiredAction(TypedDict, total=False): 

62 action_type: str 

63 detection_message: str 

64 

65 

66class _CatoRedactedMessage(TypedDict): 

67 role: NotRequired[str] 

68 content: str | None 

69 

70 

71class _CatoRedactedChat(TypedDict, total=False): 

72 all_redacted_messages: Sequence[_CatoRedactedMessage] 

73 

74 

75class _CatoAnalysisResult(TypedDict, total=False): 

76 policy_drill_down: Mapping[str, object] 

77 

78 

79class _CatoAnalyzeResponse(TypedDict): 

80 required_action: NotRequired[_CatoRequiredAction | None] 

81 analysis_result: NotRequired[_CatoAnalysisResult] 

82 redacted_chat: NotRequired[_CatoRedactedChat] 

83 

84 

85class _CatoOutputRedaction(TypedDict): 

86 redacted_output: str 

87 

88 

89class _CatoStreamMessage(TypedDict, total=False): 

90 verified_chunk: Mapping[str, object] 

91 done: bool 

92 blocking_message: str 

93 

94 

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 ] 

103 

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) 

130 

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} 

144 

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 

151 

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 

157 

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 ) 

176 

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 ) 

193 

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 

213 

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 [] 

223 

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") 

236 

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() 

248 

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)) 

265 

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 

289 

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 

325 

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) 

338 

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 

380 

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 

401 

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"] 

409 

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 

427 

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) 

456 

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 

464 

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) 

478 

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 ) 

511 

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 

528 

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 

536 

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) 

540 

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 

563 

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) 

570 

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 

591 

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 

613 

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 

621 

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) 

650 

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 

654 

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 

665 

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})) 

678 

679 @staticmethod 

680 def get_config_model() -> type["GuardrailConfigModel"] | None: 

681 from litellm.types.proxy.guardrails.guardrail_hooks.cato_networks import ( 

682 CatoNetworksGuardrailConfigModel, 

683 ) 

684 

685 return CatoNetworksGuardrailConfigModel