mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(proxy): run key project ownership right before the key row write
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
f4c161136a
commit
1d24f7fdbc
3 changed files with 78 additions and 20 deletions
|
|
@ -1553,13 +1553,6 @@ async def _common_key_generation_helper(
|
|||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
if data.project_id is not None and prisma_client is not None:
|
||||
await _check_key_project_team(
|
||||
project_id=data.project_id,
|
||||
key_team_id=data.team_id,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
response = await generate_key_helper_fn(request_type="key", **data_json, table_name="key", llm_router=llm_router)
|
||||
|
||||
response["soft_budget"] = data.soft_budget # include the user-input soft budget in the response
|
||||
|
|
@ -4922,6 +4915,13 @@ async def generate_key_helper_fn(
|
|||
# the LiteLLM_VerificationToken table will increase in size if we don't do this check
|
||||
return user_data
|
||||
|
||||
if project_id is not None:
|
||||
await _check_key_project_team(
|
||||
project_id=project_id,
|
||||
key_team_id=team_id,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
## CREATE KEY
|
||||
verbose_proxy_logger.debug(
|
||||
"prisma_client: Creating Key= %s",
|
||||
|
|
@ -5742,6 +5742,13 @@ async def _execute_virtual_key_regeneration(
|
|||
prisma_client=prisma_client,
|
||||
)
|
||||
update_data.update(update_values)
|
||||
if data is not None:
|
||||
await _check_key_project_team_on_mutation(
|
||||
data=data,
|
||||
existing_key_row=key_in_db,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
jsonified_update_data: Final[Mapping[str, object]] = prisma_client.jsonify_object(data=update_data)
|
||||
|
||||
# Snapshot before the token update: the FK cascade rewrites mapping rows to the new hash,
|
||||
|
|
@ -5766,13 +5773,6 @@ async def _execute_virtual_key_regeneration(
|
|||
grace_period=data.grace_period if data else None,
|
||||
)
|
||||
|
||||
if data is not None:
|
||||
await _check_key_project_team_on_mutation(
|
||||
data=data,
|
||||
existing_key_row=key_in_db,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
updated_token: Final[LiteLLM_VerificationToken | None] = await _prisma_table(
|
||||
VerificationTokenRepository(prisma_client)
|
||||
).update(
|
||||
|
|
|
|||
|
|
@ -33,6 +33,20 @@ def _key_rows(key: str) -> list[dict[str, JsonValue]]:
|
|||
)
|
||||
|
||||
|
||||
def _deleted_key_rows(key: str) -> list[dict[str, JsonValue]]:
|
||||
return read_rows(
|
||||
'SELECT token FROM "LiteLLM_DeletedVerificationToken" WHERE token = %s',
|
||||
(sha256(key.encode()).hexdigest(),),
|
||||
)
|
||||
|
||||
|
||||
def _deprecated_key_rows(key: str) -> list[dict[str, JsonValue]]:
|
||||
return read_rows(
|
||||
'SELECT token FROM "LiteLLM_DeprecatedVerificationToken" WHERE token = %s',
|
||||
(sha256(key.encode()).hexdigest(),),
|
||||
)
|
||||
|
||||
|
||||
def _cli_session_token(
|
||||
user_id: str,
|
||||
team_id: str | None,
|
||||
|
|
@ -289,13 +303,19 @@ def test_key_regenerate_rejects_foreign_project_without_changing_key(ownership_g
|
|||
key: Final = string_value(JSON_OBJECT.validate_json(generated.content)["key"])
|
||||
scenario.cleanups.callback(scenario.delete_key, key)
|
||||
before: Final = _key_rows(key)
|
||||
before_deleted: Final = _deleted_key_rows(key)
|
||||
before_deprecated: Final = _deprecated_key_rows(key)
|
||||
assert len(before) == 1
|
||||
assert before_deleted == []
|
||||
assert before_deprecated == []
|
||||
response: Final = ownership_gateway.request(
|
||||
"POST", f"/key/{key}/regenerate", {"project_id": project_b}
|
||||
"POST", f"/key/{key}/regenerate", {"project_id": project_b, "grace_period": "1h"}
|
||||
)
|
||||
_discard_unexpected_key(ownership_gateway, response)
|
||||
assert response.status_code == 400, response.text
|
||||
assert _key_rows(key) == before
|
||||
assert _deleted_key_rows(key) == before_deleted
|
||||
assert _deprecated_key_rows(key) == before_deprecated
|
||||
|
||||
|
||||
def test_key_generation_rejects_missing_project_without_writing_key(ownership_gateway: Gateway) -> None:
|
||||
|
|
|
|||
|
|
@ -20671,6 +20671,37 @@ async def test_key_generation_organization_membership_error_precedes_project_own
|
|||
assert error.value.detail == "Caller is not a member of organization_id=org-not-member"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_generation_premium_permission_error_precedes_project_ownership(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
user_api_key_cache: Final = await _cache_with_project(
|
||||
_OWNED_PROJECT, [], team_id=_OWNERSHIP_PROJECT_TEAM
|
||||
)
|
||||
mock_prisma_client: Final = _configure_key_endpoints(monkeypatch, user_api_key_cache)
|
||||
monkeypatch.setattr(litellm, "default_key_generate_params", None)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", False)
|
||||
|
||||
with pytest.raises(HTTPException) as error:
|
||||
await _common_key_generation_helper(
|
||||
data=GenerateKeyRequest(
|
||||
project_id=_OWNED_PROJECT,
|
||||
team_id=_OWNERSHIP_KEY_TEAM,
|
||||
permissions={"get_spend_routes": True},
|
||||
),
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-admin",
|
||||
),
|
||||
litellm_changed_by=None,
|
||||
team_table=None,
|
||||
)
|
||||
|
||||
assert error.value.status_code == 500
|
||||
assert error.value.detail == {"error": "Internal Server Error."}
|
||||
mock_prisma_client.insert_data.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_generation_duplicate_alias_error_precedes_project_ownership(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
|
|
@ -21175,6 +21206,12 @@ async def test_regenerate_checks_project_team_ownership(
|
|||
mock_prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock(
|
||||
return_value=LiteLLM_ProjectTable(project_id=_OWNED_PROJECT, team_id="team-b")
|
||||
)
|
||||
deleted_history_table: Final = MagicMock()
|
||||
deleted_history_table.create_many = AsyncMock()
|
||||
mock_prisma_client.db.litellm_deletedverificationtoken = deleted_history_table
|
||||
deprecated_key_table: Final = MagicMock()
|
||||
deprecated_key_table.upsert = AsyncMock()
|
||||
mock_prisma_client.db.litellm_deprecatedverificationtoken = deprecated_key_table
|
||||
user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id="team-b")
|
||||
|
||||
async def regenerate() -> None:
|
||||
|
|
@ -21184,10 +21221,6 @@ async def test_regenerate_checks_project_team_ownership(
|
|||
new_callable=AsyncMock,
|
||||
return_value="sk-newtoken1234ab12",
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints._insert_deprecated_key",
|
||||
new_callable=AsyncMock,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object",
|
||||
new_callable=AsyncMock,
|
||||
|
|
@ -21198,7 +21231,10 @@ async def test_regenerate_checks_project_team_ownership(
|
|||
key_in_db=existing_key,
|
||||
hashed_api_key="abc123",
|
||||
key="abc123",
|
||||
data=RegenerateKeyRequest(project_id=_OWNED_PROJECT),
|
||||
data=RegenerateKeyRequest(
|
||||
project_id=_OWNED_PROJECT,
|
||||
grace_period="1h" if expected_status is not None else None,
|
||||
),
|
||||
user_api_key_dict=_make_regenerate_user_api_key_dict(),
|
||||
litellm_changed_by=None,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
|
|
@ -21210,6 +21246,8 @@ async def test_regenerate_checks_project_team_ownership(
|
|||
await regenerate()
|
||||
assert error.value.status_code == expected_status
|
||||
assert "belongs to team team-b, but the key belongs to team-a" in str(error.value.detail)
|
||||
deleted_history_table.create_many.assert_not_awaited()
|
||||
deprecated_key_table.upsert.assert_not_awaited()
|
||||
else:
|
||||
await regenerate()
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue