Spaces:
Build error
Build error
File size: 8,738 Bytes
6505c43 2a0de14 6505c43 2a0de14 6505c43 2a0de14 6505c43 2a0de14 6505c43 2a0de14 6505c43 5f8e891 6505c43 5f8e891 2a0de14 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 | """Tests for the one-function compress() API and integrations."""
import json
import pytest
from headroom.compress import CompressResult, compress
from headroom.hooks import CompressionHooks
try:
from starlette.applications import Starlette
from starlette.requests import Request
from starlette.responses import JSONResponse
from starlette.routing import Route
from starlette.testclient import TestClient
from headroom.integrations.asgi import CompressionMiddleware
HAS_STARLETTE = True
except ImportError:
HAS_STARLETTE = False
# =============================================================================
# Tests: compress() function
# =============================================================================
class TestCompressFunction:
def test_empty_messages(self):
result = compress([], model="test")
assert result.messages == []
assert result.tokens_saved == 0
def test_small_messages_passthrough(self):
"""Small messages below compression threshold pass through unchanged."""
messages = [{"role": "user", "content": "hello"}]
result = compress(messages, model="gpt-4o")
assert result.messages[0]["content"] == "hello"
assert result.tokens_saved == 0
def test_returns_compress_result(self):
result = compress([{"role": "user", "content": "hi"}])
assert isinstance(result, CompressResult)
assert hasattr(result, "messages")
assert hasattr(result, "tokens_saved")
assert hasattr(result, "compression_ratio")
assert hasattr(result, "transforms_applied")
def test_large_tool_output_compressed(self):
"""Large JSON tool output should be compressed."""
big_data = json.dumps(
[
{"id": i, "status": "active", "name": f"item_{i}", "value": i * 17}
for i in range(200)
]
)
messages = [
{"role": "user", "content": "What are the top items?"},
{"role": "tool", "content": big_data, "tool_call_id": "call_1"},
]
result = compress(messages, model="gpt-4o")
assert result.tokens_after <= result.tokens_before
assert len(result.messages) == 2
def test_optimize_false_passthrough(self):
"""optimize=False returns messages unchanged."""
messages = [{"role": "user", "content": "hello world " * 100}]
result = compress(messages, optimize=False)
assert result.messages is messages
assert result.tokens_saved == 0
def test_with_custom_hooks(self):
"""Hooks are called when provided."""
calls = []
class TrackingHooks(CompressionHooks):
def pre_compress(self, messages, ctx):
calls.append(("pre", len(messages)))
return messages
def compute_biases(self, messages, ctx):
calls.append(("biases", len(messages)))
return {}
def post_compress(self, event):
calls.append(("post", event.tokens_saved))
big_data = json.dumps([{"id": i, "status": "active"} for i in range(100)])
messages = [
{"role": "user", "content": "analyze"},
{"role": "tool", "content": big_data, "tool_call_id": "c1"},
]
compress(messages, hooks=TrackingHooks())
assert any(c[0] == "pre" for c in calls)
assert any(c[0] == "biases" for c in calls)
class TestCompressResultFields:
def test_fields_populated(self):
big_data = json.dumps([{"id": i, "type": "log"} for i in range(100)])
messages = [
{"role": "user", "content": "summarize"},
{"role": "tool", "content": big_data, "tool_call_id": "c1"},
]
result = compress(messages, model="claude-sonnet-4-5-20250929")
assert result.tokens_before > 0
assert result.tokens_after >= 0
assert result.tokens_saved >= 0
assert 0.0 <= result.compression_ratio <= 1.0
# =============================================================================
# Tests: ASGI CompressionMiddleware (requires starlette)
# =============================================================================
def _make_asgi_app(middleware_kwargs=None):
"""Create a test ASGI app with CompressionMiddleware."""
async def chat_endpoint(request: Request) -> JSONResponse:
body = await request.json()
return JSONResponse(
{
"model": "gpt-4o",
"choices": [{"message": {"content": "response"}}],
"usage": {"prompt_tokens": 10, "completion_tokens": 5},
"_message_count": len(body.get("messages", [])),
}
)
async def health(request: Request) -> JSONResponse:
return JSONResponse({"status": "ok"})
app = Starlette(
routes=[
Route("/health", health),
Route("/v1/chat/completions", chat_endpoint, methods=["POST"]),
Route("/v1/messages", chat_endpoint, methods=["POST"]),
]
)
app.add_middleware(CompressionMiddleware, **(middleware_kwargs or {}))
return app
@pytest.mark.skipif(not HAS_STARLETTE, reason="starlette not installed")
class TestASGIMiddleware:
def test_non_llm_paths_passthrough(self):
app = _make_asgi_app()
client = TestClient(app)
resp = client.get("/health")
assert resp.status_code == 200
assert resp.json()["status"] == "ok"
def test_small_messages_passthrough(self):
app = _make_asgi_app()
client = TestClient(app)
resp = client.post(
"/v1/chat/completions",
json={"model": "gpt-4o", "messages": [{"role": "user", "content": "hi"}]},
)
assert resp.status_code == 200
def test_large_messages_compressed(self):
"""Large tool output should be compressed by middleware."""
app = _make_asgi_app()
client = TestClient(app)
big_data = json.dumps([{"id": i, "status": "active"} for i in range(200)])
resp = client.post(
"/v1/chat/completions",
json={
"model": "gpt-4o",
"messages": [
{"role": "user", "content": "analyze"},
{"role": "tool", "content": big_data, "tool_call_id": "c1"},
],
},
)
assert resp.status_code == 200
def test_anthropic_path(self):
"""Works with Anthropic /v1/messages path."""
app = _make_asgi_app()
client = TestClient(app)
resp = client.post(
"/v1/messages",
json={
"model": "claude-sonnet-4-5-20250929",
"messages": [{"role": "user", "content": "hello"}],
},
)
assert resp.status_code == 200
def test_get_requests_passthrough(self):
"""GET requests to LLM paths pass through."""
app = _make_asgi_app()
client = TestClient(app)
resp = client.get("/v1/chat/completions")
assert resp.status_code in (200, 405)
# =============================================================================
# Tests: LiteLLM Callback
# =============================================================================
class TestLiteLLMCallback:
def test_callback_imports(self):
"""Verify the callback can be imported."""
from headroom.integrations.litellm_callback import HeadroomCallback
callback = HeadroomCallback()
assert callback.total_tokens_saved == 0
def test_callback_compresses_messages(self):
"""Callback compresses messages in pre_call_hook."""
import asyncio
from headroom.integrations.litellm_callback import HeadroomCallback
callback = HeadroomCallback()
big_data = json.dumps([{"id": i, "status": "active"} for i in range(200)])
data = {
"model": "gpt-4o",
"messages": [
{"role": "user", "content": "analyze"},
{"role": "tool", "content": big_data, "tool_call_id": "c1"},
],
}
result = asyncio.run(callback.async_pre_call_hook("key", data, "completion"))
assert result is data
def test_callback_ignores_non_completion(self):
"""Non-completion calls are passed through."""
import asyncio
from headroom.integrations.litellm_callback import HeadroomCallback
callback = HeadroomCallback()
data = {"messages": [{"role": "user", "content": "hi"}]}
result = asyncio.run(callback.async_pre_call_hook("key", data, "embedding"))
assert result is data
|