From dbdfcf9b8c40eabaa92558d947dc4e0713087e62 Mon Sep 17 00:00:00 2001 From: shivam Date: Tue, 21 Jul 2026 02:18:23 +0000 Subject: [PATCH 1/2] 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> --- .../key_management_endpoints.py | 21 ++++ .../test_key_management_endpoints.py | 108 ++++++++++++++++++ 2 files changed, 129 insertions(+) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 01f4e040e58..e0602d5d073 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -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, diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index dffca3093fa..ccd9fb1febc 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -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 From 0eb952894d4dba0cde920434a0f3167455538e11 Mon Sep 17 00:00:00 2001 From: shivam Date: Tue, 21 Jul 2026 02:24:45 +0000 Subject: [PATCH 2/2] style: ruff format key_management_endpoints Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../management_endpoints/key_management_endpoints.py | 8 ++------ 1 file changed, 2 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index e0602d5d073..47a638e2b4d 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -1671,9 +1671,7 @@ 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: +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 @@ -1791,9 +1789,7 @@ 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 - ) + 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,