mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(proxy): gate v1 team filter and honor key allowlists
Only apply get_all_team_and_direct_access_models for admin or user-bound keys, then intersect with key/team model restrictions to avoid empty lists for service tokens and metadata leaks for restricted keys. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
72a960173a
commit
99c58a944f
2 changed files with 169 additions and 33 deletions
|
|
@ -12405,6 +12405,61 @@ def _deployment_matches_allowed_model_names(
|
|||
)
|
||||
|
||||
|
||||
def _get_v1_model_info_allowed_model_names(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
llm_router: Router,
|
||||
) -> Optional[Set[str]]:
|
||||
"""Return key/team allowlisted public model names, or None if unrestricted."""
|
||||
model_access_groups = llm_router.get_model_access_groups()
|
||||
proxy_model_list = llm_router.get_model_names()
|
||||
key_models = get_key_models(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_model_list=proxy_model_list,
|
||||
model_access_groups=model_access_groups,
|
||||
)
|
||||
team_models = get_team_models(
|
||||
team_models=user_api_key_dict.team_models,
|
||||
proxy_model_list=proxy_model_list,
|
||||
model_access_groups=model_access_groups,
|
||||
)
|
||||
if not key_models and not team_models:
|
||||
return None
|
||||
return set(
|
||||
get_complete_model_list(
|
||||
key_models=key_models,
|
||||
team_models=team_models,
|
||||
proxy_model_list=proxy_model_list,
|
||||
user_model=user_model,
|
||||
infer_model_from_keys=general_settings.get("infer_model_from_keys", False),
|
||||
llm_router=llm_router,
|
||||
return_wildcard_routes=False,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _filter_v1_model_info_deployments(
|
||||
all_models: List[dict],
|
||||
allowed_model_names: Optional[Set[str]],
|
||||
) -> List[dict]:
|
||||
if allowed_model_names is None:
|
||||
return all_models
|
||||
return [
|
||||
model
|
||||
for model in all_models
|
||||
if _deployment_matches_allowed_model_names(model, allowed_model_names)
|
||||
]
|
||||
|
||||
|
||||
def _should_apply_v1_team_access_filter(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> bool:
|
||||
"""Team membership filtering requires a resolvable user or admin role."""
|
||||
return (
|
||||
user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
|
||||
or user_api_key_dict.user_id is not None
|
||||
)
|
||||
|
||||
|
||||
def _translate_model_name_for_response(model: dict) -> dict:
|
||||
"""For team-scoped DB rows, replace `model_name` with the public name
|
||||
in `model_info.team_public_model_name` before returning. The DB column
|
||||
|
|
@ -12579,46 +12634,25 @@ async def model_info_v1( # noqa: PLR0915
|
|||
# use internal routing keys (model_name_{team_id}_{uuid}) and were omitted
|
||||
# when v1 resolved models only via public model_name strings.
|
||||
all_models: List[dict] = copy.deepcopy(llm_router.model_list)
|
||||
allowed_model_names = _get_v1_model_info_allowed_model_names(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
|
||||
if prisma_client is not None:
|
||||
if prisma_client is not None and _should_apply_v1_team_access_filter(
|
||||
user_api_key_dict=user_api_key_dict
|
||||
):
|
||||
all_models = await get_all_team_and_direct_access_models(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
prisma_client=prisma_client,
|
||||
llm_router=llm_router,
|
||||
all_models=all_models,
|
||||
)
|
||||
else:
|
||||
model_access_groups = llm_router.get_model_access_groups()
|
||||
proxy_model_list = llm_router.get_model_names()
|
||||
key_models = get_key_models(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_model_list=proxy_model_list,
|
||||
model_access_groups=model_access_groups,
|
||||
)
|
||||
team_models = get_team_models(
|
||||
team_models=user_api_key_dict.team_models,
|
||||
proxy_model_list=proxy_model_list,
|
||||
model_access_groups=model_access_groups,
|
||||
)
|
||||
if key_models or team_models:
|
||||
allowed_model_names = set(
|
||||
get_complete_model_list(
|
||||
key_models=key_models,
|
||||
team_models=team_models,
|
||||
proxy_model_list=proxy_model_list,
|
||||
user_model=user_model,
|
||||
infer_model_from_keys=general_settings.get(
|
||||
"infer_model_from_keys", False
|
||||
),
|
||||
llm_router=llm_router,
|
||||
return_wildcard_routes=False,
|
||||
)
|
||||
)
|
||||
all_models = [
|
||||
model
|
||||
for model in all_models
|
||||
if _deployment_matches_allowed_model_names(model, allowed_model_names)
|
||||
]
|
||||
|
||||
all_models = _filter_v1_model_info_deployments(
|
||||
all_models=all_models,
|
||||
allowed_model_names=allowed_model_names,
|
||||
)
|
||||
|
||||
all_models = [
|
||||
_translate_model_name_for_response(
|
||||
|
|
|
|||
|
|
@ -212,3 +212,105 @@ async def test_model_info_v1_list_path_translates_team_model_name(monkeypatch):
|
|||
m for m in resp["data"] if m["model_name"] == "team-claude-sonnet"
|
||||
)
|
||||
assert team_model["model_info"]["access_via_team_ids"] == ["team-abc-123"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_info_v1_no_user_id_with_db_skips_team_access_filter(monkeypatch):
|
||||
"""Service/CI keys without user_id must not hit the team-membership filter."""
|
||||
deployment = {
|
||||
"model_name": "gpt-4",
|
||||
"litellm_params": {"model": "gpt-4"},
|
||||
"model_info": {"id": "global-id-1", "db_model": False},
|
||||
}
|
||||
router = MagicMock()
|
||||
router.model_list = [deployment]
|
||||
router.get_model_names.return_value = ["gpt-4"]
|
||||
router.get_model_access_groups.return_value = {}
|
||||
|
||||
get_team_access = AsyncMock()
|
||||
monkeypatch.setattr(ps, "user_model", None)
|
||||
monkeypatch.setattr(ps, "llm_model_list", router.model_list)
|
||||
monkeypatch.setattr(ps, "llm_router", router)
|
||||
monkeypatch.setattr(ps, "prisma_client", MagicMock())
|
||||
monkeypatch.setattr(ps, "get_all_team_and_direct_access_models", get_team_access)
|
||||
monkeypatch.setattr(
|
||||
ps, "_enrich_model_info_with_litellm_data", lambda model, **kw: model
|
||||
)
|
||||
import litellm.proxy.agent_endpoints.model_list_helpers as mlh
|
||||
|
||||
monkeypatch.setattr(
|
||||
mlh,
|
||||
"append_agents_to_model_info",
|
||||
AsyncMock(side_effect=lambda models, **kw: models),
|
||||
)
|
||||
|
||||
caller = UserAPIKeyAuth(
|
||||
user_id=None,
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
models=[],
|
||||
team_models=[],
|
||||
)
|
||||
resp = await ps.model_info_v1(user_api_key_dict=caller, litellm_model_id=None)
|
||||
|
||||
assert [m["model_name"] for m in resp["data"]] == ["gpt-4"]
|
||||
get_team_access.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_info_v1_restricted_key_filters_after_team_enrichment(monkeypatch):
|
||||
"""Key-level model allowlists must apply after DB team-access enrichment."""
|
||||
team_row = _team_row()
|
||||
global_row = {
|
||||
"model_name": "gpt-4",
|
||||
"litellm_params": {"model": "gpt-4"},
|
||||
"model_info": {"id": "global-id-1", "db_model": False},
|
||||
}
|
||||
router = MagicMock()
|
||||
router.model_list = [team_row, global_row]
|
||||
router.get_model_names.return_value = ["gpt-4", "team-claude-sonnet"]
|
||||
router.get_model_access_groups.return_value = {}
|
||||
|
||||
monkeypatch.setattr(ps, "user_model", None)
|
||||
monkeypatch.setattr(ps, "llm_model_list", router.model_list)
|
||||
monkeypatch.setattr(ps, "llm_router", router)
|
||||
monkeypatch.setattr(ps, "prisma_client", MagicMock())
|
||||
monkeypatch.setattr(
|
||||
ps,
|
||||
"get_all_team_and_direct_access_models",
|
||||
AsyncMock(
|
||||
side_effect=lambda all_models, **kwargs: [
|
||||
{
|
||||
**m,
|
||||
"model_info": {
|
||||
**m.get("model_info", {}),
|
||||
**(
|
||||
{"access_via_team_ids": ["team-abc-123"]}
|
||||
if m.get("model_info", {}).get("team_id")
|
||||
else {"direct_access": True}
|
||||
),
|
||||
},
|
||||
}
|
||||
for m in all_models
|
||||
]
|
||||
),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
ps, "_enrich_model_info_with_litellm_data", lambda model, **kw: model
|
||||
)
|
||||
import litellm.proxy.agent_endpoints.model_list_helpers as mlh
|
||||
|
||||
monkeypatch.setattr(
|
||||
mlh,
|
||||
"append_agents_to_model_info",
|
||||
AsyncMock(side_effect=lambda models, **kw: models),
|
||||
)
|
||||
|
||||
caller = UserAPIKeyAuth(
|
||||
user_id="user-1",
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
models=["gpt-4"],
|
||||
team_models=[],
|
||||
)
|
||||
resp = await ps.model_info_v1(user_api_key_dict=caller, litellm_model_id=None)
|
||||
|
||||
assert [m["model_name"] for m in resp["data"]] == ["gpt-4"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue