diff --git a/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py b/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py index e8b5fab626f..e08d277788f 100644 --- a/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py +++ b/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py @@ -50,7 +50,7 @@ from litellm.repositories.table_repositories import ( ManagedFileRepository, ManagedObjectRepository, ) -from litellm.types.llms.openai import OpenAIFileObject +from litellm.types.llms.openai import BATCH_GUARDRAIL_RESPONSE_FIELD, OpenAIFileObject from litellm.types.passthrough_endpoints.managed_id_rewriter import ( ManagedFileIdReader, ManagedFileIdWriter, @@ -980,6 +980,7 @@ def _serialize_file_list_item(row: ManagedFileRow) -> dict[str, JsonValue]: file_object: Final = _parse_file_object(row.file_object) if isinstance(file_object, dict): item.update(file_object) + item.pop(BATCH_GUARDRAIL_RESPONSE_FIELD, None) item["id"] = row.unified_file_id # managed ID always wins over stored raw id return item diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 50e47071012..e7a3f825455 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -67,8 +67,10 @@ from pydantic import ( Discriminator, Field, PrivateAttr, + SerializerFunctionWrapHandler, field_serializer, field_validator, + model_serializer, ) from typing_extensions import ( NotRequired, @@ -315,6 +317,9 @@ class BatchGuardrailReport(BaseModel): """Every record that was redacted or dropped, in file order.""" +BATCH_GUARDRAIL_RESPONSE_FIELD: Final = "litellm_batch_guardrail" + + class OpenAIFileObject(BaseModel): id: str """The file identifier, which can be referenced in the API endpoints.""" @@ -363,6 +368,17 @@ class OpenAIFileObject(BaseModel): _hidden_params: dict = {"response_cost": 0.0} # no cost for writing a file + @model_serializer(mode="wrap") + def _omit_absent_batch_guardrail( # noqa: ANN202 # annotating it replaces the model's serialization schema + self, handler: SerializerFunctionWrapHandler + ): + serialized: Final[Mapping[str, object]] = handler(self) + if self.litellm_batch_guardrail is not None: + return serialized + return { # mutable-ok: pydantic's json serializer rejects a mapping that is not a dict + key: value for key, value in serialized.items() if key != BATCH_GUARDRAIL_RESPONSE_FIELD + } + def __contains__(self, key) -> bool: # Define custom behavior for the 'in' operator return hasattr(self, key) diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py index c15ba5bcedb..1b16e036ad7 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py @@ -4166,6 +4166,66 @@ def test_batch_upload_redacts_per_record(monkeypatch, llm_router: Router): ProxyLogging._callback_capabilities_cache.clear() +PLAIN_UPLOAD_RESPONSE_BODY = { + "id": "dummy-id", + "object": "file", + "bytes": 0, + "created_at": 1234567890, + "filename": "batch.jsonl", + "purpose": "batch", + "status": "uploaded", + "expires_at": None, + "status_details": None, +} + + +def test_create_file_omits_batch_guardrail_field_when_no_guardrail_configured(monkeypatch, llm_router: Router): + """An upload no guardrail is configured for serialises the plain OpenAI file shape.""" + forwarded_calls = _setup_batch_upload_endpoint(monkeypatch, llm_router) + try: + response = client.post( + "/v1/files", + files={"file": ("batch.jsonl", VALID_BATCH_LINE, "application/jsonl")}, + data={"purpose": "batch"}, + headers={"Authorization": "Bearer test-key"}, + ) + finally: + _teardown_batch_upload_endpoint() + + assert response.status_code == 200, response.text + assert len(forwarded_calls) == 1 + assert response.json() == PLAIN_UPLOAD_RESPONSE_BODY + + +def test_create_file_omits_batch_guardrail_field_when_guardrail_made_no_changes(monkeypatch, llm_router: Router): + """A guardrail that runs and changes nothing leaves the response the plain OpenAI file shape.""" + import litellm + from litellm.integrations.custom_guardrail import CustomGuardrail + from litellm.proxy.utils import ProxyLogging + + class _Passthrough(CustomGuardrail): + async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type): + return data + + forwarded_calls = _setup_batch_upload_endpoint(monkeypatch, llm_router) + monkeypatch.setattr(litellm, "callbacks", [_Passthrough(guardrail_name="noop", default_on=True)]) + ProxyLogging._callback_capabilities_cache.clear() + try: + response = client.post( + "/v1/files", + files={"file": ("batch.jsonl", VALID_BATCH_LINE, "application/jsonl")}, + data={"purpose": "batch"}, + headers={"Authorization": "Bearer test-key"}, + ) + finally: + _teardown_batch_upload_endpoint() + ProxyLogging._callback_capabilities_cache.clear() + + assert response.status_code == 200, response.text + assert len(forwarded_calls) == 1 + assert response.json() == PLAIN_UPLOAD_RESPONSE_BODY + + def test_batch_upload_closes_the_spools_it_opened(monkeypatch, llm_router: Router): """The scan and the rewrite each open a spool; the request owns both and must not leak them.""" import json as _json diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_managed_id_rewriter.py b/tests/test_litellm/proxy/pass_through_endpoints/test_managed_id_rewriter.py index 9c16c52f589..f5bec4a2585 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_managed_id_rewriter.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_managed_id_rewriter.py @@ -126,3 +126,24 @@ async def test_list_files_limit_above_batch_cap_still_served(): assert result is not None assert [item["id"] for item in result["data"]] == [managed_id] + + +@pytest.mark.asyncio +async def test_list_files_drops_batch_guardrail_key_persisted_by_an_older_proxy(): + """Rows written before the response serializer dropped the key still carry an explicit null.""" + managed_id = new_managed_id("openai", "file-abc") + row = _file_row(managed_id) + row.file_object = {**row.file_object, "litellm_batch_guardrail": None} + pc = _prisma_client(file_rows=[row]) + + result = await list_passthrough_ids_from_db( + provider="openai", + route="/openai/v1/files", + user_api_key_dict=_user(), + prisma_client=pc, + query_params={}, + ) + + assert result is not None + assert "litellm_batch_guardrail" not in result["data"][0] + assert result["data"][0]["filename"] == "test.jsonl" diff --git a/tests/test_litellm/types/llms/test_types_llms_openai.py b/tests/test_litellm/types/llms/test_types_llms_openai.py index e5e5c0183a0..3966677e928 100644 --- a/tests/test_litellm/types/llms/test_types_llms_openai.py +++ b/tests/test_litellm/types/llms/test_types_llms_openai.py @@ -450,3 +450,75 @@ def test_openai_file_object_accepts_pending_status(): status="pending", ) assert file_obj.status == "pending" + + +class TestOpenAIFileObjectBatchGuardrailSerialization: + """The proxy-only `litellm_batch_guardrail` key must reach the wire only when something set it.""" + + @staticmethod + def _file_object(**overrides): + from litellm.types.llms.openai import OpenAIFileObject + + return OpenAIFileObject( + id="file-123", + object="file", + bytes=1024, + created_at=1677610602, + filename="batch.jsonl", + purpose="batch", + status="uploaded", + **overrides, + ) + + @staticmethod + def _report(): + from litellm.types.llms.openai import BatchGuardrailRecord, BatchGuardrailReport + + return BatchGuardrailReport( + submitted_records=3, + modified_records=(BatchGuardrailRecord(line=2, custom_id="dirty", action="redacted"),), + ) + + @pytest.mark.parametrize("mode", ["python", "json"]) + def test_key_absent_when_unset(self, mode): + assert "litellm_batch_guardrail" not in self._file_object().model_dump(mode=mode) + + @pytest.mark.parametrize("mode", ["python", "json"]) + def test_key_present_when_set(self, mode): + dumped = self._file_object(litellm_batch_guardrail=self._report()).model_dump(mode=mode) + assert dumped["litellm_batch_guardrail"]["submitted_records"] == 3 + + def test_nested_nulls_of_a_set_report_survive(self): + """`exclude_none=True` was rejected as the fix because it would strip these.""" + dumped = self._file_object(litellm_batch_guardrail=self._report()).model_dump(mode="json") + assert dumped["litellm_batch_guardrail"]["modified_records"] == [ + {"line": 2, "custom_id": "dirty", "action": "redacted", "guardrail": None} + ] + + def test_by_alias_dump_also_omits_the_key(self): + """Tripwire: the serializer filters a literal key name, which an added alias would bypass.""" + assert "litellm_batch_guardrail" not in self._file_object().model_dump(mode="json", by_alias=True) + + def test_other_optional_fields_still_serialize_as_null(self): + dumped = self._file_object().model_dump(mode="json") + assert dumped["expires_at"] is None + assert dumped["status_details"] is None + + def test_round_trip_of_a_set_report_is_lossless(self): + from litellm.types.llms.openai import OpenAIFileObject + + original = self._file_object(litellm_batch_guardrail=self._report()) + assert OpenAIFileObject(**original.model_dump()) == original + + def test_serialization_json_schema_still_describes_the_model(self): + """A return annotation on the wrap serializer would collapse this to a bare object.""" + from litellm.types.llms.openai import OpenAIFileObject + + schema = OpenAIFileObject.model_json_schema(mode="serialization") + assert "litellm_batch_guardrail" in schema["properties"] + + def test_key_omitted_inside_a_file_list_page(self): + from litellm.types.llms.openai import FileListPage + + page = FileListPage(object="list", data=[self._file_object()], has_more=False) + assert "litellm_batch_guardrail" not in page.model_dump(mode="json")["data"][0]