mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
* fix(model_armor): sanitize error details by default Generated with AI Co-Authored-By: Claude Code * fix(model_armor): sanitize handler-raised HTTP errors and redact scanned content in guardrail logging The async HTTP handler raises MaskedHTTPStatusError on any non-2xx via raise_for_status, so the non-200 branch in make_model_armor_request never ran against a live API and the raw upstream body reached callers and logs. Catch the raised error and build the sanitized detail from the response status Replace the empty-dict guardrail logging payload with field-level redaction of the keys that echo scanned content (text, sanitizedText, findings) so guardrail traces keep filter states and block reasons while scanned content stays out Restore the upstream status code in the sanitized error detail, read guardrail metadata from the same key the hooks write, and keep guardrail_status within its typed literal values * fix(model_armor): bound redactor recursion depth and allowlist it in the recursion detector _redact_scanned_content walks provider JSON bounded by _REDACT_MAX_DEPTH=20 and fails closed by returning the redaction sentinel at the cap * fix(model_armor): honor fail_on_error for upstream API failures API failures now raise a dedicated ModelArmorAPIError so hooks can tell them apart from content-block HTTPExceptions; fail_on_error=False lets the request proceed on a Model Armor outage again while fail-closed configs get the same sanitized 400 as before Also addresses review notes: sanitize_error_detail constructor annotation matches the nullable config field, redaction is owned by the metadata write sites so _process_response no longer re-applies it, and the request and response debug log branches move into helpers * test(model_armor): cover fail_on_error routing on during-call, post-call, streaming, and file-scan paths * chore: remove accidentally committed pytest cache files * fix(model_armor): keep sanitize_error_detail coerced across in-memory config reloads update_in_memory_litellm_params assigns raw LitellmParams fields, so a hot reloaded config carrying an explicit null would silently disable sanitization; re-apply the only-explicit-False-opts-out coercion after the update * fix(model_armor): redact matched malicious URIs and reuse the shared recursion depth constant maliciousUriMatchedItems echoes the caller-supplied URL including path and query, so it joins the scanned-content key set; the redactor depth cap now comes from DEFAULT_MAX_RECURSE_DEPTH in litellm constants instead of a local literal * fix(model_armor): keep API failures out of the intervention trace status Fail-closed upstream failures re-raise ModelArmorAPIError instead of converting to HTTPException(400), so the shared guardrail logging keeps recording them as guardrail_failed_to_respond while content blocks stay guardrail_intervened. Callers see the same 500 shape as before this PR, with the sanitized message * chore(model_armor): drop explanatory comment per repository comment policy --------- Co-authored-by: eugene-yao-zocdoc <eugene.yao@zocdoc.com>
142 lines
6.8 KiB
Python
142 lines
6.8 KiB
Python
import ast
|
|
import os
|
|
|
|
IGNORE_FUNCTIONS = [
|
|
"_format_type",
|
|
"_remove_additional_properties",
|
|
"_remove_strict_from_schema",
|
|
"filter_schema_fields",
|
|
"text_completion",
|
|
"_check_for_os_environ_vars",
|
|
"clean_message",
|
|
"unpack_defs",
|
|
"convert_anyof_null_to_nullable", # has a set max depth
|
|
"add_object_type",
|
|
"strip_field",
|
|
"_transform_prompt",
|
|
"mask_dict",
|
|
"_serialize", # we now set a max depth for this
|
|
"_sanitize_request_body_for_spend_logs_payload", # testing added for circular reference
|
|
"_sanitize_value", # testing added for circular reference
|
|
"set_schema_property_ordering", # testing added for infinite recursion
|
|
"process_items", # testing added for infinite recursion + max depth set.
|
|
"_can_object_call_model", # max depth set.
|
|
"encode_unserializable_types", # max depth set.
|
|
"filter_value_from_dict", # max depth set.
|
|
"normalize_json_schema_types", # max depth set.
|
|
"_extract_fields_recursive", # max depth set.
|
|
"_remove_json_schema_refs", # max depth set.,
|
|
"_convert_schema_types", # max depth set.,
|
|
"_fix_enum_empty_strings", # max depth set.,
|
|
"get_access_token", # max depth set.,
|
|
"_redact_base64", # max depth set.
|
|
"_contains_vision_content", # max depth set.
|
|
"_read_all_bytes", # max depth set.
|
|
"_fix_enum_types", # max depth set.
|
|
"_collect_argument_paths", # max depth set.
|
|
"_split_text", # max depth set.
|
|
"_mask_sequence", # max depth set.
|
|
"_walk_payload", # max depth set (DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER).
|
|
"_delete_nested_value_custom", # max depth set (bounded by number of path segments).
|
|
"filter_exceptions_from_params", # max depth set (default 20) to prevent infinite recursion.
|
|
"__getattr__", # lazy loading pattern in litellm/__init__.py with proper caching to prevent infinite recursion.
|
|
"_validate_inheritance_chain", # max depth set (default 100) to prevent infinite recursion in policy inheritance validation.
|
|
"_basic_json_schema_validate", # max depth set.
|
|
"extract_text_from_a2a_message", # max depth set (default 10) to prevent infinite recursion in A2A message parsing.
|
|
"_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.
|
|
"_redact_sensitive_litellm_params", # max depth set (default 10).
|
|
"_redact_secret_values_in_obj", # max depth set (default 10, _REDACT_SECRET_MAX_DEPTH); fails closed by returning "REDACTED" at the cap.
|
|
"_resolve", # OCI: $ref resolver bounded by `resolving_stack` cycle guard.
|
|
"resolve_oci_schema_anyof", # OCI: bounded by JSON-schema tree depth (no cycles possible in well-formed input).
|
|
"sanitize_oci_schema", # OCI: bounded by JSON-schema tree depth.
|
|
"_freeze_for_dedupe", # OTEL: max depth set (default 16, _FREEZE_MAX_DEPTH); fails closed by returning repr(value) at the cap.
|
|
"apply_json_merge_patch", # max depth set (_MAX_MERGE_DEPTH=64); fails closed by raising ValueError at the cap.
|
|
"_filter_mcp_argument_value", # max depth set (DEFAULT_MAX_RECURSE_DEPTH); fails closed by blocking the MCP call at the cap.
|
|
"_redact_scanned_content", # max depth set (DEFAULT_MAX_RECURSE_DEPTH); fails closed by returning "[REDACTED]" at the cap.
|
|
]
|
|
|
|
|
|
class RecursiveFunctionFinder(ast.NodeVisitor):
|
|
def __init__(self):
|
|
self.recursive_functions = []
|
|
self.ignored_recursive_functions = []
|
|
|
|
def visit_FunctionDef(self, node):
|
|
# Check if the function calls itself
|
|
if any(self._is_recursive_call(node, call) for call in ast.walk(node)):
|
|
if node.name in IGNORE_FUNCTIONS:
|
|
self.ignored_recursive_functions.append(node.name)
|
|
else:
|
|
self.recursive_functions.append(node.name)
|
|
self.generic_visit(node)
|
|
|
|
def _is_recursive_call(self, func_node, call_node):
|
|
# Check if the call node is a function call
|
|
if not isinstance(call_node, ast.Call):
|
|
return False
|
|
|
|
# Case 1: Direct function call (e.g., my_func())
|
|
if isinstance(call_node.func, ast.Name) and call_node.func.id == func_node.name:
|
|
return True
|
|
|
|
# Case 2: Method call with self (e.g., self.my_func())
|
|
if isinstance(call_node.func, ast.Attribute) and isinstance(
|
|
call_node.func.value, ast.Name
|
|
):
|
|
return (
|
|
call_node.func.value.id == "self"
|
|
and call_node.func.attr == func_node.name
|
|
)
|
|
|
|
return False
|
|
|
|
|
|
def find_recursive_functions_in_file(file_path):
|
|
with open(file_path, "r") as file:
|
|
tree = ast.parse(file.read(), filename=file_path)
|
|
finder = RecursiveFunctionFinder()
|
|
finder.visit(tree)
|
|
return finder.recursive_functions, finder.ignored_recursive_functions
|
|
|
|
|
|
def find_recursive_functions_in_directory(directory):
|
|
recursive_functions = {}
|
|
ignored_recursive_functions = {}
|
|
for root, _, files in os.walk(directory):
|
|
for file in files:
|
|
print("file: ", file)
|
|
if file.endswith(".py"):
|
|
file_path = os.path.join(root, file)
|
|
functions, ignored = find_recursive_functions_in_file(file_path)
|
|
if functions:
|
|
recursive_functions[file_path] = functions
|
|
if ignored:
|
|
ignored_recursive_functions[file_path] = ignored
|
|
return recursive_functions, ignored_recursive_functions
|
|
|
|
|
|
if __name__ == "__main__":
|
|
# Example usage
|
|
# raise exception if any recursive functions are found, except for the ignored ones
|
|
# this is used in the CI/CD pipeline to prevent recursive functions from being merged
|
|
|
|
directory_path = "./litellm"
|
|
recursive_functions, ignored_recursive_functions = (
|
|
find_recursive_functions_in_directory(directory_path)
|
|
)
|
|
print("UNIGNORED RECURSIVE FUNCTIONS: ", recursive_functions)
|
|
print("IGNORED RECURSIVE FUNCTIONS: ", ignored_recursive_functions)
|
|
|
|
if len(recursive_functions) > 0:
|
|
# raise exception if any recursive functions are found
|
|
for file, functions in recursive_functions.items():
|
|
print(
|
|
f"🚨 Unignored recursive functions found in {file}: {functions}. THIS IS REALLY BAD, it has caused CPU Usage spikes in the past. Only keep this if it's ABSOLUTELY necessary."
|
|
)
|
|
file, functions = list(recursive_functions.items())[0]
|
|
raise Exception(
|
|
f"🚨 Unignored recursive functions found include {file}: {functions}. THIS IS REALLY BAD, it has caused CPU Usage spikes in the past. Only keep this if it's ABSOLUTELY necessary."
|
|
)
|