diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index cb8b9ec0395..2f392c48e9d 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -641,7 +641,9 @@ def _sanitize_request_body_for_spend_logs_payload( return {k: _sanitize_value(v) for k, v in request_body.items()} -def _convert_to_json_serializable_dict(obj: Any) -> Any: +def _convert_to_json_serializable_dict( + obj: Any, visited: Optional[set] = None, max_depth: int = 20 +) -> Any: """ Convert object to JSON-serializable dict, handling Pydantic models safely. @@ -650,23 +652,55 @@ def _convert_to_json_serializable_dict(obj: Any) -> Any: Args: obj: Object to convert (dict, list, Pydantic model, or primitive) + visited: Set of object IDs to track circular references + max_depth: Maximum recursion depth to prevent infinite recursion Returns: JSON-serializable version of the object """ - if isinstance(obj, BaseModel): - # Use Pydantic's model_dump() instead of pickle - return obj.model_dump() - elif isinstance(obj, dict): - return {k: _convert_to_json_serializable_dict(v) for k, v in obj.items()} - elif isinstance(obj, list): - return [_convert_to_json_serializable_dict(item) for item in obj] - elif hasattr(obj, "__dict__"): - # Handle objects with __dict__ attribute - return _convert_to_json_serializable_dict(obj.__dict__) - else: - # Primitives (str, int, float, bool, None) pass through - return obj + if max_depth <= 0: + # Return a placeholder if max depth is exceeded + return "" + + if visited is None: + visited = set() + + # Get the object's memory address to track visited objects + obj_id = id(obj) + if obj_id in visited: + # Circular reference detected, return placeholder + return "" + + # Only track mutable objects (dict, list, objects with __dict__) + if isinstance(obj, (dict, list)) or hasattr(obj, "__dict__"): + visited.add(obj_id) + + try: + if isinstance(obj, BaseModel): + # Use Pydantic's model_dump() instead of pickle + result = obj.model_dump() + # Recursively process the dumped dict + return _convert_to_json_serializable_dict(result, visited, max_depth - 1) + elif isinstance(obj, dict): + return { + k: _convert_to_json_serializable_dict(v, visited, max_depth - 1) + for k, v in obj.items() + } + elif isinstance(obj, list): + return [ + _convert_to_json_serializable_dict(item, visited, max_depth - 1) + for item in obj + ] + elif hasattr(obj, "__dict__"): + # Handle objects with __dict__ attribute + return _convert_to_json_serializable_dict(obj.__dict__, visited, max_depth - 1) + else: + # Primitives (str, int, float, bool, None) pass through + return obj + finally: + # Remove from visited set when done processing this object + if obj_id in visited: + visited.remove(obj_id) def _get_proxy_server_request_for_spend_logs_payload( diff --git a/tests/code_coverage_tests/recursive_detector.py b/tests/code_coverage_tests/recursive_detector.py index 71e7798b09e..d6bf1941a08 100644 --- a/tests/code_coverage_tests/recursive_detector.py +++ b/tests/code_coverage_tests/recursive_detector.py @@ -42,6 +42,7 @@ IGNORE_FUNCTIONS = [ "_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. ]