mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(proxy): drop budget and permission rows written for a rejected cross-team key
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
bbbd0f62b3
commit
5e19a24738
3 changed files with 181 additions and 8 deletions
|
|
@ -1379,6 +1379,10 @@ async def _common_key_generation_helper(
|
|||
)
|
||||
_budget_id = getattr(_budget, "budget_id", None)
|
||||
|
||||
created_budget_id: Final[str | None] = (
|
||||
_budget_id if prisma_client is not None and data.soft_budget is not None else None
|
||||
)
|
||||
|
||||
# ADD METADATA FIELDS
|
||||
# Set Management Endpoint Metadata Fields
|
||||
for field in LiteLLM_ManagementEndpoint_MetadataFields_Premium:
|
||||
|
|
@ -1484,10 +1488,16 @@ async def _common_key_generation_helper(
|
|||
for _op_field, _op_default_value in _default_object_permission.items():
|
||||
_caller_object_permission.setdefault(_op_field, _op_default_value)
|
||||
|
||||
should_create_object_permission: Final = prisma_client is not None and isinstance(
|
||||
data_json.get("object_permission"), dict
|
||||
)
|
||||
data_json = await _set_object_permission(
|
||||
data_json=data_json,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
created_object_permission_id: Final[str | None] = (
|
||||
cast(str | None, data_json.get("object_permission_id")) if should_create_object_permission else None
|
||||
)
|
||||
|
||||
_validate_key_alias_format(key_alias=data_json.get("key_alias", None))
|
||||
|
||||
|
|
@ -1555,7 +1565,19 @@ async def _common_key_generation_helper(
|
|||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
response = await generate_key_helper_fn(request_type="key", **data_json, table_name="key", llm_router=llm_router)
|
||||
try:
|
||||
response = await generate_key_helper_fn(
|
||||
request_type="key", **data_json, table_name="key", llm_router=llm_router
|
||||
)
|
||||
except KeyProjectTeamMismatchError:
|
||||
if prisma_client is not None:
|
||||
if created_object_permission_id is not None:
|
||||
await ObjectPermissionRepository(prisma_client).table.delete(
|
||||
where={"object_permission_id": created_object_permission_id}
|
||||
)
|
||||
if created_budget_id is not None:
|
||||
await BudgetRepository(prisma_client).table.delete(where={"budget_id": created_budget_id})
|
||||
raise
|
||||
|
||||
response["soft_budget"] = data.soft_budget # include the user-input soft budget in the response
|
||||
|
||||
|
|
@ -1831,7 +1853,7 @@ async def _check_key_project_team(
|
|||
if project_obj.team_id is None or project_obj.team_id == key_team_id:
|
||||
return
|
||||
|
||||
raise HTTPException(
|
||||
raise KeyProjectTeamMismatchError(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": (
|
||||
|
|
@ -1843,6 +1865,10 @@ async def _check_key_project_team(
|
|||
)
|
||||
|
||||
|
||||
class KeyProjectTeamMismatchError(HTTPException):
|
||||
pass
|
||||
|
||||
|
||||
async def _check_key_project_team_on_mutation(
|
||||
data: UpdateKeyRequest | RegenerateKeyRequest,
|
||||
existing_key_row: LiteLLM_VerificationToken,
|
||||
|
|
|
|||
|
|
@ -89,6 +89,13 @@ def _object_permission_table_rows() -> list[dict[str, JsonValue]]:
|
|||
)
|
||||
|
||||
|
||||
def _budget_table_rows() -> list[dict[str, JsonValue]]:
|
||||
return read_rows(
|
||||
'SELECT to_jsonb(b) AS row FROM "LiteLLM_BudgetTable" AS b ORDER BY budget_id',
|
||||
(),
|
||||
)
|
||||
|
||||
|
||||
def _budget_rows(budget_id: str) -> list[dict[str, JsonValue]]:
|
||||
return read_rows(
|
||||
'SELECT to_jsonb(b) AS row FROM "LiteLLM_BudgetTable" AS b WHERE budget_id = %s',
|
||||
|
|
@ -337,13 +344,25 @@ def test_service_account_generate_rejects_foreign_team_project_without_writing_k
|
|||
team_b: Final = scenario.team(models=[model])
|
||||
project_b: Final = scenario.project(team_b, models=[model])
|
||||
alias: Final = f"service-account-{uuid4().hex}"
|
||||
vector_store: Final = f"vector-store-{uuid4().hex}"
|
||||
budget_rows_before: Final = _budget_table_rows()
|
||||
object_permission_rows_before: Final = _object_permission_table_rows()
|
||||
response: Final = ownership_gateway.request(
|
||||
"POST",
|
||||
"/key/service-account/generate",
|
||||
{"team_id": team_a, "project_id": project_b, "key_alias": alias, "models": [model]},
|
||||
{
|
||||
"team_id": team_a,
|
||||
"project_id": project_b,
|
||||
"key_alias": alias,
|
||||
"models": [model],
|
||||
"soft_budget": 3.5,
|
||||
"object_permission": {"vector_stores": [vector_store]},
|
||||
},
|
||||
)
|
||||
_discard_unexpected_key(ownership_gateway, response)
|
||||
assert response.status_code == 400, response.text
|
||||
assert _budget_table_rows() == budget_rows_before
|
||||
assert _object_permission_table_rows() == object_permission_rows_before
|
||||
assert (
|
||||
read_rows(
|
||||
'SELECT token FROM "LiteLLM_VerificationToken" WHERE project_id = %s AND key_alias = %s',
|
||||
|
|
@ -1082,17 +1101,33 @@ def test_key_generate_rejects_foreign_team_project(ownership_gateway: Gateway) -
|
|||
team_a: Final = scenario.team(models=[model])
|
||||
team_b: Final = scenario.team(models=[model])
|
||||
project_b: Final = scenario.project(team_b, models=[model])
|
||||
alias: Final = f"cross-team-{uuid4().hex}"
|
||||
vector_store: Final = f"vector-store-{uuid4().hex}"
|
||||
budget_rows_before: Final = _budget_table_rows()
|
||||
object_permission_rows_before: Final = _object_permission_table_rows()
|
||||
cross_team: Final = ownership_gateway.request(
|
||||
"POST",
|
||||
"/key/generate",
|
||||
{"team_id": team_a, "project_id": project_b, "models": [model]},
|
||||
{
|
||||
"team_id": team_a,
|
||||
"project_id": project_b,
|
||||
"key_alias": alias,
|
||||
"models": [model],
|
||||
"soft_budget": 3.5,
|
||||
"object_permission": {"vector_stores": [vector_store]},
|
||||
},
|
||||
)
|
||||
_discard_unexpected_key(ownership_gateway, cross_team)
|
||||
assert cross_team.status_code == 400, cross_team.text
|
||||
assert read_rows(
|
||||
'SELECT token FROM "LiteLLM_VerificationToken" WHERE project_id = %s AND team_id = %s',
|
||||
(project_b, team_a),
|
||||
) == []
|
||||
assert _budget_table_rows() == budget_rows_before
|
||||
assert _object_permission_table_rows() == object_permission_rows_before
|
||||
assert (
|
||||
read_rows(
|
||||
'SELECT token FROM "LiteLLM_VerificationToken" WHERE project_id = %s AND team_id = %s',
|
||||
(project_b, team_a),
|
||||
)
|
||||
== []
|
||||
)
|
||||
|
||||
|
||||
def test_key_generate_rejects_missing_team_for_owned_project(ownership_gateway: Gateway) -> None:
|
||||
|
|
|
|||
|
|
@ -56,6 +56,7 @@ from litellm.proxy.management_endpoints.key_management_endpoints import (
|
|||
_check_project_key_limits,
|
||||
_check_team_key_limits,
|
||||
_common_key_generation_helper,
|
||||
KeyProjectTeamMismatchError,
|
||||
_effective_key_after_update,
|
||||
_effective_key_for_generate,
|
||||
_enforce_custom_key_policy,
|
||||
|
|
@ -1823,6 +1824,117 @@ async def test_generate_service_account_works_with_team_id():
|
|||
)
|
||||
|
||||
|
||||
def _identity_json_object(value: dict[str, object]) -> dict[str, object]:
|
||||
return value
|
||||
|
||||
|
||||
def _key_generation_prisma_client(
|
||||
project_record: dict[str, str] | None,
|
||||
) -> tuple[MagicMock, MagicMock, MagicMock]:
|
||||
prisma_client: Final = MagicMock()
|
||||
prisma_client.jsonify_object.side_effect = _identity_json_object
|
||||
|
||||
budget_table: Final = MagicMock()
|
||||
budget_table.create = AsyncMock(return_value=MagicMock(budget_id="created-budget-id"))
|
||||
budget_table.delete = AsyncMock()
|
||||
object_permission_table: Final = MagicMock()
|
||||
object_permission_table.create = AsyncMock(
|
||||
return_value=MagicMock(object_permission_id="created-object-permission-id")
|
||||
)
|
||||
object_permission_table.delete = AsyncMock()
|
||||
|
||||
prisma_client.db.litellm_budgettable = budget_table
|
||||
prisma_client.db.litellm_objectpermissiontable = object_permission_table
|
||||
prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
|
||||
prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock(return_value=project_record)
|
||||
prisma_client.insert_data = AsyncMock()
|
||||
return prisma_client, budget_table, object_permission_table
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rejected_key_generation_deletes_created_budget_and_default_permission(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
prisma_client, budget_table, object_permission_table = _key_generation_prisma_client(
|
||||
{"project_id": "project-b", "team_id": "team-b"}
|
||||
)
|
||||
monkeypatch.setattr(litellm, "key_generation_settings", None)
|
||||
monkeypatch.setattr(
|
||||
litellm,
|
||||
"default_key_generate_params",
|
||||
{"object_permission": {"vector_stores": ["default-vector-store"]}},
|
||||
)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", False)
|
||||
|
||||
with pytest.raises(KeyProjectTeamMismatchError) as exc:
|
||||
await _common_key_generation_helper(
|
||||
data=GenerateKeyRequest(
|
||||
budget_id="caller-budget-id",
|
||||
project_id="project-b",
|
||||
soft_budget=3.5,
|
||||
team_id="team-a",
|
||||
),
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-admin",
|
||||
user_id="admin",
|
||||
),
|
||||
litellm_changed_by=None,
|
||||
team_table=None,
|
||||
)
|
||||
|
||||
assert exc.value.status_code == 400
|
||||
budget_table.create.assert_awaited_once()
|
||||
budget_table.delete.assert_awaited_once_with(where={"budget_id": "created-budget-id"})
|
||||
object_permission_table.create.assert_awaited_once_with(data={"vector_stores": ["default-vector-store"]})
|
||||
object_permission_table.delete.assert_awaited_once_with(
|
||||
where={"object_permission_id": "created-object-permission-id"}
|
||||
)
|
||||
prisma_client.insert_data.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_generation_does_not_delete_rows_for_other_http_exceptions(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
prisma_client, budget_table, object_permission_table = _key_generation_prisma_client(None)
|
||||
monkeypatch.setattr(litellm, "key_generation_settings", None)
|
||||
monkeypatch.setattr(
|
||||
litellm,
|
||||
"default_key_generate_params",
|
||||
{"object_permission": {"vector_stores": ["default-vector-store"]}},
|
||||
)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", False)
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await _common_key_generation_helper(
|
||||
data=GenerateKeyRequest(
|
||||
project_id="project-b",
|
||||
router_settings={"weights": {"gpt-4": {"unknown-deployment": 1.0}}},
|
||||
soft_budget=3.5,
|
||||
team_id="team-a",
|
||||
),
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-admin",
|
||||
user_id="admin",
|
||||
),
|
||||
litellm_changed_by=None,
|
||||
team_table=None,
|
||||
)
|
||||
|
||||
assert type(exc.value) is HTTPException
|
||||
assert exc.value.status_code == 400
|
||||
budget_table.create.assert_awaited_once()
|
||||
budget_table.delete.assert_not_awaited()
|
||||
object_permission_table.create.assert_awaited_once_with(data={"vector_stores": ["default-vector-store"]})
|
||||
object_permission_table.delete.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_generate_key_throttle_rejected_for_non_admin():
|
||||
"""Security regression: a non-admin creating a key must not be able to set
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue