mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-05 10:24:03 +00:00
Merge pull request #26390 from BerriAI/litellm_guardrail_param_masking
[Fix] Guardrail param handling in list and submission endpoints
This commit is contained in:
@@ -3517,6 +3517,8 @@ def _get_masked_values(
|
||||
mask_all_values: bool = False,
|
||||
unmasked_length: int = 4,
|
||||
number_of_asterisks: Optional[int] = 4,
|
||||
_depth: int = 0,
|
||||
_max_depth: int = 20,
|
||||
) -> dict:
|
||||
"""
|
||||
Internal debugging helper function
|
||||
@@ -3533,38 +3535,49 @@ def _get_masked_values(
|
||||
"key",
|
||||
"secret",
|
||||
"vertex_credentials",
|
||||
"credentials",
|
||||
"password",
|
||||
"passwd",
|
||||
]
|
||||
|
||||
def _mask_value(v: Any) -> Any:
|
||||
if isinstance(v, dict):
|
||||
if _depth >= _max_depth:
|
||||
return v
|
||||
return _get_masked_values(
|
||||
v,
|
||||
ignore_sensitive_values=ignore_sensitive_values,
|
||||
mask_all_values=mask_all_values,
|
||||
unmasked_length=unmasked_length,
|
||||
number_of_asterisks=number_of_asterisks,
|
||||
_depth=_depth + 1,
|
||||
_max_depth=_max_depth,
|
||||
)
|
||||
if not isinstance(v, str):
|
||||
return v
|
||||
if len(v) <= unmasked_length:
|
||||
return "*****"
|
||||
if number_of_asterisks is not None:
|
||||
return (
|
||||
v[: unmasked_length // 2]
|
||||
+ "*" * number_of_asterisks
|
||||
+ v[-unmasked_length // 2 :]
|
||||
)
|
||||
return (
|
||||
v[: unmasked_length // 2]
|
||||
+ "*" * (len(v) - unmasked_length)
|
||||
+ v[-unmasked_length // 2 :]
|
||||
)
|
||||
|
||||
return {
|
||||
k: (
|
||||
# If ignore_sensitive_values is True, or if this key doesn't contain sensitive keywords, return original value
|
||||
v
|
||||
if ignore_sensitive_values
|
||||
or not any(
|
||||
sensitive_keyword in k.lower()
|
||||
for sensitive_keyword in sensitive_keywords
|
||||
)
|
||||
else (
|
||||
# Apply masking to sensitive keys
|
||||
(
|
||||
v[: unmasked_length // 2]
|
||||
+ "*" * number_of_asterisks
|
||||
+ v[-unmasked_length // 2 :]
|
||||
)
|
||||
if (
|
||||
isinstance(v, str)
|
||||
and len(v) > unmasked_length
|
||||
and number_of_asterisks is not None
|
||||
)
|
||||
else (
|
||||
(
|
||||
v[: unmasked_length // 2]
|
||||
+ "*" * (len(v) - unmasked_length)
|
||||
+ v[-unmasked_length // 2 :]
|
||||
)
|
||||
if (isinstance(v, str) and len(v) > unmasked_length)
|
||||
else ("*****" if isinstance(v, str) else v)
|
||||
)
|
||||
)
|
||||
else _mask_value(v)
|
||||
)
|
||||
for k, v in sensitive_object.items()
|
||||
}
|
||||
|
||||
@@ -136,10 +136,11 @@ async def list_guardrails():
|
||||
@router.get(
|
||||
"/v2/guardrails/list",
|
||||
tags=["Guardrails"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=ListGuardrailsResponse,
|
||||
)
|
||||
async def list_guardrails_v2():
|
||||
async def list_guardrails_v2(
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
List the guardrails that are available in the database using GuardrailRegistry
|
||||
|
||||
@@ -179,13 +180,29 @@ async def list_guardrails_v2():
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail="Prisma client not initialized")
|
||||
|
||||
is_admin = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
|
||||
|
||||
try:
|
||||
guardrails = await GUARDRAIL_REGISTRY.get_all_guardrails_from_db(
|
||||
prisma_client=prisma_client
|
||||
)
|
||||
|
||||
excluded_guardrail_ids: set = set()
|
||||
if not is_admin:
|
||||
caller_team_ids = await _get_user_team_ids(user_api_key_dict)
|
||||
allowed: List[Guardrail] = []
|
||||
for g in guardrails:
|
||||
g_team_id = g.get("team_id")
|
||||
if g_team_id is None or g_team_id in caller_team_ids:
|
||||
allowed.append(g)
|
||||
else:
|
||||
gid = g.get("guardrail_id")
|
||||
if gid:
|
||||
excluded_guardrail_ids.add(gid)
|
||||
guardrails = allowed
|
||||
|
||||
guardrail_configs: List[GuardrailInfoResponse] = []
|
||||
seen_guardrail_ids = set()
|
||||
seen_guardrail_ids: set = excluded_guardrail_ids.copy()
|
||||
for guardrail in guardrails:
|
||||
litellm_params: Optional[Union[LitellmParams, dict]] = guardrail.get(
|
||||
"litellm_params"
|
||||
@@ -221,34 +238,39 @@ async def list_guardrails_v2():
|
||||
# get guardrails initialized on litellm config.yaml
|
||||
in_memory_guardrails = IN_MEMORY_GUARDRAIL_HANDLER.list_in_memory_guardrails()
|
||||
for guardrail in in_memory_guardrails:
|
||||
# only add guardrails that are not in DB guardrail list already
|
||||
if guardrail.get("guardrail_id") not in seen_guardrail_ids:
|
||||
in_memory_litellm_params_raw = guardrail.get("litellm_params")
|
||||
in_memory_litellm_params_dict = (
|
||||
in_memory_litellm_params_raw.model_dump(exclude_none=True)
|
||||
if isinstance(in_memory_litellm_params_raw, LitellmParams)
|
||||
else in_memory_litellm_params_raw
|
||||
) or {}
|
||||
masked_in_memory_litellm_params = _get_masked_values(
|
||||
in_memory_litellm_params_dict,
|
||||
unmasked_length=4,
|
||||
number_of_asterisks=4,
|
||||
gid = guardrail.get("guardrail_id")
|
||||
if gid in seen_guardrail_ids:
|
||||
continue
|
||||
if not is_admin:
|
||||
g_team_id = guardrail.get("team_id")
|
||||
if g_team_id is not None and g_team_id not in caller_team_ids:
|
||||
continue
|
||||
in_memory_litellm_params_raw = guardrail.get("litellm_params")
|
||||
in_memory_litellm_params_dict = (
|
||||
in_memory_litellm_params_raw.model_dump(exclude_none=True)
|
||||
if isinstance(in_memory_litellm_params_raw, LitellmParams)
|
||||
else in_memory_litellm_params_raw
|
||||
) or {}
|
||||
masked_in_memory_litellm_params = _get_masked_values(
|
||||
in_memory_litellm_params_dict,
|
||||
unmasked_length=4,
|
||||
number_of_asterisks=4,
|
||||
)
|
||||
masked_in_memory_litellm_params_typed = (
|
||||
BaseLitellmParams(**masked_in_memory_litellm_params)
|
||||
if masked_in_memory_litellm_params
|
||||
else None
|
||||
)
|
||||
guardrail_configs.append(
|
||||
GuardrailInfoResponse(
|
||||
guardrail_id=guardrail.get("guardrail_id"),
|
||||
guardrail_name=guardrail.get("guardrail_name"),
|
||||
litellm_params=masked_in_memory_litellm_params_typed,
|
||||
guardrail_info=dict(guardrail.get("guardrail_info") or {}),
|
||||
guardrail_definition_location="config",
|
||||
)
|
||||
masked_in_memory_litellm_params_typed = (
|
||||
BaseLitellmParams(**masked_in_memory_litellm_params)
|
||||
if masked_in_memory_litellm_params
|
||||
else None
|
||||
)
|
||||
guardrail_configs.append(
|
||||
GuardrailInfoResponse(
|
||||
guardrail_id=guardrail.get("guardrail_id"),
|
||||
guardrail_name=guardrail.get("guardrail_name"),
|
||||
litellm_params=masked_in_memory_litellm_params_typed,
|
||||
guardrail_info=dict(guardrail.get("guardrail_info") or {}),
|
||||
guardrail_definition_location="config",
|
||||
)
|
||||
)
|
||||
seen_guardrail_ids.add(guardrail.get("guardrail_id"))
|
||||
)
|
||||
seen_guardrail_ids.add(gid)
|
||||
|
||||
return ListGuardrailsResponse(guardrails=guardrail_configs)
|
||||
except Exception as e:
|
||||
@@ -751,15 +773,21 @@ async def _get_user_team_ids(user_api_key_dict: UserAPIKeyAuth) -> List[str]:
|
||||
|
||||
|
||||
def _row_to_submission_item(row: Any) -> GuardrailSubmissionItem:
|
||||
from litellm.litellm_core_utils.litellm_logging import _get_masked_values
|
||||
|
||||
guardrail_info = _parse_json_field(row.guardrail_info) or {}
|
||||
team_guardrail = row.team_id is not None
|
||||
raw_params = _parse_json_field(row.litellm_params) or {}
|
||||
masked_params = _get_masked_values(
|
||||
raw_params, unmasked_length=4, number_of_asterisks=4
|
||||
)
|
||||
return GuardrailSubmissionItem(
|
||||
guardrail_id=row.guardrail_id,
|
||||
guardrail_name=row.guardrail_name,
|
||||
status=row.status or "active",
|
||||
team_id=row.team_id,
|
||||
team_guardrail=team_guardrail,
|
||||
litellm_params=_parse_json_field(row.litellm_params),
|
||||
litellm_params=masked_params,
|
||||
guardrail_info=guardrail_info,
|
||||
submitted_by_user_id=guardrail_info.get("submitted_by_user_id"),
|
||||
submitted_by_email=guardrail_info.get("submitted_by_email"),
|
||||
@@ -966,6 +994,7 @@ async def approve_guardrail_submission(
|
||||
"guardrail_name": row.guardrail_name,
|
||||
"litellm_params": litellm_params,
|
||||
"guardrail_info": guardrail_info or {},
|
||||
"team_id": row.team_id,
|
||||
}
|
||||
try:
|
||||
IN_MEMORY_GUARDRAIL_HANDLER.initialize_guardrail(
|
||||
|
||||
@@ -45,6 +45,7 @@ IGNORE_FUNCTIONS = [
|
||||
"_convert_to_json_serializable_dict", # max depth set (default 20) and circular reference protection to prevent infinite recursion.
|
||||
"dict", # max depth set. _LiteLLMParamsDictView.dict() calls builtin dict(), not itself.
|
||||
"_read_image_bytes", # max depth set.
|
||||
"_get_masked_values", # max depth set (default 20) to prevent infinite recursion while masking nested sensitive config dicts.
|
||||
]
|
||||
|
||||
|
||||
|
||||
@@ -140,7 +140,8 @@ async def test_list_guardrails_v2_with_db_and_config(
|
||||
mock_in_memory_handler,
|
||||
)
|
||||
|
||||
response = await list_guardrails_v2()
|
||||
admin_auth = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
response = await list_guardrails_v2(user_api_key_dict=admin_auth)
|
||||
|
||||
assert len(response.guardrails) == 2
|
||||
|
||||
@@ -194,7 +195,8 @@ async def test_list_guardrails_v2_masks_sensitive_data_in_db_guardrails(mocker):
|
||||
mock_in_memory_handler,
|
||||
)
|
||||
|
||||
response = await list_guardrails_v2()
|
||||
admin_auth = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
response = await list_guardrails_v2(user_api_key_dict=admin_auth)
|
||||
|
||||
assert len(response.guardrails) == 1
|
||||
guardrail = response.guardrails[0]
|
||||
@@ -248,7 +250,8 @@ async def test_list_guardrails_v2_masks_sensitive_data_in_config_guardrails(mock
|
||||
mock_in_memory_handler,
|
||||
)
|
||||
|
||||
response = await list_guardrails_v2()
|
||||
admin_auth = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
response = await list_guardrails_v2(user_api_key_dict=admin_auth)
|
||||
|
||||
assert len(response.guardrails) == 1
|
||||
guardrail = response.guardrails[0]
|
||||
|
||||
Reference in New Issue
Block a user