From 051ad73f5c97981e5f3ec38a7e053f903222591e Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Mon, 13 Apr 2026 08:30:09 -0700 Subject: [PATCH] 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 --- litellm/__init__.py | 1 + litellm/compression/__init__.py | 3 + litellm/compression/compress.py | 212 ++++++++++++++ litellm/compression/content_detection.py | 45 +++ litellm/compression/message_stubbing.py | 77 +++++ litellm/compression/retrieval_tool.py | 35 +++ litellm/compression/scoring/__init__.py | 4 + litellm/compression/scoring/bm25.py | 106 +++++++ .../compression/scoring/embedding_scorer.py | 90 ++++++ litellm/types/compression.py | 14 + tests/test_compression.py | 267 ++++++++++++++++++ 11 files changed, 854 insertions(+) create mode 100644 litellm/compression/__init__.py create mode 100644 litellm/compression/compress.py create mode 100644 litellm/compression/content_detection.py create mode 100644 litellm/compression/message_stubbing.py create mode 100644 litellm/compression/retrieval_tool.py create mode 100644 litellm/compression/scoring/__init__.py create mode 100644 litellm/compression/scoring/bm25.py create mode 100644 litellm/compression/scoring/embedding_scorer.py create mode 100644 litellm/types/compression.py create mode 100644 tests/test_compression.py diff --git a/litellm/__init__.py b/litellm/__init__.py index 8087e3f5311..8b0da380fd0 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -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 ( diff --git a/litellm/compression/__init__.py b/litellm/compression/__init__.py new file mode 100644 index 00000000000..11c5eaf84ef --- /dev/null +++ b/litellm/compression/__init__.py @@ -0,0 +1,3 @@ +from litellm.compression.compress import compress + +__all__ = ["compress"] diff --git a/litellm/compression/compress.py b/litellm/compression/compress.py new file mode 100644 index 00000000000..b738744e73c --- /dev/null +++ b/litellm/compression/compress.py @@ -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, + ) diff --git a/litellm/compression/content_detection.py b/litellm/compression/content_detection.py new file mode 100644 index 00000000000..0655a42daf5 --- /dev/null +++ b/litellm/compression/content_detection.py @@ -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" diff --git a/litellm/compression/message_stubbing.py b/litellm/compression/message_stubbing.py new file mode 100644 index 00000000000..e7d54097a61 --- /dev/null +++ b/litellm/compression/message_stubbing.py @@ -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} diff --git a/litellm/compression/retrieval_tool.py b/litellm/compression/retrieval_tool.py new file mode 100644 index 00000000000..1ee24784a63 --- /dev/null +++ b/litellm/compression/retrieval_tool.py @@ -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"], + }, + }, + } diff --git a/litellm/compression/scoring/__init__.py b/litellm/compression/scoring/__init__.py new file mode 100644 index 00000000000..78bb434d17a --- /dev/null +++ b/litellm/compression/scoring/__init__.py @@ -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"] diff --git a/litellm/compression/scoring/bm25.py b/litellm/compression/scoring/bm25.py new file mode 100644 index 00000000000..6aacb1c0d97 --- /dev/null +++ b/litellm/compression/scoring/bm25.py @@ -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 diff --git a/litellm/compression/scoring/embedding_scorer.py b/litellm/compression/scoring/embedding_scorer.py new file mode 100644 index 00000000000..2f5f3e66f67 --- /dev/null +++ b/litellm/compression/scoring/embedding_scorer.py @@ -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 diff --git a/litellm/types/compression.py b/litellm/types/compression.py new file mode 100644 index 00000000000..01d5a6dd4d6 --- /dev/null +++ b/litellm/types/compression.py @@ -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] diff --git a/tests/test_compression.py b/tests/test_compression.py new file mode 100644 index 00000000000..2443bb4c3f5 --- /dev/null +++ b/tests/test_compression.py @@ -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