Skip to content

Commit bcd6f17

Browse files
committed
feat: add DEFAULT_MODEL environment variable
1 parent 30f80ee commit bcd6f17

2 files changed

Lines changed: 61 additions & 0 deletions

File tree

app/api/routes/generate.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -305,6 +305,10 @@ def _get_default_model(provider: str) -> str:
305305
"""
306306
Get default model for provider based on available models.
307307
"""
308+
# Check if DEFAULT_MODEL is configured
309+
if settings.DEFAULT_MODEL:
310+
return settings.DEFAULT_MODEL
311+
308312
models = model_registry.get_models()
309313

310314
# Filter models for this provider

tests/test_generate.py

Lines changed: 57 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,57 @@
1+
from unittest.mock import MagicMock, patch
2+
3+
4+
def test_get_default_model_uses_env_config():
5+
"""
6+
Test that DEFAULT_MODEL env var is used.
7+
"""
8+
with patch("app.api.routes.generate.settings") as mock_settings:
9+
mock_settings.DEFAULT_MODEL = "gemini/gemini-2.0-flash-exp-image-generation"
10+
11+
from app.api.routes.generate import _get_default_model
12+
13+
result = _get_default_model("litellm")
14+
assert result == "gemini/gemini-2.0-flash-exp-image-generation"
15+
16+
17+
def test_get_default_model_fallback_to_registry():
18+
"""
19+
Test fallback to registry when DEFAULT_MODEL is not set.
20+
"""
21+
with patch("app.api.routes.generate.settings") as mock_settings:
22+
mock_settings.DEFAULT_MODEL = None
23+
24+
with patch("app.api.routes.generate.model_registry") as mock_registry:
25+
mock_model = MagicMock()
26+
mock_model.id = "dall-e-3"
27+
mock_model.provider = "openai"
28+
mock_registry.get_models.return_value = [mock_model]
29+
30+
from app.api.routes.generate import _get_default_model
31+
32+
result = _get_default_model("litellm")
33+
assert result == "dall-e-3"
34+
35+
36+
def test_get_default_model_hardcoded_fallback():
37+
"""
38+
Test hardcoded fallback when no models in registry.
39+
"""
40+
with patch("app.api.routes.generate.settings") as mock_settings:
41+
mock_settings.DEFAULT_MODEL = None
42+
43+
with patch("app.api.routes.generate.model_registry") as mock_registry:
44+
mock_registry.get_models.return_value = []
45+
46+
from app.api.routes.generate import _get_default_model
47+
48+
# LiteLLM/OpenAI provider defaults to dall-e-3
49+
result = _get_default_model("litellm")
50+
assert result == "dall-e-3"
51+
52+
result = _get_default_model("openai")
53+
assert result == "dall-e-3"
54+
55+
# Gemini provider defaults to gemini model
56+
result = _get_default_model("gemini")
57+
assert result == "gemini-2.0-flash-preview-image-generation"

0 commit comments

Comments
 (0)