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
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1from __future__ import annotations
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
12import httpx
13from fastapi import HTTPException
14from httpx import Response as HttpxResponse
15from pydantic import TypeAdapter
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
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
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])
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)
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)
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.
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
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.
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
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
143def _is_retrieve_tool_name(name: str | None) -> bool:
144 """Match the retrieve tool whether called directly or via the MCP gateway.
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 )
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 )
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
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)
184def _raw_retrieve_call_ids(messages: object) -> frozenset[str]:
185 """Retrieve-tool call ids read from the request's own, untranslated messages.
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 )
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.
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.
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 )
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.
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.
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 )
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.
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 ]
284def _build_compress_failure_detail(status_code: int, body: str) -> dict[str, object]:
285 """Build error details for failed /v1/compress responses.
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}
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 )
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()
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 }
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
362def has_headroom_retrieve_tool(tools: object) -> bool:
363 return has_tool_with_name(tools, HEADROOM_RETRIEVE_TOOL_NAME)
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 ]
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.
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 }
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)
410def _is_anthropic_messages_response(response: object) -> bool:
411 return isinstance(get_attribute_or_key(response, "content", None), list)
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.
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]
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.
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
477class HeadroomGuardrail(CustomGuardrail):
478 records_own_guardrail_information: ClassVar[bool] = True
479 server_fulfilled_tool_names: ClassVar[frozenset[str]] = frozenset({HEADROOM_RETRIEVE_TOOL_NAME})
481 @classmethod
482 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
483 return [
484 GuardrailEventHooks.pre_call,
485 GuardrailEventHooks.post_call,
486 ]
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 )
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"
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
540 @staticmethod
541 def _resolve_timeout(timeout: float | None) -> httpx.Timeout:
542 """Budget for one call to the compression service, unset meaning the default.
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))
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)
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 }
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})
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
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
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 )
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 )
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 )
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 )
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 )
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 )
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))
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
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}]"
744 if raw_response.status_code == 404:
745 return f"[Headroom: hash={hash_value} not found or expired]"
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}]"
755 try:
756 body: Final[object] = raw_response.json()
757 except ValueError:
758 return raw_response.text
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
765 return str(body)
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
778 if self._should_bypass(request_data):
779 verbose_proxy_logger.debug("Headroom: %s header set; skipping compression", BYPASS_HEADER)
780 return inputs
782 if request_data.get("background"):
783 verbose_proxy_logger.debug("Headroom: background request; skipping compression")
784 return inputs
786 structured_messages: Final = inputs.get("structured_messages")
787 if not _is_object_list(structured_messages) or not structured_messages:
788 return inputs
790 messages: Final = [m for m in structured_messages if _is_str_object_dict(m)]
791 if not messages:
792 return inputs
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
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()
820 from litellm.proxy.common_utils.callback_utils import (
821 add_guardrail_to_applied_guardrails_header,
822 )
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
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 )
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)
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]
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)
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]
878 return {**inputs, "structured_messages": compressed, "tools": merged_tools} # pyright: ignore[reportReturnType]
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 }
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, {}
912 tool_calls: Final = _extract_headroom_tool_calls(response)
913 if not tool_calls:
914 return False, {}
916 return True, {"tool_calls": tool_calls}
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", [])
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()
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))
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
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 }
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
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 )
1000 @staticmethod
1001 def get_config_model() -> type[GuardrailConfigModel[object]] | None:
1002 from litellm.types.proxy.guardrails.guardrail_hooks.headroom import (
1003 HeadroomGuardrailConfigModel,
1004 )
1006 return HeadroomGuardrailConfigModel