mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-27 01:22:18 +00:00
fix(proxy): authorize key model aliases the same way as team aliases (#43049)
This commit is contained in:
parent
c7f15709f9
commit
1d51a8dfc3
6 changed files with 564 additions and 17 deletions
|
|
@ -89,6 +89,7 @@ from litellm.proxy.common_utils.http_parsing_utils import (
|
|||
_safe_get_request_headers,
|
||||
_safe_get_request_query_params,
|
||||
)
|
||||
from litellm.proxy.common_utils.model_listing_utils import alias_map
|
||||
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
|
||||
from litellm.proxy.common_utils.user_api_key_cache import (
|
||||
END_USER_RESTRICTED_REGISTRY_OVERFLOW_SENTINEL,
|
||||
|
|
@ -722,6 +723,7 @@ async def _run_project_checks(
|
|||
model=_model,
|
||||
project_object=project_object,
|
||||
llm_router=llm_router,
|
||||
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
|
||||
)
|
||||
|
||||
if not skip_budget_checks:
|
||||
|
|
@ -1018,6 +1020,7 @@ async def common_checks(
|
|||
team_object=team_object,
|
||||
llm_router=llm_router,
|
||||
team_model_aliases=(valid_token.team_model_aliases if valid_token else None),
|
||||
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
|
||||
)
|
||||
except ProxyException as team_denial:
|
||||
if team_denial.type != ProxyErrorTypes.team_model_access_denied:
|
||||
|
|
@ -1027,6 +1030,7 @@ async def common_checks(
|
|||
valid_token=valid_token,
|
||||
team_object=team_object,
|
||||
llm_router=llm_router,
|
||||
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
|
||||
):
|
||||
raise
|
||||
|
||||
|
|
@ -1043,6 +1047,7 @@ async def common_checks(
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
team_membership=loaded_team_membership,
|
||||
team_membership_loaded=team_membership_loaded,
|
||||
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
|
||||
)
|
||||
|
||||
# Require trace id for agent keys when agent has require_trace_id_on_calls_by_agent
|
||||
|
|
@ -1081,6 +1086,7 @@ async def common_checks(
|
|||
model=_model,
|
||||
llm_router=llm_router,
|
||||
user_object=user_object,
|
||||
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
|
||||
)
|
||||
|
||||
# 1.1 - 2.2 - 3.0.2 - 3.0.3: Project checks (blocked, model access, budget)
|
||||
|
|
@ -4349,6 +4355,7 @@ def _can_object_call_model(
|
|||
models: list[str],
|
||||
team_model_aliases: dict[str, str] | None = None,
|
||||
team_id: str | None = None,
|
||||
key_model_aliases: Mapping[str, str] | None = None,
|
||||
object_type: Literal["user", "team", "key", "org", "project", "agent"] = "user",
|
||||
fallback_depth: int = 0,
|
||||
) -> Literal[True]:
|
||||
|
|
@ -4378,6 +4385,7 @@ def _can_object_call_model(
|
|||
models=models,
|
||||
team_model_aliases=team_model_aliases,
|
||||
team_id=team_id,
|
||||
key_model_aliases=key_model_aliases,
|
||||
object_type=object_type,
|
||||
fallback_depth=fallback_depth + 1,
|
||||
)
|
||||
|
|
@ -4386,13 +4394,32 @@ def _can_object_call_model(
|
|||
from litellm.router_strategy.complexity_router.context_compaction import native_compaction_parent
|
||||
|
||||
compaction_parent: Final = native_compaction_parent(model)
|
||||
potential_models: Final = [model, compaction_parent] if compaction_parent is not None else [model]
|
||||
if model in litellm.model_alias_map:
|
||||
potential_models.append(litellm.model_alias_map[model])
|
||||
elif llm_router and model in llm_router.model_group_alias:
|
||||
_model: Final = llm_router._get_model_from_alias(model)
|
||||
if _model:
|
||||
potential_models.append(_model)
|
||||
global_or_router_alias_target: Final = (
|
||||
litellm.model_alias_map[model]
|
||||
if model in litellm.model_alias_map
|
||||
else (
|
||||
llm_router._get_model_from_alias(model)
|
||||
if llm_router is not None and model in llm_router.model_group_alias
|
||||
else None
|
||||
)
|
||||
)
|
||||
after_team_alias: Final = team_model_aliases.get(model, model) if team_model_aliases else model
|
||||
after_key_alias: Final = (
|
||||
key_model_aliases.get(after_team_alias, after_team_alias) if key_model_aliases else after_team_alias
|
||||
)
|
||||
after_global_alias: Final = litellm.model_alias_map.get(after_key_alias, after_key_alias)
|
||||
dispatched_model: Final = (
|
||||
key_model_aliases.get(after_global_alias, after_global_alias) if key_model_aliases else after_global_alias
|
||||
)
|
||||
key_alias_applied: Final = after_key_alias != after_team_alias or dispatched_model != after_global_alias
|
||||
potential_models: Final = (
|
||||
(dispatched_model,)
|
||||
if key_alias_applied
|
||||
else (
|
||||
*((model, compaction_parent) if compaction_parent is not None else (model,)),
|
||||
*((global_or_router_alias_target,) if global_or_router_alias_target else ()),
|
||||
)
|
||||
)
|
||||
|
||||
## check model access for alias + underlying model - allow if either is in allowed models
|
||||
for m in potential_models:
|
||||
|
|
@ -4418,6 +4445,35 @@ def _can_object_call_model(
|
|||
)
|
||||
|
||||
|
||||
def _resolve_team_alias(
|
||||
model: str | list[str],
|
||||
team_model_aliases: dict[str, str] | None,
|
||||
team_id: str | None,
|
||||
llm_router: Router | None,
|
||||
) -> str | list[str]:
|
||||
if not team_model_aliases:
|
||||
return model
|
||||
if isinstance(model, str):
|
||||
return _live_team_alias_target(model, team_model_aliases, team_id, llm_router)
|
||||
return [ # mutable-ok: _can_object_call_model takes list[str]
|
||||
_live_team_alias_target(name, team_model_aliases, team_id, llm_router) for name in model
|
||||
]
|
||||
|
||||
|
||||
def _live_team_alias_target(
|
||||
model: str, team_model_aliases: dict[str, str], team_id: str | None, llm_router: Router | None
|
||||
) -> str:
|
||||
target: Final = team_model_aliases.get(model)
|
||||
if target is None:
|
||||
return model
|
||||
deleted_team_deployment: Final = (
|
||||
llm_router is not None
|
||||
and target.startswith(f"model_name_{team_id}_")
|
||||
and target not in llm_router.model_name_to_deployment_indices
|
||||
)
|
||||
return model if deleted_team_deployment else target
|
||||
|
||||
|
||||
async def _check_agent_access_group_model_access(
|
||||
model: str | list[str] | None, # mutable-ok: _can_object_call_model and the client message helper take list[str]
|
||||
valid_token: UserAPIKeyAuth | None,
|
||||
|
|
@ -4438,12 +4494,14 @@ async def _check_agent_access_group_model_access(
|
|||
param="model",
|
||||
code=status.HTTP_403_FORBIDDEN,
|
||||
)
|
||||
dispatched: Final = _resolve_team_alias(model, valid_token.team_model_aliases, valid_token.team_id, llm_router)
|
||||
return _can_object_call_model(
|
||||
model=model,
|
||||
model=dispatched,
|
||||
llm_router=llm_router,
|
||||
models=sorted(ceiling.models),
|
||||
team_id=valid_token.team_id,
|
||||
object_type="agent",
|
||||
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -4471,12 +4529,14 @@ async def _check_agent_caller_model_access(
|
|||
if caller_auth is None:
|
||||
return
|
||||
caller_team: Final = await load_team(valid_token)
|
||||
caller_key_model_aliases: Final = key_model_aliases_for_auth_check(valid_token)
|
||||
if caller_team is not None:
|
||||
await can_team_access_model(
|
||||
model=model,
|
||||
team_object=caller_team,
|
||||
llm_router=llm_router,
|
||||
prisma_client=prisma_client,
|
||||
key_model_aliases=caller_key_model_aliases,
|
||||
)
|
||||
await _check_team_member_model_access(
|
||||
model=model,
|
||||
|
|
@ -4486,12 +4546,18 @@ async def _check_agent_caller_model_access(
|
|||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
key_model_aliases=caller_key_model_aliases,
|
||||
)
|
||||
return
|
||||
caller_user: Final = await load_user(valid_token)
|
||||
if caller_user is None:
|
||||
return
|
||||
await can_user_call_model(model=model, llm_router=llm_router, user_object=caller_user)
|
||||
await can_user_call_model(
|
||||
model=model,
|
||||
llm_router=llm_router,
|
||||
user_object=caller_user,
|
||||
key_model_aliases=caller_key_model_aliases,
|
||||
)
|
||||
|
||||
|
||||
def _model_in_team_aliases(model: str, team_model_aliases: dict[str, str] | None = None) -> bool:
|
||||
|
|
@ -4512,6 +4578,10 @@ def _model_in_team_aliases(model: str, team_model_aliases: dict[str, str] | None
|
|||
return False
|
||||
|
||||
|
||||
def key_model_aliases_for_auth_check(valid_token: UserAPIKeyAuth | None) -> Mapping[str, str] | None:
|
||||
return alias_map(valid_token.aliases) if valid_token is not None and valid_token.aliases else None
|
||||
|
||||
|
||||
def _resolve_key_models_for_auth_check(valid_token: UserAPIKeyAuth) -> list[str]:
|
||||
"""
|
||||
Expand key model sentinels before auth checks.
|
||||
|
|
@ -4831,6 +4901,7 @@ async def can_key_call_model(
|
|||
models=key_models,
|
||||
team_model_aliases=valid_token.team_model_aliases,
|
||||
team_id=valid_token.team_id,
|
||||
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
|
||||
object_type="key",
|
||||
)
|
||||
except ProxyException:
|
||||
|
|
@ -4848,6 +4919,7 @@ async def can_key_call_model(
|
|||
models=models_from_groups,
|
||||
team_model_aliases=valid_token.team_model_aliases,
|
||||
team_id=valid_token.team_id,
|
||||
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
|
||||
object_type="key",
|
||||
)
|
||||
raise
|
||||
|
|
@ -4906,6 +4978,7 @@ async def can_key_call_resolved_model(
|
|||
team_object=team_object,
|
||||
llm_router=llm_router,
|
||||
team_model_aliases=valid_token.team_model_aliases,
|
||||
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
|
||||
)
|
||||
except ProxyException as team_denial:
|
||||
if team_denial.type != ProxyErrorTypes.team_model_access_denied:
|
||||
|
|
@ -4915,6 +4988,7 @@ async def can_key_call_resolved_model(
|
|||
valid_token=valid_token,
|
||||
team_object=team_object,
|
||||
llm_router=llm_router,
|
||||
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
|
||||
):
|
||||
raise
|
||||
|
||||
|
|
@ -4927,6 +5001,7 @@ async def can_key_call_resolved_model(
|
|||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
|
||||
)
|
||||
|
||||
if valid_token.project_id is not None:
|
||||
|
|
@ -4941,6 +5016,7 @@ async def can_key_call_resolved_model(
|
|||
model=model,
|
||||
project_object=project_object,
|
||||
llm_router=llm_router,
|
||||
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -4968,6 +5044,7 @@ async def can_team_access_model(
|
|||
team_object: LiteLLM_TeamTable | None,
|
||||
llm_router: Router | None,
|
||||
team_model_aliases: dict[str, str] | None = None,
|
||||
key_model_aliases: Mapping[str, str] | None = None,
|
||||
prisma_client: DatabaseClient | None = None,
|
||||
) -> Literal[True]:
|
||||
"""
|
||||
|
|
@ -4983,6 +5060,7 @@ async def can_team_access_model(
|
|||
models=team_object.models if team_object else [],
|
||||
team_model_aliases=team_model_aliases,
|
||||
team_id=team_object.team_id if team_object else None,
|
||||
key_model_aliases=key_model_aliases,
|
||||
object_type="team",
|
||||
)
|
||||
except ProxyException:
|
||||
|
|
@ -5000,6 +5078,7 @@ async def can_team_access_model(
|
|||
models=list(dict.fromkeys([*(team_object.models if team_object else []), *models_from_groups])),
|
||||
team_model_aliases=team_model_aliases,
|
||||
team_id=team_object.team_id if team_object else None,
|
||||
key_model_aliases=key_model_aliases,
|
||||
object_type="team",
|
||||
)
|
||||
raise
|
||||
|
|
@ -5058,6 +5137,7 @@ async def _key_access_group_grants_model(
|
|||
valid_token: UserAPIKeyAuth | None,
|
||||
team_object: LiteLLM_TeamTable | None,
|
||||
llm_router: Router | None,
|
||||
key_model_aliases: Mapping[str, str] | None = None,
|
||||
) -> bool:
|
||||
"""
|
||||
Returns True if the key's `access_group_ids` expand to models that grant
|
||||
|
|
@ -5078,6 +5158,7 @@ async def _key_access_group_grants_model(
|
|||
models=authorized_models,
|
||||
team_model_aliases=valid_token.team_model_aliases if valid_token else None,
|
||||
team_id=valid_token.team_id if valid_token else None,
|
||||
key_model_aliases=key_model_aliases,
|
||||
object_type="key",
|
||||
)
|
||||
return True
|
||||
|
|
@ -5089,6 +5170,7 @@ def can_project_access_model(
|
|||
model: str | list[str],
|
||||
project_object: LiteLLM_ProjectTable,
|
||||
llm_router: Router | None,
|
||||
key_model_aliases: Mapping[str, str] | None = None,
|
||||
) -> Literal[True]:
|
||||
"""
|
||||
Returns True if the project can access a specific model.
|
||||
|
|
@ -5099,6 +5181,7 @@ def can_project_access_model(
|
|||
model=model,
|
||||
llm_router=llm_router,
|
||||
models=project_object.models if project_object else [],
|
||||
key_model_aliases=key_model_aliases,
|
||||
object_type="project",
|
||||
)
|
||||
|
||||
|
|
@ -5107,6 +5190,7 @@ async def can_user_call_model(
|
|||
model: str | list[str],
|
||||
llm_router: Router | None,
|
||||
user_object: LiteLLM_UserTable | None,
|
||||
key_model_aliases: Mapping[str, str] | None = None,
|
||||
) -> Literal[True]:
|
||||
if user_object is None:
|
||||
return True
|
||||
|
|
@ -5128,6 +5212,7 @@ async def can_user_call_model(
|
|||
model=model,
|
||||
llm_router=llm_router,
|
||||
models=user_object.models,
|
||||
key_model_aliases=key_model_aliases,
|
||||
object_type="user",
|
||||
)
|
||||
|
||||
|
|
@ -5682,6 +5767,7 @@ async def _check_team_member_model_access(
|
|||
proxy_logging_obj: ProxyLogging,
|
||||
team_membership: LiteLLM_TeamMembership | None = None,
|
||||
team_membership_loaded: bool = False,
|
||||
key_model_aliases: Mapping[str, str] | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Check if a team member's per-member model scope allows access to the requested model.
|
||||
|
|
@ -5717,6 +5803,7 @@ async def _check_team_member_model_access(
|
|||
models=member_allowed_models,
|
||||
object_type="team",
|
||||
team_id=team_object.team_id,
|
||||
key_model_aliases=key_model_aliases,
|
||||
)
|
||||
except ProxyException:
|
||||
internal_message: Final = (
|
||||
|
|
|
|||
|
|
@ -66,6 +66,7 @@ from litellm.proxy.auth.auth_checks import (
|
|||
get_user_object,
|
||||
is_valid_fallback_model,
|
||||
jwt_key_mapping_cache_key,
|
||||
key_model_aliases_for_auth_check,
|
||||
resolve_and_validate_end_user_id,
|
||||
resolve_default_end_user_budget,
|
||||
)
|
||||
|
|
@ -469,6 +470,7 @@ async def _check_key_model_budget_with_fallback(
|
|||
models=valid_token.team_models,
|
||||
team_model_aliases=valid_token.team_model_aliases,
|
||||
team_id=valid_token.team_id,
|
||||
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
|
||||
object_type="team",
|
||||
)
|
||||
except ProxyException:
|
||||
|
|
|
|||
|
|
@ -180,7 +180,7 @@ def caller_alias_maps(
|
|||
return CallerAliases((team_aliases, key_aliases), (team_aliases, key_aliases, litellm.model_alias_map, key_aliases))
|
||||
|
||||
|
||||
def _alias_map(aliases: object) -> Mapping[str, str]:
|
||||
def alias_map(aliases: object) -> Mapping[str, str]:
|
||||
try:
|
||||
entries: Final = _ALIAS_ENTRIES.validate_python(aliases, strict=True)
|
||||
except ValidationError:
|
||||
|
|
@ -204,7 +204,7 @@ def alias_target(model_id: str, aliases: CallerAliases, listed: Container[str] =
|
|||
already `listed` keeps its own row, so it is never rewritten."""
|
||||
if model_id in listed:
|
||||
return None
|
||||
return _rewrite(model_id, tuple(_alias_map(alias_map) for alias_map in aliases.rewrite))
|
||||
return _rewrite(model_id, tuple(alias_map(raw) for raw in aliases.rewrite))
|
||||
|
||||
|
||||
def alias_listing_entries(
|
||||
|
|
@ -213,8 +213,8 @@ def alias_listing_entries(
|
|||
) -> tuple[tuple[str, str], ...]:
|
||||
"""`entries` plus one `(alias, lookup_id)` row per key or team alias whose target is
|
||||
listed. An alias colliding with a listed id keeps the listed entry."""
|
||||
maps: Final = tuple(_alias_map(alias_map) for alias_map in aliases.rewrite)
|
||||
own: Final = tuple(_alias_map(alias_map) for alias_map in aliases.own)
|
||||
maps: Final = tuple(alias_map(raw) for raw in aliases.rewrite)
|
||||
own: Final = tuple(alias_map(raw) for raw in aliases.own)
|
||||
lookup_by_response: Final = MappingProxyType(dict(entries))
|
||||
lookup_ids: Final = frozenset(lookup_by_response.values())
|
||||
targets: Final = MappingProxyType(
|
||||
|
|
|
|||
|
|
@ -0,0 +1,72 @@
|
|||
import uuid
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
|
||||
from tests.integration._support.client import Gateway, eventually, object_value, string_value
|
||||
|
||||
|
||||
def _listed_model_ids(response: httpx.Response) -> frozenset[str]:
|
||||
entries: Final = response.json()["data"]
|
||||
assert isinstance(entries, list), response.text
|
||||
return frozenset(string_value(object_value(entry)["id"]) for entry in entries)
|
||||
|
||||
|
||||
def _listed_and_callable(gateway: Gateway, key: str, model: str, alias: str) -> None:
|
||||
"""Every id /v1/models lists for this key must be callable by the same key."""
|
||||
response: Final = eventually(
|
||||
lambda: gateway.request("GET", "/v1/models", key=key),
|
||||
lambda value: value.status_code == 200 and _listed_model_ids(value) == frozenset({model, alias}),
|
||||
return_last_on_timeout=True,
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
listed: Final = _listed_model_ids(response)
|
||||
assert listed == frozenset({model, alias}), response.text
|
||||
for model_id in sorted(listed):
|
||||
called: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{"model": model_id, "messages": [{"role": "user", "content": "ping"}]},
|
||||
key=key,
|
||||
)
|
||||
assert called.status_code == 200, f"listed id {model_id} is not callable: {called.status_code} {called.text}"
|
||||
|
||||
|
||||
def test_key_alias_listed_by_v1_models_is_callable(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
model: Final = scenario.model()
|
||||
alias: Final = f"integration-alias-{uuid.uuid4().hex}"
|
||||
key: Final = scenario.key(models=[model], aliases={alias: model})
|
||||
_listed_and_callable(gateway, key, model, alias)
|
||||
|
||||
|
||||
def test_team_key_alias_listed_by_v1_models_is_callable(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
model: Final = scenario.model()
|
||||
team_id: Final = scenario.team(models=[model])
|
||||
alias: Final = f"integration-alias-{uuid.uuid4().hex}"
|
||||
key: Final = scenario.key(team_id=team_id, aliases={alias: model})
|
||||
_listed_and_callable(gateway, key, model, alias)
|
||||
|
||||
|
||||
def test_key_alias_to_model_outside_key_allowlist_is_hidden_and_denied(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
allowed: Final = scenario.model()
|
||||
hidden: Final = scenario.model()
|
||||
alias: Final = f"integration-alias-{uuid.uuid4().hex}"
|
||||
key: Final = scenario.key(models=[allowed], aliases={alias: hidden})
|
||||
response: Final = eventually(
|
||||
lambda: gateway.request("GET", "/v1/models", key=key),
|
||||
lambda value: value.status_code == 200 and _listed_model_ids(value) == frozenset({allowed}),
|
||||
return_last_on_timeout=True,
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert _listed_model_ids(response) == frozenset({allowed}), response.text
|
||||
called: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{"model": alias, "messages": [{"role": "user", "content": "ping"}]},
|
||||
key=key,
|
||||
)
|
||||
assert called.status_code == 403, called.text
|
||||
assert "key_model_access_denied" in called.text, called.text
|
||||
|
|
@ -834,12 +834,12 @@ async def test_key_model_list(model_access, model_access_level, model_endpoint):
|
|||
assert len(model_list["data"]) > 0
|
||||
if model_access == "gpt-3.5-turbo":
|
||||
if model_endpoint == "/v1/models":
|
||||
assert (
|
||||
len(model_list["data"]) == 1
|
||||
), "model_access={}, model_access_level={}".format(
|
||||
assert {entry["id"] for entry in model_list["data"]} == {
|
||||
model_access,
|
||||
"mistral-7b",
|
||||
}, "generate_key sets alias mistral-7b -> gpt-3.5-turbo, so /v1/models lists both; model_access={}, model_access_level={}".format(
|
||||
model_access, model_access_level
|
||||
)
|
||||
assert model_list["data"][0]["id"] == model_access
|
||||
elif model_endpoint == "/model/info":
|
||||
assert isinstance(model_list["data"], list)
|
||||
assert len(model_list["data"]) == 1
|
||||
|
|
|
|||
|
|
@ -1783,6 +1783,336 @@ def test_can_object_call_model_access_via_alias_only():
|
|||
assert result is True
|
||||
|
||||
|
||||
def test_can_object_call_model_key_alias_to_allowed_target_is_allowed():
|
||||
"""A key alias whose target is on the key allowlist resolves like a team alias."""
|
||||
from litellm.proxy.auth.auth_checks import _can_object_call_model
|
||||
|
||||
result = _can_object_call_model(
|
||||
model="mistral-7b",
|
||||
llm_router=None,
|
||||
models=["gpt-4o-mini"],
|
||||
key_model_aliases={"mistral-7b": "gpt-4o-mini"},
|
||||
object_type="key",
|
||||
fallback_depth=0,
|
||||
)
|
||||
|
||||
assert result is True
|
||||
|
||||
|
||||
def test_can_object_call_model_key_alias_to_disallowed_target_is_denied():
|
||||
"""A key alias whose target is outside the key allowlist stays denied."""
|
||||
from litellm.proxy._types import ProxyErrorTypes, ProxyException
|
||||
from litellm.proxy.auth.auth_checks import _can_object_call_model
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
_can_object_call_model(
|
||||
model="mistral-7b",
|
||||
llm_router=None,
|
||||
models=["gpt-4o-mini"],
|
||||
key_model_aliases={"mistral-7b": "gpt-4"},
|
||||
object_type="key",
|
||||
fallback_depth=0,
|
||||
)
|
||||
|
||||
assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied
|
||||
assert exc_info.value.code == "403"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_can_team_access_model_honors_key_alias():
|
||||
"""A key on a team can call a model through its own alias when the target is on the team allowlist."""
|
||||
from litellm.proxy.auth.auth_checks import can_team_access_model
|
||||
|
||||
team_object = LiteLLM_TeamTable(
|
||||
team_id="team-123",
|
||||
models=["gpt-4o-mini"],
|
||||
)
|
||||
|
||||
assert (
|
||||
await can_team_access_model(
|
||||
model="mistral-7b",
|
||||
team_object=team_object,
|
||||
llm_router=None,
|
||||
key_model_aliases={"mistral-7b": "gpt-4o-mini"},
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await can_team_access_model(
|
||||
model="mistral-7b",
|
||||
team_object=team_object,
|
||||
llm_router=None,
|
||||
)
|
||||
|
||||
assert exc_info.value.type == ProxyErrorTypes.team_model_access_denied
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_can_key_call_model_honors_key_alias():
|
||||
"""The real key entry point resolves a key alias to its target before the allowlist check."""
|
||||
from litellm.proxy.auth.auth_checks import can_key_call_model
|
||||
|
||||
allowed_token = UserAPIKeyAuth(
|
||||
api_key="sk-test",
|
||||
models=["gpt-4o-mini"],
|
||||
aliases={"mistral-7b": "gpt-4o-mini"},
|
||||
)
|
||||
|
||||
assert (
|
||||
await can_key_call_model(
|
||||
model="mistral-7b",
|
||||
llm_model_list=None,
|
||||
valid_token=allowed_token,
|
||||
llm_router=None,
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
denied_token = UserAPIKeyAuth(
|
||||
api_key="sk-test",
|
||||
models=["gpt-4o-mini"],
|
||||
aliases={"mistral-7b": "gpt-4"},
|
||||
)
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await can_key_call_model(
|
||||
model="mistral-7b",
|
||||
llm_model_list=None,
|
||||
valid_token=denied_token,
|
||||
llm_router=None,
|
||||
)
|
||||
|
||||
assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied
|
||||
|
||||
|
||||
def test_can_object_call_model_key_alias_applies_before_global_alias(monkeypatch):
|
||||
"""The key alias rewrite precedes the global one at dispatch, so the key target is authorized."""
|
||||
from litellm.proxy.auth.auth_checks import _can_object_call_model
|
||||
|
||||
monkeypatch.setattr(litellm, "model_alias_map", {"foo": "bar"})
|
||||
|
||||
assert (
|
||||
_can_object_call_model(
|
||||
model="foo",
|
||||
llm_router=None,
|
||||
models=["baz"],
|
||||
key_model_aliases={"foo": "baz"},
|
||||
object_type="key",
|
||||
fallback_depth=0,
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
_can_object_call_model(
|
||||
model="foo",
|
||||
llm_router=None,
|
||||
models=["bar"],
|
||||
key_model_aliases={"foo": "baz"},
|
||||
object_type="key",
|
||||
fallback_depth=0,
|
||||
)
|
||||
|
||||
assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied
|
||||
|
||||
|
||||
def test_can_object_call_model_key_alias_matches_global_rewritten_name(monkeypatch):
|
||||
"""A key alias on the globally rewritten name resolves the same way the request chain does."""
|
||||
from litellm.proxy.auth.auth_checks import _can_object_call_model
|
||||
|
||||
monkeypatch.setattr(litellm, "model_alias_map", {"foo": "bar"})
|
||||
|
||||
assert (
|
||||
_can_object_call_model(
|
||||
model="foo",
|
||||
llm_router=None,
|
||||
models=["baz"],
|
||||
key_model_aliases={"bar": "baz"},
|
||||
object_type="key",
|
||||
fallback_depth=0,
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
|
||||
def test_can_object_call_model_chained_alias_requires_final_target(monkeypatch):
|
||||
"""When a key alias fires on the globally rewritten name, only the final target is dispatched."""
|
||||
from litellm.proxy.auth.auth_checks import _can_object_call_model
|
||||
|
||||
monkeypatch.setattr(litellm, "model_alias_map", {"foo": "bar"})
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
_can_object_call_model(
|
||||
model="foo",
|
||||
llm_router=None,
|
||||
models=["bar"],
|
||||
key_model_aliases={"bar": "baz"},
|
||||
object_type="key",
|
||||
fallback_depth=0,
|
||||
)
|
||||
|
||||
assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied
|
||||
|
||||
assert (
|
||||
_can_object_call_model(
|
||||
model="foo",
|
||||
llm_router=None,
|
||||
models=["baz"],
|
||||
key_model_aliases={"bar": "baz"},
|
||||
object_type="key",
|
||||
fallback_depth=0,
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
|
||||
def test_can_object_call_model_key_alias_name_alone_is_not_enough():
|
||||
"""A key that may call the alias name but not its target cannot call the alias."""
|
||||
from litellm.proxy.auth.auth_checks import _can_object_call_model
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
_can_object_call_model(
|
||||
model="bar",
|
||||
llm_router=None,
|
||||
models=["bar"],
|
||||
key_model_aliases={"bar": "baz"},
|
||||
object_type="key",
|
||||
fallback_depth=0,
|
||||
)
|
||||
|
||||
assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied
|
||||
|
||||
assert (
|
||||
_can_object_call_model(
|
||||
model="bar",
|
||||
llm_router=None,
|
||||
models=["baz"],
|
||||
key_model_aliases={"bar": "baz"},
|
||||
object_type="key",
|
||||
fallback_depth=0,
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
|
||||
def test_can_object_call_model_team_alias_applies_before_key_alias():
|
||||
"""A key alias on the raw name loses to the team alias that rewrites it first at dispatch."""
|
||||
from litellm.proxy.auth.auth_checks import _can_object_call_model
|
||||
|
||||
assert (
|
||||
_can_object_call_model(
|
||||
model="foo",
|
||||
llm_router=None,
|
||||
models=["bar"],
|
||||
team_model_aliases={"foo": "bar"},
|
||||
key_model_aliases={"foo": "baz"},
|
||||
object_type="key",
|
||||
fallback_depth=0,
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
|
||||
def test_can_object_call_model_key_alias_on_team_alias_target():
|
||||
"""A key alias on the team-rewritten name resolves like the dispatch chain does."""
|
||||
from litellm.proxy.auth.auth_checks import _can_object_call_model
|
||||
|
||||
assert (
|
||||
_can_object_call_model(
|
||||
model="foo",
|
||||
llm_router=None,
|
||||
models=["baz"],
|
||||
team_model_aliases={"foo": "bar"},
|
||||
key_model_aliases={"bar": "baz"},
|
||||
object_type="key",
|
||||
fallback_depth=0,
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
_can_object_call_model(
|
||||
model="foo",
|
||||
llm_router=None,
|
||||
models=["bar"],
|
||||
team_model_aliases={"foo": "bar"},
|
||||
key_model_aliases={"bar": "baz"},
|
||||
object_type="key",
|
||||
fallback_depth=0,
|
||||
)
|
||||
|
||||
assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_can_user_call_model_honors_key_alias():
|
||||
"""A personal-scope key alias resolves to its target before the user allowlist check."""
|
||||
from litellm.proxy.auth.auth_checks import can_user_call_model
|
||||
|
||||
user_object = LiteLLM_UserTable(user_id="test-user", models=["gpt-4o-mini"])
|
||||
|
||||
assert (
|
||||
await can_user_call_model(
|
||||
model="mistral-7b",
|
||||
llm_router=None,
|
||||
user_object=user_object,
|
||||
key_model_aliases={"mistral-7b": "gpt-4o-mini"},
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await can_user_call_model(
|
||||
model="mistral-7b",
|
||||
llm_router=None,
|
||||
user_object=user_object,
|
||||
)
|
||||
|
||||
assert exc_info.value.type == ProxyErrorTypes.user_model_access_denied
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_team_member_model_access_honors_key_alias():
|
||||
"""A key alias resolves against the member allowlist, not just the raw alias name."""
|
||||
from litellm.proxy._types import LiteLLM_TeamMembership
|
||||
from litellm.proxy.auth.auth_checks import _check_team_member_model_access
|
||||
|
||||
membership = LiteLLM_TeamMembership(
|
||||
user_id="alice",
|
||||
team_id="team-a",
|
||||
litellm_budget_table=LiteLLM_BudgetTable(allowed_models=["gpt-4o-mini"]),
|
||||
)
|
||||
|
||||
await _check_team_member_model_access(
|
||||
model="mistral-7b",
|
||||
team_object=LiteLLM_TeamTable(team_id="team-a"),
|
||||
valid_token=UserAPIKeyAuth(token="sk-test", user_id="alice", team_id="team-a"),
|
||||
llm_router=None,
|
||||
prisma_client=None,
|
||||
user_api_key_cache=UserApiKeyCache(),
|
||||
proxy_logging_obj=MagicMock(),
|
||||
team_membership=membership,
|
||||
team_membership_loaded=True,
|
||||
key_model_aliases={"mistral-7b": "gpt-4o-mini"},
|
||||
)
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await _check_team_member_model_access(
|
||||
model="mistral-7b",
|
||||
team_object=LiteLLM_TeamTable(team_id="team-a"),
|
||||
valid_token=UserAPIKeyAuth(token="sk-test", user_id="alice", team_id="team-a"),
|
||||
llm_router=None,
|
||||
prisma_client=None,
|
||||
user_api_key_cache=UserApiKeyCache(),
|
||||
proxy_logging_obj=MagicMock(),
|
||||
team_membership=membership,
|
||||
team_membership_loaded=True,
|
||||
)
|
||||
|
||||
assert exc_info.value.type == ProxyErrorTypes.team_model_access_denied
|
||||
|
||||
|
||||
def test_can_object_call_model_access_via_underlying_model_only():
|
||||
"""
|
||||
Test that a key can access a model via underlying model even when using an alias.
|
||||
|
|
@ -9139,6 +9469,50 @@ async def test_agent_access_groups_cap_models_even_when_key_allows_them():
|
|||
assert asked == ["agent-1", "agent-1"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_access_group_ceiling_admits_the_key_alias_target():
|
||||
agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", models=["gpt-5"])
|
||||
agent_key.aliases = {"fast": "gpt-5"}
|
||||
resolve, asked = _agent_model_ceiling_resolver(frozenset({"gpt-5"}))
|
||||
|
||||
assert await _check_agent_access_group_model_access("fast", agent_key, None, resolve) is True
|
||||
assert asked == ["agent-1"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_access_group_ceiling_checks_the_team_alias_target():
|
||||
agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", team_id="team-1", models=["gpt-5"])
|
||||
agent_key.team_model_aliases = {"foo": "gpt-5"}
|
||||
resolve, asked = _agent_model_ceiling_resolver(frozenset({"gpt-5"}))
|
||||
|
||||
assert await _check_agent_access_group_model_access("foo", agent_key, None, resolve) is True
|
||||
assert asked == ["agent-1"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_access_group_ceiling_denies_a_team_alias_outside_the_ceiling():
|
||||
agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", team_id="team-1", models=["gpt-5"])
|
||||
agent_key.team_model_aliases = {"foo": "claude-sonnet-4-5"}
|
||||
resolve, _ = _agent_model_ceiling_resolver(frozenset({"gpt-5"}))
|
||||
|
||||
with pytest.raises(ModelAccessDeniedProxyException) as exc:
|
||||
await _check_agent_access_group_model_access("foo", agent_key, None, resolve)
|
||||
assert exc.value.type == ProxyErrorTypes.agent_model_access_denied
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_access_group_ceiling_keeps_the_name_for_a_deleted_team_deployment():
|
||||
from litellm.router import Router
|
||||
|
||||
agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", team_id="team-1", models=["foo"])
|
||||
agent_key.team_model_aliases = {"foo": "model_name_team-1_deadbeef"}
|
||||
router: Final = Router(model_list=[])
|
||||
resolve, asked = _agent_model_ceiling_resolver(frozenset({"foo"}))
|
||||
|
||||
assert await _check_agent_access_group_model_access("foo", agent_key, router, resolve) is True
|
||||
assert asked == ["agent-1"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_access_groups_naming_no_model_deny_every_model():
|
||||
agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", models=[])
|
||||
|
|
@ -9437,6 +9811,18 @@ async def test_agent_key_acting_for_a_teamless_user_is_capped_at_that_users_mode
|
|||
assert asked == ["team:None", "user:alice", "team:None", "user:alice"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_key_alias_resolves_against_the_echoed_teams_models():
|
||||
agent_key: Final = _agent_key_acting_for(user_id="alice", team_id="team-a")
|
||||
agent_key.aliases = {"foo": "bar"}
|
||||
load_team, load_user, asked = _caller_loaders(LiteLLM_TeamTable(team_id="team-a", models=["bar"]), None)
|
||||
cache: Final = await _cache_with_membership("alice", "team-a", allowed_models=None)
|
||||
|
||||
await _check_caller_models(agent_key, "foo", load_team, load_user, cache)
|
||||
|
||||
assert asked == ["team:team-a"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_key_without_an_echoed_caller_keeps_its_own_models():
|
||||
agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", models=["gpt-5", "claude-sonnet"])
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue