litellm/prompt_compression/headroom/main.py
Krrish Dholakia 65b8660c80 feat(prompt-compression): add headroom compression plugin
Microservice that wraps headroom-ai as a LiteLLM pre_call guardrail.
Compresses tool outputs and JSON arrays before they reach the LLM
provider using headroom's smart_crusher (JSON dedup + schema compression)
and kompress (prose/log compression). Returns compressed structured_messages
via the generic guardrail API, which required adding structured_messages
support to GenericGuardrailAPIResponse.
2026-06-25 19:49:26 -07:00

96 lines
2.9 KiB
Python

from __future__ import annotations
import logging
import os
from typing import Annotated, Any, Optional
from fastapi import FastAPI, Header, HTTPException, status
from headroom import compress
from pydantic import BaseModel, Field
logger = logging.getLogger(__name__)
app = FastAPI(title="Headroom Guardrail", version="0.1.0")
_API_KEY = os.environ.get("GUARDRAIL_API_KEY")
class GuardrailRequest(BaseModel):
input_type: str
texts: Optional[list[str]] = None
images: Optional[list[str]] = None
structured_messages: Optional[list[dict[str, Any]]] = None
tools: Optional[list[dict[str, Any]]] = None
tool_calls: Optional[list[dict[str, Any]]] = None
model: Optional[str] = None
litellm_call_id: Optional[str] = None
litellm_trace_id: Optional[str] = None
request_data: Optional[dict[str, Any]] = None
additional_provider_specific_params: Optional[dict[str, Any]] = Field(
default=None
)
class GuardrailResponse(BaseModel):
action: str
texts: Optional[list[str]] = None
images: Optional[list[str]] = None
structured_messages: Optional[list[dict[str, Any]]] = None
blocked_reason: Optional[str] = None
def _resolve_model(request: GuardrailRequest) -> str:
if request.model:
return request.model
return os.environ.get("HEADROOM_DEFAULT_MODEL", "gpt-4o-mini")
@app.get("/health")
def health() -> dict[str, str]:
return {"status": "ok"}
@app.post("/beta/litellm_basic_guardrail_api", response_model=GuardrailResponse)
async def guardrail(
request: GuardrailRequest,
x_api_key: Annotated[Optional[str], Header()] = None,
) -> GuardrailResponse:
if _API_KEY and x_api_key != _API_KEY:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Unauthorized")
if request.input_type != "request":
return GuardrailResponse(action="NONE")
messages = request.structured_messages
if not messages:
return GuardrailResponse(action="NONE")
model = _resolve_model(request)
try:
result = compress(messages, model=model, compress_user_messages=True, protect_recent=0)
except Exception:
logger.exception(
"headroom compress failed (call_id=%s); passing through unchanged",
request.litellm_call_id,
)
return GuardrailResponse(action="NONE")
compressed_messages: list[dict[str, Any]] = result.messages # type: ignore[attr-defined]
tokens_saved = getattr(result, "tokens_saved", 0)
logger.info(
"headroom compressed call_id=%s model=%s tokens_saved=%s compression_ratio=%s",
request.litellm_call_id,
model,
tokens_saved,
getattr(result, "compression_ratio", "?"),
)
if not tokens_saved:
return GuardrailResponse(action="NONE")
return GuardrailResponse(
action="GUARDRAIL_INTERVENED",
structured_messages=compressed_messages,
)