Skip to content
This repository was archived by the owner on Jun 3, 2026. It is now read-only.

Commit 821a3ff

Browse files
committed
Refactor src/models/registry.py to reduce code duplication
1 parent 43ba731 commit 821a3ff

2 files changed

Lines changed: 104 additions & 20 deletions

File tree

src/models/registry.py

Lines changed: 14 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@
1010

1111
from __future__ import annotations
1212

13+
import importlib
1314
import logging
1415
from typing import Optional
1516

@@ -20,43 +21,36 @@
2021

2122
logger = logging.getLogger("xmem.models")
2223

24+
25+
def _build_from_module(module_name: str, func_name: str, **kwargs) -> BaseChatModel:
26+
module = importlib.import_module(f"src.models.{module_name}")
27+
factory_fn = getattr(module, func_name)
28+
return factory_fn(**kwargs)
29+
30+
2331
_BUILDERS = {
24-
"gemini": lambda **kw: _build_gemini(**kw),
25-
"claude": lambda **kw: _build_claude(**kw),
26-
"openai": lambda **kw: _build_openai(**kw),
32+
"gemini": lambda **kw: _build_from_module("gemini", "build_gemini_model", **kw),
33+
"claude": lambda **kw: _build_from_module("claude", "build_claude_model", **kw),
34+
"openai": lambda **kw: _build_from_module("openai", "build_openai_model", **kw),
2735
}
2836

37+
2938
_KEY_MAP = {
3039
"gemini": lambda: settings.gemini_api_key,
3140
"claude": lambda: settings.claude_api_key,
3241
"openai": lambda: settings.openai_api_key,
3342
}
3443

3544

36-
def _build_gemini(**kw) -> BaseChatModel:
37-
from src.models.gemini import build_gemini_model
38-
return build_gemini_model(**kw)
39-
40-
41-
def _build_claude(**kw) -> BaseChatModel:
42-
from src.models.claude import build_claude_model
43-
return build_claude_model(**kw)
44-
45-
46-
def _build_openai(**kw) -> BaseChatModel:
47-
from src.models.openai import build_openai_model
48-
return build_openai_model(**kw)
49-
50-
5145
def get_model(
5246
provider: Optional[Provider] = None,
5347
model_name: Optional[str] = None,
5448
temperature: Optional[float] = None,
5549
) -> BaseChatModel:
5650
"""Build and return a chat model.
5751
58-
If *provider* is None the first provider from ``settings.fallback_order``
59-
whose API key is configured will be used. Raises ``RuntimeError`` if no
52+
If *provider* is None the first provider from settings.fallback_order
53+
whose API key is configured will be used. Raises RuntimeError if no
6054
provider can be initialised.
6155
"""
6256
kw: dict = {}

tests/unit/models/test_registry.py

Lines changed: 90 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,90 @@
1+
import pytest
2+
from unittest.mock import MagicMock, patch
3+
import sys
4+
import importlib
5+
6+
7+
@pytest.fixture
8+
def mock_modules():
9+
with patch.dict(
10+
sys.modules,
11+
{
12+
"src.models.gemini": MagicMock(),
13+
"src.models.claude": MagicMock(),
14+
"src.models.openai": MagicMock(),
15+
},
16+
):
17+
yield
18+
19+
20+
def test_build_gemini(mock_modules):
21+
from src.models import registry
22+
23+
importlib.reload(registry) # Reload to ensure _BUILDERS is fresh if needed
24+
25+
mock_builder = MagicMock()
26+
sys.modules["src.models.gemini"].build_gemini_model = mock_builder
27+
28+
# We are testing the _BUILDERS map effectively
29+
registry._BUILDERS["gemini"](temperature=0.5)
30+
31+
mock_builder.assert_called_once_with(temperature=0.5)
32+
33+
34+
def test_build_claude(mock_modules):
35+
from src.models import registry
36+
37+
importlib.reload(registry)
38+
39+
mock_builder = MagicMock()
40+
sys.modules["src.models.claude"].build_claude_model = mock_builder
41+
42+
registry._BUILDERS["claude"](model_name="claude-3")
43+
44+
mock_builder.assert_called_once_with(model_name="claude-3")
45+
46+
47+
def test_build_openai(mock_modules):
48+
from src.models import registry
49+
50+
importlib.reload(registry)
51+
52+
mock_builder = MagicMock()
53+
sys.modules["src.models.openai"].build_openai_model = mock_builder
54+
55+
registry._BUILDERS["openai"]()
56+
57+
mock_builder.assert_called_once()
58+
59+
60+
def test_get_model_specific_provider(mock_modules):
61+
from src.models import registry
62+
63+
importlib.reload(registry)
64+
65+
mock_builder = MagicMock()
66+
sys.modules["src.models.gemini"].build_gemini_model = mock_builder
67+
68+
registry.get_model("gemini", temperature=0.7)
69+
70+
mock_builder.assert_called_once_with(temperature=0.7)
71+
72+
73+
def test_get_model_fallback(mock_modules):
74+
from src.models import registry
75+
76+
importlib.reload(registry)
77+
78+
# Mock settings.fallback_order and API keys
79+
# Note: registry imports settings, so we need to patch it in registry
80+
with patch("src.models.registry.settings") as mock_settings:
81+
mock_settings.fallback_order = ["openai", "gemini"]
82+
mock_settings.openai_api_key = "sk-test"
83+
mock_settings.gemini_api_key = None
84+
85+
mock_openai_builder = MagicMock()
86+
sys.modules["src.models.openai"].build_openai_model = mock_openai_builder
87+
88+
registry.get_model()
89+
90+
mock_openai_builder.assert_called_once()

0 commit comments

Comments
 (0)