codeBOKER commited on
Commit
4ce08fd
·
1 Parent(s): b5a802c

route group message extraction through AIOrchestrator for HuggingFace fallback on Groq rate limits

Browse files
app/ai/orchestrator.py CHANGED
@@ -27,6 +27,42 @@ class AIOrchestrator:
27
  self.temperature = temperature
28
  self.max_tool_iterations = max_tool_iterations
29
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
30
  async def generate_reply(
31
  self,
32
  *,
 
27
  self.temperature = temperature
28
  self.max_tool_iterations = max_tool_iterations
29
 
30
+ async def chat(
31
+ self,
32
+ *,
33
+ messages: list[dict[str, Any]],
34
+ tools: list[dict[str, Any]] | None = None,
35
+ tool_choice: str | dict[str, Any] | None = None,
36
+ temperature: float = 0.2,
37
+ ) -> AIProviderResponse:
38
+ try:
39
+ return await self.primary.chat(
40
+ messages,
41
+ tools=tools,
42
+ tool_choice=tool_choice,
43
+ temperature=temperature,
44
+ )
45
+ except InvalidToolCallGenerationError:
46
+ logger.warning("Primary provider generated invalid tool call; retrying once")
47
+ try:
48
+ return await self.primary.chat(
49
+ messages,
50
+ tools=tools,
51
+ tool_choice=tool_choice,
52
+ temperature=max(temperature - 0.2, 0.1),
53
+ )
54
+ except RetryableProviderError:
55
+ logger.warning("Primary retry failed; falling back to Hugging Face")
56
+ except RetryableProviderError:
57
+ logger.warning("Primary provider failed; falling back to Hugging Face")
58
+
59
+ return await self.fallback.chat(
60
+ messages,
61
+ tools=tools,
62
+ tool_choice=tool_choice,
63
+ temperature=temperature,
64
+ )
65
+
66
  async def generate_reply(
67
  self,
68
  *,
app/services/container.py CHANGED
@@ -44,7 +44,7 @@ class ServiceContainer:
44
  group_message = GroupMessageService(
45
  repository=repository,
46
  embeddings=embeddings,
47
- provider=ai.primary,
48
  settings=settings,
49
  )
50
  admin = AdminService(repository=repository, embeddings=embeddings, settings=settings)
 
44
  group_message = GroupMessageService(
45
  repository=repository,
46
  embeddings=embeddings,
47
+ ai=ai,
48
  settings=settings,
49
  )
50
  admin = AdminService(repository=repository, embeddings=embeddings, settings=settings)
app/services/group_message_service.py CHANGED
@@ -4,7 +4,7 @@ from datetime import date
4
  from pathlib import Path
5
  from typing import Any
6
 
7
- from app.ai.providers import ChatProvider
8
  from app.config import Settings
9
  from app.database.supabase import SupabaseRepository
10
  from app.models.domain import ExtractedTrip, WhatsAppInboundMessage
@@ -24,12 +24,12 @@ class GroupMessageService:
24
  *,
25
  repository: SupabaseRepository,
26
  embeddings: JinaEmbeddingService,
27
- provider: ChatProvider,
28
  settings: Settings,
29
  ) -> None:
30
  self.repository = repository
31
  self.embeddings = embeddings
32
- self.provider = provider
33
  self.settings = settings
34
 
35
  async def handle_group_message(
@@ -153,7 +153,7 @@ class GroupMessageService:
153
  prompt = prompt_template.format(current_datetime=dt.isoformat())
154
 
155
  try:
156
- response = await self.provider.chat(
157
  messages=[
158
  {"role": "system", "content": prompt},
159
  {"role": "user", "content": text},
 
4
  from pathlib import Path
5
  from typing import Any
6
 
7
+ from app.ai.orchestrator import AIOrchestrator
8
  from app.config import Settings
9
  from app.database.supabase import SupabaseRepository
10
  from app.models.domain import ExtractedTrip, WhatsAppInboundMessage
 
24
  *,
25
  repository: SupabaseRepository,
26
  embeddings: JinaEmbeddingService,
27
+ ai: AIOrchestrator,
28
  settings: Settings,
29
  ) -> None:
30
  self.repository = repository
31
  self.embeddings = embeddings
32
+ self.ai = ai
33
  self.settings = settings
34
 
35
  async def handle_group_message(
 
153
  prompt = prompt_template.format(current_datetime=dt.isoformat())
154
 
155
  try:
156
+ response = await self.ai.chat(
157
  messages=[
158
  {"role": "system", "content": prompt},
159
  {"role": "user", "content": text},
tests/test_ai_orchestrator.py CHANGED
@@ -137,3 +137,70 @@ async def test_ai_reports_invalid_tool_arguments_to_model():
137
 
138
  assert reply == "Please share the question again."
139
  assert "Invalid tool arguments" in primary.calls[1]["messages"][-1]["content"]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
137
 
138
  assert reply == "Please share the question again."
139
  assert "Invalid tool arguments" in primary.calls[1]["messages"][-1]["content"]
140
+
141
+
142
+ @pytest.mark.asyncio
143
+ async def test_chat_falls_back_when_primary_rate_limited():
144
+ primary = ScriptedProvider("groq", [RetryableProviderError("rate limited")])
145
+ fallback = ScriptedProvider("huggingface", [AIProviderResponse(content="fallback reply")])
146
+ orchestrator = AIOrchestrator(
147
+ primary=primary,
148
+ fallback=fallback,
149
+ temperature=0.2,
150
+ max_tool_iterations=3,
151
+ )
152
+
153
+ response = await orchestrator.chat(
154
+ messages=[{"role": "user", "content": "test"}],
155
+ )
156
+
157
+ assert response.content == "fallback reply"
158
+ assert len(primary.calls) == 1
159
+ assert len(fallback.calls) == 1
160
+
161
+
162
+ @pytest.mark.asyncio
163
+ async def test_chat_returns_primary_on_success():
164
+ primary = ScriptedProvider("groq", [AIProviderResponse(content="primary reply")])
165
+ fallback = ScriptedProvider("huggingface", [AIProviderResponse(content="fallback")])
166
+ orchestrator = AIOrchestrator(
167
+ primary=primary,
168
+ fallback=fallback,
169
+ temperature=0.2,
170
+ max_tool_iterations=3,
171
+ )
172
+
173
+ response = await orchestrator.chat(
174
+ messages=[{"role": "user", "content": "test"}],
175
+ )
176
+
177
+ assert response.content == "primary reply"
178
+ assert len(primary.calls) == 1
179
+ assert len(fallback.calls) == 0
180
+
181
+
182
+ @pytest.mark.asyncio
183
+ async def test_chat_retries_invalid_tool_call_before_fallback():
184
+ primary = ScriptedProvider(
185
+ "groq",
186
+ [
187
+ InvalidToolCallGenerationError("bad tool"),
188
+ AIProviderResponse(content="primary retry reply"),
189
+ ],
190
+ )
191
+ fallback = ScriptedProvider("huggingface", [AIProviderResponse(content="fallback")])
192
+ orchestrator = AIOrchestrator(
193
+ primary=primary,
194
+ fallback=fallback,
195
+ temperature=0.4,
196
+ max_tool_iterations=3,
197
+ )
198
+
199
+ response = await orchestrator.chat(
200
+ messages=[{"role": "user", "content": "test"}],
201
+ temperature=0.4,
202
+ )
203
+
204
+ assert response.content == "primary retry reply"
205
+ assert [call["temperature"] for call in primary.calls] == [0.4, 0.2]
206
+ assert fallback.calls == []
tests/test_group_message_service.py CHANGED
@@ -96,7 +96,7 @@ async def test_non_trip_message_is_ignored(settings: Settings) -> None:
96
  service = GroupMessageService(
97
  repository=repo,
98
  embeddings=embeddings,
99
- provider=provider,
100
  settings=settings,
101
  )
102
 
@@ -117,7 +117,7 @@ async def test_trip_ad_from_new_driver_creates_all_entities(settings: Settings)
117
  service = GroupMessageService(
118
  repository=repo,
119
  embeddings=embeddings,
120
- provider=provider,
121
  settings=settings,
122
  )
123
 
@@ -153,7 +153,7 @@ async def test_trip_ad_from_existing_customer_is_discarded(settings: Settings) -
153
  service = GroupMessageService(
154
  repository=repo,
155
  embeddings=embeddings,
156
- provider=provider,
157
  settings=settings,
158
  )
159
 
@@ -180,7 +180,7 @@ async def test_existing_unregistered_driver_adds_new_trip(settings: Settings) ->
180
  service = GroupMessageService(
181
  repository=repo,
182
  embeddings=embeddings,
183
- provider=provider,
184
  settings=settings,
185
  )
186
 
@@ -221,7 +221,7 @@ async def test_existing_unregistered_driver_duplicate_trip_is_skipped(settings:
221
  service = GroupMessageService(
222
  repository=repo,
223
  embeddings=embeddings,
224
- provider=provider,
225
  settings=settings,
226
  )
227
 
@@ -250,7 +250,7 @@ async def test_incomplete_fields_cause_skip(settings: Settings) -> None:
250
  service = GroupMessageService(
251
  repository=repo,
252
  embeddings=embeddings,
253
- provider=provider,
254
  settings=settings,
255
  )
256
 
@@ -278,7 +278,7 @@ async def test_missing_phone_cause_skip(settings: Settings) -> None:
278
  service = GroupMessageService(
279
  repository=repo,
280
  embeddings=embeddings,
281
- provider=provider,
282
  settings=settings,
283
  )
284
 
@@ -297,7 +297,7 @@ async def test_duplicate_group_message_is_deduplicated(settings: Settings) -> No
297
  service = GroupMessageService(
298
  repository=repo,
299
  embeddings=embeddings,
300
- provider=provider,
301
  settings=settings,
302
  )
303
 
@@ -320,7 +320,7 @@ async def test_phone_normalization(settings: Settings) -> None:
320
  service = GroupMessageService(
321
  repository=repo,
322
  embeddings=embeddings,
323
- provider=provider,
324
  settings=settings,
325
  )
326
 
@@ -341,7 +341,7 @@ async def test_invalid_json_response_cause_skip(settings: Settings) -> None:
341
  service = GroupMessageService(
342
  repository=repo,
343
  embeddings=embeddings,
344
- provider=provider,
345
  settings=settings,
346
  )
347
 
@@ -361,7 +361,7 @@ async def test_markdown_fenced_json_is_parsed(settings: Settings) -> None:
361
  service = GroupMessageService(
362
  repository=repo,
363
  embeddings=embeddings,
364
- provider=provider,
365
  settings=settings,
366
  )
367
 
@@ -382,7 +382,7 @@ async def test_trip_ad_with_missing_car_type_uses_unknown(settings: Settings) ->
382
  service = GroupMessageService(
383
  repository=repo,
384
  embeddings=embeddings,
385
- provider=provider,
386
  settings=settings,
387
  )
388
 
@@ -405,7 +405,7 @@ async def test_trip_ad_with_missing_name_uses_none(settings: Settings) -> None:
405
  service = GroupMessageService(
406
  repository=repo,
407
  embeddings=embeddings,
408
- provider=provider,
409
  settings=settings,
410
  )
411
 
 
96
  service = GroupMessageService(
97
  repository=repo,
98
  embeddings=embeddings,
99
+ ai=provider,
100
  settings=settings,
101
  )
102
 
 
117
  service = GroupMessageService(
118
  repository=repo,
119
  embeddings=embeddings,
120
+ ai=provider,
121
  settings=settings,
122
  )
123
 
 
153
  service = GroupMessageService(
154
  repository=repo,
155
  embeddings=embeddings,
156
+ ai=provider,
157
  settings=settings,
158
  )
159
 
 
180
  service = GroupMessageService(
181
  repository=repo,
182
  embeddings=embeddings,
183
+ ai=provider,
184
  settings=settings,
185
  )
186
 
 
221
  service = GroupMessageService(
222
  repository=repo,
223
  embeddings=embeddings,
224
+ ai=provider,
225
  settings=settings,
226
  )
227
 
 
250
  service = GroupMessageService(
251
  repository=repo,
252
  embeddings=embeddings,
253
+ ai=provider,
254
  settings=settings,
255
  )
256
 
 
278
  service = GroupMessageService(
279
  repository=repo,
280
  embeddings=embeddings,
281
+ ai=provider,
282
  settings=settings,
283
  )
284
 
 
297
  service = GroupMessageService(
298
  repository=repo,
299
  embeddings=embeddings,
300
+ ai=provider,
301
  settings=settings,
302
  )
303
 
 
320
  service = GroupMessageService(
321
  repository=repo,
322
  embeddings=embeddings,
323
+ ai=provider,
324
  settings=settings,
325
  )
326
 
 
341
  service = GroupMessageService(
342
  repository=repo,
343
  embeddings=embeddings,
344
+ ai=provider,
345
  settings=settings,
346
  )
347
 
 
361
  service = GroupMessageService(
362
  repository=repo,
363
  embeddings=embeddings,
364
+ ai=provider,
365
  settings=settings,
366
  )
367
 
 
382
  service = GroupMessageService(
383
  repository=repo,
384
  embeddings=embeddings,
385
+ ai=provider,
386
  settings=settings,
387
  )
388
 
 
405
  service = GroupMessageService(
406
  repository=repo,
407
  embeddings=embeddings,
408
+ ai=provider,
409
  settings=settings,
410
  )
411