This commit is contained in:
Lukas 2026-09-27 16:19:45 -04:00 • committed by GitHub
commit 8f0218e54e
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 329 additions and 23 deletions

View file

@ -1263,6 +1263,19 @@ async def _common_key_generation_helper(
# check if user set upperbound key/generate params on config.yaml
_enforce_upperbound_key_params(data, fill_defaults=True)
# Checked after the defaults, because default_key_generate_params can supply
# team_id and the project's owner is checked against the key's final team.
if data.project_id is not None and prisma_client is not None:
from litellm.proxy.proxy_server import user_api_key_cache
await _check_project_key_limits(
project_id=data.project_id,
data=data,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
key_team_id=data.team_id,
)
# Delegated-authority ceiling (GHSA-q775-qw9r-2r4g): a non-admin caller
# cannot grant a key a higher budget than their own authority.
# UI session personal keys are capped by user_max_budget when it is available.
@ -1759,10 +1772,14 @@ async def _check_project_key_limits(
data: GenerateKeyRequest | UpdateKeyRequest,
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
key_team_id: str | None = None,
) -> None:
"""
Validate that key's models and budget respect its project's limits.
Validate that the key belongs to the project's team, and that its models
and budget respect the project's limits.
- The project's owning team must be the key's team. A project with no team
has no owner to protect, so it is not restricted
- Key models must be a subset of project models, except the all-team-models / all-proxy-models
sentinels, which inherit a parent scope and are narrowed by the project at request time
- Key max_budget must be <= project max_budget
@ -1779,6 +1796,16 @@ async def _check_project_key_limits(
detail={"error": f"Project not found, project_id={project_id}"},
)
if project_obj.team_id is not None and project_obj.team_id != key_team_id:
raise HTTPException(
status_code=403,
detail={
"error": f"Project {project_id} belongs to team {project_obj.team_id}, "
f"but the key belongs to {key_team_id if key_team_id is not None else 'no team'}. "
"A project can only be attached to keys of the team that owns it."
},
)
# Validate key models are a subset of project models
if data.models and len(project_obj.models) > 0:
for m in data.models:
@ -1804,6 +1831,33 @@ async def _check_project_key_limits(
)
# Touching any of these can change the project a key is under, the team it is on, or what the project must allow.
_PROJECT_LIMIT_FIELDS: Final = frozenset({"project_id", "team_id", "models", "max_budget"})
async def _check_project_key_limits_on_mutation(
data: UpdateKeyRequest | RegenerateKeyRequest,
existing_key_row: LiteLLM_VerificationToken,
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
) -> None:
"""Run _check_project_key_limits against the key as the mutation leaves it."""
if not data.model_fields_set & _PROJECT_LIMIT_FIELDS:
return
project_id: Final = data.project_id if "project_id" in data.model_fields_set else existing_key_row.project_id
if project_id is None:
return
await _check_project_key_limits(
project_id=project_id,
data=data,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
key_team_id=(data.team_id if "team_id" in data.model_fields_set else existing_key_row.team_id),
)
def check_org_key_model_specific_limits(
keys: Sequence[LiteLLM_VerificationToken],
org_table: LiteLLM_OrganizationTable,
@ -2129,15 +2183,6 @@ async def generate_key_fn(
prisma_client=prisma_client,
)
# Validate key against project limits if project_id is set
if data.project_id is not None:
await _check_project_key_limits(
project_id=data.project_id,
data=data,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
)
return await _common_key_generation_helper(
data=data,
user_api_key_dict=user_api_key_dict,
@ -3267,17 +3312,12 @@ async def _validate_update_key_data(
access_group_ids=data.access_group_ids,
)
# Validate key against project limits if project_id is being set
_project_id_to_check: Final = (
data.project_id if "project_id" in data.model_fields_set else existing_key_row.project_id
await _check_project_key_limits_on_mutation(
data=data,
existing_key_row=existing_key_row,
prisma_client=checked_prisma_client,
user_api_key_cache=user_api_key_cache,
)
if _project_id_to_check is not None and (data.models is not None or data.max_budget is not None):
await _check_project_key_limits(
project_id=_project_id_to_check,
data=data,
prisma_client=checked_prisma_client,
user_api_key_cache=user_api_key_cache,
)
# When the caller asks to change the key's organization_id, require that
# they are a member of (or a proxy admin over) the target organization.
@ -5611,6 +5651,12 @@ async def _execute_virtual_key_regeneration(
await _enforce_custom_key_update_policy(hook=_custom_key_update_hook(proxy_server), data=update_request)
# Enforce upperbound key params on regenerate (don't fill defaults)
_enforce_upperbound_key_params(data, fill_defaults=False)
await _check_project_key_limits_on_mutation(
data=data,
existing_key_row=key_in_db,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
)
non_default_values = await prepare_key_update_data(
data=data, existing_key_row=key_in_db, prisma_client=prisma_client, llm_router=llm_router
)

View file

@ -47,6 +47,7 @@ from litellm.litellm_core_utils.duration_parser import duration_in_seconds
from litellm.proxy.management_endpoints.key_management_endpoints import (
_check_org_key_limits,
_check_project_key_limits,
_check_project_key_limits_on_mutation,
_check_team_key_limits,
_common_key_generation_helper,
_effective_key_after_update,
@ -20059,11 +20060,13 @@ async def test_regenerate_key_repoints_live_membership_not_the_key_row_it_read(
) == ["attached-model"]
async def _cache_with_project(project_id: str, project_models: list[str]) -> UserApiKeyCache:
async def _cache_with_project(
project_id: str, project_models: list[str], team_id: str | None = "team-lit-5823"
) -> UserApiKeyCache:
user_api_key_cache = UserApiKeyCache()
await user_api_key_cache.async_set_cache(
key=project_cache_key(project_id),
value=LiteLLM_ProjectTableCachedObj(project_id=project_id, team_id="team-lit-5823", models=project_models),
value=LiteLLM_ProjectTableCachedObj(project_id=project_id, team_id=team_id, models=project_models),
model_type=LiteLLM_ProjectTableCachedObj,
)
return user_api_key_cache
@ -20081,6 +20084,7 @@ async def test_check_project_key_limits_accepts_inherited_model_sentinels(reques
data=request_cls(key="sk-lit-5823", models=[sentinel]),
prisma_client=MagicMock(),
user_api_key_cache=user_api_key_cache,
key_team_id="team-lit-5823",
)
@ -20099,6 +20103,7 @@ async def test_check_project_key_limits_still_rejects_real_model_outside_project
data=request_cls(key="sk-lit-5823", models=key_models),
prisma_client=MagicMock(),
user_api_key_cache=user_api_key_cache,
key_team_id="team-lit-5823",
)
assert exc_info.value.status_code == 400
@ -20467,7 +20472,10 @@ async def test_project_detachment_preserves_omission_and_other_key_fields():
@pytest.mark.asyncio
async def test_project_detachment_uses_effective_project_for_validation(project_id: str | None):
existing: Final = LiteLLM_VerificationToken(token="project-detach-token", project_id="project-orbit")
cache: Final = await _cache_with_project("project-orbit", ["model-orbit"])
# An unowned project: this cell is about WHICH project the validation uses,
# so the ownership gate (#41089) must not be what it measures. Giving the
# key a team instead would pull the whole team lookup into a MagicMock db.
cache: Final = await _cache_with_project("project-orbit", ["model-orbit"], team_id=None)
data: Final = UpdateKeyRequest(key=existing.token, project_id=project_id, models=["model-other"])
if project_id is None:
await _validate_update_key_data(
@ -20502,6 +20510,258 @@ async def test_key_creator_cannot_detach_project_without_admin_access():
assert "Only proxy admins, team admins, or org admins" in str(exc.value.detail)
# --- Tests: a project may only be attached to keys of the team that owns it ---
_OWNED_PROJECT: Final = "proj-owned-1"
@pytest.mark.parametrize(
"project_team_id, key_team_id, expected_status, expected_in_detail",
[
# A key on another team must not point at this project.
("team-b", "team-a", 403, ["team-b", "team-a"]),
# The issue's step 5: no team at all still charges a team's project.
("team-b", None, 403, ["no team"]),
# Accept controls. Without these the check is a wall, not a boundary.
("team-a", "team-a", None, []),
# A project with no owning team has nobody to protect.
(None, "team-a", None, []),
(None, None, None, []),
],
)
@pytest.mark.asyncio
async def test_a_project_is_only_for_keys_of_its_own_team(
project_team_id, key_team_id, expected_status, expected_in_detail
):
user_api_key_cache: Final = await _cache_with_project(
_OWNED_PROJECT, [], team_id=project_team_id
)
async def check():
await _check_project_key_limits(
project_id=_OWNED_PROJECT,
data=GenerateKeyRequest(team_id=key_team_id),
prisma_client=MagicMock(),
user_api_key_cache=user_api_key_cache,
key_team_id=key_team_id,
)
if expected_status is None:
await check()
# The project was read and accepted, rather than never reached.
assert await user_api_key_cache.async_get_cache(key=project_cache_key(_OWNED_PROJECT))
return
with pytest.raises(HTTPException) as exc:
await check()
assert exc.value.status_code == expected_status
for fragment in expected_in_detail:
assert fragment in str(exc.value.detail)
@pytest.mark.asyncio
async def test_a_foreign_project_is_refused_as_foreign_not_as_a_model_problem():
"""A 403 reads as a tenancy problem; the 400 would read as a configuration one."""
user_api_key_cache: Final = await _cache_with_project(
_OWNED_PROJECT, ["gpt-4o-mini"], team_id="team-b"
)
with pytest.raises(HTTPException) as exc:
await _check_project_key_limits(
project_id=_OWNED_PROJECT,
data=GenerateKeyRequest(team_id="team-a", models=["gpt-4o"]),
prisma_client=MagicMock(),
user_api_key_cache=user_api_key_cache,
key_team_id="team-a",
)
assert exc.value.status_code == 403
@pytest.mark.parametrize(
"existing_team_id, data_kwargs, expected_status",
[
# Moves only the project: the key's stored team is what counts.
("team-a", {"project_id": _OWNED_PROJECT}, 403),
# Moves only the team: the project stays attached, so it still counts.
("team-b", {"team_id": "team-a"}, 403),
# Moves the key onto the project's own team. Accept control.
(None, {"team_id": "team-b"}, None),
# Touches neither, so there is nothing to re-check.
("team-a", {"key_alias": "renamed"}, None),
],
)
@pytest.mark.asyncio
async def test_an_update_checks_the_project_the_mutation_leaves_attached(
existing_team_id, data_kwargs, expected_status
):
existing: Final = LiteLLM_VerificationToken(
token="sk-hash", team_id=existing_team_id, project_id=_OWNED_PROJECT
)
user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id="team-b")
async def check():
await _check_project_key_limits_on_mutation(
data=UpdateKeyRequest(key="sk-x", **data_kwargs),
existing_key_row=existing,
prisma_client=MagicMock(),
user_api_key_cache=user_api_key_cache,
)
if expected_status is None:
await check()
assert await user_api_key_cache.async_get_cache(key=project_cache_key(_OWNED_PROJECT))
return
with pytest.raises(HTTPException) as exc:
await check()
assert exc.value.status_code == expected_status
@pytest.mark.asyncio
async def test_an_update_that_touches_none_of_the_fields_does_not_read_the_project():
existing: Final = LiteLLM_VerificationToken(
token="sk-hash", team_id="team-a", project_id=_OWNED_PROJECT
)
prisma_client: Final = MagicMock()
# The cache is empty, so any lookup would have to reach the database.
await _check_project_key_limits_on_mutation(
data=UpdateKeyRequest(key="sk-x", key_alias="renamed"),
existing_key_row=existing,
prisma_client=prisma_client,
user_api_key_cache=UserApiKeyCache(),
)
assert prisma_client.mock_calls == []
@pytest.mark.parametrize(
"key_team_id, expected_status, expected_updates",
[
# /key/regenerate wrote project_id and team_id with no project check.
("team-a", 403, 0),
# Accept control: a key already on the project's team regenerates.
("team-b", None, 1),
],
)
@pytest.mark.asyncio
async def test_regenerate_checks_the_project_it_attaches(
key_team_id, expected_status, expected_updates
):
from litellm.proxy._types import RegenerateKeyRequest
from litellm.proxy.management_endpoints.key_management_endpoints import (
_execute_virtual_key_regeneration,
)
existing_key: Final = LiteLLM_VerificationToken(token="abc123", team_id=key_team_id)
mock_prisma_client: Final = _make_regenerate_mock_prisma()
user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id="team-b")
async def regenerate():
with (
patch( # test-quality-ok: a fresh random token is a side effect of regeneration, not the project check under test; the file's other regenerate cells stub the same seam
"litellm.proxy.management_endpoints.key_management_endpoints.get_new_token",
new_callable=AsyncMock,
return_value="sk-newtoken1234ab12",
),
patch( # test-quality-ok: the deprecated-key row is a side effect of regeneration; stubbing it keeps the DB assertion below about the key update alone
"litellm.proxy.management_endpoints.key_management_endpoints._insert_deprecated_key",
new_callable=AsyncMock,
),
patch( # test-quality-ok: the cache delete is a side effect of regeneration and would evict the seeded project this cell reads
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object",
new_callable=AsyncMock,
),
):
await _execute_virtual_key_regeneration(
prisma_client=mock_prisma_client,
key_in_db=existing_key,
hashed_api_key="abc123",
key="abc123",
data=RegenerateKeyRequest(project_id=_OWNED_PROJECT),
user_api_key_dict=_make_regenerate_user_api_key_dict(),
litellm_changed_by=None,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=MagicMock(),
)
if expected_status is None:
await regenerate()
else:
with pytest.raises(HTTPException) as exc:
await regenerate()
assert exc.value.status_code == expected_status
# A refused regenerate must not reach the DB update.
assert (
mock_prisma_client.db.litellm_verificationtoken.update.await_count == expected_updates
)
def _make_generate_mock_prisma():
"""Mock prisma client shaped for _common_key_generation_helper."""
mock_prisma_client = AsyncMock()
mock_prisma_client.insert_data = AsyncMock(
return_value=MagicMock(
token="hashed_token_123", litellm_budget_table=None, object_permission=None
)
)
mock_prisma_client.db = MagicMock()
mock_prisma_client.db.litellm_verificationtoken = MagicMock()
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=None)
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma_client.db.litellm_verificationtoken.count = AsyncMock(return_value=0)
mock_prisma_client.db.litellm_verificationtoken.update = AsyncMock(
return_value=MagicMock(
token="hashed_token_123", litellm_budget_table=None, object_permission=None
)
)
return mock_prisma_client
@pytest.mark.parametrize(
"project_team_id, expected_status",
[
# default_key_generate_params fills team_id after the request is parsed,
# so the key ends up on team-b and may use its projects.
("team-b", None),
# ... and still may not use another team's.
("team-c", 403),
],
)
@pytest.mark.asyncio
async def test_a_team_supplied_by_defaults_is_what_the_project_is_checked_against(
monkeypatch, project_team_id, expected_status
):
monkeypatch.setattr(
"litellm.proxy.proxy_server.prisma_client", _make_generate_mock_prisma()
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.user_api_key_cache",
await _cache_with_project(_OWNED_PROJECT, [], team_id=project_team_id),
)
monkeypatch.setattr(litellm, "default_key_generate_params", {"team_id": "team-b"})
async def generate():
return await _common_key_generation_helper(
data=GenerateKeyRequest(project_id=_OWNED_PROJECT),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1234", user_id="1234"
),
litellm_changed_by=None,
team_table=None,
)
if expected_status is None:
response = await generate()
assert response.team_id == "team-b"
return
with pytest.raises(HTTPException) as exc:
await generate()
assert exc.value.status_code == expected_status
@pytest.mark.asyncio
async def test_bulk_update_team_keys_runs_custom_key_policy_per_key(monkeypatch):
from litellm.types.proxy.management_endpoints.key_management_endpoints import (