diff --git a/litellm/proxy/credential_endpoints/endpoints.py b/litellm/proxy/credential_endpoints/endpoints.py index 2d05270e2ed..b9abe916aed 100644 --- a/litellm/proxy/credential_endpoints/endpoints.py +++ b/litellm/proxy/credential_endpoints/endpoints.py @@ -2,7 +2,7 @@ CRUD endpoints for storing reusable credentials. """ -from typing import Optional +from typing import Any, Optional from fastapi import APIRouter, Depends, HTTPException, Request, Response, Path @@ -12,7 +12,10 @@ from litellm.litellm_core_utils.credential_accessor import CredentialAccessor from litellm.litellm_core_utils.litellm_logging import _get_masked_values from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper +from litellm.proxy.common_utils.encrypt_decrypt_utils import ( + decrypt_value_helper, + encrypt_value_helper, +) from litellm.proxy.utils import handle_exception_on_proxy, jsonify_object from litellm.types.utils import CreateCredentialItem, CredentialItem @@ -40,6 +43,88 @@ class CredentialHelperUtils: ) +def _is_masked_credential_value( + *, + key: str, + submitted_value: Any, + stored_value: Any, + unmasked_length: int = 4, + number_of_asterisks: int = 4, +) -> bool: + if not isinstance(submitted_value, str) or not isinstance(stored_value, str): + return False + + masked_value = _get_masked_values( + {key: stored_value}, + unmasked_length=unmasked_length, + number_of_asterisks=number_of_asterisks, + )[key] + return submitted_value == masked_value + + +def _looks_like_masked_credential_value( + *, + key: str, + submitted_value: Any, + unmasked_length: int = 4, + number_of_asterisks: int = 4, +) -> bool: + if not isinstance(submitted_value, str) or "*" not in submitted_value: + return False + + return _is_masked_credential_value( + key=key, + submitted_value=submitted_value, + stored_value=submitted_value, + unmasked_length=unmasked_length, + number_of_asterisks=number_of_asterisks, + ) + + +def _preserve_masked_credential_values( + db_credential: CredentialItem, + updated_patch: CredentialItem, +) -> CredentialItem: + if not updated_patch.credential_values: + return updated_patch + + existing_credential_values = db_credential.credential_values or {} + filtered_credential_values = {} + + for key, value in updated_patch.credential_values.items(): + existing_value = existing_credential_values.get(key) + decrypted_existing_value = existing_value + if isinstance(existing_value, str): + decrypted_existing_value = decrypt_value_helper( + value=existing_value, + key=key, + exception_type="debug", + return_original_value=False, + ) + if decrypted_existing_value is None: + if _looks_like_masked_credential_value( + key=key, + submitted_value=value, + ): + continue + decrypted_existing_value = existing_value + + if _is_masked_credential_value( + key=key, + submitted_value=value, + stored_value=decrypted_existing_value, + ): + continue + + filtered_credential_values[key] = value + + return CredentialItem( + credential_name=updated_patch.credential_name, + credential_values=filtered_credential_values, + credential_info=updated_patch.credential_info, + ) + + @router.post( "/credentials", dependencies=[Depends(user_api_key_auth)], @@ -331,6 +416,7 @@ async def update_credential( ) if db_credential is None: raise HTTPException(status_code=404, detail="Credential not found in DB.") + credential = _preserve_masked_credential_values(db_credential, credential) merged_credential = update_db_credential(db_credential, credential) credential_object_jsonified = jsonify_object(merged_credential.model_dump()) await prisma_client.db.litellm_credentialstable.update( diff --git a/tests/test_litellm/proxy/management_endpoints/credential_endpoints/test_credential_update.py b/tests/test_litellm/proxy/management_endpoints/credential_endpoints/test_credential_update.py new file mode 100644 index 00000000000..bb703c5fa22 --- /dev/null +++ b/tests/test_litellm/proxy/management_endpoints/credential_endpoints/test_credential_update.py @@ -0,0 +1,114 @@ +from unittest.mock import patch + +from litellm.proxy.credential_endpoints.endpoints import ( + _preserve_masked_credential_values, + update_db_credential, +) +from litellm.types.utils import CredentialItem + + +def test_preserve_masked_credential_value_keeps_existing_secret(): + db_credential = CredentialItem( + credential_name="openai-prod", + credential_values={ + "api_key": "encrypted-stored-api-key", + "api_base": "https://old.example", + }, + credential_info={}, + ) + updated_patch = CredentialItem( + credential_name="openai-prod", + credential_values={ + "api_key": "sk****34", + "api_base": "https://new.example", + }, + credential_info={}, + ) + + with ( + patch( + "litellm.proxy.credential_endpoints.endpoints.decrypt_value_helper", + return_value="sk-live-secret-1234", + ), + patch( + "litellm.proxy.credential_endpoints.endpoints.encrypt_value_helper", + side_effect=lambda value, new_encryption_key=None: f"encrypted:{value}", + ), + ): + sanitized_patch = _preserve_masked_credential_values( + db_credential, updated_patch + ) + merged_credential = update_db_credential(db_credential, sanitized_patch) + + assert sanitized_patch.credential_values == {"api_base": "https://new.example"} + assert merged_credential.credential_values["api_key"] == "encrypted-stored-api-key" + assert merged_credential.credential_values["api_base"] == ( + "encrypted:https://new.example" + ) + + +def test_preserve_masked_credential_value_allows_real_secret_rotation(): + db_credential = CredentialItem( + credential_name="openai-prod", + credential_values={"api_key": "encrypted-stored-api-key"}, + credential_info={}, + ) + updated_patch = CredentialItem( + credential_name="openai-prod", + credential_values={"api_key": "sk-new-secret-5678"}, + credential_info={}, + ) + + with ( + patch( + "litellm.proxy.credential_endpoints.endpoints.decrypt_value_helper", + return_value="sk-live-secret-1234", + ), + patch( + "litellm.proxy.credential_endpoints.endpoints.encrypt_value_helper", + side_effect=lambda value, new_encryption_key=None: f"encrypted:{value}", + ), + ): + sanitized_patch = _preserve_masked_credential_values( + db_credential, updated_patch + ) + merged_credential = update_db_credential(db_credential, sanitized_patch) + + assert sanitized_patch.credential_values == {"api_key": "sk-new-secret-5678"} + assert merged_credential.credential_values["api_key"] == ( + "encrypted:sk-new-secret-5678" + ) + + +def test_preserve_masked_credential_value_on_decryption_failure(): + db_credential = CredentialItem( + credential_name="openai-prod", + credential_values={"api_key": "encrypted-with-old-master-key"}, + credential_info={}, + ) + updated_patch = CredentialItem( + credential_name="openai-prod", + credential_values={"api_key": "sk****34"}, + credential_info={}, + ) + + with ( + patch( + "litellm.proxy.credential_endpoints.endpoints.decrypt_value_helper", + return_value=None, + ), + patch( + "litellm.proxy.credential_endpoints.endpoints.encrypt_value_helper", + side_effect=lambda value, new_encryption_key=None: f"encrypted:{value}", + ), + ): + sanitized_patch = _preserve_masked_credential_values( + db_credential, updated_patch + ) + merged_credential = update_db_credential(db_credential, sanitized_patch) + + assert sanitized_patch.credential_values == {} + assert ( + merged_credential.credential_values["api_key"] + == "encrypted-with-old-master-key" + )