route group message extraction through AIOrchestrator for HuggingFace fallback on Groq rate limits
Browse files- app/ai/orchestrator.py +36 -0
- app/services/container.py +1 -1
- app/services/group_message_service.py +4 -4
- tests/test_ai_orchestrator.py +67 -0
- tests/test_group_message_service.py +13 -13
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 |
-
|
| 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.
|
| 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 |
-
|
| 28 |
settings: Settings,
|
| 29 |
) -> None:
|
| 30 |
self.repository = repository
|
| 31 |
self.embeddings = embeddings
|
| 32 |
-
self.
|
| 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.
|
| 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 |
-
|
| 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 |
-
|
| 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 |
-
|
| 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 |
-
|
| 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 |
-
|
| 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 |
-
|
| 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 |
-
|
| 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 |
-
|
| 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 |
-
|
| 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 |
-
|
| 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 |
-
|
| 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 |
-
|
| 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 |
-
|
| 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 |
|