diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 08b3238bc6b..b54f23cdd1e 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -4123,6 +4123,12 @@ LiteLLM_ManagementEndpoint_MetadataFields_Premium = [ "allowed_passthrough_routes", ] +# Metadata keys that are immutable once set: preserved when an update omits them, +# and rejected (400) when an update tries to change them. +LiteLLM_Reserved_Metadata_Fields = [ + "service_account_id", +] + class ProviderBudgetResponseObject(LiteLLMPydanticObjectBase): """ diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 1fdf62cdb5b..2ebe1a86f96 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -1523,6 +1523,22 @@ def prepare_metadata_fields( casted_metadata = cast(dict, non_default_values["metadata"]) + # Reserved metadata fields are immutable once set. Preserve the existing value + # when omitted, reject any explicit attempt to change it (including null). + for reserved_field in LiteLLM_Reserved_Metadata_Fields: + existing_value = existing_metadata.get(reserved_field) + if existing_value is None: + continue + if casted_metadata is None or ( + reserved_field in casted_metadata + and casted_metadata[reserved_field] != existing_value + ): + raise HTTPException( + status_code=400, + detail=f"{reserved_field} is immutable once set and cannot be changed via update.", + ) + casted_metadata[reserved_field] = existing_value + data_json = data.model_dump(exclude_unset=True, exclude_none=True) try: diff --git a/tests/proxy_unit_tests/test_proxy_utils.py b/tests/proxy_unit_tests/test_proxy_utils.py index 8703331ea83..feca4251a45 100644 --- a/tests/proxy_unit_tests/test_proxy_utils.py +++ b/tests/proxy_unit_tests/test_proxy_utils.py @@ -686,12 +686,13 @@ async def test_proxy_config_update_from_db(): @pytest.mark.asyncio async def test_prepare_key_update_data(): - from litellm.proxy._types import UpdateKeyRequest + from litellm.proxy._types import LiteLLM_VerificationToken, UpdateKeyRequest from litellm.proxy.management_endpoints.key_management_endpoints import ( prepare_key_update_data, ) - existing_key_row = MagicMock() + existing_key_row = MagicMock(spec=LiteLLM_VerificationToken) + existing_key_row.metadata = {} data = UpdateKeyRequest(key="test_key", models=["gpt-4"], duration="120s") updated_data = await prepare_key_update_data(data, existing_key_row) assert "expires" in updated_data 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 8e6fd78a6f6..647be49e0d1 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 @@ -1211,6 +1211,131 @@ async def test_update_service_account_works_with_team_id(): await prepare_key_update_data(data=data, existing_key_row=existing_key) +@pytest.mark.asyncio +async def test_update_preserves_service_account_id_when_metadata_replaced(): + """ + Regression: /key/update wholesale-replaced metadata, silently dropping + service_account_id. The pre-call check then treated the key as a regular + key and bypassed service_account_settings.enforced_params. + """ + data = UpdateKeyRequest( + key="sk-1", + metadata={"unrelated": "value"}, + team_id="IJ", + ) + existing_key = LiteLLM_VerificationToken( + token="hashed", + team_id="IJ", + metadata={"service_account_id": "sa-123"}, + ) + + result = await prepare_key_update_data(data=data, existing_key_row=existing_key) + + assert result["metadata"]["service_account_id"] == "sa-123" + assert result["metadata"]["unrelated"] == "value" + + +@pytest.mark.asyncio +async def test_update_rejects_service_account_id_overwrite(): + """ + Once assigned, a key's service_account_id is an identity marker — rebinding + it via update would break spend attribution. Reject rather than silently + ignore so scripted callers surface the bug. + """ + data = UpdateKeyRequest( + key="sk-1", + metadata={"service_account_id": "sa-new"}, + team_id="IJ", + ) + existing_key = LiteLLM_VerificationToken( + token="hashed", + team_id="IJ", + metadata={"service_account_id": "sa-old"}, + ) + + with pytest.raises(HTTPException) as exc_info: + await prepare_key_update_data(data=data, existing_key_row=existing_key) + assert exc_info.value.status_code == 400 + + +@pytest.mark.asyncio +async def test_update_allows_matching_service_account_id(): + """Resending the same value (e.g. UI round-trip) is a no-op, not a conflict.""" + data = UpdateKeyRequest( + key="sk-1", + metadata={"service_account_id": "sa-123", "other": "value"}, + team_id="IJ", + ) + existing_key = LiteLLM_VerificationToken( + token="hashed", + team_id="IJ", + metadata={"service_account_id": "sa-123"}, + ) + + result = await prepare_key_update_data(data=data, existing_key_row=existing_key) + + assert result["metadata"]["service_account_id"] == "sa-123" + assert result["metadata"]["other"] == "value" + + +@pytest.mark.asyncio +async def test_update_rejects_explicit_null_service_account_id(): + """ + Explicit null is an attempt to clear — not an omission. Silently ignoring + it would let a caller think they cleared the field when they didn't, so + treat it the same as any other rebind attempt and return 400. + """ + data = UpdateKeyRequest( + key="sk-1", + metadata={"service_account_id": None}, + team_id="IJ", + ) + existing_key = LiteLLM_VerificationToken( + token="hashed", + team_id="IJ", + metadata={"service_account_id": "sa-old"}, + ) + + with pytest.raises(HTTPException) as exc_info: + await prepare_key_update_data(data=data, existing_key_row=existing_key) + assert exc_info.value.status_code == 400 + + +@pytest.mark.asyncio +async def test_update_rejects_whole_metadata_null_on_service_account_key(): + """ + `metadata: null` on a key with a reserved field would have dereferenced None + inside the reserved-field loop and returned a 500. Must surface as a 400 + since the effect would be to clear an immutable field. + """ + data = UpdateKeyRequest(key="sk-1", metadata=None, team_id="IJ") + existing_key = LiteLLM_VerificationToken( + token="hashed", + team_id="IJ", + metadata={"service_account_id": "sa-123"}, + ) + + with pytest.raises(HTTPException) as exc_info: + await prepare_key_update_data(data=data, existing_key_row=existing_key) + assert exc_info.value.status_code == 400 + + +@pytest.mark.asyncio +async def test_update_without_metadata_still_preserves_existing(): + """Omitting metadata entirely must not drop existing metadata fields.""" + data = UpdateKeyRequest(key="sk-1", max_budget=100) + existing_key = LiteLLM_VerificationToken( + token="hashed", + team_id="IJ", + metadata={"service_account_id": "sa-123", "other": "kept"}, + ) + + result = await prepare_key_update_data(data=data, existing_key_row=existing_key) + + assert result["metadata"]["service_account_id"] == "sa-123" + assert result["metadata"]["other"] == "kept" + + @pytest.mark.asyncio async def test_prepare_key_update_data_duration_never_expires(): """Test that duration="-1" sets expires to None (never expires).""" @@ -8007,6 +8132,7 @@ async def test_update_key_non_budget_fields_allowed_for_internal_user(monkeypatc mock_existing_key.max_budget = 10.0 mock_existing_key.key_alias = None mock_existing_key.models = [] + mock_existing_key.metadata = {} mock_existing_key.model_dump.return_value = { "token": test_hashed_token, "user_id": "internal_user", @@ -8089,9 +8215,7 @@ async def test_update_key_non_budget_rejects_cross_user_modification(monkeypatch ) mock_prisma_client = AsyncMock() - test_hashed_token = ( - "cafebabe" * 8 - ) + test_hashed_token = "cafebabe" * 8 mock_existing_key = MagicMock() mock_existing_key.token = test_hashed_token