mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
test: add unit tests for _get_team_deployments filtering logic
Tests cover: matching deployments, wrong team_id filtering, string-encoded model_info, empty results, invalid model_info, and mixed deployment filtering. Also updates MockPrismaClient.find_many to support the new startswith query. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
parent
8e6300f0bf
commit
566a04126f
1 changed files with 98 additions and 33 deletions
|
|
@ -22,6 +22,7 @@ from litellm.proxy._types import (
|
|||
)
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
ModelManagementAuthChecks,
|
||||
_get_team_deployments,
|
||||
clear_cache,
|
||||
)
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
|
@ -53,44 +54,23 @@ class MockPrismaClient:
|
|||
)
|
||||
return None
|
||||
|
||||
async def find_many(self, where):
|
||||
# Filter sibling deployments by team_id if where clause specifies it
|
||||
async def find_many(self, where=None):
|
||||
# Filter sibling deployments based on where clause
|
||||
if not self.sibling_deployments:
|
||||
return []
|
||||
|
||||
# Extract team_id from where clause if present
|
||||
team_id_filter = None
|
||||
if where and "model_info" in where:
|
||||
model_info_filter = where["model_info"]
|
||||
if isinstance(model_info_filter, dict) and "path" in model_info_filter:
|
||||
if (
|
||||
model_info_filter["path"] == ["team_id"]
|
||||
and "equals" in model_info_filter
|
||||
):
|
||||
team_id_filter = model_info_filter["equals"]
|
||||
results = self.sibling_deployments
|
||||
|
||||
# Filter deployments by team_id if specified
|
||||
if team_id_filter:
|
||||
# Support model_name startswith filter (used by _get_team_deployments)
|
||||
if where and "model_name" in where:
|
||||
model_name_filter = where["model_name"]
|
||||
if isinstance(model_name_filter, dict) and "startswith" in model_name_filter:
|
||||
prefix = model_name_filter["startswith"]
|
||||
results = [
|
||||
d for d in results if d.model_name.startswith(prefix)
|
||||
]
|
||||
|
||||
def _get_team_id(model_info):
|
||||
if isinstance(model_info, dict):
|
||||
return model_info.get("team_id")
|
||||
if isinstance(model_info, str):
|
||||
try:
|
||||
parsed = json.loads(model_info)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
if isinstance(parsed, dict):
|
||||
return parsed.get("team_id")
|
||||
return None
|
||||
|
||||
return [
|
||||
d
|
||||
for d in self.sibling_deployments
|
||||
if _get_team_id(d.model_info) == team_id_filter
|
||||
]
|
||||
|
||||
return self.sibling_deployments
|
||||
return results
|
||||
|
||||
@property
|
||||
def litellm_teamtable(self):
|
||||
|
|
@ -1276,3 +1256,88 @@ class TestAddAndDeleteModelLifecycle:
|
|||
user_api_key_dict=admin_user,
|
||||
)
|
||||
assert str(exc_info.value.code) == "400"
|
||||
|
||||
|
||||
class TestGetTeamDeployments:
|
||||
"""Tests for _get_team_deployments which filters by model_name prefix + Python-side team_id check."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_matching_team_deployments(self):
|
||||
"""Deployments with matching model_name prefix and team_id are returned."""
|
||||
team_id = "team_abc"
|
||||
dep = MagicMock()
|
||||
dep.model_name = f"model_name_{team_id}_uuid1"
|
||||
dep.model_info = {"team_id": team_id, "team_public_model_name": "gpt-4"}
|
||||
|
||||
prisma_client = MockPrismaClient(sibling_deployments=[dep])
|
||||
result = await _get_team_deployments(team_id, prisma_client)
|
||||
assert len(result) == 1
|
||||
assert result[0] is dep
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_filters_out_wrong_team_id_in_model_info(self):
|
||||
"""A deployment whose model_name matches but model_info.team_id differs is excluded."""
|
||||
team_id = "team_abc"
|
||||
dep = MagicMock()
|
||||
dep.model_name = f"model_name_{team_id}_uuid1"
|
||||
dep.model_info = {"team_id": "other_team"}
|
||||
|
||||
prisma_client = MockPrismaClient(sibling_deployments=[dep])
|
||||
result = await _get_team_deployments(team_id, prisma_client)
|
||||
assert len(result) == 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_handles_string_encoded_model_info(self):
|
||||
"""Legacy rows with JSON-string model_info are parsed and filtered correctly."""
|
||||
team_id = "team_abc"
|
||||
dep = MagicMock()
|
||||
dep.model_name = f"model_name_{team_id}_uuid1"
|
||||
dep.model_info = json.dumps({"team_id": team_id})
|
||||
|
||||
prisma_client = MockPrismaClient(sibling_deployments=[dep])
|
||||
result = await _get_team_deployments(team_id, prisma_client)
|
||||
assert len(result) == 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_empty_when_no_deployments(self):
|
||||
"""Returns empty list when no deployments exist."""
|
||||
prisma_client = MockPrismaClient(sibling_deployments=[])
|
||||
result = await _get_team_deployments("team_abc", prisma_client)
|
||||
assert result == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_skips_rows_with_invalid_model_info(self):
|
||||
"""Rows with non-dict, non-parseable model_info are skipped."""
|
||||
team_id = "team_abc"
|
||||
dep = MagicMock()
|
||||
dep.model_name = f"model_name_{team_id}_uuid1"
|
||||
dep.model_info = "not-valid-json"
|
||||
|
||||
prisma_client = MockPrismaClient(sibling_deployments=[dep])
|
||||
result = await _get_team_deployments(team_id, prisma_client)
|
||||
assert len(result) == 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_multiple_deployments_mixed_filtering(self):
|
||||
"""Only deployments with correct prefix AND team_id are returned."""
|
||||
team_id = "team_abc"
|
||||
|
||||
# Matches both prefix and team_id
|
||||
dep1 = MagicMock()
|
||||
dep1.model_name = f"model_name_{team_id}_uuid1"
|
||||
dep1.model_info = {"team_id": team_id}
|
||||
|
||||
# Matches prefix but wrong team_id
|
||||
dep2 = MagicMock()
|
||||
dep2.model_name = f"model_name_{team_id}_uuid2"
|
||||
dep2.model_info = {"team_id": "wrong_team"}
|
||||
|
||||
# Different prefix entirely (won't be returned by mock's startswith filter)
|
||||
dep3 = MagicMock()
|
||||
dep3.model_name = "model_name_other_team_uuid3"
|
||||
dep3.model_info = {"team_id": "other_team"}
|
||||
|
||||
prisma_client = MockPrismaClient(sibling_deployments=[dep1, dep2, dep3])
|
||||
result = await _get_team_deployments(team_id, prisma_client)
|
||||
assert len(result) == 1
|
||||
assert result[0] is dep1
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue