mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(credentials): preserve masked values on update
This commit is contained in:
parent
73e9071311
commit
b195cec77e
2 changed files with 202 additions and 2 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
Loading…
Add table
Reference in a new issue