Skip to content

Commit b21a2be

Browse files
committed
feat(litellm): add direct provider fallback for unsupported features
LiteLLM doesn't support all provider features (e.g., Gemini aspect ratios other than 1:1). New DIRECT_PROVIDER_FALLBACK env var enables automatic fallback to native provider APIs when needed. - Add DIRECT_PROVIDER_FALLBACK config option (default: false) - Implement fallback logic in LiteLLMService for Gemini aspect ratios - Add helper methods for Gemini model detection and normalization - Add 14 tests for fallback logic
1 parent b9f280a commit b21a2be

5 files changed

Lines changed: 280 additions & 0 deletions

File tree

.env.example

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -53,6 +53,17 @@ MARKDOWN_EMBED_IMAGES=false
5353
# Bearer token for API authentication (leave empty to disable auth)
5454
API_BEARER_TOKEN=
5555

56+
# -----------------------------------------------------------------------------
57+
# Provider Fallback
58+
# -----------------------------------------------------------------------------
59+
# Use direct provider API when LiteLLM doesn't support a feature (default: false)
60+
# When enabled, requests using unsupported features will bypass LiteLLM and
61+
# use the provider's native API directly. Requires the respective API key.
62+
#
63+
# Currently applies to:
64+
# - Gemini aspect ratios other than 1:1 (requires GEMINI_API_KEY)
65+
DIRECT_PROVIDER_FALLBACK=false
66+
5667
# -----------------------------------------------------------------------------
5768
# Model Configuration
5869
# -----------------------------------------------------------------------------

README.md

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -127,6 +127,18 @@ OPENAI_API_KEY=sk-...
127127
GEMINI_API_KEY=...
128128
```
129129

130+
### Direct Provider Fallback
131+
132+
LiteLLM doesn't support all provider features. Enable `DIRECT_PROVIDER_FALLBACK` to automatically use the native provider API when needed:
133+
134+
```env
135+
DIRECT_PROVIDER_FALLBACK=true
136+
GEMINI_API_KEY=... # Required for Gemini fallback
137+
```
138+
139+
Currently applies to:
140+
- **Gemini aspect ratios**: 16:9, 9:16, 4:3, 3:4 (LiteLLM only supports 1:1)
141+
130142
## Supported Models
131143
132144
### OpenAI

app/core/config.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -57,6 +57,9 @@ class Settings(BaseSettings):
5757
FILTER_IMAGE_MODELS: bool = True # Only return image generation models from LiteLLM
5858
DEFAULT_MODEL: str | None = None # Default model for image generation
5959

60+
# Provider Fallback
61+
DIRECT_PROVIDER_FALLBACK: bool = False # Use direct provider API for unsupported LiteLLM features
62+
6063
# Server Configuration
6164
HOST: str = "0.0.0.0"
6265
PORT: int = 8000

app/services/litellm_service.py

Lines changed: 75 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -58,6 +58,16 @@ async def generate_image(
5858
"""
5959
logger.info(f"Generating {n} image(s) with {model} via LiteLLM")
6060

61+
# Fallback to direct provider API for unsupported features
62+
if self._should_use_direct_provider(model, aspect_ratio):
63+
return await self._generate_via_direct_provider(
64+
prompt=prompt,
65+
model=model,
66+
aspect_ratio=aspect_ratio,
67+
quality=quality,
68+
n=n,
69+
)
70+
6171
# Get model capabilities to adjust parameters
6272
model_info = model_registry.get_model(model)
6373
size = self._get_size(aspect_ratio)
@@ -168,6 +178,71 @@ def _get_size(self, aspect_ratio: str) -> str:
168178
"""
169179
return self.ASPECT_RATIO_SIZES.get(aspect_ratio, "1024x1024")
170180

181+
def _should_use_direct_provider(self, model: str, aspect_ratio: str) -> bool:
182+
"""
183+
Check if we should use direct provider API instead of LiteLLM.
184+
185+
Currently applies to:
186+
- Gemini models with non-square aspect ratios (LiteLLM doesn't support this)
187+
"""
188+
if not settings.DIRECT_PROVIDER_FALLBACK:
189+
return False
190+
191+
# Gemini: aspect ratios other than 1:1 not supported via LiteLLM
192+
if self._is_gemini_model(model) and aspect_ratio != "1:1":
193+
if settings.gemini_available:
194+
return True
195+
logger.warning(
196+
f"DIRECT_PROVIDER_FALLBACK enabled but GEMINI_API_KEY not set. "
197+
f"Falling back to LiteLLM (aspect_ratio={aspect_ratio} may not work)."
198+
)
199+
200+
return False
201+
202+
def _is_gemini_model(self, model: str) -> bool:
203+
"""Check if model is a Gemini/Imagen model."""
204+
model_lower = model.lower()
205+
return "gemini" in model_lower or "imagen" in model_lower
206+
207+
def _normalize_gemini_model(self, model: str) -> str:
208+
"""
209+
Normalize Gemini model name for direct API.
210+
211+
LiteLLM uses "gemini/model-name", direct API uses "model-name".
212+
"""
213+
if model.startswith("gemini/"):
214+
return model[7:]
215+
return model
216+
217+
async def _generate_via_direct_provider(
218+
self,
219+
prompt: str,
220+
model: str,
221+
aspect_ratio: str,
222+
quality: str,
223+
n: int,
224+
) -> list[str]:
225+
"""
226+
Generate images via direct provider API.
227+
228+
Used when LiteLLM doesn't support a required feature.
229+
"""
230+
if self._is_gemini_model(model):
231+
from app.services.gemini_service import get_gemini_service
232+
233+
logger.info(f"Using direct Gemini API for {model} (aspect_ratio={aspect_ratio})")
234+
gemini_service = get_gemini_service()
235+
return await gemini_service.generate_image(
236+
prompt=prompt,
237+
model=self._normalize_gemini_model(model),
238+
aspect_ratio=aspect_ratio,
239+
quality=quality,
240+
n=n,
241+
)
242+
243+
# Add other providers here as needed
244+
raise ValueError(f"No direct provider fallback available for model: {model}")
245+
171246

172247
def get_litellm_service() -> LiteLLMService:
173248
"""

tests/test_litellm_service.py

Lines changed: 179 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,179 @@
1+
from unittest.mock import AsyncMock, MagicMock, patch
2+
3+
import pytest
4+
5+
from app.services.litellm_service import LiteLLMService
6+
7+
8+
class TestIsGeminiModel:
9+
"""Tests for _is_gemini_model helper."""
10+
11+
@pytest.fixture
12+
def service(self):
13+
with patch("app.services.litellm_service.settings") as mock_settings:
14+
mock_settings.LITELLM_BASE_URL = "http://localhost:4000"
15+
mock_settings.LITELLM_API_KEY = "test-key"
16+
return LiteLLMService()
17+
18+
def test_gemini_model(self, service):
19+
assert service._is_gemini_model("gemini-2.0-flash-preview-image-generation") is True
20+
assert service._is_gemini_model("gemini/gemini-2.0-flash-exp") is True
21+
assert service._is_gemini_model("GEMINI-MODEL") is True
22+
23+
def test_imagen_model(self, service):
24+
assert service._is_gemini_model("imagen-3.0-generate-002") is True
25+
assert service._is_gemini_model("vertex/imagen-3") is True
26+
27+
def test_non_gemini_model(self, service):
28+
assert service._is_gemini_model("dall-e-3") is False
29+
assert service._is_gemini_model("gpt-image-1") is False
30+
assert service._is_gemini_model("stable-diffusion") is False
31+
32+
33+
class TestNormalizeGeminiModel:
34+
"""Tests for _normalize_gemini_model helper."""
35+
36+
@pytest.fixture
37+
def service(self):
38+
with patch("app.services.litellm_service.settings") as mock_settings:
39+
mock_settings.LITELLM_BASE_URL = "http://localhost:4000"
40+
mock_settings.LITELLM_API_KEY = "test-key"
41+
return LiteLLMService()
42+
43+
def test_removes_gemini_prefix(self, service):
44+
assert service._normalize_gemini_model("gemini/gemini-2.0-flash-exp") == "gemini-2.0-flash-exp"
45+
46+
def test_keeps_model_without_prefix(self, service):
47+
assert service._normalize_gemini_model("gemini-2.0-flash-exp") == "gemini-2.0-flash-exp"
48+
49+
def test_keeps_other_prefixes(self, service):
50+
assert service._normalize_gemini_model("vertex/imagen-3") == "vertex/imagen-3"
51+
52+
53+
class TestShouldUseDirectProvider:
54+
"""Tests for _should_use_direct_provider logic."""
55+
56+
@pytest.fixture
57+
def service(self):
58+
with patch("app.services.litellm_service.settings") as mock_settings:
59+
mock_settings.LITELLM_BASE_URL = "http://localhost:4000"
60+
mock_settings.LITELLM_API_KEY = "test-key"
61+
return LiteLLMService()
62+
63+
def test_disabled_when_fallback_false(self, service):
64+
with patch("app.services.litellm_service.settings") as mock_settings:
65+
mock_settings.DIRECT_PROVIDER_FALLBACK = False
66+
assert service._should_use_direct_provider("gemini/model", "16:9") is False
67+
68+
def test_disabled_for_square_aspect_ratio(self, service):
69+
with patch("app.services.litellm_service.settings") as mock_settings:
70+
mock_settings.DIRECT_PROVIDER_FALLBACK = True
71+
mock_settings.gemini_available = True
72+
assert service._should_use_direct_provider("gemini/model", "1:1") is False
73+
74+
def test_enabled_for_gemini_non_square(self, service):
75+
with patch("app.services.litellm_service.settings") as mock_settings:
76+
mock_settings.DIRECT_PROVIDER_FALLBACK = True
77+
mock_settings.gemini_available = True
78+
assert service._should_use_direct_provider("gemini/model", "16:9") is True
79+
assert service._should_use_direct_provider("gemini/model", "9:16") is True
80+
assert service._should_use_direct_provider("imagen-3", "4:3") is True
81+
82+
def test_disabled_when_gemini_key_missing(self, service, caplog):
83+
with patch("app.services.litellm_service.settings") as mock_settings:
84+
mock_settings.DIRECT_PROVIDER_FALLBACK = True
85+
mock_settings.gemini_available = False
86+
assert service._should_use_direct_provider("gemini/model", "16:9") is False
87+
assert "GEMINI_API_KEY not set" in caplog.text
88+
89+
def test_disabled_for_non_gemini_models(self, service):
90+
with patch("app.services.litellm_service.settings") as mock_settings:
91+
mock_settings.DIRECT_PROVIDER_FALLBACK = True
92+
mock_settings.gemini_available = True
93+
assert service._should_use_direct_provider("dall-e-3", "16:9") is False
94+
assert service._should_use_direct_provider("gpt-image-1", "9:16") is False
95+
96+
97+
class TestGenerateImageWithFallback:
98+
"""Tests for generate_image with direct provider fallback."""
99+
100+
@pytest.fixture
101+
def service(self):
102+
with patch("app.services.litellm_service.settings") as mock_settings:
103+
mock_settings.LITELLM_BASE_URL = "http://localhost:4000"
104+
mock_settings.LITELLM_API_KEY = "test-key"
105+
return LiteLLMService()
106+
107+
@pytest.mark.asyncio
108+
async def test_uses_direct_gemini_for_non_square(self, service):
109+
with patch("app.services.litellm_service.settings") as mock_settings:
110+
mock_settings.DIRECT_PROVIDER_FALLBACK = True
111+
mock_settings.gemini_available = True
112+
113+
mock_gemini_service = MagicMock()
114+
mock_gemini_service.generate_image = AsyncMock(return_value=["http://example.com/image.png"])
115+
116+
with patch("app.services.gemini_service.get_gemini_service", return_value=mock_gemini_service):
117+
result = await service.generate_image(
118+
prompt="test prompt",
119+
model="gemini/gemini-2.0-flash-exp",
120+
aspect_ratio="16:9",
121+
quality="standard",
122+
n=1,
123+
)
124+
125+
assert result == ["http://example.com/image.png"]
126+
mock_gemini_service.generate_image.assert_called_once_with(
127+
prompt="test prompt",
128+
model="gemini-2.0-flash-exp", # Prefix removed
129+
aspect_ratio="16:9",
130+
quality="standard",
131+
n=1,
132+
)
133+
134+
@pytest.mark.asyncio
135+
async def test_uses_litellm_for_square_aspect_ratio(self, service, mock_openai_response):
136+
with patch("app.services.litellm_service.settings") as mock_settings:
137+
mock_settings.DIRECT_PROVIDER_FALLBACK = True
138+
mock_settings.gemini_available = True
139+
140+
with patch.object(service.client.images, "generate", return_value=mock_openai_response):
141+
with patch("app.services.litellm_service.model_registry") as mock_registry:
142+
mock_registry.get_model.return_value = None
143+
144+
with patch("app.services.litellm_service.storage_service") as mock_storage:
145+
mock_storage.save_image = AsyncMock(return_value="http://example.com/image.png")
146+
147+
result = await service.generate_image(
148+
prompt="test prompt",
149+
model="gemini/gemini-2.0-flash-exp",
150+
aspect_ratio="1:1",
151+
quality="standard",
152+
n=1,
153+
)
154+
155+
assert result == ["http://example.com/image.png"]
156+
service.client.images.generate.assert_called_once()
157+
158+
@pytest.mark.asyncio
159+
async def test_uses_litellm_when_fallback_disabled(self, service, mock_openai_response):
160+
with patch("app.services.litellm_service.settings") as mock_settings:
161+
mock_settings.DIRECT_PROVIDER_FALLBACK = False
162+
163+
with patch.object(service.client.images, "generate", return_value=mock_openai_response):
164+
with patch("app.services.litellm_service.model_registry") as mock_registry:
165+
mock_registry.get_model.return_value = None
166+
167+
with patch("app.services.litellm_service.storage_service") as mock_storage:
168+
mock_storage.save_image = AsyncMock(return_value="http://example.com/image.png")
169+
170+
result = await service.generate_image(
171+
prompt="test prompt",
172+
model="gemini/gemini-2.0-flash-exp",
173+
aspect_ratio="16:9",
174+
quality="standard",
175+
n=1,
176+
)
177+
178+
assert result == ["http://example.com/image.png"]
179+
service.client.images.generate.assert_called_once()

0 commit comments

Comments
 (0)