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()