From 7a5a24e092ec31796160d70140eb77ee60a7d07e Mon Sep 17 00:00:00 2001 From: Harshit Jain Date: Mon, 23 Feb 2026 21:01:02 +0530 Subject: [PATCH] fix: trigger pre-call guardrail hooks for batch file --- .../batch_guardrail_utils.py | 56 +++++++++ litellm/proxy/batches_endpoints/endpoints.py | 4 + .../openai_files_endpoints/files_endpoints.py | 11 ++ .../test_batch_guardrail_utils.py | 107 ++++++++++++++++++ .../test_files_endpoint.py | 44 +++++++ 5 files changed, 222 insertions(+) create mode 100644 litellm/proxy/batches_endpoints/batch_guardrail_utils.py create mode 100644 tests/test_litellm/proxy/batches_endpoints/test_batch_guardrail_utils.py diff --git a/litellm/proxy/batches_endpoints/batch_guardrail_utils.py b/litellm/proxy/batches_endpoints/batch_guardrail_utils.py new file mode 100644 index 00000000000..0c6f63118cf --- /dev/null +++ b/litellm/proxy/batches_endpoints/batch_guardrail_utils.py @@ -0,0 +1,56 @@ +import json +from typing import Optional +from litellm.proxy._types import UserAPIKeyAuth +from litellm.types.utils import CallTypesLiteral +from litellm.proxy.utils import ProxyLogging + +def _get_call_type_from_endpoint(endpoint: str) -> CallTypesLiteral: + """Map batch JSONL endpoint to CallTypesLiteral.""" + if "chat/completions" in endpoint: + return "acompletion" + elif "embeddings" in endpoint: + return "aembedding" + else: + return "acompletion" # default fallback + +async def run_pre_call_guardrails_on_batch_file( + file_content: bytes, + proxy_logging_obj: ProxyLogging, + user_api_key_dict: UserAPIKeyAuth, +) -> None: + """ + Parse a batch JSONL file and run pre-call guardrails on each request. + + Raises an exception if any guardrail rejects a request. + """ + lines = file_content.decode("utf-8").strip().splitlines() + + for line_num, line in enumerate(lines, start=1): + if not line.strip(): + continue + + try: + json_obj = json.loads(line) + except json.JSONDecodeError: + continue + + body = json_obj.get("body", {}) + + if not body or "messages" not in body: + continue + + # Determine call_type from the endpoint in the JSONL line + endpoint = json_obj.get("url", "") + call_type = _get_call_type_from_endpoint(endpoint) + + try: + # Run pre_call_hook (which triggers all guardrails) + await proxy_logging_obj.pre_call_hook( + user_api_key_dict=user_api_key_dict, + data=body, + call_type=call_type, + ) + except Exception as e: + custom_id = json_obj.get("custom_id", f"line_{line_num}") + # Reraise exception to abort the upload, appending item information + raise Exception(f"Guardrail rejected batch item {custom_id} (line {line_num}): {str(e)}") from e diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index 1c9ba6cb248..0d2e8dae492 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -63,6 +63,10 @@ async def create_batch( # noqa: PLR0915 Create large batches of API requests for asynchronous processing. This is the equivalent of POST https://api.openai.com/v1/batch Supports Identical Params as: https://platform.openai.com/docs/api-reference/batch + + Note: Guardrails for batch content are executed at file upload time + (`/v1/files` with `purpose="batch"`). They are NOT executed here since + the actual file contents/messages are not available in this request. Example Curl ``` diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index ec6e9733344..c7254553e1b 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -461,6 +461,17 @@ async def create_file( # noqa: PLR0915 **data ) + # Run pre-call guardrails on batch file content + if purpose == "batch": + from litellm.proxy.batches_endpoints.batch_guardrail_utils import ( + run_pre_call_guardrails_on_batch_file, + ) + await run_pre_call_guardrails_on_batch_file( + file_content=file_content, + proxy_logging_obj=proxy_logging_obj, + user_api_key_dict=user_api_key_dict, + ) + response = await route_create_file( llm_router=llm_router, _create_file_request=_create_file_request, diff --git a/tests/test_litellm/proxy/batches_endpoints/test_batch_guardrail_utils.py b/tests/test_litellm/proxy/batches_endpoints/test_batch_guardrail_utils.py new file mode 100644 index 00000000000..ccb443fc41f --- /dev/null +++ b/tests/test_litellm/proxy/batches_endpoints/test_batch_guardrail_utils.py @@ -0,0 +1,107 @@ +import json +import pytest +from unittest.mock import AsyncMock, MagicMock +from litellm.proxy.batches_endpoints.batch_guardrail_utils import run_pre_call_guardrails_on_batch_file, _get_call_type_from_endpoint +from litellm.proxy._types import UserAPIKeyAuth + +@pytest.fixture +def mock_proxy_logging_obj(): + mock_obj = MagicMock() + mock_obj.pre_call_hook = AsyncMock() + return mock_obj + +@pytest.fixture +def mock_user_api_key_dict(): + return UserAPIKeyAuth() + +@pytest.mark.asyncio +async def test_guardrail_runs_on_batch_file_messages(mock_proxy_logging_obj, mock_user_api_key_dict): + # Create a batch JSONL file with messages + req1 = {"custom_id": "request-1", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "gpt-3.5-turbo", "messages": [{"role": "system", "content": "You are a helpful assistant."}]}} + req2 = {"custom_id": "request-2", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "Hello!"}]}} + + file_content = f"{json.dumps(req1)}\n{json.dumps(req2)}\n".encode("utf-8") + + await run_pre_call_guardrails_on_batch_file( + file_content=file_content, + proxy_logging_obj=mock_proxy_logging_obj, + user_api_key_dict=mock_user_api_key_dict + ) + + # Verify pre_call_hook is called for each line + assert mock_proxy_logging_obj.pre_call_hook.call_count == 2 + + # Check args of the first call + call_args_1 = mock_proxy_logging_obj.pre_call_hook.call_args_list[0].kwargs + assert call_args_1["call_type"] == "acompletion" + assert "messages" in call_args_1["data"] + assert call_args_1["data"]["messages"][0]["content"] == "You are a helpful assistant." + + # Check args of the second call + call_args_2 = mock_proxy_logging_obj.pre_call_hook.call_args_list[1].kwargs + assert call_args_2["call_type"] == "acompletion" + assert "messages" in call_args_2["data"] + assert call_args_2["data"]["messages"][0]["content"] == "Hello!" + +@pytest.mark.asyncio +async def test_guardrail_rejection_raises_exception(mock_proxy_logging_obj, mock_user_api_key_dict): + # Configure mock to raise an exception simulating a guardrail block + mock_proxy_logging_obj.pre_call_hook.side_effect = Exception("Content blocked by policy") + + req = {"custom_id": "bad-request-1", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "gpt-4", "messages": [{"role": "user", "content": "bad stuff"}]}} + file_content = f"{json.dumps(req)}\n".encode("utf-8") + + # Verify rejection propagates and includes the custom_id and line info + with pytest.raises(Exception) as exc_info: + await run_pre_call_guardrails_on_batch_file( + file_content=file_content, + proxy_logging_obj=mock_proxy_logging_obj, + user_api_key_dict=mock_user_api_key_dict + ) + + assert "bad-request-1" in str(exc_info.value) + assert "line 1" in str(exc_info.value) + assert "Content blocked by policy" in str(exc_info.value) + +@pytest.mark.asyncio +async def test_endpoint_to_call_type_mapping(): + assert _get_call_type_from_endpoint("/v1/chat/completions") == "acompletion" + assert _get_call_type_from_endpoint("/v1/embeddings") == "aembedding" + assert _get_call_type_from_endpoint("/v1/audio/transcriptions") == "acompletion" # Default fallback + assert _get_call_type_from_endpoint("/custom/endpoint") == "acompletion" + +@pytest.mark.asyncio +async def test_malformed_jsonl_handling(mock_proxy_logging_obj, mock_user_api_key_dict): + # Test file with some valid lines and some malformed/empty lines + req1 = {"custom_id": "request-1", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "gpt-3.5", "messages": [{"role": "user", "content": "Hi"}]}} + + file_content = f"{json.dumps(req1)}\nthis is not json\n\n{{\"not_complete\": true \n".encode("utf-8") + + # Should complete without error, skipping malformed lines + await run_pre_call_guardrails_on_batch_file( + file_content=file_content, + proxy_logging_obj=mock_proxy_logging_obj, + user_api_key_dict=mock_user_api_key_dict + ) + + # Only the first valid line should have triggered a call + assert mock_proxy_logging_obj.pre_call_hook.call_count == 1 + +@pytest.mark.asyncio +async def test_skips_lines_without_messages(mock_proxy_logging_obj, mock_user_api_key_dict): + # Request without messages + req1 = {"custom_id": "req-1", "method": "POST", "url": "/v1/models", "body": {}} + # Request with messages + req2 = {"custom_id": "req-2", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "gpt-4", "messages": []}} + + file_content = f"{json.dumps(req1)}\n{json.dumps(req2)}\n".encode("utf-8") + + await run_pre_call_guardrails_on_batch_file( + file_content=file_content, + proxy_logging_obj=mock_proxy_logging_obj, + user_api_key_dict=mock_user_api_key_dict + ) + + # Only req2 should trigger the hook since req1 has no messages + assert mock_proxy_logging_obj.pre_call_hook.call_count == 1 + 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 9c6182493d0..fb6aab1680e 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 @@ -1134,3 +1134,47 @@ def test_create_file_with_deep_nested_litellm_metadata( assert captured_litellm_metadata["config"]["database"]["port"] == "5432" assert "cache" in captured_litellm_metadata["config"] assert captured_litellm_metadata["config"]["cache"]["enabled"] == "true" + + +def test_create_batch_file_with_failing_guardrail( + mocker: MockerFixture, monkeypatch, llm_router: Router +): + """ + Test that uploading a batch file triggers guardrails and fails if the guardrail raises an exception. + """ + proxy_logging_obj = ProxyLogging( + user_api_key_cache=DualCache(default_in_memory_ttl=1) + ) + proxy_logging_obj._add_proxy_hooks(llm_router) + + # Mock pre_call_hook to raise an exception simulating a blocked request + mock_pre_call_hook = mocker.patch.object( + proxy_logging_obj, + "pre_call_hook", + side_effect=Exception("Guardrail blocked this batch item") + ) + + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", llm_router) + monkeypatch.setattr( + "litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj + ) + + test_file_content = b'{"custom_id": "req-1", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "Bad stuff"}]}}' + test_file = ("nested.jsonl", test_file_content, "application/jsonl") + + response = client.post( + "/v1/files", + files={"file": test_file}, + data={ + "purpose": "batch", + }, + headers={"Authorization": "Bearer test-key"}, + ) + + # Verify the guardrail was called + assert mock_pre_call_hook.call_count == 1 + + # Verify request failed + assert response.status_code == 500 + error_detail = response.json() + assert "Guardrail blocked this batch item" in error_detail["error"]["message"]