61 lines
2.3 KiB
Python
61 lines
2.3 KiB
Python
from __future__ import annotations
|
|
|
|
import sys
|
|
import unittest
|
|
from pathlib import Path
|
|
|
|
|
|
ROOT = Path(__file__).resolve().parents[2]
|
|
sys.path.insert(0, str(ROOT / "backend"))
|
|
|
|
from app.settings import ProviderConfig, ProviderModel, Settings # noqa: E402
|
|
|
|
|
|
class SettingsModelSelectionTests(unittest.TestCase):
|
|
def test_implicit_author_uses_configured_default_model(self) -> None:
|
|
provider = ProviderConfig(
|
|
"author",
|
|
"Author",
|
|
"https://example.invalid/v1",
|
|
"key",
|
|
(ProviderModel("fast"), ProviderModel("reliable")),
|
|
)
|
|
settings = Settings(
|
|
task_root=ROOT / "tmp-tasks",
|
|
conversation_root=ROOT / "tmp-conversations",
|
|
library_root=ROOT / "backend" / "cdsl_library",
|
|
engine_root=ROOT / "backend" / "engine" / "cdsl_engine",
|
|
llm_base_url=provider.base_url,
|
|
llm_api_key=provider.api_key,
|
|
llm_model="reliable",
|
|
llm_timeout_s=1,
|
|
default_provider_id="author",
|
|
providers=(provider,),
|
|
)
|
|
|
|
resolved_provider, resolved_model = settings.resolve_model(None, None)
|
|
|
|
self.assertEqual(resolved_provider.id, "author")
|
|
self.assertEqual(resolved_model.id, "reliable")
|
|
|
|
def test_explicit_provider_without_model_uses_its_first_model(self) -> None:
|
|
default = ProviderConfig("default", "Default", "https://default.invalid", "key", (ProviderModel("default-model"),))
|
|
alternate = ProviderConfig("alternate", "Alternate", "https://alternate.invalid", "key", (ProviderModel("alternate-first"), ProviderModel("alternate-second")))
|
|
settings = Settings(
|
|
task_root=ROOT / "tmp-tasks",
|
|
conversation_root=ROOT / "tmp-conversations",
|
|
library_root=ROOT / "backend" / "cdsl_library",
|
|
engine_root=ROOT / "backend" / "engine" / "cdsl_engine",
|
|
llm_base_url=default.base_url,
|
|
llm_api_key=default.api_key,
|
|
llm_model="default-model",
|
|
llm_timeout_s=1,
|
|
default_provider_id="default",
|
|
providers=(default, alternate),
|
|
)
|
|
|
|
resolved_provider, resolved_model = settings.resolve_model("alternate", None)
|
|
|
|
self.assertEqual(resolved_provider.id, "alternate")
|
|
self.assertEqual(resolved_model.id, "alternate-first")
|