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:
shivam 2026-07-21 02:18:23 +00:00
parent 2ea09a9969
commit dbdfcf9b8c
2 changed files with 129 additions and 0 deletions

View file

@ -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,

View file

@ -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