mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(key_management): merge service_account_id into service-account key metadata
The /key/service-account/generate endpoint relied on the caller (the Admin UI) to place service_account_id in metadata, so raw-API callers had to manage the reserved field themselves and a client that replaced metadata could drop the caller's keys before custom_key_generate validation. Inject service_account_id server-side by merging it into the incoming metadata (deriving it from an explicit value or the key alias) before the custom hook runs, so validation sees the combined result and caller keys are preserved. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
2ea09a9969
commit
dbdfcf9b8c
2 changed files with 129 additions and 0 deletions
|
|
@ -1671,6 +1671,23 @@ async def generate_key_fn(
|
|||
raise handle_exception_on_proxy(e)
|
||||
|
||||
|
||||
def _ensure_service_account_id_in_metadata(
|
||||
metadata: dict | None, key_alias: str | None
|
||||
) -> dict | None:
|
||||
"""Merge the LiteLLM-internal ``service_account_id`` into caller metadata.
|
||||
|
||||
Keys created via ``/key/service-account/generate`` are identified by
|
||||
``metadata["service_account_id"]``. Deriving it here (from an explicit value
|
||||
or the key alias) and merging keeps it transparent to callers and preserves
|
||||
any user-supplied keys, so ``custom_key_generate`` validates the combined
|
||||
result instead of losing the caller's metadata.
|
||||
"""
|
||||
service_account_id = (metadata or {}).get("service_account_id") or key_alias
|
||||
if service_account_id is None:
|
||||
return metadata
|
||||
return {**(metadata or {}), "service_account_id": service_account_id}
|
||||
|
||||
|
||||
@router.post(
|
||||
"/key/service-account/generate",
|
||||
tags=["key management"],
|
||||
|
|
@ -1774,6 +1791,10 @@ async def generate_service_account_key_fn(
|
|||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
data.metadata = _ensure_service_account_id_in_metadata(
|
||||
metadata=data.metadata, key_alias=data.key_alias
|
||||
)
|
||||
|
||||
await validate_team_id_used_in_service_account_request(
|
||||
team_id=data.team_id,
|
||||
prisma_client=prisma_client,
|
||||
|
|
|
|||
|
|
@ -37,6 +37,7 @@ from litellm.proxy.management_endpoints.key_management_endpoints import (
|
|||
_check_team_key_limits,
|
||||
_common_key_generation_helper,
|
||||
_enforce_upperbound_key_params,
|
||||
_ensure_service_account_id_in_metadata,
|
||||
_get_and_validate_existing_key,
|
||||
_list_key_helper,
|
||||
_persist_deleted_verification_tokens,
|
||||
|
|
@ -52,6 +53,7 @@ from litellm.proxy.management_endpoints.key_management_endpoints import (
|
|||
delete_verification_tokens,
|
||||
generate_key_fn,
|
||||
generate_key_helper_fn,
|
||||
generate_service_account_key_fn,
|
||||
key_aliases,
|
||||
key_generation_check,
|
||||
list_keys,
|
||||
|
|
@ -1652,6 +1654,112 @@ async def test_generate_service_account_works_with_team_id():
|
|||
)
|
||||
|
||||
|
||||
def test_ensure_service_account_id_in_metadata_merges_without_overwrite():
|
||||
"""LIT-4635: injecting service_account_id must keep the caller's metadata."""
|
||||
result = _ensure_service_account_id_in_metadata(
|
||||
metadata={
|
||||
"bn-product-id": "12345",
|
||||
"bn-contact": "user@example.com",
|
||||
"bn-environment": "prod",
|
||||
},
|
||||
key_alias="deletemetoo",
|
||||
)
|
||||
assert result == {
|
||||
"bn-product-id": "12345",
|
||||
"bn-contact": "user@example.com",
|
||||
"bn-environment": "prod",
|
||||
"service_account_id": "deletemetoo",
|
||||
}
|
||||
|
||||
|
||||
def test_ensure_service_account_id_prefers_explicit_metadata_value():
|
||||
result = _ensure_service_account_id_in_metadata(
|
||||
metadata={"service_account_id": "explicit", "team": "core"},
|
||||
key_alias="deletemetoo",
|
||||
)
|
||||
assert result == {"service_account_id": "explicit", "team": "core"}
|
||||
|
||||
|
||||
def test_ensure_service_account_id_derives_from_key_alias_when_metadata_empty():
|
||||
assert _ensure_service_account_id_in_metadata(
|
||||
metadata=None, key_alias="deleteme"
|
||||
) == {"service_account_id": "deleteme"}
|
||||
|
||||
|
||||
def test_ensure_service_account_id_noop_when_nothing_to_derive():
|
||||
assert _ensure_service_account_id_in_metadata(metadata=None, key_alias=None) is None
|
||||
assert _ensure_service_account_id_in_metadata(
|
||||
metadata={"a": 1}, key_alias=None
|
||||
) == {"a": 1}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_service_account_generate_gives_custom_hook_merged_metadata():
|
||||
"""LIT-4635 regression: custom_key_generate must receive the caller's
|
||||
metadata plus the injected service_account_id, even when the caller never
|
||||
sets service_account_id itself, so required-metadata validation passes."""
|
||||
received_metadata: dict = {}
|
||||
|
||||
async def recording_custom_key_generate(data):
|
||||
received_metadata.update(data.metadata or {})
|
||||
return {"decision": True}
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", AsyncMock()),
|
||||
patch("litellm.proxy.proxy_server.llm_router", None),
|
||||
patch("litellm.proxy.proxy_server.premium_user", False),
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.user_custom_key_generate",
|
||||
recording_custom_key_generate,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.validate_team_id_used_in_service_account_request",
|
||||
AsyncMock(return_value=True),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.get_team_object",
|
||||
AsyncMock(side_effect=Exception("no team table")),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints._enforce_unique_key_alias",
|
||||
AsyncMock(return_value=None),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn",
|
||||
AsyncMock(
|
||||
return_value={
|
||||
"key": "sk-test-key",
|
||||
"expires": None,
|
||||
"user_id": None,
|
||||
"team_id": "IJ",
|
||||
}
|
||||
),
|
||||
),
|
||||
):
|
||||
await generate_service_account_key_fn(
|
||||
data=GenerateKeyRequest(
|
||||
key_alias="deletemetoo",
|
||||
team_id="IJ",
|
||||
metadata={
|
||||
"bn-product-id": "12345",
|
||||
"bn-contact": "user@example.com",
|
||||
"bn-environment": "prod",
|
||||
},
|
||||
),
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1"
|
||||
),
|
||||
litellm_changed_by=None,
|
||||
)
|
||||
|
||||
assert received_metadata == {
|
||||
"bn-product-id": "12345",
|
||||
"bn-contact": "user@example.com",
|
||||
"bn-environment": "prod",
|
||||
"service_account_id": "deletemetoo",
|
||||
}
|
||||
|
||||
|
||||
@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