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

371 statements  

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

1from __future__ import annotations 

2 

3import json 

4import math 

5import re 

6import time 

7import uuid 

8from collections.abc import Mapping, Sequence 

9from dataclasses import dataclass 

10from typing import TYPE_CHECKING, ClassVar, Final, Literal, TypeGuard 

11 

12import httpx 

13from fastapi import HTTPException 

14from httpx import Response as HttpxResponse 

15from pydantic import TypeAdapter 

16 

17import litellm 

18from litellm._logging import verbose_proxy_logger 

19from litellm.compression.compress import get_protected_indices 

20from litellm.constants import HTTP_HANDLER_CONNECT_TIMEOUT_SECONDS 

21from litellm.integrations.custom_guardrail import ( 

22 CustomGuardrail, 

23 log_guardrail_information, 

24) 

25from litellm.litellm_core_utils.prompt_templates.factory import ( 

26 get_attribute_or_key, 

27 get_tool_calls_from_response, 

28 group_tool_exchanges, 

29 has_tool_with_name, 

30) 

31from litellm.llms.custom_httpx.http_handler import ( 

32 get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] 

33 httpxSpecialProvider, 

34) 

35from litellm.proxy.guardrails.guardrail_hooks.content_text import ( 

36 assistant_text_from_response, 

37 content_to_text, 

38 is_all_text_parts, 

39 merge_rewritten_text_parts, 

40) 

41from litellm.proxy.spend_tracking.compression_savings import HEADROOM_GUARDRAIL_PROVIDER 

42from litellm.secret_managers.main import get_secret_str 

43from litellm.types.guardrails import GuardrailEventHooks, Mode 

44from litellm.types.integrations.custom_logger import ( 

45 HEADROOM_CONVERTED_STREAM_KEY, 

46 AgenticLoopPlan, 

47 AgenticLoopRequestPatch, 

48) 

49from litellm.types.utils import CallTypes, GenericGuardrailAPIInputs 

50 

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

52 from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj 

53 from litellm.llms.base_llm.anthropic_messages.transformation import BaseAnthropicMessagesConfig 

54 from litellm.types.guardrails import LitellmParams 

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

56 

57BYPASS_HEADER: Final = "x-headroom-bypass" 

58_STREAM_CONVERTIBLE_CALL_TYPES: Final = frozenset( 

59 (CallTypes.completion, CallTypes.acompletion, CallTypes.responses, CallTypes.aresponses) 

60) 

61# The shared GuardrailCallback client carries no per-call bound, so without this a 

62# stalled service holds the caller's request and a pooled connection for 600s or more. 

63_COMPRESS_TIMEOUT_SECONDS: Final = 60.0 

64HEADROOM_RETRIEVE_TOOL_NAME: Final = "headroom_retrieve" 

65_HASH_PATTERN: Final = re.compile(r"[a-f0-9]{12,24}") 

66_HASH_CACHE_TTL_SECONDS: Final = 15 * 60 

67# Narrows the base class's bare-dict ``request_data`` at the boundary so its 

68# untranslated messages can be read with concrete types (values pass through by 

69# reference, so this is a shallow top-level reconstruction). 

70_REQUEST_DATA_ADAPTER: Final = TypeAdapter(dict[str, object]) 

71 

72 

73def _is_str_object_dict(value: object) -> TypeGuard[dict[str, object]]: # guard-ok: isinstance narrows correctly; predicate is trivially correct # fmt: skip 

74 return isinstance(value, dict) 

75 

76 

77def _is_object_list(value: object) -> TypeGuard[list[object]]: # guard-ok: isinstance narrows correctly; predicate is trivially correct # fmt: skip 

78 return isinstance(value, list) 

79 

80 

81def _flatten_messages_for_compression(messages: list[dict[str, object]]) -> list[dict[str, object]]: 

82 """Collapse all-text list-of-parts content to plain strings for /v1/compress. 

83 

84 The compression service's transforms only rewrite string content and skip 

85 the OpenAI list-of-parts shape, which is what every Anthropic-format 

86 request translates to. Only rows whose parts are ALL text are flattened: 

87 cache_control breakpoints are positional (each caches the prefix ending 

88 at its part), so merging text across a non-text part would move a later 

89 breakpoint to the other side of it. Rows with non-text parts are sent 

90 unchanged and pass through the service untouched. 

91 """ 

92 flattened: Final[list[dict[str, object]]] = [] 

93 for msg in messages: 

94 content = msg.get("content") 

95 if is_all_text_parts(content): 

96 text = content_to_text(content) 

97 if text: 

98 flattened.append({**msg, "content": text}) 

99 continue 

100 flattened.append(msg) 

101 return flattened 

102 

103 

104def _restore_content_shapes( 

105 originals: list[dict[str, object]], returned: list[dict[str, object]] 

106) -> list[dict[str, object]]: 

107 """Write compressed text back into each original row's content shape. 

108 

109 Rows are matched positionally; the pairing is only trusted when the 

110 service kept the row count and every role lines up. If it restructured 

111 the conversation (e.g. dropped rows), its output is adopted as-is, which 

112 is the pre-flattening behavior. 

113 """ 

114 if len(returned) != len(originals): 

115 return returned 

116 for orig, ret in zip(originals, returned): 

117 if orig.get("role") != ret.get("role"): 

118 return returned 

119 restored: Final[list[dict[str, object]]] = [] 

120 for orig, ret in zip(originals, returned): 

121 orig_content = orig.get("content") 

122 ret_content = ret.get("content") 

123 if isinstance(orig_content, list) and isinstance(ret_content, str): 

124 if ret_content == content_to_text(orig_content): 

125 # Untouched row: keep the exact original parts, including 

126 # per-part fields like cache_control on later text parts. 

127 restored.append({**ret, "content": orig_content}) 

128 else: 

129 restored.append({**ret, "content": merge_rewritten_text_parts(orig_content, ret_content)}) 

130 else: 

131 restored.append(ret) 

132 return restored 

133 

134 

135def _tool_call_name(tool_call: Mapping[str, object]) -> str | None: 

136 function: Final = tool_call.get("function") 

137 if not _is_str_object_dict(function): 

138 return None 

139 name: Final = function.get("name") 

140 return name if isinstance(name, str) else None 

141 

142 

143def _is_retrieve_tool_name(name: str | None) -> bool: 

144 """Match the retrieve tool whether called directly or via the MCP gateway. 

145 

146 Server-side the tool is ``headroom_retrieve``; exposed through LiteLLM's MCP 

147 gateway a client calls it as ``mcp__<server>__headroom_retrieve``. 

148 """ 

149 return name is not None and ( 

150 name == HEADROOM_RETRIEVE_TOOL_NAME or name.endswith(f"__{HEADROOM_RETRIEVE_TOOL_NAME}") 

151 ) 

152 

153 

154def _retrieve_call_ids_in_message(message: Mapping[str, object]) -> frozenset[str]: 

155 if message.get("role") != "assistant": 

156 return frozenset() 

157 tool_calls: Final = message.get("tool_calls") 

158 if not _is_object_list(tool_calls): 

159 return frozenset() 

160 return frozenset( 

161 str(tool_call["id"]) 

162 for tool_call in tool_calls 

163 if _is_str_object_dict(tool_call) and tool_call.get("id") and _is_retrieve_tool_name(_tool_call_name(tool_call)) 

164 ) 

165 

166 

167def _anthropic_tool_use_retrieve_id(block: object) -> str | None: 

168 if not _is_str_object_dict(block) or block.get("type") != "tool_use": 

169 return None 

170 name: Final = block.get("name") 

171 call_id: Final = block.get("id") 

172 if isinstance(name, str) and call_id is not None and _is_retrieve_tool_name(name): 

173 return str(call_id) 

174 return None 

175 

176 

177def _anthropic_retrieve_ids_in_message(message: Mapping[str, object]) -> frozenset[str]: 

178 content: Final = message.get("content") 

179 if not _is_object_list(content): 

180 return frozenset() 

181 return frozenset(call_id for block in content if (call_id := _anthropic_tool_use_retrieve_id(block)) is not None) 

182 

183 

184def _raw_retrieve_call_ids(messages: object) -> frozenset[str]: 

185 """Retrieve-tool call ids read from the request's own, untranslated messages. 

186 

187 The guardrail otherwise scans an OpenAI-translated view where a tool name 

188 over 64 chars is truncated to ``{prefix}_{hash}``, which drops the 

189 ``__headroom_retrieve`` suffix a long ``mcp__<server>__`` prefix pushes past 

190 the limit. Tool-call ids are never truncated, so pairing the tool result to 

191 an id read from the original request keeps the match intact. Both wire 

192 shapes are handled: OpenAI ``tool_calls`` and Anthropic ``tool_use`` blocks. 

193 """ 

194 if not _is_object_list(messages): 

195 return frozenset() 

196 return frozenset( 

197 call_id 

198 for message in messages 

199 if _is_str_object_dict(message) 

200 for call_id in _retrieve_call_ids_in_message(message) | _anthropic_retrieve_ids_in_message(message) 

201 ) 

202 

203 

204def _retrieval_result_indices( 

205 messages: Sequence[Mapping[str, object]], extra_retrieve_call_ids: frozenset[str] = frozenset() 

206) -> frozenset[int]: 

207 """Indices of tool-result rows that carry ``headroom_retrieve`` output. 

208 

209 When the retrieve tool is exposed to a client that runs its own tool loop 

210 (the LiteLLM MCP gateway path), the client executes the call and sends the 

211 recovered original content back as a tool result on the next turn. That 

212 content is exactly what a prior compression stubbed, so compressing it again 

213 re-derives the identical content hash: a no-op that strands the model on the 

214 marker and loops the agent. Hold those rows back so the expansion survives. 

215 

216 ``extra_retrieve_call_ids`` carries ids recovered from the untruncated 

217 request so the pairing survives tool-name truncation (see 

218 ``_raw_retrieve_call_ids``). 

219 """ 

220 retrieve_call_ids: Final = extra_retrieve_call_ids | frozenset( 

221 call_id for message in messages for call_id in _retrieve_call_ids_in_message(message) 

222 ) 

223 if not retrieve_call_ids: 

224 return frozenset() 

225 return frozenset( 

226 index 

227 for index, message in enumerate(messages) 

228 if message.get("role") in ("tool", "function") and str(message.get("tool_call_id")) in retrieve_call_ids 

229 ) 

230 

231 

232def _protected_indices( 

233 messages: Sequence[Mapping[str, object]], extra_retrieve_call_ids: frozenset[str] = frozenset() 

234) -> frozenset[int]: 

235 """Indices headroom must not send to the compression service. 

236 

237 ``get_protected_indices`` is litellm's own compression policy: the system 

238 rows, the last user row, the last assistant row. Rows carrying just-retrieved 

239 ``headroom_retrieve`` output are added so re-compression can't collapse them 

240 back to the marker they were expanded from. The union is expanded over whole 

241 tool exchanges the way ``compress()`` expands it, so a protected assistant 

242 tool call cannot end up answered by a marker standing in for the result the 

243 model just asked for. 

244 

245 Every assistant row is then withheld without expanding its tool exchange: 

246 the service protects assistant text blocks but has no gate for assistant 

247 strings, and the Anthropic adapter hands assistant blocks over as strings, 

248 so the model's own earlier tables came back rewritten and it imitated the 

249 shape. The tool results those turns asked for stay compressible. 

250 """ 

251 protected: Final = frozenset(get_protected_indices(messages)) | _retrieval_result_indices( 

252 messages, extra_retrieve_call_ids 

253 ) 

254 return ( 

255 protected 

256 | frozenset( 

257 index 

258 for group in group_tool_exchanges(messages) 

259 if any(member in protected for member in group) 

260 for index in group 

261 ) 

262 | frozenset(index for index, message in enumerate(messages) if message.get("role") == "assistant") 

263 ) 

264 

265 

266def _restore_protected_messages( 

267 messages: Sequence[dict[str, object]], 

268 compressed: Sequence[dict[str, object]], 

269 protected_indices: frozenset[int], 

270) -> Sequence[dict[str, object]]: 

271 """Put the rows that were held back from compression at their original positions. 

272 

273 Requires one returned row per row actually sent, which ``_call_compress`` 

274 enforces; a service that changed the row count is treated as a failure 

275 there, because a reshaped conversation cannot be re-interleaved. 

276 """ 

277 sent_positions: Final = tuple(index for index in range(len(messages)) if index not in protected_indices) 

278 compressed_by_index: Final = dict(zip(sent_positions, compressed)) 

279 return [ 

280 messages[index] if index in protected_indices else compressed_by_index[index] for index in range(len(messages)) 

281 ] 

282 

283 

284def _build_compress_failure_detail(status_code: int, body: str) -> dict[str, object]: 

285 """Build error details for failed /v1/compress responses. 

286 

287 Adds troubleshooting hints for known deployment-related errors while 

288 preserving the upstream status code and response body. 

289 """ 

290 if status_code == 404: 

291 return { 

292 "status_code": status_code, 

293 "body": body, 

294 "hint": ( 

295 "The Headroom compression endpoint returned HTTP 404. " 

296 "Verify that the configured Headroom endpoint is correct and that " 

297 "the compression endpoint is available. If you are using a " 

298 "self-hosted deployment, some deployments require enabling remote " 

299 "compression (for example, HEADROOM_COMPRESS_ALLOW_REMOTE=1)." 

300 ), 

301 } 

302 return {"status_code": status_code, "body": body} 

303 

304 

305def _read_ccr_hashes(body: Mapping[str, object]) -> frozenset[str]: 

306 ccr_hashes: Final = body.get("ccr_hashes") 

307 if not isinstance(ccr_hashes, list): 

308 return frozenset() 

309 return frozenset( 

310 hash_value.lower() 

311 for hash_value in ccr_hashes 

312 if isinstance(hash_value, str) and _HASH_PATTERN.fullmatch(hash_value.lower()) 

313 ) 

314 

315 

316@dataclass(frozen=True, slots=True) 

317class _CompressResult: 

318 messages: list[dict[str, object]] 

319 succeeded: bool 

320 stats: dict[str, object] 

321 ccr_hashes: frozenset[str] = frozenset() 

322 

323 

324def _build_headroom_retrieve_tool() -> dict[str, object]: 

325 return { 

326 "type": "function", 

327 "function": { 

328 "name": HEADROOM_RETRIEVE_TOOL_NAME, 

329 "description": ( 

330 "Retrieve original content that was compressed by Headroom. " 

331 "Call this when you encounter a compression marker containing a hash." 

332 ), 

333 "parameters": { 

334 "type": "object", 

335 "properties": { 

336 "hash": { 

337 "type": "string", 

338 "description": "The hex hash from the compression marker.", 

339 }, 

340 "query": { 

341 "type": "string", 

342 "description": "Optional search query for BM25-ranked retrieval.", 

343 }, 

344 }, 

345 "required": ["hash"], 

346 }, 

347 }, 

348 } 

349 

350 

351def _resolve_call_id(logging_obj: object, request_state: dict[str, object]) -> str | None: 

352 """Resolve the litellm_call_id shared by a request's pre-call hook and its 

353 agentic-loop hooks, so CCR hash validation can be scoped per call instead 

354 of trusting any hash-shaped string that shows up in message text.""" 

355 logging_call_id: Final = getattr(logging_obj, "litellm_call_id", None) 

356 if isinstance(logging_call_id, str) and logging_call_id: 

357 return logging_call_id 

358 kwargs_call_id: Final = request_state.get("litellm_call_id") 

359 return kwargs_call_id if isinstance(kwargs_call_id, str) else None 

360 

361 

362def has_headroom_retrieve_tool(tools: object) -> bool: 

363 return has_tool_with_name(tools, HEADROOM_RETRIEVE_TOOL_NAME) 

364 

365 

366def _extract_headroom_tool_calls(response: object) -> list[dict[str, object]]: 

367 return [ 

368 {"id": tc["id"], "type": "function", "name": tc["name"], "arguments": tc["arguments"]} 

369 for tc in get_tool_calls_from_response(response) 

370 if tc["name"] == HEADROOM_RETRIEVE_TOOL_NAME 

371 ] 

372 

373 

374def _build_assistant_message_from_response( 

375 response: object, 

376 retrieved: Sequence[tuple[dict[str, object], str]], 

377) -> dict[str, object]: 

378 """Rebuild the chat-completions assistant turn for the retrieval follow-up. 

379 

380 Only the ``headroom_retrieve`` calls are echoed, each answered by a tool 

381 result below. Other tool calls made in the same turn are omitted on purpose: 

382 the follow-up re-runs the model with the recovered content so it re-plans 

383 them. Echoing them would leave tool_calls with no matching tool result and 

384 the provider would reject the request. 

385 """ 

386 return { 

387 "role": "assistant", 

388 "content": assistant_text_from_response(response), 

389 "tool_calls": [ 

390 { 

391 "id": tool_call.get("id"), 

392 "type": "function", 

393 "function": { 

394 "name": tool_call.get("name"), 

395 "arguments": json.dumps(tool_call.get("arguments", {})), 

396 }, 

397 } 

398 for tool_call, _ in retrieved 

399 ], 

400 } 

401 

402 

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

404 # Real response objects can be plain dicts at runtime (e.g. TypedDict-based 

405 # response types), so getattr alone would silently miss the key -- use the 

406 # same dict-or-object accessor as the tool-call extractors. 

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

408 

409 

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

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

412 

413 

414def _build_anthropic_followup_messages( 

415 response: object, 

416 retrieved: list[tuple[dict[str, object], str]], 

417) -> list[dict[str, object]]: 

418 """Build Anthropic Messages API follow-up messages for a tool round-trip. 

419 

420 Anthropic requires the tool_use block to be echoed back in an assistant 

421 message, paired with a tool_result block in a user message keyed by the 

422 same tool_use_id -- it does not accept chat-style tool-role messages. Any 

423 text the model wrote alongside the tool call is preserved, so its reasoning 

424 survives into the follow-up turn. 

425 """ 

426 text: Final = assistant_text_from_response(response) 

427 assistant_message: Final[dict[str, object]] = { 

428 "role": "assistant", 

429 "content": ([{"type": "text", "text": text}] if text else []) 

430 + [ 

431 { 

432 "type": "tool_use", 

433 "id": tool_call.get("id"), 

434 "name": tool_call.get("name"), 

435 "input": tool_call.get("arguments", {}), 

436 } 

437 for tool_call, _ in retrieved 

438 ], 

439 } 

440 user_message: Final[dict[str, object]] = { 

441 "role": "user", 

442 "content": [ 

443 {"type": "tool_result", "tool_use_id": tool_call.get("id"), "content": content} 

444 for tool_call, content in retrieved 

445 ], 

446 } 

447 return [assistant_message, user_message] 

448 

449 

450def _build_responses_followup_items( 

451 response: object, 

452 retrieved: list[tuple[dict[str, object], str]], 

453) -> list[dict[str, object]]: 

454 """Build Responses API input items for a tool round-trip. 

455 

456 The Responses API does not accept chat-style assistant/tool messages as 

457 follow-up input; it requires the model's function_call to be echoed back 

458 paired with a function_call_output keyed by the same call_id. Any text the 

459 model wrote alongside the tool call is preserved. 

460 """ 

461 text: Final = assistant_text_from_response(response) 

462 items: Final[list[dict[str, object]]] = [{"role": "assistant", "content": text}] if text else [] 

463 for tool_call, content in retrieved: 

464 call_id = tool_call.get("id") 

465 items.append( 

466 { 

467 "type": "function_call", 

468 "call_id": call_id, 

469 "name": tool_call.get("name"), 

470 "arguments": json.dumps(tool_call.get("arguments", {})), 

471 } 

472 ) 

473 items.append({"type": "function_call_output", "call_id": call_id, "output": content}) 

474 return items 

475 

476 

477class HeadroomGuardrail(CustomGuardrail): 

478 records_own_guardrail_information: ClassVar[bool] = True 

479 server_fulfilled_tool_names: ClassVar[frozenset[str]] = frozenset({HEADROOM_RETRIEVE_TOOL_NAME}) 

480 

481 @classmethod 

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

483 return [ 

484 GuardrailEventHooks.pre_call, 

485 GuardrailEventHooks.post_call, 

486 ] 

487 

488 def __init__( 

489 self, 

490 api_base: str | None = None, 

491 api_key: str | None = None, 

492 model: str | None = None, 

493 guardrail_name: str | None = None, 

494 event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None = None, 

495 default_on: bool = False, 

496 unreachable_fallback: str | None = None, 

497 timeout: float | None = None, 

498 ccr_retrieval: bool = True, 

499 ): 

500 self.headroom_api_base = (api_base or get_secret_str("HEADROOM_API_BASE") or "").rstrip("/") 

501 if not self.headroom_api_base: 

502 raise ValueError( 

503 "Headroom guardrail requires an API base URL. " 

504 "Set `api_base` in the guardrail config or HEADROOM_API_BASE env var." 

505 ) 

506 self.headroom_api_key = api_key or get_secret_str("HEADROOM_API_KEY") 

507 self.headroom_model = model 

508 self.unreachable_fallback: Literal["fail_closed", "fail_open"] = ( 

509 "fail_open" if unreachable_fallback == "fail_open" else "fail_closed" 

510 ) 

511 self.timeout: httpx.Timeout = self._resolve_timeout(timeout) 

512 self.ccr_retrieval = ccr_retrieval 

513 self.async_handler = get_async_httpx_client( 

514 llm_provider=httpxSpecialProvider.GuardrailCallback, 

515 ) 

516 self._issued_hashes_by_call_id: dict[str, tuple[frozenset[str], float]] = {} 

517 super().__init__( # pyright: ignore[reportUnknownMemberType] 

518 guardrail_name=guardrail_name, 

519 event_hook=event_hook, 

520 default_on=default_on, 

521 supported_event_hooks=list(self.get_supported_event_hooks()), 

522 ) 

523 

524 def _should_bypass(self, request_data: dict) -> bool: 

525 psr: Final = request_data.get("proxy_server_request") 

526 if not _is_str_object_dict(psr): 

527 return False 

528 headers: Final = psr.get("headers") 

529 if not _is_str_object_dict(headers): 

530 return False 

531 value: Final = headers.get(BYPASS_HEADER) 

532 return str(value).lower() == "true" 

533 

534 def _request_headers(self) -> dict[str, str]: 

535 headers: Final[dict[str, str]] = {"Content-Type": "application/json"} 

536 if self.headroom_api_key: 

537 headers["Authorization"] = f"Bearer {self.headroom_api_key}" 

538 return headers 

539 

540 @staticmethod 

541 def _resolve_timeout(timeout: float | None) -> httpx.Timeout: 

542 """Budget for one call to the compression service, unset meaning the default. 

543 

544 Zero, negative and non-finite values are rejected instead of passed through: 

545 httpx accepts them, and the transport then reads 0 and inf as no deadline at 

546 all and a negative one as a deadline already past. 

547 """ 

548 rejected: Final = timeout is not None and not (math.isfinite(timeout) and timeout > 0) 

549 if rejected: 

550 verbose_proxy_logger.warning( 

551 "Headroom: ignoring unusable timeout %s, using %s seconds", 

552 timeout, 

553 _COMPRESS_TIMEOUT_SECONDS, 

554 ) 

555 seconds: Final = _COMPRESS_TIMEOUT_SECONDS if timeout is None or rejected else timeout 

556 return httpx.Timeout(timeout=seconds, connect=min(seconds, HTTP_HANDLER_CONNECT_TIMEOUT_SECONDS)) 

557 

558 def update_in_memory_litellm_params(self, litellm_params: LitellmParams) -> None: 

559 """Re-resolve the timeout, which the base implementation would otherwise null out.""" 

560 super().update_in_memory_litellm_params(litellm_params) 

561 self.timeout = self._resolve_timeout(litellm_params.timeout) 

562 

563 def _prune_expired_hashes(self) -> None: 

564 now: Final = time.monotonic() 

565 self._issued_hashes_by_call_id = { 

566 call_id: (hashes, expiry) 

567 for call_id, (hashes, expiry) in self._issued_hashes_by_call_id.items() 

568 if expiry > now 

569 } 

570 

571 def _handle_compress_failure( 

572 self, 

573 messages: list[dict[str, object]], 

574 error: str, 

575 detail: dict[str, object], 

576 ) -> list[dict[str, object]]: 

577 if self.unreachable_fallback == "fail_open": 

578 verbose_proxy_logger.critical( 

579 "Headroom: %s; fail_open configured, forwarding request uncompressed. detail=%s", 

580 error, 

581 detail, 

582 ) 

583 return messages 

584 raise HTTPException(status_code=502, detail={"error": error, **detail}) 

585 

586 async def _call_compress( 

587 self, 

588 messages: list[dict[str, object]], 

589 model: str | None, 

590 ) -> _CompressResult: 

591 payload: Final[dict[str, object]] = {"messages": messages} 

592 if model: 

593 payload["model"] = model 

594 

595 try: 

596 raw_response: HttpxResponse = await self.async_handler.post( # pyright: ignore[reportUnknownMemberType] # AsyncHTTPHandler.post is untyped 

597 url=f"{self.headroom_api_base}/v1/compress", 

598 json=payload, 

599 headers=self._request_headers(), 

600 timeout=self.timeout, 

601 ) 

602 except httpx.HTTPStatusError as e: 

603 return _CompressResult( 

604 self._handle_compress_failure( 

605 messages, 

606 "Headroom compression service returned an error", 

607 _build_compress_failure_detail(e.response.status_code, e.response.text), 

608 ), 

609 False, 

610 {}, 

611 ) 

612 except (httpx.ConnectError, httpx.TimeoutException, httpx.TransportError, litellm.Timeout) as e: 

613 return _CompressResult( 

614 self._handle_compress_failure( 

615 messages, 

616 "Headroom compression service unreachable", 

617 {"detail": str(e)}, 

618 ), 

619 False, 

620 {}, 

621 ) 

622 response: Final[HttpxResponse] = raw_response 

623 

624 if response.status_code != 200: 

625 return _CompressResult( 

626 self._handle_compress_failure( 

627 messages, 

628 "Headroom compression service returned an error", 

629 _build_compress_failure_detail(response.status_code, response.text), 

630 ), 

631 False, 

632 {}, 

633 ) 

634 

635 try: 

636 body: Final[object] = response.json() 

637 except ValueError: 

638 return _CompressResult( 

639 self._handle_compress_failure( 

640 messages, 

641 "Headroom compression service returned non-JSON response", 

642 {"body": response.text[:500]}, 

643 ), 

644 False, 

645 {}, 

646 ) 

647 if not _is_str_object_dict(body): 

648 return _CompressResult( 

649 self._handle_compress_failure( 

650 messages, 

651 "Headroom compression service returned unexpected response shape", 

652 {"body": response.text[:500]}, 

653 ), 

654 False, 

655 {}, 

656 ) 

657 

658 compressed_messages: Final = body.get("messages") 

659 if not _is_object_list(compressed_messages): 

660 return _CompressResult( 

661 self._handle_compress_failure( 

662 messages, 

663 "Headroom compression service response missing 'messages'", 

664 {"body": response.text}, 

665 ), 

666 False, 

667 {}, 

668 ) 

669 

670 filtered: Final = [item for item in compressed_messages if _is_str_object_dict(item)] 

671 if not filtered: 

672 return _CompressResult( 

673 self._handle_compress_failure( 

674 messages, 

675 "Headroom compression service returned empty message list", 

676 {"body": response.text}, 

677 ), 

678 False, 

679 {}, 

680 ) 

681 

682 if len(filtered) != len(messages): 

683 # Rows are matched positionally when the never-compressed messages 

684 # are put back, so a reshaped conversation cannot be applied at all. 

685 return _CompressResult( 

686 self._handle_compress_failure( 

687 messages, 

688 "Headroom compression service changed the message count", 

689 {"sent": len(messages), "returned": len(filtered)}, 

690 ), 

691 False, 

692 {}, 

693 ) 

694 

695 verbose_proxy_logger.debug( 

696 "Headroom: compressed %s tokens -> %s tokens (ratio %.2f)", 

697 body.get("tokens_before", "?"), 

698 body.get("tokens_after", "?"), 

699 body.get("compression_ratio", 0), 

700 ) 

701 

702 stats: Final = { 

703 key: body[key] 

704 for key in ( 

705 "tokens_before", 

706 "tokens_after", 

707 "tokens_saved", 

708 "compression_ratio", 

709 "transforms_applied", 

710 ) 

711 if key in body 

712 } 

713 tokens_before: Final = stats.get("tokens_before") 

714 tokens_after: Final = stats.get("tokens_after") 

715 if ( 

716 "tokens_saved" not in stats 

717 and isinstance(tokens_before, (int, float)) 

718 and not isinstance(tokens_before, bool) 

719 and isinstance(tokens_after, (int, float)) 

720 and not isinstance(tokens_after, bool) 

721 ): 

722 # Spend tracking (extract_compression_saved_tokens) reads only 

723 # tokens_saved, which the live compression service omits; derive it 

724 # so savings are counted, but let a service-sent value win. 

725 stats["tokens_saved"] = tokens_before - tokens_after 

726 return _CompressResult(filtered, True, stats, _read_ccr_hashes(body)) 

727 

728 async def _call_retrieve(self, hash_value: str, query: str | None = None) -> str: 

729 params: Final[dict[str, str]] = {} 

730 if query: 

731 params["query"] = query 

732 

733 try: 

734 raw_response: HttpxResponse = await self.async_handler.get( # pyright: ignore[reportUnknownMemberType] # AsyncHTTPHandler.get is untyped 

735 url=f"{self.headroom_api_base}/v1/retrieve/{hash_value}", 

736 params=params, 

737 headers=self._request_headers(), 

738 timeout=self.timeout, 

739 ) 

740 except (httpx.ConnectError, httpx.TimeoutException, httpx.TransportError, litellm.Timeout) as e: 

741 verbose_proxy_logger.warning("Headroom: retrieve failed for hash=%s: %s", hash_value, e) 

742 return f"[Headroom: retrieval failed for hash={hash_value}]" 

743 

744 if raw_response.status_code == 404: 

745 return f"[Headroom: hash={hash_value} not found or expired]" 

746 

747 if raw_response.status_code != 200: 

748 verbose_proxy_logger.warning( 

749 "Headroom: retrieve returned %s for hash=%s", 

750 raw_response.status_code, 

751 hash_value, 

752 ) 

753 return f"[Headroom: retrieval error {raw_response.status_code} for hash={hash_value}]" 

754 

755 try: 

756 body: Final[object] = raw_response.json() 

757 except ValueError: 

758 return raw_response.text 

759 

760 if _is_str_object_dict(body): 

761 original_content: Final = body.get("original_content") 

762 if isinstance(original_content, str): 

763 return original_content 

764 

765 return str(body) 

766 

767 @log_guardrail_information 

768 async def apply_guardrail( 

769 self, 

770 inputs: GenericGuardrailAPIInputs, 

771 request_data: dict, 

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

773 logging_obj: LiteLLMLoggingObj | None = None, 

774 ) -> GenericGuardrailAPIInputs: 

775 if input_type != "request": 

776 return inputs 

777 

778 if self._should_bypass(request_data): 

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

780 return inputs 

781 

782 if request_data.get("background"): 

783 verbose_proxy_logger.debug("Headroom: background request; skipping compression") 

784 return inputs 

785 

786 structured_messages: Final = inputs.get("structured_messages") 

787 if not _is_object_list(structured_messages) or not structured_messages: 

788 return inputs 

789 

790 messages: Final = [m for m in structured_messages if _is_str_object_dict(m)] 

791 if not messages: 

792 return inputs 

793 

794 # The last user message is the instruction the model is being asked to 

795 # act on, so replacing it with a marker means the model answers a 

796 # retrieval result instead of the request. Protected rows are held back 

797 # from the payload rather than pinned after the fact, so their tokens 

798 # are not counted as savings we never apply; the Anthropic write-back 

799 # discards a compressed system prompt outright. Keep it that way unless 

800 # /v1/compress grows a field for sending the live turn as the retrieval 

801 # query without compressing it: query-aware compression reads the newest 

802 # user message, so it is withheld here at some cost to history ranking. 

803 # request_data is a bare dict on the base signature; narrow it before 

804 # reading the untranslated messages so long tool names can be recovered. 

805 raw_messages: Final = _REQUEST_DATA_ADAPTER.validate_python(request_data).get("messages") 

806 raw_retrieve_call_ids: Final = _raw_retrieve_call_ids(raw_messages) 

807 protected_indices: Final = _protected_indices(messages, raw_retrieve_call_ids) 

808 compressible: Final = [m for i, m in enumerate(messages) if i not in protected_indices] 

809 if not compressible: 

810 return inputs 

811 

812 model: Final = self.headroom_model or request_data.get("model") 

813 start_time: Final = time.time() 

814 result: Final = await self._call_compress( 

815 messages=_flatten_messages_for_compression(compressible), 

816 model=model if isinstance(model, str) else None, 

817 ) 

818 end_time: Final = time.time() 

819 

820 from litellm.proxy.common_utils.callback_utils import ( 

821 add_guardrail_to_applied_guardrails_header, 

822 ) 

823 

824 if not result.succeeded: 

825 self.add_standard_logging_guardrail_information_to_request_data( 

826 guardrail_json_response={"error": "headroom compression unavailable; request forwarded uncompressed"}, 

827 request_data=request_data, 

828 guardrail_status="guardrail_failed_to_respond", 

829 guardrail_provider=HEADROOM_GUARDRAIL_PROVIDER, 

830 start_time=start_time, 

831 end_time=end_time, 

832 duration=end_time - start_time, 

833 ) 

834 add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name) 

835 # Hand back the caller's own inputs object. Translation handlers 

836 # detect "the guardrail rewrote the messages" by identity, so 

837 # returning a rebuilt copy sends an unchanged request through the 

838 # write-back and restructures it for nothing. 

839 return inputs 

840 

841 compressed: Final = _restore_protected_messages( 

842 messages=messages, 

843 compressed=_restore_content_shapes(originals=compressible, returned=result.messages), 

844 protected_indices=protected_indices, 

845 ) 

846 

847 self.add_standard_logging_guardrail_information_to_request_data( 

848 guardrail_json_response=result.stats, 

849 request_data=request_data, 

850 guardrail_status="success", 

851 guardrail_provider=HEADROOM_GUARDRAIL_PROVIDER, 

852 start_time=start_time, 

853 end_time=end_time, 

854 duration=end_time - start_time, 

855 ) 

856 add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name) 

857 

858 hashes: Final = result.ccr_hashes if self.ccr_retrieval else frozenset() 

859 if not hashes: 

860 return {**inputs, "structured_messages": compressed} # pyright: ignore[reportReturnType] 

861 

862 self._prune_expired_hashes() 

863 call_id = _resolve_call_id(logging_obj, request_data) 

864 if not call_id: 

865 call_id = str(uuid.uuid4()) 

866 request_data["litellm_call_id"] = call_id 

867 self._issued_hashes_by_call_id[call_id] = (frozenset(hashes), time.monotonic() + _HASH_CACHE_TTL_SECONDS) 

868 

869 existing_tools: Final = inputs.get("tools") 

870 retrieve_tool: Final = _build_headroom_retrieve_tool() 

871 if isinstance(existing_tools, list) and not has_headroom_retrieve_tool(existing_tools): 

872 merged_tools: list[object] = list(existing_tools) + [retrieve_tool] 

873 elif existing_tools is None: 

874 merged_tools = [retrieve_tool] 

875 else: 

876 merged_tools = list(existing_tools) if isinstance(existing_tools, list) else [retrieve_tool] 

877 

878 return {**inputs, "structured_messages": compressed, "tools": merged_tools} # pyright: ignore[reportReturnType] 

879 

880 async def async_pre_call_deployment_hook( 

881 self, 

882 kwargs: dict[str, object], 

883 call_type: CallTypes | None, 

884 ) -> dict[str, object] | None: # mutable-ok: overrides CustomLogger hook whose contract is a plain dict 

885 base_result: Final = await super().async_pre_call_deployment_hook(kwargs, call_type) 

886 effective: Final = base_result if base_result is not None else kwargs 

887 if call_type not in _STREAM_CONVERTIBLE_CALL_TYPES: 

888 return base_result 

889 if not effective.get("stream") or effective.get("background"): 

890 return base_result 

891 if not has_headroom_retrieve_tool(effective.get("tools")): 

892 return base_result 

893 return { # mutable-ok: the hook contract is a plain dict the router merges into the request kwargs 

894 **effective, 

895 "stream": False, 

896 HEADROOM_CONVERTED_STREAM_KEY: True, 

897 } 

898 

899 async def async_should_run_agentic_loop( 

900 self, 

901 response: object, 

902 model: str, 

903 messages: list[dict], 

904 tools: list[dict] | None, 

905 stream: bool, 

906 custom_llm_provider: str, 

907 kwargs: dict, 

908 ) -> tuple[bool, dict]: 

909 if not has_headroom_retrieve_tool(tools): 

910 return False, {} 

911 

912 tool_calls: Final = _extract_headroom_tool_calls(response) 

913 if not tool_calls: 

914 return False, {} 

915 

916 return True, {"tool_calls": tool_calls} 

917 

918 async def async_build_agentic_loop_plan( 

919 self, 

920 tools: dict, 

921 model: str, 

922 messages: list[dict], 

923 response: object, 

924 anthropic_messages_provider_config: BaseAnthropicMessagesConfig | None, 

925 anthropic_messages_optional_request_params: dict, 

926 logging_obj: LiteLLMLoggingObj | None, 

927 stream: bool, 

928 kwargs: dict, 

929 ) -> AgenticLoopPlan: 

930 tool_calls: Final[list[dict[str, object]]] = tools.get("tool_calls", []) 

931 

932 self._prune_expired_hashes() 

933 call_id: Final = _resolve_call_id(logging_obj, kwargs) 

934 valid_hashes = self._issued_hashes_by_call_id.get(call_id, (frozenset(), 0.0))[0] if call_id else frozenset() 

935 

936 retrieved: Final[list[tuple[dict[str, object], str]]] = [] 

937 for tc in tool_calls: 

938 arguments = tc.get("arguments", {}) 

939 raw_hash = arguments.get("hash", "") if isinstance(arguments, dict) else "" 

940 hash_value = str(raw_hash).lower() 

941 query = arguments.get("query") if isinstance(arguments, dict) else None 

942 # A hash is only honored if it was issued by *this request's own* 

943 # Headroom /v1/compress call, scoped by litellm_call_id. Scoping by 

944 # message text alone is forgeable -- an attacker can plant a 

945 # hash-shaped string in their own prompt, and a hash issued for one 

946 # request would validate for any other request that echoes it back. 

947 if hash_value not in valid_hashes: 

948 verbose_proxy_logger.warning( 

949 "Headroom CCR: rejecting hash=%s not produced by current request compression", 

950 hash_value, 

951 ) 

952 content = f"[Headroom: hash={hash_value} was not produced by the current request]" 

953 else: 

954 content = await self._call_retrieve( 

955 hash_value=hash_value, 

956 query=str(query) if query else None, 

957 ) 

958 verbose_proxy_logger.debug("Headroom CCR: retrieved hash=%s (%d chars)", hash_value, len(content)) 

959 retrieved.append((tc, content)) 

960 

961 if _is_responses_api_response(response): 

962 follow_up_messages = list(messages) + _build_responses_followup_items(response, retrieved) 

963 elif _is_anthropic_messages_response(response): 

964 follow_up_messages = list(messages) + _build_anthropic_followup_messages(response, retrieved) 

965 else: 

966 assistant_message: Final = _build_assistant_message_from_response(response, retrieved) 

967 tool_results: Final = [ 

968 {"role": "tool", "tool_call_id": tc.get("id"), "content": content} for tc, content in retrieved 

969 ] 

970 follow_up_messages = list(messages) + [assistant_message] + tool_results 

971 

972 max_tokens: Final[int | None] = anthropic_messages_optional_request_params.get("max_tokens") or kwargs.get( 

973 "max_tokens" 

974 ) 

975 optional_params_without_max_tokens: Final = { 

976 k: v for k, v in anthropic_messages_optional_request_params.items() if k != "max_tokens" 

977 } 

978 

979 full_model_name = model 

980 if logging_obj is not None: 

981 agentic_params: Final = getattr(logging_obj, "model_call_details", {}).get("agentic_loop_params", {}) 

982 candidate: Final = agentic_params.get("model", model) 

983 if isinstance(candidate, str) and candidate: 

984 full_model_name = candidate 

985 

986 return AgenticLoopPlan( 

987 run_agentic_loop=True, 

988 request_patch=AgenticLoopRequestPatch( 

989 model=full_model_name, 

990 messages=follow_up_messages, 

991 max_tokens=max_tokens, 

992 optional_params=optional_params_without_max_tokens, 

993 kwargs={ 

994 k: v for k, v in kwargs.items() if not k.startswith("_headroom") and k != "litellm_logging_obj" 

995 }, 

996 ), 

997 metadata={"tool_type": "headroom_ccr"}, 

998 ) 

999 

1000 @staticmethod 

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

1002 from litellm.types.proxy.guardrails.guardrail_hooks.headroom import ( 

1003 HeadroomGuardrailConfigModel, 

1004 ) 

1005 

1006 return HeadroomGuardrailConfigModel