fix: test cases on team scoped model

This commit is contained in:
Harshit28j 2026-03-05 20:20:33 +05:30
parent 303eb78667
commit 36b13b5332
3 changed files with 456 additions and 48 deletions

View file

@ -126,7 +126,10 @@ def _calculate_key_rotation_time(rotation_interval: str) -> datetime:
def _set_key_rotation_fields(
data: dict, auto_rotate: bool, rotation_interval: Optional[str], existing_key_alias: Optional[str] = None
data: dict,
auto_rotate: bool,
rotation_interval: Optional[str],
existing_key_alias: Optional[str] = None,
) -> None:
"""
Helper function to set rotation fields in key data if auto_rotate is enabled.
@ -948,7 +951,16 @@ async def _validate_key_models_against_effective_team_models(
team_member_models=member_models,
)
# 3. Fallback: if effective models empty but team has models, use team.models
# 3. Cap effective models to team.models (same intersection the runtime
# auth check applies) so the key is never stamped with models the
# runtime would block.
if (
team_table.models
and SpecialModelNames.all_proxy_models.value not in team_table.models
):
effective_models = list(set(effective_models) & set(team_table.models))
# 4. Fallback: if effective models empty but team has models, use team.models
if not effective_models:
if team_table.models:
effective_models = team_table.models
@ -960,7 +972,7 @@ async def _validate_key_models_against_effective_team_models(
},
)
# 4. Step 6b: If data.models is empty, default to effective models
# 5. If data.models is empty, default to effective models
if not data.models:
data.models = effective_models
else:
@ -3156,7 +3168,10 @@ async def delete_verification_tokens(
hashed_token = hash_token(cast(str, key))
user_api_key_cache.delete_cache(hashed_token)
return {"deleted_keys": deleted_tokens, "failed_tokens": failed_tokens}, _keys_being_deleted
return {
"deleted_keys": deleted_tokens,
"failed_tokens": failed_tokens,
}, _keys_being_deleted
def _transform_verification_tokens_to_deleted_records(
@ -3259,7 +3274,7 @@ async def delete_key_aliases(
)
async def _rotate_master_key( # noqa: PLR0915
async def _rotate_master_key( # noqa: PLR0915
prisma_client: PrismaClient,
user_api_key_dict: UserAPIKeyAuth,
current_master_key: str,
@ -3474,6 +3489,8 @@ async def _insert_deprecated_key(
"Failed to insert deprecated key for grace period: %s",
deprecated_err,
)
async def _execute_virtual_key_regeneration(
*,
prisma_client: PrismaClient,
@ -4037,8 +4054,7 @@ def _get_member_team_ids_from_objects(
team.team_id
for team in team_objects
if any(
member.user_id is not None
and member.user_id == user_api_key_dict.user_id
member.user_id is not None and member.user_id == user_api_key_dict.user_id
for member in team.members_with_roles
)
]
@ -4282,9 +4298,7 @@ async def key_aliases(
where_sql = " AND ".join(where_parts)
count_sql = (
f'SELECT COUNT(*) AS count FROM "LiteLLM_VerificationToken" WHERE {where_sql}'
)
count_sql = f'SELECT COUNT(*) AS count FROM "LiteLLM_VerificationToken" WHERE {where_sql}'
count_rows = await prisma_client.db.query_raw(count_sql, *query_params)
total_count = int(count_rows[0]["count"]) if count_rows else 0
@ -4299,7 +4313,9 @@ async def key_aliases(
f" LIMIT ${limit_idx} OFFSET ${offset_idx}"
)
alias_rows = await prisma_client.db.query_raw(aliases_sql, *aliases_params)
aliases: List[str] = [row["key_alias"] for row in alias_rows if row.get("key_alias")]
aliases: List[str] = [
row["key_alias"] for row in alias_rows if row.get("key_alias")
]
total_pages = -(-total_count // size) if total_count > 0 else 0
verbose_proxy_logger.debug(

View file

@ -1744,25 +1744,23 @@ async def _process_team_members(
updated_users: List[LiteLLM_UserTable] = []
updated_team_memberships: List[LiteLLM_TeamMembership] = []
# Always validate member models ⊆ team.models regardless of feature flag.
# This prevents storing out-of-bounds data that could become effective
# if the flag is enabled later.
if data.models is not None:
from litellm.proxy.management_endpoints.common_utils import (
_is_team_model_overrides_enabled,
)
if _is_team_model_overrides_enabled():
if (
complete_team_data.models
and SpecialModelNames.all_proxy_models.value
not in complete_team_data.models
):
invalid = set(data.models) - set(complete_team_data.models)
if invalid:
raise HTTPException(
status_code=400,
detail={
"error": f"Models {list(invalid)} not in team's allowed models: {complete_team_data.models}"
},
)
if (
complete_team_data.models
and SpecialModelNames.all_proxy_models.value
not in complete_team_data.models
):
invalid = set(data.models) - set(complete_team_data.models)
if invalid:
raise HTTPException(
status_code=400,
detail={
"error": f"Models {list(invalid)} not in team's allowed models: {complete_team_data.models}"
},
)
default_team_budget_id = (
complete_team_data.metadata.get("team_member_budget_id")
@ -2426,26 +2424,20 @@ async def team_member_update(
identified_budget_id = tm.budget_id
break
# Always validate member models ⊆ team.models regardless of feature flag.
if data.models is not None:
from litellm.proxy.management_endpoints.common_utils import (
_is_team_model_overrides_enabled,
)
if _is_team_model_overrides_enabled():
# Validate models are within team's allowed set (team.models)
if (
existing_team_row.models
and SpecialModelNames.all_proxy_models.value
not in existing_team_row.models
):
invalid = set(data.models) - set(existing_team_row.models)
if invalid:
raise HTTPException(
status_code=400,
detail={
"error": f"Models {list(invalid)} not in team's allowed models: {existing_team_row.models}"
},
)
if (
existing_team_row.models
and SpecialModelNames.all_proxy_models.value not in existing_team_row.models
):
invalid = set(data.models) - set(existing_team_row.models)
if invalid:
raise HTTPException(
status_code=400,
detail={
"error": f"Models {list(invalid)} not in team's allowed models: {existing_team_row.models}"
},
)
### upsert new budget
async with prisma_client.db.tx() as tx:

View file

@ -0,0 +1,400 @@
"""
Unit tests for team-scoped model overrides.
Tests cover:
- compute_effective_team_models (union logic)
- can_team_access_model with overrides (runtime enforcement)
- default_models team.models validation
- member models team.models validation
- _validate_key_models_against_effective_team_models (key creation)
"""
import os
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
from litellm.proxy._types import (
LiteLLM_TeamTable,
UserAPIKeyAuth,
)
from litellm.proxy.auth.auth_checks import (
can_team_access_model,
compute_effective_team_models,
)
# ── compute_effective_team_models ────────────────────────────────────────────
class TestComputeEffectiveTeamModels:
def test_union_of_defaults_and_member(self):
result = compute_effective_team_models(
team_default_models=["gpt-4o"],
team_member_models=["claude-sonnet"],
)
assert set(result) == {"gpt-4o", "claude-sonnet"}
def test_deduplicates(self):
result = compute_effective_team_models(
team_default_models=["gpt-4o", "claude-sonnet"],
team_member_models=["claude-sonnet"],
)
assert sorted(result) == sorted(["gpt-4o", "claude-sonnet"])
def test_none_defaults(self):
result = compute_effective_team_models(
team_default_models=None,
team_member_models=["claude-sonnet"],
)
assert result == ["claude-sonnet"]
def test_none_member(self):
result = compute_effective_team_models(
team_default_models=["gpt-4o"],
team_member_models=None,
)
assert result == ["gpt-4o"]
def test_both_none(self):
result = compute_effective_team_models(
team_default_models=None,
team_member_models=None,
)
assert result == []
def test_both_empty(self):
result = compute_effective_team_models(
team_default_models=[],
team_member_models=[],
)
assert result == []
# ── can_team_access_model (runtime check) ────────────────────────────────────
class TestCanTeamAccessModelOverrides:
@pytest.mark.asyncio
@patch.dict(os.environ, {"LITELLM_TEAM_MODEL_OVERRIDES": "true"})
async def test_allowed_model_in_defaults(self):
"""Model in default_models should be allowed."""
team_object = LiteLLM_TeamTable(
team_id="team-1",
models=["gpt-4o", "gpt-4o-mini", "claude-sonnet"],
)
valid_token = UserAPIKeyAuth(
token="test-token",
team_id="team-1",
team_default_models=["gpt-4o"],
team_member_models=None,
)
mock_router = MagicMock()
mock_router.get_model_group_info.return_value = None
result = await can_team_access_model(
model="gpt-4o",
team_object=team_object,
llm_router=mock_router,
team_model_aliases=None,
valid_token=valid_token,
)
assert result is True
@pytest.mark.asyncio
@patch.dict(os.environ, {"LITELLM_TEAM_MODEL_OVERRIDES": "true"})
async def test_allowed_model_in_member_override(self):
"""Model in member models should be allowed."""
team_object = LiteLLM_TeamTable(
team_id="team-1",
models=["gpt-4o", "gpt-4o-mini", "claude-sonnet"],
)
valid_token = UserAPIKeyAuth(
token="test-token",
team_id="team-1",
team_default_models=["gpt-4o"],
team_member_models=["claude-sonnet"],
)
mock_router = MagicMock()
mock_router.get_model_group_info.return_value = None
result = await can_team_access_model(
model="claude-sonnet",
team_object=team_object,
llm_router=mock_router,
team_model_aliases=None,
valid_token=valid_token,
)
assert result is True
@pytest.mark.asyncio
@patch.dict(os.environ, {"LITELLM_TEAM_MODEL_OVERRIDES": "true"})
async def test_blocked_model_not_in_effective(self):
"""Model in team.models but NOT in effective models should be blocked."""
team_object = LiteLLM_TeamTable(
team_id="team-1",
models=["gpt-4o", "gpt-4o-mini", "claude-sonnet"],
)
valid_token = UserAPIKeyAuth(
token="test-token",
team_id="team-1",
team_default_models=["gpt-4o"],
team_member_models=["claude-sonnet"],
)
mock_router = MagicMock()
mock_router.get_model_group_info.return_value = None
with pytest.raises(Exception):
await can_team_access_model(
model="gpt-4o-mini",
team_object=team_object,
llm_router=mock_router,
team_model_aliases=None,
valid_token=valid_token,
)
@pytest.mark.asyncio
@patch.dict(os.environ, {"LITELLM_TEAM_MODEL_OVERRIDES": "true"})
async def test_effective_models_intersected_with_team_models(self):
"""Even if default_models has out-of-bounds model, runtime should block it."""
team_object = LiteLLM_TeamTable(
team_id="team-1",
models=["gpt-4o"], # team only allows gpt-4o
)
valid_token = UserAPIKeyAuth(
token="test-token",
team_id="team-1",
team_default_models=[
"gpt-4o",
"claude-sonnet",
], # claude-sonnet is out of bounds
team_member_models=None,
)
mock_router = MagicMock()
mock_router.get_model_group_info.return_value = None
# claude-sonnet should be blocked even though it's in default_models
with pytest.raises(Exception):
await can_team_access_model(
model="claude-sonnet",
team_object=team_object,
llm_router=mock_router,
team_model_aliases=None,
valid_token=valid_token,
)
@pytest.mark.asyncio
@patch.dict(os.environ, {"LITELLM_TEAM_MODEL_OVERRIDES": "false"})
async def test_flag_off_uses_team_models(self):
"""When flag is off, should use team.models as before."""
team_object = LiteLLM_TeamTable(
team_id="team-1",
models=["gpt-4o", "gpt-4o-mini"],
)
valid_token = UserAPIKeyAuth(
token="test-token",
team_id="team-1",
team_default_models=["gpt-4o"],
team_member_models=None,
)
mock_router = MagicMock()
mock_router.get_model_group_info.return_value = None
# gpt-4o-mini should be allowed because flag is off, team.models is used
result = await can_team_access_model(
model="gpt-4o-mini",
team_object=team_object,
llm_router=mock_router,
team_model_aliases=None,
valid_token=valid_token,
)
assert result is True
@pytest.mark.asyncio
@patch.dict(os.environ, {"LITELLM_TEAM_MODEL_OVERRIDES": "true"})
async def test_no_overrides_configured_uses_team_models(self):
"""Team with no default_models/member_models uses team.models unchanged."""
team_object = LiteLLM_TeamTable(
team_id="team-1",
models=["gpt-4o", "gpt-4o-mini"],
)
valid_token = UserAPIKeyAuth(
token="test-token",
team_id="team-1",
team_default_models=None,
team_member_models=None,
)
mock_router = MagicMock()
mock_router.get_model_group_info.return_value = None
result = await can_team_access_model(
model="gpt-4o-mini",
team_object=team_object,
llm_router=mock_router,
team_model_aliases=None,
valid_token=valid_token,
)
assert result is True
# ── _validate_key_models_against_effective_team_models ───────────────────────
class TestValidateKeyModelsAgainstEffective:
@pytest.mark.asyncio
@patch.dict(os.environ, {"LITELLM_TEAM_MODEL_OVERRIDES": "true"})
async def test_empty_data_models_gets_effective(self):
"""Key with no models should inherit effective models."""
from litellm.proxy.management_endpoints.key_management_endpoints import (
_validate_key_models_against_effective_team_models,
)
mock_prisma = MagicMock()
mock_membership = MagicMock()
mock_membership.models = ["claude-sonnet"]
mock_prisma.db.litellm_teammembership.find_unique = AsyncMock(
return_value=mock_membership
)
team_table = MagicMock()
team_table.default_models = ["gpt-4o"]
team_table.models = ["gpt-4o", "claude-sonnet", "gpt-4o-mini"]
data = MagicMock()
data.models = []
await _validate_key_models_against_effective_team_models(
team_id="team-1",
user_id="user-1",
data=data,
team_table=team_table,
prisma_client=mock_prisma,
)
assert set(data.models) == {"gpt-4o", "claude-sonnet"}
@pytest.mark.asyncio
@patch.dict(os.environ, {"LITELLM_TEAM_MODEL_OVERRIDES": "true"})
async def test_effective_models_capped_to_team_models(self):
"""Key effective models should be intersected with team.models."""
from litellm.proxy.management_endpoints.key_management_endpoints import (
_validate_key_models_against_effective_team_models,
)
mock_prisma = MagicMock()
mock_membership = MagicMock()
mock_membership.models = ["claude-sonnet"] # out-of-bounds
mock_prisma.db.litellm_teammembership.find_unique = AsyncMock(
return_value=mock_membership
)
team_table = MagicMock()
team_table.default_models = ["gpt-4o"]
team_table.models = ["gpt-4o"] # team only allows gpt-4o
data = MagicMock()
data.models = []
await _validate_key_models_against_effective_team_models(
team_id="team-1",
user_id="user-1",
data=data,
team_table=team_table,
prisma_client=mock_prisma,
)
# claude-sonnet should be capped out
assert data.models == ["gpt-4o"]
@pytest.mark.asyncio
@patch.dict(os.environ, {"LITELLM_TEAM_MODEL_OVERRIDES": "true"})
async def test_disallowed_model_in_key_raises(self):
"""Key requesting model outside effective set should raise 403."""
from litellm.proxy.management_endpoints.key_management_endpoints import (
_validate_key_models_against_effective_team_models,
)
mock_prisma = MagicMock()
mock_membership = MagicMock()
mock_membership.models = ["claude-sonnet"]
mock_prisma.db.litellm_teammembership.find_unique = AsyncMock(
return_value=mock_membership
)
team_table = MagicMock()
team_table.default_models = ["gpt-4o"]
team_table.models = ["gpt-4o", "claude-sonnet", "gpt-4o-mini"]
data = MagicMock()
data.models = ["gpt-4o-mini"] # not in effective models
with pytest.raises(HTTPException) as exc_info:
await _validate_key_models_against_effective_team_models(
team_id="team-1",
user_id="user-1",
data=data,
team_table=team_table,
prisma_client=mock_prisma,
)
assert exc_info.value.status_code == 403
@pytest.mark.asyncio
@patch.dict(os.environ, {"LITELLM_TEAM_MODEL_OVERRIDES": "true"})
async def test_no_overrides_skips_validation(self):
"""Teams without default_models or member models skip override validation."""
from litellm.proxy.management_endpoints.key_management_endpoints import (
_validate_key_models_against_effective_team_models,
)
mock_prisma = MagicMock()
mock_membership = MagicMock()
mock_membership.models = []
mock_prisma.db.litellm_teammembership.find_unique = AsyncMock(
return_value=mock_membership
)
team_table = MagicMock()
team_table.default_models = []
team_table.models = ["gpt-4o", "gpt-4o-mini"]
data = MagicMock()
data.models = ["gpt-4o-mini"]
# Should return without modifying data.models
await _validate_key_models_against_effective_team_models(
team_id="team-1",
user_id="user-1",
data=data,
team_table=team_table,
prisma_client=mock_prisma,
)
assert data.models == ["gpt-4o-mini"]
@pytest.mark.asyncio
@patch.dict(os.environ, {"LITELLM_TEAM_MODEL_OVERRIDES": "false"})
async def test_flag_off_skips_entirely(self):
"""When flag is off, validation is skipped entirely."""
from litellm.proxy.management_endpoints.key_management_endpoints import (
_validate_key_models_against_effective_team_models,
)
mock_prisma = MagicMock()
team_table = MagicMock()
data = MagicMock()
data.models = ["anything"]
await _validate_key_models_against_effective_team_models(
team_id="team-1",
user_id="user-1",
data=data,
team_table=team_table,
prisma_client=mock_prisma,
)
# Should be unchanged — no validation happened
assert data.models == ["anything"]