mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-20 00:11:50 +00:00
Merge 0eb952894d into af6dc1db08
This commit is contained in:
commit
953b97dfd2
2 changed files with 125 additions and 0 deletions
|
|
@ -1936,6 +1936,21 @@ 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"],
|
||||
|
|
@ -2041,6 +2056,8 @@ 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,
|
||||
|
|
|
|||
|
|
@ -38,6 +38,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,
|
||||
|
|
@ -53,6 +54,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,
|
||||
|
|
@ -1728,6 +1730,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