File size: 4,048 Bytes
e0c9c7c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
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()