diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index c09369b9b2c..ba131ecac4c 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -1,6 +1,6 @@ { "reportAny": { - "limit": 22344 + "limit": 22343 }, "reportArgumentType": { "limit": 2578 diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index c4a350fb285..1b1e3c70110 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -38,6 +38,8 @@ from litellm.proxy._types import ( CommonProxyErrors, LitellmDataForBackendLLMCall, LitellmUserRoles, + ProxyErrorTypes, + ProxyException, SpecialHeaders, TeamCallbackMetadata, UserAPIKeyAuth, @@ -348,6 +350,36 @@ def reject_url_valued_destination(field: str, value: str) -> None: ) +_METADATA_JSON_TYPE_NAMES: Final[Mapping[type, str]] = MappingProxyType( + {bool: "a boolean", int: "an integer", float: "a number", str: "a string", list: "an array"} +) + + +def _invalid_metadata_type_error(field: str, value: object) -> ProxyException: + received_type: Final = _METADATA_JSON_TYPE_NAMES.get(type(value), f"a {type(value).__name__}") + return ProxyException( + message=f"Invalid type for '{field}': expected an object, but got {received_type} instead.", + type=ProxyErrorTypes.bad_request_error, + param=field, + code=400, + ) + + +def _normalized_metadata_object(field: str, value: object) -> Mapping[str, Any]: + """Return ``value`` as a metadata object or raise a 400 like OpenAI does. + + A JSON string that parses to an object is accepted because multipart/form-data + and ``extra_body`` callers can only send metadata as a string. The caller pops + the raw value from the request body before validating so the failure-logging + hooks that inspect the body afterwards don't crash on it and mask the 400 as a 500. + """ + if isinstance(value, dict): + return value + if isinstance(value, str) and isinstance((parsed := safe_json_loads(value)), dict): + return parsed + raise _invalid_metadata_type_error(field=field, value=value) + + def _strip_untrusted_request_header_controls( headers: Any, *, @@ -1572,6 +1604,10 @@ async def add_litellm_data_to_request( continue data.pop(_internal_key, None) _reject_url_valued_destinations(data) + for _metadata_field in ("metadata", "litellm_metadata"): + if (_raw_metadata := data.get(_metadata_field)) is not None: + data.pop(_metadata_field) + data[_metadata_field] = _normalized_metadata_object(_metadata_field, _raw_metadata) # Strip spoofable auth metadata from user-supplied metadata dict _user_metadata = data.get("metadata") if isinstance(_user_metadata, dict): @@ -1711,29 +1747,10 @@ async def add_litellm_data_to_request( verbose_proxy_logger.debug("receiving data: %s", data) - # Parse metadata if it's a string (e.g., from multipart/form-data) - if "metadata" in data and data["metadata"] is not None: - if isinstance(data["metadata"], str): - data["metadata"] = safe_json_loads(data["metadata"]) - if not isinstance(data["metadata"], dict): - verbose_proxy_logger.warning( - "Failed to parse 'metadata' as JSON dict. Received value: %s", data["metadata"] - ) - # requester_metadata is snapshotted AFTER the strip below so - # downstream consumers (e.g. PANW guardrail reading user_ip / - # profile_id) don't see attacker-injected admin slots preserved in - # the deepcopy. - - # Parse litellm_metadata if it's a string (e.g., from multipart/form-data or extra_body) - if "litellm_metadata" in data and data["litellm_metadata"] is not None: - if isinstance(data["litellm_metadata"], str): - parsed_litellm_metadata: Final = safe_json_loads(data["litellm_metadata"]) - if not isinstance(parsed_litellm_metadata, dict): - verbose_proxy_logger.warning( - "Failed to parse 'litellm_metadata' as JSON dict. Received value: %s", data["litellm_metadata"] - ) - else: - data["litellm_metadata"] = parsed_litellm_metadata + # requester_metadata is snapshotted AFTER the strip below so + # downstream consumers (e.g. PANW guardrail reading user_ip / + # profile_id) don't see attacker-injected admin slots preserved in + # the deepcopy. # Strip internal pipeline state and admin-injection slots from user input. # Runs AFTER the string-to-dict parse above so JSON-string metadata (sent diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index 22d3119d457..15396a95632 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -201,7 +201,7 @@ "limit": 58 }, "SIM102": { - "limit": 319 + "limit": 317 }, "SIM103": { "limit": 119 diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index e31058f402e..f6b801d3122 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -11,7 +11,7 @@ from pydantic import ValidationError as PydanticValidationError from starlette.datastructures import Headers import litellm -from litellm.proxy._types import AddTeamCallback, TeamCallbackMetadata, UserAPIKeyAuth +from litellm.proxy._types import AddTeamCallback, ProxyException, TeamCallbackMetadata, UserAPIKeyAuth from litellm.proxy.litellm_pre_call_utils import ( KeyAndTeamLoggingSettings, LiteLLMProxyRequestSetup, @@ -417,6 +417,69 @@ async def test_add_litellm_data_to_request_string_metadata_does_not_crash(): assert updated["metadata"].get("generation_name") == "test" +def _batches_request_mock() -> MagicMock: + request_mock = MagicMock(spec=Request) + request_mock.url = MagicMock() + request_mock.url.__str__.return_value = "http://localhost/v1/batches" + request_mock.url.path = "/v1/batches" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = {"Content-Type": "application/json"} + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + return request_mock + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "field,value,received_type", + [ + ("metadata", "abc", "a string"), + ("litellm_metadata", "abc", "a string"), + ("metadata", 42, "an integer"), + ("litellm_metadata", [1, 2], "an array"), + ("metadata", True, "a boolean"), + ], +) +async def test_add_litellm_data_to_request_rejects_non_object_metadata(field, value, received_type): + """Regression for https://github.com/BerriAI/litellm/issues/37147: a + non-object metadata was silently dropped with a 200, and a non-object + litellm_metadata crashed later with a 500 ('str' object has no attribute + 'update'). Both must be a 400 naming the field, like OpenAI returns.""" + data = {"input_file_id": "file-abc", "endpoint": "/v1/chat/completions", field: value} + + with pytest.raises(ProxyException) as exc_info: + await add_litellm_data_to_request( + data=data, + request=_batches_request_mock(), + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"), + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert exc_info.value.code == "400" + assert exc_info.value.param == field + assert exc_info.value.message == f"Invalid type for '{field}': expected an object, but got {received_type} instead." + assert field not in data + + +@pytest.mark.asyncio +async def test_add_litellm_data_to_request_parses_json_object_string_litellm_metadata(): + data = {"input_file_id": "file-abc", "litellm_metadata": json.dumps({"cost_centre": "research"})} + + updated = await add_litellm_data_to_request( + data=data, + request=_batches_request_mock(), + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"), + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert updated["litellm_metadata"]["cost_centre"] == "research" + + @pytest.mark.asyncio async def test_add_litellm_data_to_request_proxy_server_request_body_is_post_strip(): """Regression: proxy_server_request['body'] used to be snapshotted before diff --git a/type-discipline-budget.json b/type-discipline-budget.json index ca848190a32..6d70a6aa5f4 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -30,7 +30,7 @@ "limit": 16713 }, "LIT011": { - "limit": 5591 + "limit": 5590 }, "LIT012": { "limit": 4519