From bac0506a4a94104d871dd7db753a136c92d31f70 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Tue, 3 Dec 2024 11:07:15 -0800 Subject: [PATCH] fix(key_management_endpoints.py): override metadata field value on update allow user to override tags --- docs/my-website/docs/proxy/virtual_keys.md | 17 ++++-- .../key_management_endpoints.py | 18 +----- .../test_key_management.py | 58 ++++++++----------- 3 files changed, 39 insertions(+), 54 deletions(-) diff --git a/docs/my-website/docs/proxy/virtual_keys.md b/docs/my-website/docs/proxy/virtual_keys.md index 5bbb6b2a002..2107698f322 100644 --- a/docs/my-website/docs/proxy/virtual_keys.md +++ b/docs/my-website/docs/proxy/virtual_keys.md @@ -828,8 +828,18 @@ litellm_settings: #### Spec ```python +key_generation_settings: Optional[StandardKeyGenerationConfig] = None +``` + +#### Types + +```python +class StandardKeyGenerationConfig(TypedDict, total=False): + team_key_generation: TeamUIKeyGenerationConfig + personal_key_generation: PersonalUIKeyGenerationConfig + class TeamUIKeyGenerationConfig(TypedDict): - allowed_team_member_roles: List[str] + allowed_team_member_roles: List[str] # either 'user' or 'admin' required_params: List[str] # require params on `/key/generate` to be set if a team key (team_id in request) is being generated @@ -838,11 +848,6 @@ class PersonalUIKeyGenerationConfig(TypedDict): required_params: List[str] # require params on `/key/generate` to be set if a personal key (no team_id in request) is being generated -class StandardKeyGenerationConfig(TypedDict, total=False): - team_key_generation: TeamUIKeyGenerationConfig - personal_key_generation: PersonalUIKeyGenerationConfig - - class LitellmUserRoles(str, enum.Enum): """ Admin Roles: diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 287de569640..c102b9054c4 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -459,7 +459,6 @@ def prepare_metadata_fields( """ Check LiteLLM_ManagementEndpoint_MetadataFields (proxy/_types.py) for fields that are allowed to be updated """ - if "metadata" not in non_default_values: # allow user to set metadata to none non_default_values["metadata"] = existing_metadata.copy() @@ -469,18 +468,8 @@ def prepare_metadata_fields( try: for k, v in data_json.items(): - if k == "model_tpm_limit" or k == "model_rpm_limit": - if k not in casted_metadata or casted_metadata[k] is None: - casted_metadata[k] = {} - casted_metadata[k].update(v) - - if k == "tags" or k == "guardrails": - if k not in casted_metadata or casted_metadata[k] is None: - casted_metadata[k] = [] - seen = set(casted_metadata[k]) - casted_metadata[k].extend( - x for x in v if x not in seen and not seen.add(x) # type: ignore - ) # prevent duplicates from being added + maintain initial order + if k in LiteLLM_ManagementEndpoint_MetadataFields: + casted_metadata[k] = v except Exception as e: verbose_proxy_logger.exception( @@ -498,10 +487,9 @@ def prepare_key_update_data( ): data_json: dict = data.model_dump(exclude_unset=True) data_json.pop("key", None) - _metadata_fields = ["model_rpm_limit", "model_tpm_limit", "guardrails", "tags"] non_default_values = {} for k, v in data_json.items(): - if k in _metadata_fields: + if k in LiteLLM_ManagementEndpoint_MetadataFields: continue non_default_values[k] = v diff --git a/tests/proxy_admin_ui_tests/test_key_management.py b/tests/proxy_admin_ui_tests/test_key_management.py index 7a2764e3f28..b4fd2735e16 100644 --- a/tests/proxy_admin_ui_tests/test_key_management.py +++ b/tests/proxy_admin_ui_tests/test_key_management.py @@ -695,45 +695,37 @@ def test_personal_key_generation_check(): ) -def test_prepare_metadata_fields(): +@pytest.mark.parametrize( + "update_request_data, non_default_values, existing_metadata, expected_result", + [ + ( + {"metadata": {"test": "new"}}, + {"metadata": {"test": "new"}}, + {"test": "test"}, + {"metadata": {"test": "new"}}, + ), + ( + {"tags": ["new_tag"]}, + {}, + {"tags": ["old_tag"]}, + {"metadata": {"tags": ["new_tag"]}}, + ), + ], +) +def test_prepare_metadata_fields( + update_request_data, non_default_values, existing_metadata, expected_result +): from litellm.proxy.management_endpoints.key_management_endpoints import ( prepare_metadata_fields, ) - new_metadata = {"test": "new"} - old_metadata = {"test": "test"} - args = { "data": UpdateKeyRequest( - key_alias=None, - duration=None, - models=[], - spend=None, - max_budget=None, - user_id=None, - team_id=None, - max_parallel_requests=None, - metadata=new_metadata, - tpm_limit=None, - rpm_limit=None, - budget_duration=None, - allowed_cache_controls=[], - soft_budget=None, - config={}, - permissions={}, - model_max_budget={}, - send_invite_email=None, - model_rpm_limit=None, - model_tpm_limit=None, - guardrails=None, - blocked=None, - aliases={}, - key="sk-1qGQUJJTcljeaPfzgWRrXQ", - tags=None, + key="sk-1qGQUJJTcljeaPfzgWRrXQ", **update_request_data ), - "non_default_values": {"metadata": new_metadata}, - "existing_metadata": {"tags": None, **old_metadata}, + "non_default_values": non_default_values, + "existing_metadata": existing_metadata, } - non_default_values = prepare_metadata_fields(**args) - assert non_default_values == {"metadata": new_metadata} + updated_non_default_values = prepare_metadata_fields(**args) + assert updated_non_default_values == expected_result