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

1"""Compresr guardrail — query-aware, recoverable context compression. 

2 

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. 

9 

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

15 

16from __future__ import annotations 

17 

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 

27 

28import httpx 

29from fastapi import HTTPException 

30from httpx import Response as HttpxResponse 

31 

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 

62 

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 ) 

73 

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) 

107 

108 

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 

123 

124 

125def _validate_api_base(url: str) -> str: 

126 """Return ``url`` if it passes basic outbound-target checks, else raise. 

127 

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 

145 

146 

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) 

149 

150 

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) 

153 

154 

155def _replace_text_in_content(content: object, new_text: str) -> object: 

156 """Write ``new_text`` back into a ``content`` value, preserving shape. 

157 

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 

169 

170 

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 

184 

185 

186def _query_for_target(messages: list[dict[str, object]], target_idx: int, fallback: str) -> str: 

187 """Query used to compress ``messages[target_idx]``. 

188 

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 

196 

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 

221 

222 

223def _safe_int(value: object) -> int: 

224 """Parse a token-stat field defensively. 

225 

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 

235 

236 

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] 

247 

248 

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] 

253 

254 

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

258 

259 

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

268 

269 

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 ) 

276 

277 

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 } 

300 

301 

302def has_compresr_retrieve_tool(tools: object) -> bool: 

303 return has_tool_with_name(tools, COMPRESR_RETRIEVE_TOOL_NAME) 

304 

305 

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] 

318 

319 

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 ] 

326 

327 

328def _resolve_call_id(logging_obj: object) -> str | None: 

329 """The call id from the framework logging object. 

330 

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 

341 

342 

343def _caller_scope(logging_obj: object) -> str: 

344 """The caller's virtual-key hash, used to partition the recovery store. 

345 

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

371 

372 

373def _scoped_store_key(logging_obj: object) -> str | None: 

374 """Key for the recovery store: caller identity plus framework call id. 

375 

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 

387 

388 

389def _is_responses_api_response(response: object) -> bool: 

390 return isinstance(get_attribute_or_key(response, "output", None), list) 

391 

392 

393def _is_anthropic_messages_response(response: object) -> bool: 

394 return isinstance(get_attribute_or_key(response, "content", None), list) 

395 

396 

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. 

402 

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 } 

424 

425 

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] 

456 

457 

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 

481 

482 

483@dataclass 

484class _CompressionResult: 

485 """Outcome of applying compression results to a message list.""" 

486 

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 

497 

498 

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 ) 

587 

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" 

598 

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 } 

604 

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

620 

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) 

625 

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

638 

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] 

645 

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

657 

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 

675 

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 

688 

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 

716 

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 } 

752 

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 

793 

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 

811 

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 

828 

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 

854 

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 

863 

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``. 

874 

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 

917 

918 @staticmethod 

919 def _mirror_texts_channel(input_texts: object, applied: _CompressionResult) -> list[object] | None: 

920 """Compressed content mirrored into the Responses `texts` channel. 

921 

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] 

942 

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 

953 

954 if self._should_bypass(request_data): 

955 verbose_proxy_logger.debug("Compresr: %s header set; skipping compression", BYPASS_HEADER) 

956 return inputs 

957 

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 

964 

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 

979 

980 contexts: Final = [content_to_text(messages[idx].get("content")) for idx in targets] 

981 

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 

987 

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 ) 

1002 

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 

1013 

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 ) 

1036 

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 

1041 

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 

1045 

1046 self._store_originals(store_key, originals) 

1047 

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 

1052 

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} 

1069 

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 

1083 

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) 

1090 

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 

1101 

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 } 

1107 

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 

1114 

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 ) 

1126 

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 

1150 

1151 @staticmethod 

1152 def get_config_model() -> type[GuardrailConfigModel[object]] | None: 

1153 from litellm.types.proxy.guardrails.guardrail_hooks.compresr import ( 

1154 CompresrGuardrailConfigModel, 

1155 ) 

1156 

1157 return CompresrGuardrailConfigModel