Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py: 30%
176 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"""TypeSafe (Jev) relevance-based compaction guardrail.
3Instead of summarizing tool output, the guardrail asks TypeSafe's Jev model
4one yes/no question per completed tool exchange ("is this result still needed
5for the current task?") over ``POST {api_base}/v1/systemone`` and blanks the
6tool results Jev judges no longer relevant.
7"""
9from __future__ import annotations
11import asyncio
12import time
13from collections.abc import Mapping, Sequence
14from typing import TYPE_CHECKING, Annotated, Final, Literal
16import httpx
17from fastapi import HTTPException
18from httpx import Response as HttpxResponse
19from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError
21from litellm._logging import verbose_proxy_logger
22from litellm.compression.compress import get_protected_indices
23from litellm.integrations.custom_guardrail import (
24 CustomGuardrail,
25 log_guardrail_information, # pyright: ignore[reportUnknownVariableType] # decorator is untyped in custom_guardrail
26)
27from litellm.litellm_core_utils.prompt_templates.factory import group_tool_exchanges
28from litellm.llms.custom_httpx.http_handler import (
29 AsyncHTTPHandler,
30 get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # helper is untyped in http_handler
31 httpxSpecialProvider,
32)
33from litellm.proxy.guardrails.guardrail_hooks.content_text import content_to_text
34from litellm.secret_managers.main import get_secret_str
35from litellm.types.guardrails import GuardrailEventHooks, Mode
36from litellm.types.utils import GenericGuardrailAPIInputs
38if TYPE_CHECKING: 38 ↛ 39line 38 didn't jump to line 39 because the condition on line 38 was never true
39 from litellm.litellm_core_utils.litellm_logging import (
40 Logging as LiteLLMLoggingObj,
41 )
42 from litellm.types.proxy.guardrails.guardrail_hooks.typesafe import (
43 TypeSafeGuardrailConfigModel,
44 )
46DEFAULT_API_BASE: Final = "https://api.typesafe.ai"
47DEFAULT_MODEL: Final = "jev-latest"
48DEFAULT_RELEVANCE_THRESHOLD: Final = 0.2
49DEFAULT_MIN_CHARS_TO_EVALUATE: Final = 200
50DEFAULT_MAX_RESULT_CHARS_IN_STATE: Final = 4000
51_MAX_EXCHANGES_EVALUATED: Final = 200
52_JEV_TIMEOUT_SECONDS: Final = 30.0
53DROPPED_RESULT_TEXT: Final = (
54 "[Tool result removed by TypeSafe compaction: judged no longer relevant to the current task]"
55)
56_ELISION_MARKER: Final = "\n... [middle truncated] ...\n"
59_STR_OBJECT_DICT_ADAPTER: Final = TypeAdapter(dict[str, object])
60_OBJECT_LIST_ADAPTER: Final = TypeAdapter(list[object])
63def _as_str_object_dict(value: object) -> dict[str, object] | None:
64 try:
65 return _STR_OBJECT_DICT_ADAPTER.validate_python(value)
66 except ValidationError:
67 return None
70def _as_object_list(value: object) -> list[object] | None:
71 try:
72 return _OBJECT_LIST_ADAPTER.validate_python(value)
73 except ValidationError:
74 return None
77def _safe_response_text(response: HttpxResponse | None, limit: int = 500) -> str:
78 if response is None:
79 return ""
80 try:
81 text: Final = response.text
82 except httpx.DecodingError:
83 return "<undecodable response body>"
84 return (text or "")[:limit]
87class _JevNoulAnswer(BaseModel):
88 model_config = ConfigDict(frozen=True, allow_inf_nan=False)
90 type: Literal["noul"]
91 noul: Annotated[float, Field(ge=0.0, le=1.0)]
94class _JevSystemOneResponse(BaseModel):
95 model_config = ConfigDict(frozen=True)
97 answers: Mapping[str, _JevNoulAnswer]
100_JEV_RESPONSE_ADAPTER: Final = TypeAdapter(_JevSystemOneResponse)
103def _truncate_for_state(text: str, max_chars: int) -> str:
104 """Keeps the head and tail within ``max_chars`` so Jev sees both ends of a long result."""
105 if len(text) <= max_chars:
106 return text
107 if max_chars <= len(_ELISION_MARKER):
108 return text[:max_chars]
109 budget: Final = max_chars - len(_ELISION_MARKER)
110 head: Final = budget // 2
111 return text[:head] + _ELISION_MARKER + text[len(text) - (budget - head) :]
114def _question_instructions(question_id: str) -> str:
115 return (
116 f"Is tool exchange `{question_id}` in `tool_exchanges` still needed by the assistant to "
117 "complete `task`? Answer yes if its result contains information the assistant has not yet "
118 "fully used or will need again; answer no if it is off-topic, superseded, or already "
119 "incorporated into later messages."
120 )
123def _tool_call_entry(tool_call: object) -> dict[str, object] | None:
124 parsed_call = _as_str_object_dict(tool_call)
125 if parsed_call is None:
126 return None
127 function = _as_str_object_dict(parsed_call.get("function"))
128 fn = function if function is not None else parsed_call
129 return {"name": fn.get("name"), "arguments": fn.get("arguments")} # mutable-ok: serialized to JSON
132def _tool_call_entries(assistant_message: Mapping[str, object]) -> tuple[dict[str, object], ...]:
133 tool_calls: Final = _as_object_list(assistant_message.get("tool_calls"))
134 if tool_calls is None:
135 return ()
136 return tuple(entry for tool_call in tool_calls if (entry := _tool_call_entry(tool_call)) is not None)
139def _protected_indices(messages: Sequence[Mapping[str, object]]) -> frozenset[int]:
140 """``get_protected_indices`` expanded over whole tool exchanges, so the most recent exchange is never evaluated."""
141 protected: Final = frozenset(get_protected_indices(messages))
142 return protected | frozenset(
143 index
144 for group in group_tool_exchanges(messages)
145 if any(member in protected for member in group)
146 for index in group
147 )
150class TypeSafeGuardrail(CustomGuardrail):
151 def __init__(
152 self,
153 api_base: str | None = None,
154 api_key: str | None = None,
155 model: str | None = None,
156 relevance_threshold: float | None = None,
157 min_chars_to_evaluate: int | None = None,
158 max_result_chars_in_state: int | None = None,
159 unreachable_fallback: str | None = None,
160 guardrail_name: str | None = None,
161 event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None = None,
162 default_on: bool = False,
163 async_handler: AsyncHTTPHandler | None = None,
164 ) -> None:
165 raw_api_base: Final = (api_base or get_secret_str("TYPESAFE_API_BASE") or DEFAULT_API_BASE).rstrip("/")
166 self.typesafe_api_base = raw_api_base
167 self.typesafe_api_key = api_key or get_secret_str("TYPESAFE_API_KEY")
168 if not self.typesafe_api_key:
169 raise ValueError(
170 "TypeSafe guardrail requires an API key. Set `api_key` in the "
171 "guardrail config or the TYPESAFE_API_KEY env var."
172 )
173 self.jev_model = model or DEFAULT_MODEL
174 self.relevance_threshold = DEFAULT_RELEVANCE_THRESHOLD if relevance_threshold is None else relevance_threshold
175 self.min_chars_to_evaluate = (
176 DEFAULT_MIN_CHARS_TO_EVALUATE if min_chars_to_evaluate is None else min_chars_to_evaluate
177 )
178 self.max_result_chars_in_state = (
179 DEFAULT_MAX_RESULT_CHARS_IN_STATE if max_result_chars_in_state is None else max_result_chars_in_state
180 )
181 self.unreachable_fallback: Literal["fail_closed", "fail_open"] = (
182 "fail_closed" if unreachable_fallback == "fail_closed" else "fail_open"
183 )
184 self.async_handler: AsyncHTTPHandler = async_handler or get_async_httpx_client(
185 llm_provider=httpxSpecialProvider.GuardrailCallback,
186 )
187 super().__init__( # pyright: ignore[reportUnknownMemberType] # CustomGuardrail.__init__ is untyped
188 guardrail_name=guardrail_name,
189 event_hook=event_hook,
190 default_on=default_on,
191 )
193 def _handle_failure(self, error: str, log_detail: dict[str, object]) -> None:
194 """fail_open logs and returns; fail_closed raises a generic 502 (upstream bodies stay in server logs)."""
195 if self.unreachable_fallback == "fail_open":
196 verbose_proxy_logger.warning(
197 "TypeSafe: %s; fail_open configured, forwarding request uncompacted. detail=%s",
198 error,
199 log_detail,
200 )
201 return
202 verbose_proxy_logger.error("TypeSafe: %s. detail=%s", error, log_detail)
203 raise HTTPException(status_code=502, detail={"error": error}) # mutable-ok: FastAPI wants a dict detail
205 def _candidate_exchanges(self, messages: Sequence[dict[str, object]]) -> tuple[tuple[int, ...], ...]:
206 """Completed tool exchanges eligible for evaluation: unprotected, and long enough to be worth a call."""
207 protected: Final = _protected_indices(messages)
208 candidates: Final = tuple(
209 group
210 for group in group_tool_exchanges(messages)
211 if len(group) >= 2
212 and messages[group[0]].get("role") == "assistant"
213 and not any(member in protected for member in group)
214 and len(self._exchange_tool_text(messages, group)) >= self.min_chars_to_evaluate
215 )
216 return candidates[-_MAX_EXCHANGES_EVALUATED:]
218 @staticmethod
219 def _exchange_tool_text(messages: Sequence[dict[str, object]], group: tuple[int, ...]) -> str:
220 return "".join(
221 content_to_text(messages[index].get("content"))
222 for index in group[1:]
223 if messages[index].get("role") in ("tool", "function")
224 )
226 def _build_state(
227 self, messages: Sequence[dict[str, object]], candidates: tuple[tuple[int, ...], ...]
228 ) -> dict[str, object]:
229 task: Final = next(
230 (
231 content_to_text(messages[index].get("content"))
232 for index in range(len(messages) - 1, -1, -1)
233 if messages[index].get("role") == "user"
234 ),
235 "",
236 )
237 system: Final = "\n\n".join(
238 content_to_text(message.get("content")) for message in messages if message.get("role") == "system"
239 )
240 tool_exchanges: Final = { # mutable-ok: accumulated once, serialized to JSON
241 f"e{ordinal}": { # mutable-ok: serialized to JSON
242 "tool_calls": _tool_call_entries(messages[group[0]]),
243 "result": _truncate_for_state(
244 self._exchange_tool_text(messages, group), self.max_result_chars_in_state
245 ),
246 }
247 for ordinal, group in enumerate(candidates)
248 }
249 return {"task": task, "system": system, "tool_exchanges": tool_exchanges} # mutable-ok: serialized to JSON
251 async def _call_systemone(
252 self, state: dict[str, object], question_ids: Sequence[str]
253 ) -> _JevSystemOneResponse | None:
254 """Returns the response, or None when the service failed and fail_open applies."""
255 payload: Final[dict[str, object]] = { # mutable-ok: serialized to JSON by httpx
256 "model": self.jev_model,
257 "state": state,
258 "questions": { # mutable-ok: serialized to JSON
259 question_id: { # mutable-ok: serialized to JSON
260 "type": "noul",
261 "instructions": _question_instructions(question_id),
262 }
263 for question_id in question_ids
264 },
265 }
266 try:
267 raw_response: HttpxResponse = await self.async_handler.post( # pyright: ignore[reportUnknownMemberType] # AsyncHTTPHandler.post is untyped
268 url=f"{self.typesafe_api_base}/v1/systemone",
269 json=payload,
270 headers={ # mutable-ok: httpx header contract is a dict
271 "Authorization": f"Bearer {self.typesafe_api_key}",
272 "Content-Type": "application/json",
273 },
274 timeout=_JEV_TIMEOUT_SECONDS,
275 )
276 except asyncio.CancelledError:
277 raise
278 except Exception as e:
279 detail: Final[dict[str, object]] = (
280 { # mutable-ok: log detail record
281 "error_type": type(e).__name__,
282 "detail": str(e),
283 "status_code": e.response.status_code,
284 "body": _safe_response_text(e.response),
285 }
286 if isinstance(e, httpx.HTTPStatusError)
287 else {"error_type": type(e).__name__, "detail": str(e)} # mutable-ok: log detail record
288 )
289 self._handle_failure("TypeSafe evaluation service request failed", detail)
290 return None
291 if not 200 <= raw_response.status_code < 300:
292 self._handle_failure(
293 "TypeSafe evaluation service returned an error",
294 { # mutable-ok: log detail record
295 "status_code": raw_response.status_code,
296 "body": _safe_response_text(raw_response),
297 },
298 )
299 return None
300 try:
301 body: Final[object] = raw_response.json() # pyright: ignore[reportAny] # httpx Response.json() is untyped
302 except (ValueError, httpx.DecodingError, RecursionError):
303 self._handle_failure(
304 "TypeSafe evaluation service returned an unreadable response",
305 {"body": _safe_response_text(raw_response)}, # mutable-ok: log detail record
306 )
307 return None
308 try:
309 return _JEV_RESPONSE_ADAPTER.validate_python(body)
310 except ValidationError:
311 self._handle_failure(
312 "TypeSafe evaluation service returned unexpected response shape",
313 {"body": _safe_response_text(raw_response)}, # mutable-ok: log detail record
314 )
315 return None
317 @log_guardrail_information
318 async def apply_guardrail(
319 self,
320 inputs: GenericGuardrailAPIInputs,
321 request_data: dict[str, object],
322 input_type: Literal["request", "response"],
323 logging_obj: LiteLLMLoggingObj | None = None,
324 ) -> GenericGuardrailAPIInputs:
325 if input_type != "request":
326 return inputs
328 structured_messages: Final = _as_object_list(inputs.get("structured_messages"))
329 if not structured_messages:
330 return inputs
331 parsed_messages: Final = tuple(_as_str_object_dict(m) for m in structured_messages)
332 if any(m is None for m in parsed_messages):
333 return inputs
334 messages: Final = tuple(m for m in parsed_messages if m is not None)
336 candidates: Final = self._candidate_exchanges(messages)
337 if not candidates:
338 verbose_proxy_logger.debug("TypeSafe: no completed tool exchanges eligible for evaluation")
339 return inputs
341 question_ids: Final = tuple(f"e{ordinal}" for ordinal in range(len(candidates)))
342 state: Final = self._build_state(messages, candidates)
344 start_time: Final = time.monotonic()
345 response: Final = await self._call_systemone(state, question_ids)
346 end_time: Final = time.monotonic()
347 if response is None:
348 self.add_standard_logging_guardrail_information_to_request_data( # pyright: ignore[reportUnknownMemberType] # untyped base helper
349 guardrail_json_response={ # mutable-ok: must stay JSON-serializable for shared logging
350 "error": "TypeSafe evaluation unavailable; request forwarded uncompacted",
351 "model": self.jev_model,
352 },
353 request_data=request_data,
354 guardrail_status="guardrail_failed_to_respond",
355 guardrail_provider="typesafe",
356 start_time=start_time,
357 end_time=end_time,
358 duration=end_time - start_time,
359 )
360 return inputs
362 dropped_ordinals: Final = frozenset(
363 ordinal
364 for ordinal in range(len(candidates))
365 if (answer := response.answers.get(f"e{ordinal}")) is not None and answer.noul < self.relevance_threshold
366 )
367 dropped_tool_indices: Final[frozenset[int]] = frozenset(
368 index
369 for ordinal in dropped_ordinals
370 for index in candidates[ordinal][1:]
371 if messages[index].get("role") in ("tool", "function")
372 )
373 if not dropped_tool_indices:
374 verbose_proxy_logger.debug("TypeSafe: all evaluated exchanges still relevant; request unchanged")
375 return inputs
377 compacted_messages: Final = [ # mutable-ok: structured_messages contract is a list of dicts
378 {**message, "content": DROPPED_RESULT_TEXT} # mutable-ok: JSON message row
379 if index in dropped_tool_indices
380 else message
381 for index, message in enumerate(messages)
382 ]
383 chars_removed: Final = sum(
384 len(content_to_text(messages[index].get("content"))) - len(DROPPED_RESULT_TEXT)
385 for index in dropped_tool_indices
386 )
387 exchanges_dropped: Final = len(dropped_ordinals)
388 verbose_proxy_logger.info(
389 "TypeSafe: evaluated %s tool exchange(s), dropped %s, ~%s chars removed",
390 len(candidates),
391 exchanges_dropped,
392 chars_removed,
393 )
394 self.add_standard_logging_guardrail_information_to_request_data( # pyright: ignore[reportUnknownMemberType] # untyped base helper
395 guardrail_json_response={ # mutable-ok: must stay JSON-serializable for shared logging
396 "exchanges_evaluated": len(candidates),
397 "exchanges_dropped": exchanges_dropped,
398 "chars_removed": chars_removed,
399 "model": self.jev_model,
400 },
401 request_data=request_data,
402 guardrail_status="success",
403 guardrail_provider="typesafe",
404 start_time=start_time,
405 end_time=end_time,
406 duration=end_time - start_time,
407 )
408 return {**inputs, "structured_messages": compacted_messages} # pyright: ignore[reportReturnType] # mutable-ok: inputs protocol is a plain dict # plain dicts satisfy AllMessageValues at runtime
410 @staticmethod
411 def get_config_model() -> type[TypeSafeGuardrailConfigModel] | None:
412 from litellm.types.proxy.guardrails.guardrail_hooks.typesafe import (
413 TypeSafeGuardrailConfigModel,
414 )
416 return TypeSafeGuardrailConfigModel