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