Terminal / tests /test_provider_model_migrations.py
github-actions[bot]
Sync backend-only Space export 807f4ac8
c8365f5
Raw
History Blame Contribute Delete
4.05 kB
import os
import unittest
from types import SimpleNamespace
from unittest.mock import patch
from api import providers
from models.ai_client import AIClient
from models.role_router import Role, RoleRouter
class _RecordingTable:
def __init__(self):
self.calls: list[dict] = []
self._pending: dict = {}
def update(self, values: dict):
self._pending = {"values": values, "filters": []}
return self
def eq(self, field: str, value: str):
self._pending["filters"].append((field, value))
return self
def execute(self):
self.calls.append(self._pending)
return SimpleNamespace(data=[{"id": len(self.calls)}])
class _RecordingSupabase:
def __init__(self):
self.table_ref = _RecordingTable()
def table(self, name: str):
if name != "ai_providers":
raise AssertionError(f"Tabella inattesa: {name}")
return self.table_ref
class RoleRouterModelMigrationTests(unittest.TestCase):
@patch.object(AIClient, "_load_providers", return_value=[])
@patch.dict(
os.environ,
{"GROQ_API_KEY": "test-groq-key"},
clear=True,
)
def test_architect_uses_supported_groq_default(self, _load_providers):
client = RoleRouter.get_client(Role.ARCHITECT)
self.assertEqual(client.provider_name, "groq-architect")
self.assertEqual(client.default_model, "qwen/qwen3.6-27b")
@patch.object(AIClient, "_load_providers", return_value=[])
@patch.dict(
os.environ,
{"OPENROUTER_API_KEY": "test-openrouter-key"},
clear=True,
)
def test_openrouter_role_fallbacks_use_available_free_model(self, _load_providers):
architect = RoleRouter.get_client(Role.ARCHITECT)
coder = RoleRouter.get_client(Role.CODER)
self.assertEqual(architect.default_model, "openrouter/free")
self.assertEqual(coder.default_model, "openrouter/free")
@patch.object(AIClient, "_load_providers", return_value=[])
@patch.dict(
os.environ,
{"CEREBRAS_API_KEY": "test-cerebras-key"},
clear=True,
)
def test_reasoner_uses_cerebras_gpt_oss_default(self, _load_providers):
client = RoleRouter.get_client(Role.REASONER)
self.assertEqual(client.provider_name, "cerebras-reasoner")
self.assertEqual(client.default_model, "gpt-oss-120b")
class ProviderTableMigrationTests(unittest.IsolatedAsyncioTestCase):
async def test_update_models_filters_by_provider_and_never_downgrades_gpt_oss(self):
original_supabase = providers._sb
fake_supabase = _RecordingSupabase()
providers._sb = fake_supabase
try:
payload = await providers.update_provider_models(role=None)
finally:
providers._sb = original_supabase
self.assertTrue(payload["ok"])
self.assertEqual(payload["total_updated"], 8)
self.assertEqual(len(fake_supabase.table_ref.calls), 8)
expected = {
("groq", "llama-3.1-70b-versatile", "qwen/qwen3.6-27b"),
("cerebras", "llama3.1-70b", "gpt-oss-120b"),
("nvidia", "llama-3.1-405b-instruct", "meta/llama-3.3-70b-instruct"),
("openrouter", "llama-3.1-405b", "openrouter/free"),
("sambanova", "llama3-70b", "DeepSeek-V3.2"),
("gemini", "gemini-1.5-flash", "gemini-3.5-flash-lite"),
("gemini", "gemini-1.5-pro", "gemini-3.6-flash"),
("openrouter", "claude-3.5-sonnet", "openrouter/free"),
}
actual = {
(
dict(call["filters"])["name"],
dict(call["filters"])["default_model"],
call["values"]["default_model"],
)
for call in fake_supabase.table_ref.calls
}
self.assertEqual(actual, expected)
self.assertNotIn("llama-4-scout", {new for _, _, new in actual})
self.assertNotIn(("cerebras", "gpt-oss-120b", "llama-4-scout"), actual)
if __name__ == "__main__":
unittest.main()