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

1import json 

2import os 

3from collections.abc import Mapping, Sequence 

4from types import MappingProxyType 

5from typing import Any, Final 

6from urllib.parse import urlparse 

7 

8import httpx 

9import pydantic 

10from typing_extensions import TypedDict, Unpack 

11 

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 

40 

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

46 

47 

48class _CustomGuardrailOptions(TypedDict, total=False, extra_items=object): 

49 """Base-class constructor options this guardrail forwards untouched to CustomGuardrail.""" 

50 

51 

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 ) 

78 

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

81 

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 

87 

88 self.timeout = _DEFAULT_TIMEOUT if timeout is None else timeout 

89 

90 self.async_handler = get_async_httpx_client( 

91 llm_provider=httpxSpecialProvider.GuardrailCallback, 

92 ) 

93 

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 ] 

102 

103 super().__init__(**kwargs) 

104 

105 @staticmethod 

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

107 from litellm.types.proxy.guardrails.guardrail_hooks.singulr import ( 

108 SingulrGuardrailConfigModel, 

109 ) 

110 

111 return SingulrGuardrailConfigModel 

112 

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 ) 

126 

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 

134 

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 

142 

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 

161 

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 

165 

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

176 

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) 

180 

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 

192 

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 

205 

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 

214 

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 

223 

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 ) 

236 

237 images: Final = inputs.get("images") 

238 tools: Final = inputs.get("tools") 

239 

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 

243 

244 metadata: Final = self._build_metadata(request_data=request_data) 

245 

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) 

257 

258 if guardrail_resp is None: 

259 return inputs 

260 

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 

269 

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

273 

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

278 

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

286 

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) 

289 

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) 

299 

300 if guardrail_resp is None: 

301 return 

302 

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 ) 

310 

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 

316 

317 metadata: Final = self._build_metadata(request_data=request_data) 

318 

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) 

327 

328 if guardrail_resp is None: 

329 return inputs 

330 

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 ) 

338 

339 return inputs 

340 

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 ) 

360 

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 

365 

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 ) 

372 

373 assistant_message: Final = AssistantMessage( 

374 role="assistant", 

375 content=combined_texts, 

376 tool_calls=tool_calls_res, 

377 ) 

378 

379 metadata: Final = self._build_metadata(request_data=request_data) 

380 

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 ) 

390 

391 payload = singulr_resp_obj.model_dump(mode="json") 

392 guardrail_resp = await self._call_api(payload) 

393 

394 if guardrail_resp is None: 

395 return inputs 

396 

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 

405 

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

416 

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 ) 

423 

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