FreshPixels commited on
Commit
d566dad
·
verified ·
1 Parent(s): 7c6d2b9

Create tests/unit/test_llm_manager.py

Browse files
Files changed (1) hide show
  1. tests/unit/test_llm_manager.py +244 -0
tests/unit/test_llm_manager.py ADDED
@@ -0,0 +1,244 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import asyncio
4
+ import pytest
5
+ import pytest_asyncio
6
+
7
+ from llm.manager import (
8
+ AllProvidersFailedError,
9
+ FallbackPolicy,
10
+ LLMManager,
11
+ ProviderAlreadyRegisteredError,
12
+ ProviderNotFoundError,
13
+ )
14
+ from llm.providers.base_provider import (
15
+ BaseProvider,
16
+ LLMMessage,
17
+ LLMProviderError,
18
+ LLMRequest,
19
+ LLMResponse,
20
+ )
21
+
22
+
23
+ # ---------------------------------------------------------------------------
24
+ # Stub providers for testing
25
+ # ---------------------------------------------------------------------------
26
+
27
+ class _SuccessProvider(BaseProvider):
28
+ """A provider that always succeeds."""
29
+
30
+ def __init__(self, name: str = "success") -> None:
31
+ super().__init__(name=name)
32
+
33
+ async def generate(self, request: LLMRequest) -> LLMResponse:
34
+ return LLMResponse(
35
+ content=f"response from {self.name}",
36
+ model=request.model,
37
+ provider=self.name,
38
+ )
39
+
40
+
41
+ class _FailProvider(BaseProvider):
42
+ """A provider that always fails with LLMProviderError."""
43
+
44
+ def __init__(self, name: str = "fail") -> None:
45
+ super().__init__(name=name)
46
+ self.call_count: int = 0
47
+
48
+ async def generate(self, request: LLMRequest) -> LLMResponse:
49
+ self.call_count += 1
50
+ raise LLMProviderError(
51
+ f"provider {self.name} failed (attempt {self.call_count})",
52
+ provider_name=self.name,
53
+ )
54
+
55
+
56
+ class _FlakyProvider(BaseProvider):
57
+ """A provider that fails N times then succeeds."""
58
+
59
+ def __init__(self, fail_count: int = 1, name: str = "flaky") -> None:
60
+ super().__init__(name=name)
61
+ self._fail_count = fail_count
62
+ self.call_count: int = 0
63
+
64
+ async def generate(self, request: LLMRequest) -> LLMResponse:
65
+ self.call_count += 1
66
+ if self.call_count <= self._fail_count:
67
+ raise LLMProviderError(
68
+ f"transient failure (attempt {self.call_count})",
69
+ provider_name=self.name,
70
+ )
71
+ return LLMResponse(
72
+ content=f"success after {self.call_count} attempts",
73
+ model=request.model,
74
+ provider=self.name,
75
+ )
76
+
77
+
78
+ # ---------------------------------------------------------------------------
79
+ # Helper
80
+ # ---------------------------------------------------------------------------
81
+
82
+ def _make_request(model: str = "test-model") -> LLMRequest:
83
+ return LLMRequest(
84
+ messages=[LLMMessage(role="user", content="Hello")],
85
+ model=model,
86
+ )
87
+
88
+
89
+ # ---------------------------------------------------------------------------
90
+ # Tests
91
+ # ---------------------------------------------------------------------------
92
+
93
+ class TestLLMManagerRegistration:
94
+ @pytest.mark.asyncio
95
+ async def test_register_and_list(self) -> None:
96
+ manager = LLMManager(default_provider="p1")
97
+ await manager.register_provider("p1", _SuccessProvider("p1"))
98
+ assert manager.list_providers() == ["p1"]
99
+
100
+ @pytest.mark.asyncio
101
+ async def test_duplicate_registration_raises(self) -> None:
102
+ manager = LLMManager(default_provider="p1")
103
+ await manager.register_provider("p1", _SuccessProvider("p1"))
104
+ with pytest.raises(ProviderAlreadyRegisteredError):
105
+ await manager.register_provider("p1", _SuccessProvider("p1"))
106
+
107
+ @pytest.mark.asyncio
108
+ async def test_duplicate_with_overwrite(self) -> None:
109
+ manager = LLMManager(default_provider="p1")
110
+ await manager.register_provider("p1", _SuccessProvider("p1"))
111
+ await manager.register_provider("p1", _SuccessProvider("p1"), overwrite=True)
112
+ assert manager.list_providers() == ["p1"]
113
+
114
+ @pytest.mark.asyncio
115
+ async def test_has_providers(self) -> None:
116
+ manager = LLMManager()
117
+ assert not manager.has_providers()
118
+ await manager.register_provider("p1", _SuccessProvider("p1"))
119
+ assert manager.has_providers()
120
+
121
+
122
+ class TestLLMManagerGetProvider:
123
+ def test_get_default(self) -> None:
124
+ manager = LLMManager(default_provider="p1")
125
+ provider = _SuccessProvider("p1")
126
+ manager._providers["p1"] = provider
127
+ assert manager.get_provider() is provider
128
+
129
+ def test_get_by_name(self) -> None:
130
+ manager = LLMManager(default_provider="p1")
131
+ provider = _SuccessProvider("p2")
132
+ manager._providers["p2"] = provider
133
+ assert manager.get_provider("p2") is provider
134
+
135
+ def test_not_found_raises(self) -> None:
136
+ manager = LLMManager(default_provider="missing")
137
+ with pytest.raises(ProviderNotFoundError):
138
+ manager.get_provider()
139
+
140
+
141
+ class TestLLMManagerGenerate:
142
+ @pytest.mark.asyncio
143
+ async def test_successful_generate(self) -> None:
144
+ manager = LLMManager(default_provider="p1")
145
+ await manager.register_provider("p1", _SuccessProvider("p1"))
146
+ response = await manager.generate(_make_request())
147
+ assert response.provider == "p1"
148
+ assert "response" in response.content
149
+
150
+ @pytest.mark.asyncio
151
+ async def test_no_providers_raises(self) -> None:
152
+ manager = LLMManager(default_provider="missing")
153
+ with pytest.raises(ProviderNotFoundError):
154
+ await manager.generate(_make_request())
155
+
156
+ @pytest.mark.asyncio
157
+ async def test_fallback_to_second_provider(self) -> None:
158
+ policy = FallbackPolicy(
159
+ providers=["fail", "success"],
160
+ max_retries_per_provider=0,
161
+ )
162
+ manager = LLMManager(default_provider="fail", fallback_policy=policy)
163
+ fail_provider = _FailProvider("fail")
164
+ await manager.register_provider("fail", fail_provider)
165
+ await manager.register_provider("success", _SuccessProvider("success"))
166
+
167
+ response = await manager.generate(_make_request())
168
+ assert response.provider == "success"
169
+
170
+ @pytest.mark.asyncio
171
+ async def test_all_providers_fail_raises(self) -> None:
172
+ policy = FallbackPolicy(
173
+ providers=["fail1", "fail2"],
174
+ max_retries_per_provider=0,
175
+ )
176
+ manager = LLMManager(default_provider="fail1", fallback_policy=policy)
177
+ await manager.register_provider("fail1", _FailProvider("fail1"))
178
+ await manager.register_provider("fail2", _FailProvider("fail2"))
179
+
180
+ with pytest.raises(AllProvidersFailedError) as exc_info:
181
+ await manager.generate(_make_request())
182
+
183
+ assert len(exc_info.value.attempts) == 2
184
+
185
+ @pytest.mark.asyncio
186
+ async def test_retry_then_success(self) -> None:
187
+ policy = FallbackPolicy(
188
+ providers=["flaky"],
189
+ max_retries_per_provider=2,
190
+ base_delay_seconds=0.01,
191
+ max_delay_seconds=0.05,
192
+ )
193
+ manager = LLMManager(default_provider="flaky", fallback_policy=policy)
194
+ flaky = _FlakyProvider(fail_count=1, name="flaky")
195
+ await manager.register_provider("flaky", flaky)
196
+
197
+ response = await manager.generate(_make_request())
198
+ assert response.provider == "flaky"
199
+ assert flaky.call_count == 2 # 1 fail + 1 success
200
+
201
+ @pytest.mark.asyncio
202
+ async def test_retry_exhausted_then_fallback(self) -> None:
203
+ policy = FallbackPolicy(
204
+ providers=["flaky", "success"],
205
+ max_retries_per_provider=1,
206
+ base_delay_seconds=0.01,
207
+ max_delay_seconds=0.05,
208
+ )
209
+ manager = LLMManager(default_provider="flaky", fallback_policy=policy)
210
+ flaky = _FlakyProvider(fail_count=3, name="flaky") # Will never succeed within 1 retry
211
+ await manager.register_provider("flaky", flaky)
212
+ await manager.register_provider("success", _SuccessProvider("success"))
213
+
214
+ response = await manager.generate(_make_request())
215
+ assert response.provider == "success"
216
+
217
+ @pytest.mark.asyncio
218
+ async def test_explicit_provider_overrides_default(self) -> None:
219
+ manager = LLMManager(default_provider="p1")
220
+ await manager.register_provider("p1", _SuccessProvider("p1"))
221
+ await manager.register_provider("p2", _SuccessProvider("p2"))
222
+
223
+ response = await manager.generate(_make_request(), provider_name="p2")
224
+ assert response.provider == "p2"
225
+
226
+
227
+ class TestFallbackPolicy:
228
+ def test_default_values(self) -> None:
229
+ policy = FallbackPolicy()
230
+ assert policy.providers == []
231
+ assert policy.max_retries_per_provider == 2
232
+ assert policy.base_delay_seconds == 1.0
233
+ assert policy.max_delay_seconds == 30.0
234
+
235
+ def test_custom_values(self) -> None:
236
+ policy = FallbackPolicy(
237
+ providers=["openai", "claude"],
238
+ max_retries_per_provider=3,
239
+ base_delay_seconds=0.5,
240
+ max_delay_seconds=10.0,
241
+ )
242
+ assert policy.providers == ["openai", "claude"]
243
+ assert policy.max_retries_per_provider == 3
244
+