Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/onyx/onyx.py: 33%

65 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1# +-------------------------------------------------------------+ 

2# 

3# Use Onyx Guardrails for your LLM calls 

4# https://onyx.security/ 

5# 

6# +-------------------------------------------------------------+ 

7import os 

8import uuid 

9from typing import TYPE_CHECKING, Final, Literal, Optional 

10 

11import httpx 

12from fastapi import HTTPException 

13 

14from litellm._logging import verbose_proxy_logger 

15from litellm.integrations.custom_guardrail import ( 

16 CustomGuardrail, 

17 log_guardrail_information, 

18) 

19from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj 

20from litellm.llms.custom_httpx.http_handler import ( 

21 get_async_httpx_client, 

22 httpxSpecialProvider, 

23) 

24from litellm.types.guardrails import GuardrailEventHooks 

25from litellm.types.utils import GenericGuardrailAPIInputs, ModelResponse 

26 

27if TYPE_CHECKING: 27 ↛ 28line 27 didn't jump to line 28 because the condition on line 27 was never true

28 from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel 

29 

30 

31class OnyxGuardrail(CustomGuardrail): 

32 @classmethod 

33 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: 

34 return [ 

35 GuardrailEventHooks.pre_call, 

36 GuardrailEventHooks.during_call, 

37 GuardrailEventHooks.post_call, 

38 ] 

39 

40 def __init__( 

41 self, 

42 api_base: str | None = None, 

43 api_key: str | None = None, 

44 timeout: float | None = 10.0, 

45 **kwargs, 

46 ): 

47 kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks())) 

48 timeout = timeout or int(os.getenv("ONYX_TIMEOUT", 10.0)) 

49 self.async_handler = get_async_httpx_client( 

50 llm_provider=httpxSpecialProvider.GuardrailCallback, 

51 params={"timeout": httpx.Timeout(timeout=timeout, connect=5.0)}, 

52 ) 

53 self.api_base = api_base or os.getenv( 

54 "ONYX_API_BASE", 

55 "https://ai-guard.onyx.security", 

56 ) 

57 self.api_key = api_key or os.getenv("ONYX_API_KEY") 

58 if not self.api_key: 

59 raise ValueError("ONYX_API_KEY environment variable is not set") 

60 self.optional_params = kwargs 

61 super().__init__(**kwargs) 

62 verbose_proxy_logger.info("OnyxGuard initialized with server: %s", self.api_base) 

63 

64 async def _validate_with_guard_server( 

65 self, 

66 payload: object, 

67 input_type: Literal["request", "response"], 

68 conversation_id: str, 

69 ) -> dict: 

70 """ 

71 Call external Onyx Guard server for validation 

72 """ 

73 response: Final = await self.async_handler.post( 

74 f"{self.api_base}/guard/evaluate/v1/{self.api_key}/litellm", 

75 json={ 

76 "payload": payload, 

77 "input_type": input_type, 

78 "conversation_id": conversation_id, 

79 }, 

80 headers={ 

81 "Content-Type": "application/json", 

82 }, 

83 ) 

84 response.raise_for_status() 

85 result: Final = response.json() 

86 if not result.get("allowed", True): 

87 detection_message = "Unknown violation" 

88 if "violated_rules" in result: 

89 detection_message = ", ".join(result["violated_rules"]) 

90 verbose_proxy_logger.warning("Request blocked by Onyx Guard. Violations: %s.", detection_message) 

91 raise HTTPException( 

92 status_code=400, 

93 detail=f"Request blocked by Onyx Guard. Violations: {detection_message}.", 

94 ) 

95 return result 

96 

97 @log_guardrail_information 

98 async def apply_guardrail( 

99 self, 

100 inputs: GenericGuardrailAPIInputs, 

101 request_data: dict, 

102 input_type: Literal["request", "response"], 

103 logging_obj: Optional["LiteLLMLoggingObj"] = None, 

104 ) -> GenericGuardrailAPIInputs: 

105 conversation_id: Final = logging_obj.litellm_call_id if logging_obj else str(uuid.uuid4()) 

106 

107 verbose_proxy_logger.info( 

108 "Running Onyx Guard apply_guardrail hook", 

109 extra={"conversation_id": conversation_id, "input_type": input_type}, 

110 ) 

111 payload = {} 

112 if input_type == "request": 

113 payload = request_data.get("proxy_server_request", {}) 

114 else: 

115 try: 

116 response: Final = ModelResponse(**request_data) 

117 parsed: Final = response.json() 

118 payload = parsed.get("response", {}) 

119 except Exception as e: 

120 verbose_proxy_logger.error( 

121 "Error in converting request_data to ModelResponse: %s", 

122 e, 

123 extra={ 

124 "conversation_id": conversation_id, 

125 "input_type": input_type, 

126 }, 

127 ) 

128 payload = request_data 

129 

130 try: 

131 await self._validate_with_guard_server(payload, input_type, conversation_id) 

132 return inputs 

133 except HTTPException as e: 

134 raise e 

135 except Exception as e: 

136 verbose_proxy_logger.error( 

137 "Error in apply_guardrail guard: %s", 

138 e, 

139 extra={"conversation_id": conversation_id, "input_type": input_type}, 

140 ) 

141 return inputs 

142 

143 @staticmethod 

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

145 from litellm.types.proxy.guardrails.guardrail_hooks.onyx import ( 

146 OnyxGuardrailConfigModel, 

147 ) 

148 

149 return OnyxGuardrailConfigModel