feat: add litellm.compress() for BM25-based context compression

Adds a compress() utility that reduces context size for LLM calls using
BM25 relevance scoring (with optional semantic embeddings via
litellm.embedding()). Messages below a token threshold pass through
unchanged; messages above are scored, ranked, and the lowest-relevance
ones replaced with stubs. Originals are cached and a retrieval tool is
injected so the model can recover dropped content on demand.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
Krrish Dholakia 2026-04-13 08:30:09 -07:00
parent 0eae9f101e
commit 051ad73f5c
11 changed files with 854 additions and 0 deletions

View file

@ -1176,6 +1176,7 @@ from litellm.types.utils import LlmProviders
## Lazy loading this is not straightforward, will leave it here for now.
from .main import * # type: ignore
from .compression import compress
# Skills API
from .skills.main import (

View file

@ -0,0 +1,3 @@
from litellm.compression.compress import compress
__all__ = ["compress"]

View file

@ -0,0 +1,212 @@
"""
Main compress() function — orchestrates BM25/embedding scoring, message stubbing,
and retrieval tool injection.
"""
from typing import Dict, List, Optional, Set
from litellm.caching.dual_cache import DualCache
from litellm.compression.message_stubbing import extract_key, stub_message
from litellm.compression.retrieval_tool import build_retrieval_tool
from litellm.compression.scoring.bm25 import bm25_score_messages
from litellm.litellm_core_utils.token_counter import token_counter
from litellm.types.compression import CompressedResult
def _extract_last_user_message(messages: List[dict]) -> str:
"""Return the text content of the last user message."""
for msg in reversed(messages):
if msg.get("role") == "user":
content = msg.get("content", "")
if isinstance(content, str):
return content
if isinstance(content, list):
parts = []
for part in content:
if isinstance(part, dict) and part.get("type") == "text":
parts.append(part.get("text", ""))
elif isinstance(part, str):
parts.append(part)
return " ".join(parts)
return ""
def _get_protected_indices(messages: List[dict]) -> List[int]:
"""
Return indices of messages that must never be compressed:
- All system messages
- The last user message
- The last assistant message
"""
protected: List[int] = []
last_user_idx = None
last_assistant_idx = None
for i, msg in enumerate(messages):
role = msg.get("role", "")
if role == "system":
protected.append(i)
elif role == "user":
last_user_idx = i
elif role == "assistant":
last_assistant_idx = i
if last_user_idx is not None:
protected.append(last_user_idx)
if last_assistant_idx is not None:
protected.append(last_assistant_idx)
return protected
def _combine_scores(
bm25_scores: List[float],
emb_scores: List[float],
bm25_weight: float = 0.4,
) -> List[float]:
"""Weighted average of BM25 and embedding scores, with min-max normalization."""
def _normalize(scores: List[float]) -> List[float]:
min_s = min(scores) if scores else 0.0
max_s = max(scores) if scores else 0.0
rng = max_s - min_s
if rng == 0:
return [0.0] * len(scores)
return [(s - min_s) / rng for s in scores]
norm_bm25 = _normalize(bm25_scores)
norm_emb = _normalize(emb_scores)
emb_weight = 1.0 - bm25_weight
return [bm25_weight * b + emb_weight * e for b, e in zip(norm_bm25, norm_emb)]
def compress(
messages: List[dict],
model: str,
compression_trigger: int = 200_000,
compression_target: Optional[int] = None,
embedding_model: Optional[str] = None,
compression_cache: Optional[DualCache] = None,
) -> CompressedResult:
"""
Compress a list of messages by replacing low-relevance content with stubs.
Messages below ``compression_trigger`` tokens pass through unchanged.
Messages above are scored with BM25 (and optionally embeddings), ranked,
and the lowest-relevance messages are replaced with stubs. Originals are
cached and a retrieval tool is injected so the model can recover dropped
content on demand.
Parameters:
messages: The conversation messages to (potentially) compress.
model: The LLM model name — used for token counting.
compression_trigger: Only compress if input exceeds this token count.
compression_target: Target token count after compression.
Defaults to ``compression_trigger // 2``.
embedding_model: If provided, use BM25 + embeddings for scoring.
If ``None``, BM25 only.
compression_cache: Passed through to ``litellm.embedding()`` for
cross-turn caching of embedding vectors.
Returns:
A ``CompressedResult`` dict containing compressed messages, token
counts, a cache of original content, and the retrieval tool definition.
"""
if compression_target is None:
compression_target = compression_trigger // 2
original_tokens = token_counter(model=model, messages=messages)
# Pass through if below trigger
if original_tokens <= compression_trigger:
return CompressedResult(
messages=messages,
original_tokens=original_tokens,
compressed_tokens=original_tokens,
compression_ratio=0.0,
cache={},
tools=[],
)
# Extract query for relevance scoring
query = _extract_last_user_message(messages)
# Score each message
bm25_scores = bm25_score_messages(query, messages)
if embedding_model:
from litellm.compression.scoring.embedding_scorer import (
embedding_score_messages,
)
emb_scores = embedding_score_messages(
query, messages, model=embedding_model, cache=compression_cache
)
combined_scores = _combine_scores(bm25_scores, emb_scores, bm25_weight=0.4)
else:
combined_scores = bm25_scores
# Sort message indices by score descending
ranked_indices = sorted(
range(len(messages)),
key=lambda i: combined_scores[i],
reverse=True,
)
# Protected messages are never compressed
protected_indices = _get_protected_indices(messages)
kept_indices: Set[int] = set(protected_indices)
# Count tokens for protected messages
current_tokens = 0
for i in kept_indices:
current_tokens += token_counter(
model=model, text=messages[i].get("content", "") or ""
)
# Fill token budget from highest-scoring messages
for idx in ranked_indices:
if idx in kept_indices:
continue
msg_content = messages[idx].get("content", "") or ""
msg_tokens = token_counter(model=model, text=msg_content)
if current_tokens + msg_tokens <= compression_target:
kept_indices.add(idx)
current_tokens += msg_tokens
# Build compressed messages and cache
compressed_messages: List[dict] = []
cache: Dict[str, str] = {}
used_keys: Set[str] = set()
for i, msg in enumerate(messages):
if i in kept_indices:
compressed_messages.append(msg)
else:
key = extract_key(msg, fallback_index=i, used_keys=used_keys)
content = msg.get("content", "")
if isinstance(content, list):
content = " ".join(
p.get("text", "") if isinstance(p, dict) else str(p)
for p in content
)
cache[key] = content
compressed_messages.append(stub_message(msg, key))
# Build retrieval tool
tools = [build_retrieval_tool(list(cache.keys()))] if cache else []
compressed_tokens = token_counter(model=model, messages=compressed_messages)
return CompressedResult(
messages=compressed_messages,
original_tokens=original_tokens,
compressed_tokens=compressed_tokens,
compression_ratio=round(1 - (compressed_tokens / original_tokens), 4)
if original_tokens > 0
else 0.0,
cache=cache,
tools=tools,
)

View file

@ -0,0 +1,45 @@
"""
Auto-detect content type per message: code, JSON, or text.
"""
import json
import re
_CODE_KEYWORDS = re.compile(
r"\b(?:def |function |class |import |from |require\(|#include|fn |func |const |let |var |public |private |static )\b"
)
def detect_content_type(content: str) -> str:
"""
Detect whether content is code, JSON, or plain text.
Returns one of: "code", "json", "text"
"""
stripped = content.strip()
if not stripped:
return "text"
# Check JSON
if stripped[0] in ("{", "["):
try:
json.loads(stripped)
return "json"
except (json.JSONDecodeError, ValueError):
pass
# Check code indicators
# Sample first 5000 chars for performance
sample = stripped[:5000]
keyword_matches = len(_CODE_KEYWORDS.findall(sample))
lines = sample.split("\n")
indented_lines = sum(
1 for line in lines if line.startswith((" ", "\t")) and line.strip()
)
# If we see multiple code keywords or significant indentation, it's likely code
if keyword_matches >= 3 or (indented_lines > len(lines) * 0.3 and len(lines) > 5):
return "code"
return "text"

View file

@ -0,0 +1,77 @@
"""
Replace messages with compact stubs and extract human-readable keys.
"""
import re
from typing import Set
from litellm.compression.content_detection import detect_content_type
# Patterns for extracting file paths from content
_FILE_PATH_PATTERNS = [
re.compile(r"^#\s*(\S+\.\w+)", re.MULTILINE), # # filename.py
re.compile(r"^//\s*(\S+\.\w+)", re.MULTILINE), # // filename.js
re.compile(r"^File:\s*(\S+)", re.MULTILINE), # File: path/to/file
re.compile(r"^---\s*(\S+\.\w+)", re.MULTILINE), # --- filename.ext
re.compile(r"`(\S+\.\w{1,5})`"), # `filename.ext` in backticks
]
def extract_key(message: dict, fallback_index: int, used_keys: Set[str]) -> str:
"""
Extract a human-readable key for the message.
Looks for file path patterns in the content. Falls back to message_{index}.
Handles duplicates by appending _2, _3, etc.
"""
content = message.get("content", "")
if isinstance(content, list):
content = " ".join(
p.get("text", "") if isinstance(p, dict) else str(p) for p in content
)
key = None
for pattern in _FILE_PATH_PATTERNS:
match = pattern.search(content[:2000]) # Only search the beginning
if match:
# Use just the filename, not full path
path = match.group(1)
key = path.split("/")[-1]
break
if key is None:
key = f"message_{fallback_index}"
# Handle duplicates
base_key = key
counter = 2
while key in used_keys:
key = f"{base_key}_{counter}"
counter += 1
used_keys.add(key)
return key
def stub_message(message: dict, key: str) -> dict:
"""
Replace message content with a compact stub.
Returns a new message dict with the same role but content replaced
with a short description referencing the retrieval tool.
"""
content = message.get("content", "")
if isinstance(content, list):
content = " ".join(
p.get("text", "") if isinstance(p, dict) else str(p) for p in content
)
line_count = content.count("\n") + 1
content_type = detect_content_type(content)
stub_content = (
f"[Compressed: {key} — {line_count} lines, {content_type}. "
f"Use litellm_content_retrieve tool to get full content.]"
)
return {**message, "content": stub_content}

View file

@ -0,0 +1,35 @@
"""
Build the litellm_content_retrieve tool definition for the LLM.
"""
from typing import List
def build_retrieval_tool(available_keys: List[str]) -> dict:
"""
Return an OpenAI-format tool definition that lets the model
retrieve the full content of a compressed message.
"""
return {
"type": "function",
"function": {
"name": "litellm_content_retrieve",
"description": (
"Retrieve the full content of a file or message that was "
"compressed to save tokens. Use this when you need the complete "
"content to answer accurately. Available keys: "
+ ", ".join(available_keys)
),
"parameters": {
"type": "object",
"properties": {
"key": {
"type": "string",
"description": "The identifier of the content to retrieve",
"enum": available_keys,
}
},
"required": ["key"],
},
},
}

View file

@ -0,0 +1,4 @@
from litellm.compression.scoring.bm25 import bm25_score_messages
from litellm.compression.scoring.embedding_scorer import embedding_score_messages
__all__ = ["bm25_score_messages", "embedding_score_messages"]

View file

@ -0,0 +1,106 @@
"""
Pure Python BM25 (Okapi BM25) relevance scorer.
No external dependencies — uses only stdlib.
"""
import math
import re
from collections import Counter
from typing import Dict, List
def _tokenize(text: str) -> List[str]:
"""Split text into lowercase tokens on word boundaries."""
return re.findall(r"[a-z0-9_]+", text.lower())
def _extract_content(message: dict) -> str:
"""Extract text content from a message dict."""
content = message.get("content", "")
if isinstance(content, str):
return content
if isinstance(content, list):
parts = []
for part in content:
if isinstance(part, dict) and part.get("type") == "text":
parts.append(part.get("text", ""))
elif isinstance(part, str):
parts.append(part)
return " ".join(parts)
return ""
def bm25_score_messages(
query: str,
messages: List[dict],
k1: float = 1.5,
b: float = 0.75,
) -> List[float]:
"""
Score each message's relevance to the query using BM25 (Okapi BM25).
Parameters:
query: The reference text to score against (typically the last user message).
messages: List of message dicts with "content" fields.
k1: Term frequency saturation parameter.
b: Length normalization parameter.
Returns:
List of float scores, one per message. Higher = more relevant.
"""
query_terms = _tokenize(query)
if not query_terms:
return [0.0] * len(messages)
# Tokenize all documents
doc_tokens: List[List[str]] = []
for msg in messages:
doc_tokens.append(_tokenize(_extract_content(msg)))
n = len(doc_tokens)
if n == 0:
return []
# Average document length
doc_lengths = [len(dt) for dt in doc_tokens]
avgdl = sum(doc_lengths) / n if n > 0 else 1.0
# Document frequency for each term
df: Dict[str, int] = {}
for dt in doc_tokens:
seen = set(dt)
for term in seen:
df[term] = df.get(term, 0) + 1
# IDF for query terms
idf: Dict[str, float] = {}
for term in set(query_terms):
term_df = df.get(term, 0)
# Standard BM25 IDF: log((N - df + 0.5) / (df + 0.5) + 1)
idf[term] = math.log((n - term_df + 0.5) / (term_df + 0.5) + 1.0)
# Score each document
scores: List[float] = []
for i, dt in enumerate(doc_tokens):
if not dt:
scores.append(0.0)
continue
tf_counts = Counter(dt)
dl = doc_lengths[i]
score = 0.0
for term in query_terms:
if term not in idf:
continue
tf = tf_counts.get(term, 0)
if tf == 0:
continue
numerator = tf * (k1 + 1)
denominator = tf + k1 * (1 - b + b * dl / avgdl)
score += idf[term] * numerator / denominator
scores.append(score)
return scores

View file

@ -0,0 +1,90 @@
"""
Semantic scoring via litellm.embedding().
Computes cosine similarity between the query embedding and each message embedding.
"""
import math
from typing import List, Optional
from litellm.caching.dual_cache import DualCache
def _extract_content(message: dict) -> str:
"""Extract text content from a message dict."""
content = message.get("content", "")
if isinstance(content, str):
return content
if isinstance(content, list):
parts = []
for part in content:
if isinstance(part, dict) and part.get("type") == "text":
parts.append(part.get("text", ""))
elif isinstance(part, str):
parts.append(part)
return " ".join(parts)
return ""
def _truncate_text(text: str, max_chars: int = 30000) -> str:
"""Truncate long text, keeping first and last portions."""
if len(text) <= max_chars:
return text
half = max_chars // 2
return text[:half] + "\n...\n" + text[-half:]
def _cosine_similarity(a: List[float], b: List[float]) -> float:
"""Compute cosine similarity between two vectors."""
dot = sum(x * y for x, y in zip(a, b))
norm_a = math.sqrt(sum(x * x for x in a))
norm_b = math.sqrt(sum(x * x for x in b))
if norm_a == 0 or norm_b == 0:
return 0.0
return dot / (norm_a * norm_b)
def embedding_score_messages(
query: str,
messages: List[dict],
model: str,
cache: Optional[DualCache] = None,
) -> List[float]:
"""
Score each message's semantic similarity to the query using embeddings.
Parameters:
query: The reference text to score against.
messages: List of message dicts with "content" fields.
model: The embedding model to use (e.g., "text-embedding-3-small").
cache: Optional DualCache for cross-turn embedding caching.
Returns:
List of float scores (cosine similarity), one per message.
"""
import litellm
texts = [_truncate_text(query)]
for msg in messages:
texts.append(_truncate_text(_extract_content(msg)))
# Filter out empty texts — replace with a placeholder to maintain indexing
processed_texts = [t if t.strip() else "empty" for t in texts]
kwargs = {
"model": model,
"input": processed_texts,
"caching": cache is not None,
}
response = litellm.embedding(**kwargs)
# Extract embedding vectors
embeddings = [item["embedding"] for item in response.data]
query_embedding = embeddings[0]
scores: List[float] = []
for i in range(1, len(embeddings)):
scores.append(_cosine_similarity(query_embedding, embeddings[i]))
return scores

View file

@ -0,0 +1,14 @@
"""
Type definitions for litellm.compress().
"""
from typing import Dict, List, TypedDict
class CompressedResult(TypedDict):
messages: List[dict] # compressed messages (stubs replace low-relevance messages)
original_tokens: int # token count before compression
compressed_tokens: int # token count after compression
compression_ratio: float # fraction reduced, e.g. 0.6 means 60% reduction
cache: Dict[str, str] # key -> original content (for retrieval tool responses)
tools: List[dict] # [litellm_content_retrieve tool definition]

267
tests/test_compression.py Normal file
View file

@ -0,0 +1,267 @@
"""
Unit tests for litellm.compress().
"""
import os
import pytest
import litellm
from litellm.compression.scoring.bm25 import bm25_score_messages
from litellm.compression.content_detection import detect_content_type
from litellm.compression.message_stubbing import extract_key, stub_message
from litellm.compression.retrieval_tool import build_retrieval_tool
# ---------------------------------------------------------------------------
# BM25 scorer
# ---------------------------------------------------------------------------
def test_bm25_relevance_ranking():
query = "Fix the authentication bug in the login handler"
messages = [
{
"role": "user",
"content": "def login_handler(): authentication check bug fix",
},
{"role": "user", "content": "def render_template(name): css styling layout"},
{"role": "user", "content": "def verify(): authentication token bug handler"},
]
scores = bm25_score_messages(query, messages)
# Messages sharing query terms should score higher than unrelated ones
assert scores[0] > scores[1]
assert scores[2] > scores[1]
def test_bm25_empty_query():
scores = bm25_score_messages("", [{"role": "user", "content": "hello"}])
assert scores == [0.0]
def test_bm25_empty_messages():
scores = bm25_score_messages("query", [])
assert scores == []
def test_bm25_empty_content():
scores = bm25_score_messages("query", [{"role": "user", "content": ""}])
assert scores == [0.0]
# ---------------------------------------------------------------------------
# Content detection
# ---------------------------------------------------------------------------
def test_detect_code():
code = """
import os
from pathlib import Path
def main():
class Foo:
pass
return Foo()
"""
assert detect_content_type(code) == "code"
def test_detect_json():
assert detect_content_type('{"key": "value", "num": 42}') == "json"
assert detect_content_type("[1, 2, 3]") == "json"
def test_detect_text():
assert detect_content_type("This is a plain text paragraph about dogs.") == "text"
def test_detect_empty():
assert detect_content_type("") == "text"
# ---------------------------------------------------------------------------
# Message stubbing
# ---------------------------------------------------------------------------
def test_extract_key_with_filename():
msg = {"role": "user", "content": "# auth.py\ndef authenticate():\n pass"}
used: set = set()
key = extract_key(msg, fallback_index=0, used_keys=used)
assert key == "auth.py"
def test_extract_key_fallback():
msg = {"role": "user", "content": "Some random content without a filename"}
used: set = set()
key = extract_key(msg, fallback_index=5, used_keys=used)
assert key == "message_5"
def test_extract_key_duplicates():
used: set = set()
msg = {"role": "user", "content": "# auth.py\ncode here"}
k1 = extract_key(msg, fallback_index=0, used_keys=used)
k2 = extract_key(msg, fallback_index=1, used_keys=used)
assert k1 == "auth.py"
assert k2 == "auth.py_2"
def test_stub_message():
msg = {"role": "user", "content": "line1\nline2\nline3"}
stubbed = stub_message(msg, "test_key")
assert stubbed["role"] == "user"
assert "test_key" in stubbed["content"]
assert "litellm_content_retrieve" in stubbed["content"]
assert "3 lines" in stubbed["content"]
# ---------------------------------------------------------------------------
# Retrieval tool
# ---------------------------------------------------------------------------
def test_retrieval_tool_schema():
tool = build_retrieval_tool(["auth.py", "utils.py"])
assert tool["type"] == "function"
assert tool["function"]["name"] == "litellm_content_retrieve"
assert "key" in tool["function"]["parameters"]["properties"]
assert tool["function"]["parameters"]["properties"]["key"]["enum"] == [
"auth.py",
"utils.py",
]
assert tool["function"]["parameters"]["required"] == ["key"]
def test_retrieval_tool_description_lists_keys():
tool = build_retrieval_tool(["foo.py", "bar.js"])
desc = tool["function"]["description"]
assert "foo.py" in desc
assert "bar.js" in desc
# ---------------------------------------------------------------------------
# compress() — end-to-end
# ---------------------------------------------------------------------------
def test_compress_below_trigger_passthrough():
messages = [{"role": "user", "content": "hello"}]
result = litellm.compress(messages, model="gpt-4o")
assert result["messages"] == messages
assert result["cache"] == {}
assert result["tools"] == []
assert result["compression_ratio"] == 0.0
assert result["original_tokens"] == result["compressed_tokens"]
def test_compress_above_trigger():
big_messages = [
{"role": "system", "content": "You are a coding assistant."},
{
"role": "user",
"content": "# auth.py\n" + "def authenticate():\n pass\n" * 2000,
},
{
"role": "user",
"content": "# utils.py\n" + "def helper():\n pass\n" * 2000,
},
{
"role": "user",
"content": "# readme.md\n" + "This is documentation. " * 2000,
},
{"role": "user", "content": "Fix the bug in auth.py"},
]
result = litellm.compress(
big_messages,
model="gpt-4o",
compression_trigger=1000,
compression_target=500,
)
assert result["compressed_tokens"] < result["original_tokens"]
assert result["compression_ratio"] > 0
assert len(result["cache"]) > 0
assert len(result["tools"]) == 1
assert result["tools"][0]["function"]["name"] == "litellm_content_retrieve"
def test_compress_preserves_system_message():
messages = [
{"role": "system", "content": "System prompt. " * 500},
{"role": "user", "content": "Large file content. " * 5000},
{"role": "user", "content": "Fix the bug"},
]
result = litellm.compress(messages, model="gpt-4o", compression_trigger=1000)
assert result["messages"][0]["role"] == "system"
assert "System prompt" in result["messages"][0]["content"]
def test_compress_preserves_last_user_message():
messages = [
{"role": "user", "content": "Big context " * 5000},
{"role": "user", "content": "Fix the bug in auth.py"},
]
result = litellm.compress(messages, model="gpt-4o", compression_trigger=1000)
last_user = [m for m in result["messages"] if m["role"] == "user"][-1]
assert "Fix the bug in auth.py" in last_user["content"]
def test_compress_preserves_last_assistant_message():
messages = [
{"role": "user", "content": "Big context " * 5000},
{"role": "assistant", "content": "I'll help with that. " * 2000},
{"role": "user", "content": "Now fix the bug"},
]
result = litellm.compress(messages, model="gpt-4o", compression_trigger=1000)
assistant_msgs = [m for m in result["messages"] if m["role"] == "assistant"]
assert len(assistant_msgs) >= 1
# The last assistant message should be preserved (not stubbed)
last_assistant = assistant_msgs[-1]
assert "I'll help with that" in last_assistant["content"]
def test_cache_keys_match_stubs():
messages = [
{"role": "user", "content": "# auth.py\n" + "code " * 5000},
{"role": "user", "content": "Fix it"},
]
result = litellm.compress(messages, model="gpt-4o", compression_trigger=1000)
if result["tools"]:
tool_desc = result["tools"][0]["function"]["description"]
for key in result["cache"]:
assert key in tool_desc
def test_compress_default_target():
"""compression_target defaults to compression_trigger // 2."""
messages = [
{"role": "user", "content": "content " * 5000},
{"role": "user", "content": "query"},
]
result = litellm.compress(messages, model="gpt-4o", compression_trigger=2000)
# Should have compressed — target = 1000
assert result["compressed_tokens"] <= result["original_tokens"]
# ---------------------------------------------------------------------------
# Embedding scorer — integration test (skipped without API key)
# ---------------------------------------------------------------------------
@pytest.mark.skipif(not os.environ.get("OPENAI_API_KEY"), reason="Needs OPENAI_API_KEY")
def test_embedding_scorer():
result = litellm.compress(
messages=[
{"role": "user", "content": "Authentication code " * 2000},
{"role": "user", "content": "Unrelated cooking recipes " * 2000},
{"role": "user", "content": "Fix auth"},
],
model="gpt-4o",
compression_trigger=1000,
embedding_model="text-embedding-3-small",
)
assert result["compression_ratio"] > 0
assert len(result["cache"]) > 0