Skip to content

Commit af0fa95

Browse files
authored
fix model mapping (#47)
1 parent 697f26a commit af0fa95

2 files changed

Lines changed: 337 additions & 6 deletions

File tree

app/services/provider_service.py

Lines changed: 20 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -139,6 +139,19 @@ def _get_adapters(self) -> dict[str, ProviderAdapter]:
139139
ProviderService._adapters_cache = ProviderAdapterFactory.get_all_adapters()
140140
return ProviderService._adapters_cache
141141

142+
def _ensure_model_mapping_dict(self, model_mapping: Any) -> dict[str, Any]:
143+
"""Ensure model_mapping is a dictionary, handling cases where it might be a string."""
144+
if isinstance(model_mapping, dict):
145+
return model_mapping
146+
elif isinstance(model_mapping, str):
147+
try:
148+
import json
149+
return json.loads(model_mapping) if model_mapping else {}
150+
except (json.JSONDecodeError, TypeError):
151+
return {}
152+
else:
153+
return {}
154+
142155
async def _load_provider_keys(self) -> dict[str, dict[str, Any]]:
143156
"""Load all provider keys for the user synchronously, with lazy loading and caching."""
144157
if self._keys_loaded:
@@ -169,7 +182,7 @@ async def _load_provider_keys(self) -> dict[str, dict[str, Any]]:
169182

170183
keys = {}
171184
for provider_key in provider_key_records:
172-
model_mapping = provider_key.model_mapping or {}
185+
model_mapping = self._ensure_model_mapping_dict(provider_key.model_mapping or {})
173186

174187
keys[provider_key.provider_name] = {
175188
"api_key": decrypt_api_key(provider_key.encrypted_api_key),
@@ -221,7 +234,7 @@ async def _load_provider_keys_async(self) -> dict[str, dict[str, Any]]:
221234

222235
keys = {}
223236
for provider_key in provider_key_records:
224-
model_mapping = provider_key.model_mapping or {}
237+
model_mapping = self._ensure_model_mapping_dict(provider_key.model_mapping or {})
225238

226239
keys[provider_key.provider_name] = {
227240
"api_key": decrypt_api_key(provider_key.encrypted_api_key),
@@ -285,7 +298,7 @@ def _get_provider_info_with_prefix(
285298

286299
provider_data = self.provider_keys[matching_provider]
287300

288-
model_mapping = provider_data.get("model_mapping", {})
301+
model_mapping = self._ensure_model_mapping_dict(provider_data.get("model_mapping", {}))
289302
mapped_model = model_mapping.get(model_name, model_name)
290303
return (
291304
matching_provider,
@@ -308,7 +321,7 @@ def _find_provider_for_unprefixed_model(
308321

309322
# Check custom model mappings
310323
for provider_name, provider_data in sorted_providers:
311-
model_mapping = provider_data.get("model_mapping", {})
324+
model_mapping = self._ensure_model_mapping_dict(provider_data.get("model_mapping", {}))
312325
if model in model_mapping:
313326
mapped_model = model_mapping[model]
314327
return (
@@ -369,7 +382,8 @@ async def list_models(
369382

370383
# Create a cache key unique to this provider config
371384
base_url = provider_data.get("base_url", "default")
372-
cache_key = f"{base_url}:{hash(frozenset(provider_data.get('model_mapping', {}).items()))}"
385+
model_mapping = self._ensure_model_mapping_dict(provider_data.get("model_mapping", {}))
386+
cache_key = f"{base_url}:{hash(frozenset(model_mapping.items()))}"
373387

374388
# Check if we have cached models for this provider
375389
cached_models = await self.get_cached_models(provider_name, cache_key)
@@ -387,7 +401,7 @@ async def _list_models_helper(
387401
) -> list[dict[str, Any]]:
388402
try:
389403
model_names = await adapter.list_models(api_key)
390-
model_mapping = provider_data.get("model_mapping", {})
404+
model_mapping = self._ensure_model_mapping_dict(provider_data.get("model_mapping", {}))
391405
reverse_model_mapping = {v: k for k, v in model_mapping.items()}
392406
provider_models = [
393407
{
Lines changed: 317 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,317 @@
1+
import pytest
2+
from unittest.mock import AsyncMock, MagicMock
3+
from sqlalchemy.ext.asyncio import AsyncSession
4+
5+
from app.services.provider_service import ProviderService
6+
from app.models.provider_key import ProviderKey
7+
from app.models.user import User
8+
9+
10+
class TestModelMappingFix:
11+
"""Test cases for the model_mapping string-to-dict conversion fix."""
12+
13+
def test_ensure_model_mapping_dict_helper(self):
14+
"""Test the _ensure_model_mapping_dict helper method with various inputs."""
15+
# Create a ProviderService instance (db can be None for this test)
16+
ps = ProviderService(1, None)
17+
18+
# Test valid JSON string
19+
result = ps._ensure_model_mapping_dict('{"gpt-4": "gpt-4-turbo", "claude": "claude-3-opus"}')
20+
assert result == {"gpt-4": "gpt-4-turbo", "claude": "claude-3-opus"}
21+
22+
# Test empty string
23+
result = ps._ensure_model_mapping_dict("")
24+
assert result == {}
25+
26+
# Test None
27+
result = ps._ensure_model_mapping_dict(None)
28+
assert result == {}
29+
30+
# Test already valid dict
31+
test_dict = {"test": "value"}
32+
result = ps._ensure_model_mapping_dict(test_dict)
33+
assert result == test_dict
34+
assert result is test_dict # Should return the same object
35+
36+
# Test invalid JSON string
37+
result = ps._ensure_model_mapping_dict('{invalid json}')
38+
assert result == {}
39+
40+
# Test malformed JSON string
41+
result = ps._ensure_model_mapping_dict('{"key": "value",}')
42+
assert result == {}
43+
44+
# Test non-string, non-dict input
45+
result = ps._ensure_model_mapping_dict(123)
46+
assert result == {}
47+
48+
result = ps._ensure_model_mapping_dict([])
49+
assert result == {}
50+
51+
@pytest.mark.asyncio
52+
async def test_list_models_with_string_model_mapping(self):
53+
"""Test that list_models works correctly when model_mapping is stored as a string."""
54+
# Mock database session
55+
mock_db = AsyncMock(spec=AsyncSession)
56+
57+
# Mock user
58+
mock_user = MagicMock(spec=User)
59+
mock_user.id = 1
60+
61+
# Create ProviderService instance
62+
ps = ProviderService(1, mock_db)
63+
64+
# Mock the database query to return a provider key with string model_mapping
65+
mock_provider_key = MagicMock(spec=ProviderKey)
66+
mock_provider_key.provider_name = "openai"
67+
mock_provider_key.encrypted_api_key = "encrypted_key"
68+
mock_provider_key.base_url = "https://api.openai.com"
69+
# This simulates old data where model_mapping was stored as a string
70+
mock_provider_key.model_mapping = '{"gpt-4": "gpt-4-turbo"}'
71+
72+
# Mock the database query result
73+
mock_result = MagicMock()
74+
mock_result.scalars.return_value.all.return_value = [mock_provider_key]
75+
mock_db.execute.return_value = mock_result
76+
77+
# Mock the cache to return None (no cached data)
78+
ps._keys_loaded = False
79+
80+
# Mock the provider adapter
81+
mock_adapter = MagicMock()
82+
mock_adapter.list_models = AsyncMock(return_value=["gpt-4", "gpt-3.5-turbo"])
83+
84+
# Mock the adapter factory
85+
with pytest.MonkeyPatch().context() as m:
86+
m.setattr("app.services.provider_service.ProviderAdapterFactory.get_adapter_cls",
87+
lambda x: MagicMock(deserialize_api_key_config=lambda x: ("api_key", {})))
88+
m.setattr("app.services.provider_service.ProviderAdapterFactory.get_adapter",
89+
lambda x, y, z: mock_adapter)
90+
m.setattr("app.services.provider_service.decrypt_api_key", lambda x: "decrypted_key")
91+
m.setattr("app.services.provider_service.async_provider_service_cache.get",
92+
AsyncMock(return_value=None))
93+
m.setattr("app.services.provider_service.async_provider_service_cache.set",
94+
AsyncMock())
95+
96+
# Call list_models - this should not raise an error
97+
result = await ps.list_models()
98+
99+
# Verify the result
100+
assert isinstance(result, list)
101+
assert len(result) == 2 # Two models returned
102+
103+
# Verify the models have the correct structure
104+
for model in result:
105+
assert "id" in model
106+
assert "display_name" in model
107+
assert "object" in model
108+
assert "owned_by" in model
109+
assert model["object"] == "model"
110+
assert model["owned_by"] == "openai"
111+
112+
@pytest.mark.asyncio
113+
async def test_list_models_with_invalid_json_string(self):
114+
"""Test that list_models handles invalid JSON strings gracefully."""
115+
# Mock database session
116+
mock_db = AsyncMock(spec=AsyncSession)
117+
118+
# Mock user
119+
mock_user = MagicMock(spec=User)
120+
mock_user.id = 1
121+
122+
# Create ProviderService instance
123+
ps = ProviderService(1, mock_db)
124+
125+
# Mock the database query to return a provider key with invalid JSON string
126+
mock_provider_key = MagicMock(spec=ProviderKey)
127+
mock_provider_key.provider_name = "openai"
128+
mock_provider_key.encrypted_api_key = "encrypted_key"
129+
mock_provider_key.base_url = "https://api.openai.com"
130+
# This simulates corrupted data
131+
mock_provider_key.model_mapping = '{invalid json string'
132+
133+
# Mock the database query result
134+
mock_result = MagicMock()
135+
mock_result.scalars.return_value.all.return_value = [mock_provider_key]
136+
mock_db.execute.return_value = mock_result
137+
138+
# Mock the cache to return None (no cached data)
139+
ps._keys_loaded = False
140+
141+
# Mock the provider adapter
142+
mock_adapter = MagicMock()
143+
mock_adapter.list_models = AsyncMock(return_value=["gpt-4", "gpt-3.5-turbo"])
144+
145+
# Mock the adapter factory
146+
with pytest.MonkeyPatch().context() as m:
147+
m.setattr("app.services.provider_service.ProviderAdapterFactory.get_adapter_cls",
148+
lambda x: MagicMock(deserialize_api_key_config=lambda x: ("api_key", {})))
149+
m.setattr("app.services.provider_service.ProviderAdapterFactory.get_adapter",
150+
lambda x, y, z: mock_adapter)
151+
m.setattr("app.services.provider_service.decrypt_api_key", lambda x: "decrypted_key")
152+
m.setattr("app.services.provider_service.async_provider_service_cache.get",
153+
AsyncMock(return_value=None))
154+
m.setattr("app.services.provider_service.async_provider_service_cache.set",
155+
AsyncMock())
156+
157+
# Call list_models - this should not raise an error
158+
result = await ps.list_models()
159+
160+
# Verify the result
161+
assert isinstance(result, list)
162+
assert len(result) == 2 # Two models returned
163+
164+
# Since model_mapping was invalid, display_name should be the same as the model name
165+
for model in result:
166+
assert model["display_name"] == model["id"].split("/")[1]
167+
168+
@pytest.mark.asyncio
169+
async def test_list_models_with_none_model_mapping(self):
170+
"""Test that list_models works correctly when model_mapping is None."""
171+
# Mock database session
172+
mock_db = AsyncMock(spec=AsyncSession)
173+
174+
# Mock user
175+
mock_user = MagicMock(spec=User)
176+
mock_user.id = 1
177+
178+
# Create ProviderService instance
179+
ps = ProviderService(1, mock_db)
180+
181+
# Mock the database query to return a provider key with None model_mapping
182+
mock_provider_key = MagicMock(spec=ProviderKey)
183+
mock_provider_key.provider_name = "openai"
184+
mock_provider_key.encrypted_api_key = "encrypted_key"
185+
mock_provider_key.base_url = "https://api.openai.com"
186+
mock_provider_key.model_mapping = None
187+
188+
# Mock the database query result
189+
mock_result = MagicMock()
190+
mock_result.scalars.return_value.all.return_value = [mock_provider_key]
191+
mock_db.execute.return_value = mock_result
192+
193+
# Mock the cache to return None (no cached data)
194+
ps._keys_loaded = False
195+
196+
# Mock the provider adapter
197+
mock_adapter = MagicMock()
198+
mock_adapter.list_models = AsyncMock(return_value=["gpt-4", "gpt-3.5-turbo"])
199+
200+
# Mock the adapter factory
201+
with pytest.MonkeyPatch().context() as m:
202+
m.setattr("app.services.provider_service.ProviderAdapterFactory.get_adapter_cls",
203+
lambda x: MagicMock(deserialize_api_key_config=lambda x: ("api_key", {})))
204+
m.setattr("app.services.provider_service.ProviderAdapterFactory.get_adapter",
205+
lambda x, y, z: mock_adapter)
206+
m.setattr("app.services.provider_service.decrypt_api_key", lambda x: "decrypted_key")
207+
m.setattr("app.services.provider_service.async_provider_service_cache.get",
208+
AsyncMock(return_value=None))
209+
m.setattr("app.services.provider_service.async_provider_service_cache.set",
210+
AsyncMock())
211+
212+
# Call list_models - this should not raise an error
213+
result = await ps.list_models()
214+
215+
# Verify the result
216+
assert isinstance(result, list)
217+
assert len(result) == 2 # Two models returned
218+
219+
# Since model_mapping was None, display_name should be the same as the model name
220+
for model in result:
221+
assert model["display_name"] == model["id"].split("/")[1]
222+
223+
def test_get_provider_info_with_string_model_mapping(self):
224+
"""Test that _get_provider_info_with_prefix works with string model_mapping."""
225+
ps = ProviderService(1, None)
226+
227+
# Mock provider_keys with string model_mapping
228+
ps.provider_keys = {
229+
"openai": {
230+
"api_key": "test_key",
231+
"base_url": "https://api.openai.com",
232+
"model_mapping": '{"custom-gpt": "gpt-4"}'
233+
}
234+
}
235+
ps._keys_loaded = True
236+
237+
# Test that it works correctly
238+
provider_name, mapped_model, base_url = ps._get_provider_info_with_prefix(
239+
"openai", "custom-gpt", "openai/custom-gpt"
240+
)
241+
242+
assert provider_name == "openai"
243+
assert mapped_model == "gpt-4" # Should be mapped correctly
244+
assert base_url == "https://api.openai.com"
245+
246+
def test_find_provider_for_unprefixed_model_with_string_model_mapping(self):
247+
"""Test that _find_provider_for_unprefixed_model works with string model_mapping."""
248+
ps = ProviderService(1, None)
249+
250+
# Mock provider_keys with string model_mapping
251+
ps.provider_keys = {
252+
"openai": {
253+
"api_key": "test_key",
254+
"base_url": "https://api.openai.com",
255+
"model_mapping": '{"custom-gpt": "gpt-4"}'
256+
}
257+
}
258+
ps._keys_loaded = True
259+
260+
# Test that it works correctly
261+
provider_name, mapped_model, base_url = ps._find_provider_for_unprefixed_model("custom-gpt")
262+
263+
assert provider_name == "openai"
264+
assert mapped_model == "gpt-4" # Should be mapped correctly
265+
assert base_url == "https://api.openai.com"
266+
267+
def test_original_error_scenario_prevention(self):
268+
"""Test that the original 'str' object has no attribute 'items' error is prevented."""
269+
ps = ProviderService(1, None)
270+
271+
# Simulate the exact scenario that caused the original error
272+
# This would have caused the error before our fix
273+
provider_data = {
274+
"base_url": "https://api.openai.com",
275+
"model_mapping": '{"gpt-4": "gpt-4-turbo"}' # String instead of dict
276+
}
277+
278+
# This line would have failed before our fix:
279+
# cache_key = f"{base_url}:{hash(frozenset(provider_data.get('model_mapping', {}).items()))}"
280+
# Because provider_data.get('model_mapping', {}) would return a string, and strings don't have .items()
281+
282+
# Now with our fix, this should work:
283+
base_url = provider_data.get("base_url", "default")
284+
model_mapping = ps._ensure_model_mapping_dict(provider_data.get("model_mapping", {}))
285+
cache_key = f"{base_url}:{hash(frozenset(model_mapping.items()))}"
286+
287+
# Verify that no error was raised and we got a valid cache key
288+
assert isinstance(cache_key, str)
289+
assert "https://api.openai.com" in cache_key
290+
assert len(cache_key) > 0
291+
292+
def test_cache_key_generation_with_various_model_mappings(self):
293+
"""Test that cache key generation works with various model_mapping types."""
294+
ps = ProviderService(1, None)
295+
296+
test_cases = [
297+
# (model_mapping, expected_type)
298+
('{"gpt-4": "gpt-4-turbo"}', dict),
299+
('', dict),
300+
(None, dict),
301+
('{invalid json}', dict),
302+
({"valid": "dict"}, dict),
303+
]
304+
305+
for model_mapping, expected_type in test_cases:
306+
result = ps._ensure_model_mapping_dict(model_mapping)
307+
assert isinstance(result, expected_type)
308+
309+
# Test that we can call .items() on the result
310+
items = result.items()
311+
assert hasattr(items, '__iter__') # Should be iterable
312+
313+
# Test cache key generation
314+
base_url = "https://api.openai.com"
315+
cache_key = f"{base_url}:{hash(frozenset(result.items()))}"
316+
assert isinstance(cache_key, str)
317+
assert base_url in cache_key

0 commit comments

Comments
 (0)