File size: 6,538 Bytes
7a05808
 
 
e4a41fa
 
 
 
 
7a05808
 
e4a41fa
7a05808
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Tests for HeadroomClient cache optimizer integration."""

import os
import tempfile
from unittest.mock import MagicMock

import pytest

from headroom import (
    AnthropicCacheOptimizer,
    HeadroomClient,
)


@pytest.fixture
def temp_db():
    """Create a temporary database file."""
    fd, path = tempfile.mkstemp(suffix=".db")
    os.close(fd)
    yield f"sqlite:///{path}"
    if os.path.exists(path):
        os.unlink(path)


class MockTokenCounter:
    """Mock token counter for testing."""

    def count_tokens(self, text: str) -> int:
        return len(text) // 4

    def count_messages(self, messages: list) -> int:
        total = 0
        for msg in messages:
            content = msg.get("content", "")
            if isinstance(content, str):
                total += len(content) // 4
            elif isinstance(content, list):
                for block in content:
                    if isinstance(block, dict):
                        total += len(block.get("text", "")) // 4
        return total


class MockAnthropicProvider:
    """Mock Anthropic provider for testing."""

    name = "anthropic"

    def get_token_counter(self, model: str):
        return MockTokenCounter()

    def get_context_limit(self, model: str) -> int:
        return 200000


class MockOpenAIProvider:
    """Mock OpenAI provider for testing."""

    name = "openai"

    def get_token_counter(self, model: str):
        return MockTokenCounter()

    def get_context_limit(self, model: str) -> int:
        return 128000


class TestHeadroomClientCacheIntegration:
    """Test HeadroomClient cache optimizer integration."""

    def test_auto_detect_anthropic_optimizer(self, temp_db):
        """Test that Anthropic optimizer is auto-detected."""
        mock_client = MagicMock()
        provider = MockAnthropicProvider()

        client = HeadroomClient(
            original_client=mock_client,
            provider=provider,
            store_url=temp_db,
            enable_cache_optimizer=True,
        )

        assert client._cache_optimizer is not None
        assert client._cache_optimizer.name == "anthropic-cache-optimizer"

    def test_auto_detect_openai_optimizer(self, temp_db):
        """Test that OpenAI optimizer is auto-detected."""
        mock_client = MagicMock()
        provider = MockOpenAIProvider()

        client = HeadroomClient(
            original_client=mock_client,
            provider=provider,
            store_url=temp_db,
            enable_cache_optimizer=True,
        )

        assert client._cache_optimizer is not None
        assert client._cache_optimizer.name == "openai-prefix-stabilizer"

    def test_custom_optimizer(self, temp_db):
        """Test using a custom optimizer."""
        mock_client = MagicMock()
        provider = MockAnthropicProvider()
        custom_optimizer = AnthropicCacheOptimizer()

        client = HeadroomClient(
            original_client=mock_client,
            provider=provider,
            store_url=temp_db,
            cache_optimizer=custom_optimizer,
        )

        assert client._cache_optimizer is custom_optimizer

    def test_disable_cache_optimizer(self, temp_db):
        """Test disabling cache optimizer."""
        mock_client = MagicMock()
        provider = MockAnthropicProvider()

        client = HeadroomClient(
            original_client=mock_client,
            provider=provider,
            store_url=temp_db,
            enable_cache_optimizer=False,
        )

        assert client._cache_optimizer is None

    def test_semantic_cache_layer_creation(self, temp_db):
        """Test semantic cache layer is created when enabled."""
        mock_client = MagicMock()
        provider = MockAnthropicProvider()

        client = HeadroomClient(
            original_client=mock_client,
            provider=provider,
            store_url=temp_db,
            enable_cache_optimizer=True,
            enable_semantic_cache=True,
        )

        assert client._semantic_cache_layer is not None
        assert client._cache_optimizer is not None

    def test_extract_query_from_string_content(self, temp_db):
        """Test query extraction from string content."""
        mock_client = MagicMock()
        provider = MockAnthropicProvider()

        client = HeadroomClient(
            original_client=mock_client,
            provider=provider,
            store_url=temp_db,
        )

        messages = [
            {"role": "system", "content": "You are helpful."},
            {"role": "user", "content": "What is 2+2?"},
        ]

        query = client._extract_query(messages)
        assert query == "What is 2+2?"

    def test_extract_query_from_content_blocks(self, temp_db):
        """Test query extraction from content block format."""
        mock_client = MagicMock()
        provider = MockAnthropicProvider()

        client = HeadroomClient(
            original_client=mock_client,
            provider=provider,
            store_url=temp_db,
        )

        messages = [
            {"role": "system", "content": "You are helpful."},
            {
                "role": "user",
                "content": [{"type": "text", "text": "What is 2+2?"}],
            },
        ]

        query = client._extract_query(messages)
        assert query == "What is 2+2?"

    def test_extract_query_last_user_message(self, temp_db):
        """Test that query extraction uses last user message."""
        mock_client = MagicMock()
        provider = MockAnthropicProvider()

        client = HeadroomClient(
            original_client=mock_client,
            provider=provider,
            store_url=temp_db,
        )

        messages = [
            {"role": "user", "content": "First question"},
            {"role": "assistant", "content": "First answer"},
            {"role": "user", "content": "Second question"},
        ]

        query = client._extract_query(messages)
        assert query == "Second question"

    def test_config_propagation(self, temp_db):
        """Test that config is properly propagated."""
        mock_client = MagicMock()
        provider = MockAnthropicProvider()

        client = HeadroomClient(
            original_client=mock_client,
            provider=provider,
            store_url=temp_db,
            enable_cache_optimizer=True,
            enable_semantic_cache=True,
        )

        assert client._config.cache_optimizer.enabled is True
        assert client._config.cache_optimizer.enable_semantic_cache is True