Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/compresr/compresr.py: 14%
542 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"""Compresr guardrail — query-aware, recoverable context compression.
3Compresses bulky message content (tool outputs by default) through the
4Compresr API before the request reaches the LLM. Each compressed message
5carries a hash marker; a ``compresr_retrieve`` tool is injected so the model
6can fetch the original content back through the agentic loop when the
7compressed version is not enough — making compression recoverable instead
8of lossy.
10Unlike gateway-side compressors that operate on whole message lists, each
11target is compressed *query-aware*: the query sent to Compresr is the intent
12of the tool call that produced the message (``name + arguments``, resolved
13via ``tool_call_id``), falling back to the last user message.
14"""
16from __future__ import annotations
18import asyncio
19import hashlib
20import ipaddress
21import json
22import time
23from collections import Counter, OrderedDict
24from dataclasses import dataclass, field
25from typing import TYPE_CHECKING, Final, Literal, TypeGuard
26from urllib.parse import urlparse
28import httpx
29from fastapi import HTTPException
30from httpx import Response as HttpxResponse
32import litellm
33from litellm._logging import verbose_proxy_logger
34from litellm.constants import PRE_CALL_EXECUTED_GUARDRAILS_KEY
35from litellm.integrations.custom_guardrail import (
36 CustomGuardrail,
37 log_guardrail_information,
38)
39from litellm.litellm_core_utils.prompt_templates.factory import (
40 get_attribute_or_key,
41 get_tool_calls_from_response,
42 has_tool_with_name,
43)
44from litellm.llms.custom_httpx.http_handler import (
45 get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # helper is untyped in http_handler
46 httpxSpecialProvider,
47)
48from litellm.proxy._types import UserAPIKeyAuth
49from litellm.proxy.guardrails.guardrail_hooks.content_text import (
50 assistant_text_from_response,
51 content_to_text,
52 is_all_text_parts,
53 merge_rewritten_text_parts,
54)
55from litellm.secret_managers.main import get_secret_str
56from litellm.types.guardrails import GuardrailEventHooks, Mode
57from litellm.types.integrations.custom_logger import (
58 AgenticLoopPlan,
59 AgenticLoopRequestPatch,
60)
61from litellm.types.utils import GenericGuardrailAPIInputs
63if TYPE_CHECKING: 63 ↛ 64line 63 didn't jump to line 64 because the condition on line 63 was never true
64 from litellm.litellm_core_utils.litellm_logging import (
65 Logging as LiteLLMLoggingObj,
66 )
67 from litellm.llms.base_llm.anthropic_messages.transformation import (
68 BaseAnthropicMessagesConfig,
69 )
70 from litellm.types.proxy.guardrails.guardrail_hooks.base import (
71 GuardrailConfigModel,
72 )
74BYPASS_HEADER: Final = "x-compresr-bypass"
75COMPRESR_RETRIEVE_TOOL_NAME: Final = "compresr_retrieve"
76DEFAULT_API_BASE: Final = "https://api.compresr.ai"
77DEFAULT_COMPRESSION_MODEL: Final = "latte_v2"
78DEFAULT_TARGET_COMPRESSION_RATIO: Final = 0.5
79DEFAULT_MIN_CHARS_TO_COMPRESS: Final = 500
80_ORIGINALS_TTL_SECONDS: Final = 15 * 60
81_NO_SCOPE_WARNING_INTERVAL_SECONDS: Final = 15 * 60
82_MAX_TRACKED_CALLS: Final = 256
83_DEFAULT_MAX_BYTES_PER_CALL: Final = 10 * 1024 * 1024
84# Aggregate ceiling across all recovery-store entries. max_bytes_per_call only
85# bounds a single call; this caps the whole store so many calls cannot exhaust it.
86_MAX_TOTAL_STORE_BYTES: Final = 256 * 1024 * 1024
87# Max compresr_retrieve calls expanded into a single follow-up (repeats deduped).
88_MAX_RETRIEVALS_PER_LOOP: Final = 8
89# The shared client's 600s read timeout is far too long for an on-request
90# guardrail; bound the compress call so a stall hits the fail policy quickly.
91_COMPRESS_TIMEOUT_SECONDS: Final = 60.0
92_SOURCE_TAG: Final = "integration:litellm"
93# Request-content fields the compression_params passthrough must never
94# override — they carry the actual message content/queries being compressed.
95_RESERVED_COMPRESSION_PARAM_KEYS: Final = frozenset({"context", "query", "inputs"})
96_BLOCKED_METADATA_HOSTS: Final = frozenset(
97 {
98 "metadata.google.internal",
99 "metadata.goog",
100 "metadata.azure.com",
101 "metadata.azure.internal",
102 }
103)
104_BLOCKED_METADATA_IPS: Final = frozenset(
105 ipaddress.ip_address(ip) for ip in ("169.254.169.254", "fd00:ec2::254", "100.100.100.200", "168.63.129.16")
106)
109def _parse_ip_literal(host: str) -> ipaddress.IPv4Address | ipaddress.IPv6Address | None:
110 """Parse ``host`` as an IP literal, covering the alternate spellings the
111 socket layer accepts (decimal/hex single-integer IPv4, IPv4-mapped IPv6)
112 so a blocked address cannot be smuggled past a string comparison."""
113 try:
114 addr = ipaddress.ip_address(host)
115 except ValueError:
116 try:
117 addr = ipaddress.ip_address(int(host, 0))
118 except (TypeError, ValueError):
119 return None
120 if isinstance(addr, ipaddress.IPv6Address) and addr.ipv4_mapped is not None:
121 return addr.ipv4_mapped
122 return addr
125def _validate_api_base(url: str) -> str:
126 """Return ``url`` if it passes basic outbound-target checks, else raise.
128 Best-effort defense in depth for a mis/maliciously-configured ``api_base``:
129 rejects non-http(s) schemes and cloud-metadata IPs/hosts (incl. alternate IP
130 encodings); private ranges are allowed for on-prem deployments. NOT a complete
131 SSRF control — no DNS resolution, and the shared client follows redirects and
132 re-resolves DNS (TOCTOU / rebinding); ``api_base`` is trusted operator config,
133 so this is an accepted limitation.
134 """
135 parsed: Final = urlparse(url)
136 if parsed.scheme not in ("http", "https"):
137 raise ValueError(f"Compresr guardrail api_base must be http or https, got scheme={parsed.scheme!r}")
138 host: Final = (parsed.hostname or "").lower()
139 if not host:
140 raise ValueError("Compresr guardrail api_base has no host")
141 ip_literal: Final = _parse_ip_literal(host)
142 if host in _BLOCKED_METADATA_HOSTS or (ip_literal is not None and ip_literal in _BLOCKED_METADATA_IPS):
143 raise ValueError(f"Compresr guardrail api_base {host!r} is a blocked cloud-metadata host")
144 return url
147def _is_str_object_dict(value: object) -> TypeGuard[dict[str, object]]: # guard-ok: isinstance narrows correctly; predicate is trivially correct # fmt: skip
148 return isinstance(value, dict)
151def _is_object_list(value: object) -> TypeGuard[list[object]]: # guard-ok: isinstance narrows correctly; predicate is trivially correct # fmt: skip
152 return isinstance(value, list)
155def _replace_text_in_content(content: object, new_text: str) -> object:
156 """Write ``new_text`` back into a ``content`` value, preserving shape.
158 ``str`` content is replaced directly. An all-text part list collapses to a
159 single part carrying the last declared cache_control breakpoint. Anything
160 else is returned unchanged: breakpoints are positional, so one compressed
161 string cannot be written back across a non-text part without moving text
162 to the other side of it.
163 """
164 if isinstance(content, str):
165 return new_text
166 if _is_object_list(content) and is_all_text_parts(content):
167 return merge_rewritten_text_parts(content, new_text)
168 return content
171def _render_tool_intent(fn: dict[str, object]) -> str:
172 name: Final = str(fn.get("name") or "").strip()
173 args: Final = fn.get("arguments")
174 if isinstance(args, dict):
175 try:
176 args_str = json.dumps(args, separators=(",", ":"))
177 except (TypeError, ValueError):
178 args_str = str(args)
179 else:
180 args_str = str(args).strip() if args is not None else ""
181 if name and args_str:
182 return f"{name}: {args_str}"
183 return name or args_str
186def _query_for_target(messages: list[dict[str, object]], target_idx: int, fallback: str) -> str:
187 """Query used to compress ``messages[target_idx]``.
189 Tool/function outputs are compressed against the intent of the tool call
190 that produced them (found via ``tool_call_id`` on a prior assistant
191 message); everything else uses the last user message.
192 """
193 msg: Final = messages[target_idx]
194 if msg.get("role") not in ("tool", "function"):
195 return fallback
197 tool_call_id: Final = msg.get("tool_call_id")
198 fn_name: Final = msg.get("name")
199 for j in range(target_idx - 1, -1, -1):
200 prev = messages[j]
201 if prev.get("role") != "assistant":
202 continue
203 tool_calls = prev.get("tool_calls")
204 if isinstance(tool_calls, list):
205 for tc in tool_calls:
206 if not isinstance(tc, dict):
207 continue
208 if tool_call_id and tc.get("id") == tool_call_id:
209 fn = tc.get("function")
210 intent = _render_tool_intent(fn if isinstance(fn, dict) else {})
211 if intent:
212 return intent
213 # Legacy function_call fallback: require a name match, else an earlier
214 # function_call turn would attribute the wrong intent.
215 fc = prev.get("function_call")
216 if isinstance(fc, dict) and fn_name and fc.get("name") == fn_name:
217 intent = _render_tool_intent(fc)
218 if intent:
219 return intent
220 return fallback
223def _safe_int(value: object) -> int:
224 """Parse a token-stat field defensively.
226 A malformed-but-200 response must not raise here: ``_call_compress`` has
227 already returned successfully, so the fail_open/fail_closed decision is
228 behind us. A bare ``int()`` on a non-numeric field would surface as an
229 unhandled 500 even when ``fail_open`` is configured.
230 """
231 try:
232 return int(value) if value is not None else 0
233 except (TypeError, ValueError):
234 return 0
237def _safe_response_text(response: object, limit: int = 500) -> str:
238 """Read a response body for error logging without letting the read itself
239 raise. A corrupt ``Content-Encoding`` makes ``httpx``'s ``.text`` raise a
240 ``DecodingError``; if that happened while building a failure detail it would
241 turn an already-handled error into an unhandled 500."""
242 try:
243 text: Final = getattr(response, "text", "")
244 except httpx.DecodingError:
245 return "<undecodable response body>"
246 return (text or "")[:limit]
249def _content_hash(text: str) -> str:
250 # surrogatepass so a lone surrogate in untrusted content (valid via a JSON
251 # \uXXXX escape) hashes instead of raising past the fail policy.
252 return hashlib.sha256(text.encode("utf-8", "surrogatepass")).hexdigest()[:24]
255def _entry_bytes(originals: dict[str, str]) -> int:
256 """UTF-8 byte size of one recovery-store entry (surrogatepass, like _content_hash)."""
257 return sum(len(value.encode("utf-8", "surrogatepass")) for value in originals.values())
260def _display_hash(hash_value: str) -> str:
261 """Bound a model-supplied hash for logs/fallback text. A real marker hash is
262 24 hex chars; a prompt-injected ``compresr_retrieve`` call could pass a huge
263 or control-character-laden string, so strip non-printables (no forged log
264 lines / ANSI escapes) and cap length before echoing into logs and the
265 conversation."""
266 printable: Final = "".join(ch for ch in hash_value if ch.isprintable())
267 return printable if len(printable) <= 32 else f"{printable[:32]}…"
270def _recovery_marker(hash_value: str) -> str:
271 return (
272 f"\n\n[compresr hash={hash_value}: parts of this content were compressed "
273 f"away. If you need the full original, call the "
274 f"{COMPRESR_RETRIEVE_TOOL_NAME} tool with this hash.]"
275 )
278def _build_compresr_retrieve_tool() -> dict[str, object]:
279 return {
280 "type": "function",
281 "function": {
282 "name": COMPRESR_RETRIEVE_TOOL_NAME,
283 "description": (
284 "Retrieve the original, uncompressed content behind a Compresr "
285 "compression marker. Call this when a compression marker's hash "
286 "points at content you need in full."
287 ),
288 "parameters": {
289 "type": "object",
290 "properties": {
291 "hash": {
292 "type": "string",
293 "description": "The 24-character hex hash from the compression marker.",
294 },
295 },
296 "required": ["hash"],
297 },
298 },
299 }
302def has_compresr_retrieve_tool(tools: object) -> bool:
303 return has_tool_with_name(tools, COMPRESR_RETRIEVE_TOOL_NAME)
306def _merge_retrieve_tool(existing_tools: object) -> list[object] | None:
307 """The request's tools plus the retrieve tool, or None when the incoming
308 shape is not a list (leave the caller's tools untouched; markers stay
309 inert text)."""
310 if existing_tools is not None and not isinstance(existing_tools, list):
311 return None
312 retrieve_tool: Final = _build_compresr_retrieve_tool()
313 if existing_tools is None:
314 return [retrieve_tool]
315 if has_compresr_retrieve_tool(existing_tools):
316 return list(existing_tools)
317 return list(existing_tools) + [retrieve_tool]
320def _extract_compresr_tool_calls(response: object) -> list[dict[str, object]]:
321 return [
322 {"id": tc.get("id"), "type": "function", "name": tc.get("name"), "arguments": tc.get("arguments", {})}
323 for tc in get_tool_calls_from_response(response)
324 if tc.get("name") == COMPRESR_RETRIEVE_TOOL_NAME
325 ]
328def _resolve_call_id(logging_obj: object) -> str | None:
329 """The call id from the framework logging object.
331 This value ultimately derives from the client-settable ``x-litellm-call-id``
332 header and is echoed back in responses, so it is NOT a trust boundary on its
333 own — ``_scoped_store_key`` prefixes it with the caller's virtual-key hash to
334 partition the recovery store per tenant. Request-body/kwargs call ids are
335 deliberately not consulted here.
336 """
337 logging_call_id: Final = getattr(logging_obj, "litellm_call_id", None)
338 if isinstance(logging_call_id, str) and logging_call_id:
339 return logging_call_id
340 return None
343def _caller_scope(logging_obj: object) -> str:
344 """The caller's virtual-key hash, used to partition the recovery store.
346 Trust is anchored on the ``UserAPIKeyAuth`` object the proxy sets
347 server-side (``metadata.user_api_key_auth``, litellm_pre_call_utils). Its
348 ``api_key`` is the hash of the authenticated key. Both metadata spellings
349 are scanned (``/v1/messages`` and ``/v1/responses`` carry it under
350 ``litellm_metadata``), but the bare ``user_api_key`` *string* is never
351 trusted on its own: a JSON request body can place one in the client-supplied
352 ``metadata`` field, which is only sanitized on the route's canonical
353 container. Returns "" when the proxy runs without per-key auth, in which case
354 all traffic is a single trust domain and the call id alone suffices.
355 """
356 details: Final = getattr(logging_obj, "model_call_details", None)
357 if not _is_str_object_dict(details):
358 return ""
359 litellm_params: Final = details.get("litellm_params")
360 for container in (litellm_params, details):
361 if not _is_str_object_dict(container):
362 continue
363 for meta_key in ("metadata", "litellm_metadata"):
364 metadata = container.get(meta_key)
365 if not _is_str_object_dict(metadata):
366 continue
367 auth = metadata.get("user_api_key_auth")
368 if isinstance(auth, UserAPIKeyAuth) and isinstance(auth.api_key, str) and auth.api_key:
369 return auth.api_key
370 return ""
373def _scoped_store_key(logging_obj: object) -> str | None:
374 """Key for the recovery store: caller identity plus framework call id.
376 Keying on the call id alone is unsafe: it comes from the client-settable
377 ``x-litellm-call-id`` header and is echoed back in responses, so one caller
378 could read or evict another's originals by reusing the id. Prefixing the
379 unforgeable virtual-key hash binds each entry to the tenant that created it.
380 Returns None when there is no call id, which disables recovery for the call.
381 """
382 call_id: Final = _resolve_call_id(logging_obj)
383 if call_id is None:
384 return None
385 scope: Final = _caller_scope(logging_obj)
386 return f"{scope}\x00{call_id}" if scope else call_id
389def _is_responses_api_response(response: object) -> bool:
390 return isinstance(get_attribute_or_key(response, "output", None), list)
393def _is_anthropic_messages_response(response: object) -> bool:
394 return isinstance(get_attribute_or_key(response, "content", None), list)
397def _build_assistant_message_from_response(
398 response: object,
399 retrieved: list[tuple[dict[str, object], str]],
400) -> dict[str, object]:
401 """Rebuild the chat-completions assistant turn for the retrieval follow-up.
403 Only the ``compresr_retrieve`` calls are echoed, each answered by a tool
404 result below. Other tool calls made in the same turn are omitted on purpose:
405 the follow-up re-runs the model with the recovered content so it re-plans
406 them. Echoing them would leave tool_calls with no matching tool result and
407 the provider would reject the request.
408 """
409 return {
410 "role": "assistant",
411 "content": assistant_text_from_response(response),
412 "tool_calls": [
413 {
414 "id": tool_call.get("id"),
415 "type": "function",
416 "function": {
417 "name": tool_call.get("name"),
418 "arguments": json.dumps(tool_call.get("arguments", {})),
419 },
420 }
421 for tool_call, _ in retrieved
422 ],
423 }
426def _build_anthropic_followup_messages(
427 response: object,
428 retrieved: list[tuple[dict[str, object], str]],
429) -> list[dict[str, object]]:
430 """Anthropic requires the tool_use block echoed back in an assistant
431 message paired with a tool_result block keyed by the same tool_use_id. The
432 assistant text is preserved; non-retrieve tool calls are re-planned by the
433 follow-up (see _build_assistant_message_from_response)."""
434 assistant_content: Final[list[dict[str, object]]] = []
435 text: Final = assistant_text_from_response(response)
436 if text:
437 assistant_content.append({"type": "text", "text": text})
438 assistant_content.extend(
439 {
440 "type": "tool_use",
441 "id": tool_call.get("id"),
442 "name": tool_call.get("name"),
443 "input": tool_call.get("arguments", {}),
444 }
445 for tool_call, _ in retrieved
446 )
447 assistant_message: Final[dict[str, object]] = {"role": "assistant", "content": assistant_content}
448 user_message: Final[dict[str, object]] = {
449 "role": "user",
450 "content": [
451 {"type": "tool_result", "tool_use_id": tool_call.get("id"), "content": content}
452 for tool_call, content in retrieved
453 ],
454 }
455 return [assistant_message, user_message]
458def _build_responses_followup_items(
459 response: object,
460 retrieved: list[tuple[dict[str, object], str]],
461) -> list[dict[str, object]]:
462 """The Responses API requires the model's function_call echoed back paired
463 with a function_call_output keyed by the same call_id. The assistant text is
464 preserved; non-retrieve tool calls are re-planned by the follow-up."""
465 items: Final[list[dict[str, object]]] = []
466 text: Final = assistant_text_from_response(response)
467 if text:
468 items.append({"role": "assistant", "content": text})
469 for tool_call, content in retrieved:
470 call_id = tool_call.get("id")
471 items.append(
472 {
473 "type": "function_call",
474 "call_id": call_id,
475 "name": tool_call.get("name"),
476 "arguments": json.dumps(tool_call.get("arguments", {})),
477 }
478 )
479 items.append({"type": "function_call_output", "call_id": call_id, "output": content})
480 return items
483@dataclass
484class _CompressionResult:
485 """Outcome of applying compression results to a message list."""
487 compressed_messages: list[dict[str, object]]
488 originals: dict[str, str] = field(default_factory=dict)
489 # original text -> compressed text, plus the machinery the Responses `texts`
490 # mirror needs to replace only where it is unambiguous.
491 text_replacements: dict[str, str] = field(default_factory=dict)
492 replaced_text_counts: dict[str, int] = field(default_factory=dict)
493 ambiguous_texts: set[str] = field(default_factory=set)
494 messages_compressed: int = 0
495 tokens_before: int = 0
496 tokens_after: int = 0
499class CompresrGuardrail(CustomGuardrail):
500 def __init__(
501 self,
502 api_base: str | None = None,
503 api_key: str | None = None,
504 model: str | None = None,
505 target_compression_ratio: float | None = None,
506 coarse: bool | None = None,
507 min_chars_to_compress: int | None = None,
508 compress_tool_outputs: bool | None = None,
509 compress_system: bool | None = None,
510 compress_history: bool | None = None,
511 compress_last_user: bool | None = None,
512 enable_retrieval: bool | None = None,
513 guardrail_name: str | None = None,
514 event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None = None,
515 default_on: bool = False,
516 unreachable_fallback: str | None = None,
517 max_bytes_per_call: int | None = None,
518 allow_bypass_header: bool | None = None,
519 dynamic: bool | None = None,
520 dynamic_min_ratio: float | None = None,
521 dynamic_max_ratio: float | None = None,
522 compression_params: dict[str, object] | None = None,
523 ):
524 raw_api_base: Final = (api_base or get_secret_str("COMPRESR_API_BASE") or DEFAULT_API_BASE).rstrip("/")
525 self.compresr_api_base = _validate_api_base(raw_api_base)
526 self.compresr_api_key = api_key or get_secret_str("COMPRESR_API_KEY")
527 if not self.compresr_api_key:
528 raise ValueError(
529 "Compresr guardrail requires an API key. Set `api_key` in the "
530 "guardrail config or the COMPRESR_API_KEY env var."
531 )
532 self.compression_model = model or DEFAULT_COMPRESSION_MODEL
533 self.target_compression_ratio = (
534 DEFAULT_TARGET_COMPRESSION_RATIO if target_compression_ratio is None else target_compression_ratio
535 )
536 self.coarse = True if coarse is None else coarse
537 self.min_chars_to_compress = (
538 DEFAULT_MIN_CHARS_TO_COMPRESS if min_chars_to_compress is None else min_chars_to_compress
539 )
540 self.compress_tool_outputs = True if compress_tool_outputs is None else compress_tool_outputs
541 self.compress_system = False if compress_system is None else compress_system
542 self.compress_history = False if compress_history is None else compress_history
543 self.compress_last_user = False if compress_last_user is None else compress_last_user
544 self.enable_retrieval = True if enable_retrieval is None else enable_retrieval
545 self.unreachable_fallback: Literal["fail_closed", "fail_open"] = (
546 "fail_open" if unreachable_fallback == "fail_open" else "fail_closed"
547 )
548 self.max_bytes_per_call = _DEFAULT_MAX_BYTES_PER_CALL if max_bytes_per_call is None else max_bytes_per_call
549 if self.max_bytes_per_call < 0:
550 raise ValueError("max_bytes_per_call must be >= 0 (0 disables the cap; positive values enforce it)")
551 self.allow_bypass_header = False if allow_bypass_header is None else allow_bypass_header
552 # Dynamic (adaptive) compression — latte_v2 only, on by default: the server
553 # picks the ratio per input instead of honoring target_compression_ratio.
554 self.dynamic = True if dynamic is None else dynamic
555 self.dynamic_min_ratio = dynamic_min_ratio
556 self.dynamic_max_ratio = dynamic_max_ratio
557 # Passthrough of extra compression params forwarded verbatim, so a new
558 # Compresr feature works without changing this guardrail. Named fields win;
559 # request-content fields are stripped.
560 reserved_keys: Final = _RESERVED_COMPRESSION_PARAM_KEYS.intersection(compression_params or {})
561 if reserved_keys:
562 verbose_proxy_logger.warning(
563 "Compresr: ignoring reserved compression_params keys %s", sorted(reserved_keys)
564 )
565 self.compression_params: dict[str, object] = {
566 k: v for k, v in (compression_params or {}).items() if k not in _RESERVED_COMPRESSION_PARAM_KEYS
567 }
568 self.async_handler = get_async_httpx_client(
569 llm_provider=httpxSpecialProvider.GuardrailCallback,
570 )
571 self._originals_by_call_id: OrderedDict[str, tuple[dict[str, str], float]] = OrderedDict()
572 # Running byte size of the store, kept in sync to enforce the global cap cheaply.
573 self._store_total_bytes = 0
574 # Rate-limits the "recovery skipped, no auth scope" warning so an ongoing
575 # misconfiguration stays visible without flooding hot-path logs.
576 self._no_scope_warning_expiry = 0.0
577 if self.enable_retrieval:
578 verbose_proxy_logger.warning(
579 "Compresr: enable_retrieval is on; the recovery store is per-process. "
580 "For multi-worker deployments, set enable_retrieval=false or run with --workers 1."
581 )
582 super().__init__( # pyright: ignore[reportUnknownMemberType] # CustomGuardrail.__init__ is untyped
583 guardrail_name=guardrail_name,
584 event_hook=event_hook,
585 default_on=default_on,
586 )
588 def _should_bypass(self, request_data: dict) -> bool:
589 if not self.allow_bypass_header:
590 return False
591 psr: Final = request_data.get("proxy_server_request")
592 if not _is_str_object_dict(psr):
593 return False
594 headers: Final = psr.get("headers")
595 if not _is_str_object_dict(headers):
596 return False
597 return str(headers.get(BYPASS_HEADER)).lower() == "true"
599 def _request_headers(self) -> dict[str, str]:
600 return {
601 "Content-Type": "application/json",
602 "X-API-Key": self.compresr_api_key or "",
603 }
605 def _handle_compress_failure(self, error: str, log_detail: dict[str, object]) -> None:
606 """fail_open logs and returns (caller forwards uncompressed);
607 fail_closed raises. ``log_detail`` may include upstream response bodies
608 and is written only to server logs; the raised ``HTTPException`` carries
609 a generic message so a malicious ``api_base`` cannot exfiltrate response
610 bytes through the client-visible error."""
611 if self.unreachable_fallback == "fail_open":
612 verbose_proxy_logger.warning(
613 "Compresr: %s; fail_open configured, forwarding request uncompressed. detail=%s",
614 error,
615 log_detail,
616 )
617 return
618 verbose_proxy_logger.error("Compresr: %s. detail=%s", error, log_detail)
619 raise HTTPException(status_code=502, detail={"error": error})
621 def _evict_oldest(self) -> None:
622 """Drop the front (oldest) entry and decrement the running byte total."""
623 _key, (evicted, _expiry) = self._originals_by_call_id.popitem(last=False)
624 self._store_total_bytes -= _entry_bytes(evicted)
626 def _prune_originals(self) -> None:
627 # Insertion order == expiry order (shared TTL); prune from the front.
628 now: Final = time.monotonic()
629 store: Final = self._originals_by_call_id
630 while store and store[next(iter(store))][1] <= now:
631 self._evict_oldest()
632 while len(store) > _MAX_TRACKED_CALLS:
633 self._evict_oldest()
634 # Global byte budget; keep the most-recent entry so the current call's
635 # originals survive (a single call is already bounded by max_bytes_per_call).
636 while len(store) > 1 and self._store_total_bytes > _MAX_TOTAL_STORE_BYTES:
637 self._evict_oldest()
639 def _existing_originals(self, store_key: str | None) -> dict[str, str]:
640 """Originals already stored under this key, so the per-call byte budget
641 can account for an earlier turn that reused the store key."""
642 if store_key is None:
643 return {}
644 return self._originals_by_call_id.get(store_key, ({}, 0.0))[0]
646 def _store_originals(self, store_key: str, originals: dict[str, str]) -> None:
647 existing, _ = self._originals_by_call_id.get(store_key, ({}, 0.0))
648 merged: Final = self._bound_call_bytes({**existing, **originals})
649 # Keep the running total in sync: drop the overwritten entry, add the new one.
650 self._store_total_bytes += _entry_bytes(merged) - _entry_bytes(existing)
651 self._originals_by_call_id[store_key] = (
652 merged,
653 time.monotonic() + _ORIGINALS_TTL_SECONDS,
654 )
655 self._originals_by_call_id.move_to_end(store_key)
656 self._prune_originals()
658 def _bound_call_bytes(self, merged: dict[str, str]) -> dict[str, str]:
659 """Drop oldest entries (dict insertion order) until the aggregate byte
660 size fits ``self.max_bytes_per_call``. Prevents one call with many
661 large tool outputs from growing proxy memory without bound."""
662 if self.max_bytes_per_call <= 0:
663 return merged
664 total = _entry_bytes(merged)
665 if total <= self.max_bytes_per_call:
666 return merged
667 bounded: Final = dict(merged)
668 for key in list(bounded.keys()):
669 if total <= self.max_bytes_per_call:
670 break
671 total -= len(bounded[key].encode("utf-8", "surrogatepass"))
672 del bounded[key]
673 verbose_proxy_logger.warning("Compresr: originals-store byte cap hit, evicted hash=%s", key)
674 return bounded
676 def _retrieve_original(self, store_key: str | None, hash_value: str) -> str | None:
677 """Stored original for a marker hash, or None if not issued for this
678 request (unknown, expired, or from another caller's scope)."""
679 if store_key:
680 originals, expiry = self._originals_by_call_id.get(store_key, ({}, 0.0))
681 if expiry > time.monotonic() and hash_value in originals:
682 return originals[hash_value]
683 verbose_proxy_logger.warning(
684 "Compresr retrieve: rejecting hash=%s (not issued for this request, or expired)",
685 _display_hash(hash_value),
686 )
687 return None
689 def _resolve_retrievals(
690 self, store_key: str | None, tool_calls: list[dict[str, object]]
691 ) -> tuple[list[tuple[dict[str, object], str]], bool]:
692 """Resolve compresr_retrieve calls to (call, result_text) pairs, deduping
693 repeated hashes and capping the count so the follow-up cannot be amplified.
694 The bool is True iff at least one call resolved to real stored content."""
695 retrieved: Final[list[tuple[dict[str, object], str]]] = []
696 seen: Final[set[str]] = set()
697 resolved_any = False
698 for idx, tc in enumerate(tool_calls):
699 arguments = tc.get("arguments", {})
700 hash_value = str(arguments.get("hash", "")) if isinstance(arguments, dict) else ""
701 if idx >= _MAX_RETRIEVALS_PER_LOOP:
702 result = "[compresr: retrieval limit reached for this turn]"
703 elif hash_value in seen:
704 result = "[compresr: already retrieved above for this hash]"
705 else:
706 content = self._retrieve_original(store_key, hash_value)
707 if content is None:
708 result = f"[compresr: hash={_display_hash(hash_value)} not found, expired, or not issued for this request]"
709 else:
710 seen.add(hash_value)
711 resolved_any = True
712 result = content
713 verbose_proxy_logger.debug("Compresr retrieve: hash=%s -> %d chars", _display_hash(hash_value), len(result))
714 retrieved.append((tc, result))
715 return retrieved, resolved_any
717 async def _call_compress(
718 self,
719 contexts: list[str],
720 queries: list[str],
721 ) -> list[dict[str, object]] | None:
722 """Compress ``contexts`` (query-aware). Returns one result dict per
723 context, or None when the service failed and fail_open applies."""
724 common: Final[dict[str, object]] = {
725 # Passthrough first so the named fields below always win on collision.
726 **self.compression_params,
727 "compression_model_name": self.compression_model,
728 "target_compression_ratio": self.target_compression_ratio,
729 "coarse": self.coarse,
730 "dynamic": self.dynamic,
731 "source": _SOURCE_TAG,
732 }
733 # Only send the bounds the operator actually set; otherwise let the
734 # server apply its own floor/ceiling.
735 if self.dynamic_min_ratio is not None:
736 common["dynamic_min_ratio"] = self.dynamic_min_ratio
737 if self.dynamic_max_ratio is not None:
738 common["dynamic_max_ratio"] = self.dynamic_max_ratio
739 if len(contexts) == 1:
740 url = f"{self.compresr_api_base}/api/compress/question-specific/"
741 payload: dict[str, object] = {
742 "context": contexts[0],
743 "query": queries[0],
744 **common,
745 }
746 else:
747 url = f"{self.compresr_api_base}/api/compress/question-specific/batch"
748 payload = {
749 "inputs": [{"context": ctx, "query": q} for ctx, q in zip(contexts, queries)],
750 **common,
751 }
753 try:
754 raw_response: HttpxResponse = await self.async_handler.post( # pyright: ignore[reportUnknownMemberType] # AsyncHTTPHandler.post is untyped
755 url=url,
756 json=payload,
757 headers=self._request_headers(),
758 timeout=_COMPRESS_TIMEOUT_SECONDS,
759 )
760 except asyncio.CancelledError:
761 raise
762 except httpx.HTTPStatusError as e:
763 # The shared handler calls raise_for_status(), so a non-2xx reply arrives
764 # here as an error carrying the upstream body + our API key header; route
765 # it through the fail policy so none of that reaches the client.
766 resp: Final = getattr(e, "response", None)
767 self._handle_compress_failure(
768 "Compresr compression service returned an error",
769 {
770 "status_code": getattr(resp, "status_code", None),
771 "body": _safe_response_text(resp),
772 },
773 )
774 return None
775 except (httpx.RequestError, litellm.Timeout) as e:
776 # Every request-side httpx failure is a RequestError; route the whole
777 # class through the fail policy so none escapes as a 500 under fail_open.
778 # (HTTPStatusError is handled above and is not a RequestError.)
779 self._handle_compress_failure(
780 "Compresr compression service request failed",
781 {"detail": str(e)},
782 )
783 return None
784 if not 200 <= raw_response.status_code < 300:
785 self._handle_compress_failure(
786 "Compresr compression service returned an error",
787 {
788 "status_code": raw_response.status_code,
789 "body": _safe_response_text(raw_response),
790 },
791 )
792 return None
794 try:
795 body: Final[object] = raw_response.json()
796 except (ValueError, httpx.DecodingError, RecursionError):
797 # RecursionError: a deeply nested JSON body overflows the parser;
798 # route it through the fail policy rather than let it escape as a 500.
799 self._handle_compress_failure(
800 "Compresr compression service returned an unreadable response",
801 {"body": _safe_response_text(raw_response)},
802 )
803 return None
804 if not _is_str_object_dict(body) or not _is_str_object_dict(body.get("data")):
805 self._handle_compress_failure(
806 "Compresr compression service returned unexpected response shape",
807 {"body": _safe_response_text(raw_response)},
808 )
809 return None
810 data: dict[str, object] = body["data"] # pyright: ignore[reportAssignmentType] # dict-guarded above; subscript does not narrow
812 if len(contexts) == 1:
813 return [data]
814 results: Final = data.get("results")
815 if (
816 not _is_object_list(results)
817 or len(results) != len(contexts)
818 or not all(_is_str_object_dict(r) for r in results)
819 ):
820 # Anything but a 1:1 dict-per-context mapping would misalign
821 # results with their target messages.
822 self._handle_compress_failure(
823 "Compresr batch response missing or mismatched 'results'",
824 {"expected": len(contexts), "got": len(results) if _is_object_list(results) else None},
825 )
826 return None
827 return results # pyright: ignore[reportReturnType] # every element dict-checked above; list[object] does not narrow
829 def _select_targets(self, messages: list[dict[str, object]], query_idx: int | None) -> list[int]:
830 """Indices of messages whose text content should be compressed."""
831 targets: Final[list[int]] = []
832 for idx, msg in enumerate(messages):
833 if idx == query_idx and not self.compress_last_user:
834 continue
835 role = msg.get("role")
836 if role in ("tool", "function"):
837 if not self.compress_tool_outputs:
838 continue
839 elif role == "system":
840 if not self.compress_system:
841 continue
842 elif role == "user":
843 if idx != query_idx and not self.compress_history:
844 continue
845 else:
846 continue
847 content = msg.get("content")
848 if _is_object_list(content) and not is_all_text_parts(content):
849 continue
850 if len(content_to_text(content)) < self.min_chars_to_compress:
851 continue
852 targets.append(idx)
853 return targets
855 @staticmethod
856 def _extract_fallback_query(
857 messages: list[dict[str, object]],
858 ) -> tuple[str, int | None]:
859 for idx in range(len(messages) - 1, -1, -1):
860 if messages[idx].get("role") == "user":
861 return content_to_text(messages[idx].get("content")), idx
862 return "", None
864 def _apply_compression_results(
865 self,
866 messages: list[dict[str, object]],
867 targets: list[int],
868 contexts: list[str],
869 results: list[dict[str, object]],
870 recovery_enabled: bool,
871 existing_originals: dict[str, str] | None = None,
872 ) -> _CompressionResult:
873 """Write each compression result into a copy of ``messages``.
875 A result is a real compression only when it is a non-empty string that
876 differs from the original; identical text is treated as a no-op so an
877 untouched request is not needlessly rewritten downstream.
878 """
879 out: Final = _CompressionResult(compressed_messages=list(messages))
880 existing: Final = existing_originals or {}
881 cap: Final = self.max_bytes_per_call
882 # Seed with what is already stored under this store key: markers are
883 # attached only while the store (existing + this call's originals) stays
884 # within the cap, so _store_originals never has to evict a hash this call
885 # just shipped a marker for -- including on a later turn that reuses the
886 # store key. A hash already stored (or repeated here) costs no new bytes.
887 recovery_bytes = _entry_bytes(existing)
888 for target_idx, original_text, result in zip(targets, contexts, results):
889 compressed_text = result.get("compressed_context")
890 if not isinstance(compressed_text, str) or not compressed_text or compressed_text == original_text:
891 continue
892 out.messages_compressed += 1
893 if recovery_enabled:
894 hash_value = _content_hash(original_text)
895 already_stored = hash_value in existing or hash_value in out.originals
896 new_bytes = 0 if already_stored else len(original_text.encode("utf-8", "surrogatepass"))
897 if cap <= 0 or recovery_bytes + new_bytes <= cap:
898 recovery_bytes += new_bytes
899 out.originals[hash_value] = original_text
900 compressed_text += _recovery_marker(hash_value)
901 previous = out.text_replacements.get(original_text)
902 if previous is not None and previous != compressed_text:
903 # Two targets with identical text but different query-specific
904 # compressions; a value-keyed replacement cannot tell them apart.
905 out.ambiguous_texts.add(original_text)
906 else:
907 out.text_replacements[original_text] = compressed_text
908 out.replaced_text_counts[original_text] = out.replaced_text_counts.get(original_text, 0) + 1
909 original_msg = out.compressed_messages[target_idx]
910 out.compressed_messages[target_idx] = {
911 **original_msg,
912 "content": _replace_text_in_content(original_msg.get("content"), compressed_text),
913 }
914 out.tokens_before += _safe_int(result.get("original_tokens"))
915 out.tokens_after += _safe_int(result.get("compressed_tokens"))
916 return out
918 @staticmethod
919 def _mirror_texts_channel(input_texts: object, applied: _CompressionResult) -> list[object] | None:
920 """Compressed content mirrored into the Responses `texts` channel.
922 The chat/Anthropic/Responses handlers round-trip
923 ``structured_messages``; translations without that round-trip write
924 back through ``texts``, so the compressed content is mirrored there
925 too. This matches by value, so a
926 replacement is applied only when it is unambiguous: one compression per
927 text, and every occurrence in ``texts`` accounted for by a compressed
928 target. Anything else is left uncompressed rather than risk a wrong or
929 out-of-policy replacement. Returns None when nothing safe applies.
930 """
931 if not applied.text_replacements or not isinstance(input_texts, list):
932 return None
933 counts: Final = Counter(text for text in input_texts if isinstance(text, str))
934 safe: Final = {
935 text: replacement
936 for text, replacement in applied.text_replacements.items()
937 if text not in applied.ambiguous_texts and counts.get(text) == applied.replaced_text_counts.get(text)
938 }
939 if not safe:
940 return None
941 return [safe.get(text, text) if isinstance(text, str) else text for text in input_texts]
943 @log_guardrail_information
944 async def apply_guardrail(
945 self,
946 inputs: GenericGuardrailAPIInputs,
947 request_data: dict,
948 input_type: Literal["request", "response"],
949 logging_obj: LiteLLMLoggingObj | None = None,
950 ) -> GenericGuardrailAPIInputs:
951 if input_type != "request":
952 return inputs
954 if self._should_bypass(request_data):
955 verbose_proxy_logger.debug("Compresr: %s header set; skipping compression", BYPASS_HEADER)
956 return inputs
958 structured_messages: Final = inputs.get("structured_messages")
959 if not _is_object_list(structured_messages) or not structured_messages:
960 return inputs
961 messages: Final = [m for m in structured_messages if _is_str_object_dict(m)]
962 if len(messages) != len(structured_messages):
963 return inputs
965 fallback_query, query_idx = self._extract_fallback_query(messages)
966 targets: Final[list[int]] = []
967 queries: Final[list[str]] = []
968 for idx in self._select_targets(messages, query_idx):
969 query = _query_for_target(messages, idx, fallback_query)
970 # latte models require a non-empty query; leave targets we cannot
971 # derive one for uncompressed rather than erroring.
972 if not query.strip():
973 continue
974 targets.append(idx)
975 queries.append(query)
976 if not targets:
977 verbose_proxy_logger.debug("Compresr: no messages eligible for compression")
978 return inputs
980 contexts: Final = [content_to_text(messages[idx].get("content")) for idx in targets]
982 start_time: Final = time.monotonic()
983 results: Final = await self._call_compress(contexts=contexts, queries=queries)
984 end_time: Final = time.monotonic()
985 if results is None: # service failed, fail_open configured
986 return inputs
988 # Recovery needs a per-tenant scope; without per-key auth the key would fall
989 # back to the client-settable call id (cross-tenant reads), so skip it.
990 store_key: Final = _scoped_store_key(logging_obj)
991 scope: Final = _caller_scope(logging_obj)
992 recovery_enabled: Final = self.enable_retrieval and store_key is not None and bool(scope)
993 if self.enable_retrieval and not scope and time.monotonic() >= self._no_scope_warning_expiry:
994 # Surface the silent no-recovery case (compressed, but no auth scope
995 # to inject the retrieve tool), re-warning once per interval.
996 self._no_scope_warning_expiry = time.monotonic() + _NO_SCOPE_WARNING_INTERVAL_SECONDS
997 verbose_proxy_logger.warning(
998 "Compresr: enable_retrieval is on but this request has no per-key auth scope; "
999 "compressing without recovery (compresr_retrieve tool not injected). "
1000 "Configure virtual-key auth to enable recovery."
1001 )
1003 existing_originals: Final = self._existing_originals(store_key)
1004 applied: Final = self._apply_compression_results(
1005 messages, targets, contexts, results, recovery_enabled, existing_originals
1006 )
1007 if applied.messages_compressed == 0:
1008 # Nothing replaced: return the original inputs object (handlers detect
1009 # edits by identity; a fresh list forces write-back that strips Anthropic
1010 # cache_control from thinking blocks).
1011 verbose_proxy_logger.debug("Compresr: service returned no compressed content; request unchanged")
1012 return inputs
1014 stats: Final[dict[str, object]] = {
1015 "messages_compressed": applied.messages_compressed,
1016 "tokens_before": applied.tokens_before,
1017 "tokens_after": applied.tokens_after,
1018 "tokens_saved": applied.tokens_before - applied.tokens_after,
1019 "compression_model": self.compression_model,
1020 }
1021 verbose_proxy_logger.debug(
1022 "Compresr: compressed %s message(s), %s -> %s tokens",
1023 applied.messages_compressed,
1024 applied.tokens_before,
1025 applied.tokens_after,
1026 )
1027 self.add_standard_logging_guardrail_information_to_request_data(
1028 guardrail_json_response=stats,
1029 request_data=request_data,
1030 guardrail_status="success",
1031 guardrail_provider="compresr",
1032 start_time=start_time,
1033 end_time=end_time,
1034 duration=end_time - start_time,
1035 )
1037 compressed_inputs: Final[dict[str, object]] = {**inputs, "structured_messages": applied.compressed_messages}
1038 mirrored_texts: Final = self._mirror_texts_channel(inputs.get("texts"), applied)
1039 if mirrored_texts is not None:
1040 compressed_inputs["texts"] = mirrored_texts
1042 originals: Final = applied.originals
1043 if not recovery_enabled or not originals or store_key is None:
1044 return compressed_inputs # pyright: ignore[reportReturnType] # plain dicts satisfy AllMessageValues at runtime
1046 self._store_originals(store_key, originals)
1048 merged_tools: Final = _merge_retrieve_tool(inputs.get("tools"))
1049 if merged_tools is not None:
1050 compressed_inputs["tools"] = merged_tools
1051 return compressed_inputs # pyright: ignore[reportReturnType] # plain dicts satisfy AllMessageValues at runtime
1053 async def async_should_run_agentic_loop(
1054 self,
1055 response: object,
1056 model: str,
1057 messages: list[dict],
1058 tools: list[dict] | None,
1059 stream: bool,
1060 custom_llm_provider: str,
1061 kwargs: dict,
1062 ) -> tuple[bool, dict]:
1063 if not has_compresr_retrieve_tool(tools):
1064 return False, {}
1065 tool_calls: Final = _extract_compresr_tool_calls(response)
1066 if not tool_calls:
1067 return False, {}
1068 return True, {"tool_calls": tool_calls}
1070 async def async_build_agentic_loop_plan(
1071 self,
1072 tools: dict,
1073 model: str,
1074 messages: list[dict],
1075 response: object,
1076 anthropic_messages_provider_config: BaseAnthropicMessagesConfig | None,
1077 anthropic_messages_optional_request_params: dict,
1078 logging_obj: LiteLLMLoggingObj | None,
1079 stream: bool,
1080 kwargs: dict,
1081 ) -> AgenticLoopPlan:
1082 tool_calls: list[dict[str, object]] = tools.get("tool_calls", []) # pyright: ignore[reportAssignmentType] # gate hook builds this dict with list values only
1084 self._prune_originals()
1085 store_key: Final = _scoped_store_key(logging_obj)
1086 retrieved, resolved_any = self._resolve_retrievals(store_key, tool_calls)
1087 if not resolved_any:
1088 # Nothing this guardrail stored resolved; skip the extra provider round-trip.
1089 return AgenticLoopPlan(run_agentic_loop=False)
1091 if _is_responses_api_response(response):
1092 follow_up_messages = list(messages) + _build_responses_followup_items(response, retrieved)
1093 elif _is_anthropic_messages_response(response):
1094 follow_up_messages = list(messages) + _build_anthropic_followup_messages(response, retrieved)
1095 else:
1096 assistant_message: Final = _build_assistant_message_from_response(response, retrieved)
1097 tool_results: Final = [
1098 {"role": "tool", "tool_call_id": tc.get("id"), "content": content} for tc, content in retrieved
1099 ]
1100 follow_up_messages = list(messages) + [assistant_message] + tool_results
1102 anthropic_max: Final = anthropic_messages_optional_request_params.get("max_tokens")
1103 max_tokens: Final[int | None] = anthropic_max if anthropic_max is not None else kwargs.get("max_tokens")
1104 optional_params_without_max_tokens: Final = {
1105 k: v for k, v in anthropic_messages_optional_request_params.items() if k != "max_tokens"
1106 }
1108 full_model_name = model
1109 if logging_obj is not None:
1110 agentic_params: Final = getattr(logging_obj, "model_call_details", {}).get("agentic_loop_params", {})
1111 candidate: Final = agentic_params.get("model", model)
1112 if isinstance(candidate, str) and candidate:
1113 full_model_name = candidate
1115 return AgenticLoopPlan(
1116 run_agentic_loop=True,
1117 request_patch=AgenticLoopRequestPatch(
1118 model=full_model_name,
1119 messages=follow_up_messages,
1120 max_tokens=max_tokens,
1121 optional_params=optional_params_without_max_tokens,
1122 kwargs=self._sanitized_follow_up_kwargs(kwargs),
1123 ),
1124 metadata={"tool_type": "compresr_retrieve"},
1125 )
1127 def _sanitized_follow_up_kwargs(self, kwargs: dict) -> dict[str, object]:
1128 """Copy of the request kwargs for the retrieval follow-up with other
1129 guardrails' pre-call-executed markers stripped, so input guardrails
1130 re-inspect the restored originals; only this guardrail's own marker is
1131 kept, to avoid recompressing what it just retrieved."""
1132 out: Final[dict[str, object]] = {
1133 k: v for k, v in kwargs.items() if not k.startswith("_compresr") and k != "litellm_logging_obj"
1134 }
1135 own_marker: Final = self._pre_call_marker()
1136 for meta_key in ("metadata", "litellm_metadata"):
1137 meta = out.get(meta_key)
1138 if not isinstance(meta, dict):
1139 continue
1140 executed = meta.get(PRE_CALL_EXECUTED_GUARDRAILS_KEY)
1141 if not isinstance(executed, list):
1142 continue
1143 kept = [m for m in executed if own_marker is not None and m == own_marker]
1144 out[meta_key] = (
1145 {**meta, PRE_CALL_EXECUTED_GUARDRAILS_KEY: kept}
1146 if kept
1147 else {k: v for k, v in meta.items() if k != PRE_CALL_EXECUTED_GUARDRAILS_KEY}
1148 )
1149 return out
1151 @staticmethod
1152 def get_config_model() -> type[GuardrailConfigModel[object]] | None:
1153 from litellm.types.proxy.guardrails.guardrail_hooks.compresr import (
1154 CompresrGuardrailConfigModel,
1155 )
1157 return CompresrGuardrailConfigModel