From 31ecd4ce49c5d6eaa65b8323328ab3bf31ca11ca Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Thu, 27 Nov 2025 09:12:44 -0800 Subject: [PATCH] Revert "Respect custom llm provider in header" (#17211) --- litellm/proxy/batches_endpoints/endpoints.py | 6 +- .../test_openai_batches_endpoint.py | 86 +------------------ 2 files changed, 3 insertions(+), 89 deletions(-) diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index 03b9ac3dea..98492bcc2d 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -31,6 +31,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( get_original_file_id, prepare_data_with_credentials, ) + from litellm.proxy.utils import handle_exception_on_proxy, is_known_model from litellm.types.llms.openai import LiteLLMBatchCreateRequest @@ -111,10 +112,7 @@ async def create_batch( # noqa: PLR0915 is_router_model = is_known_model(model=router_model, llm_router=llm_router) custom_llm_provider = ( - provider - or data.pop("custom_llm_provider", None) - or get_custom_llm_provider_from_request_headers(request=request) - or "openai" + provider or data.pop("custom_llm_provider", None) or "openai" ) _create_batch_data = LiteLLMBatchCreateRequest(**data) input_file_id = _create_batch_data.get("input_file_id", None) diff --git a/tests/openai_endpoints_tests/test_openai_batches_endpoint.py b/tests/openai_endpoints_tests/test_openai_batches_endpoint.py index be931f2f4b..5a84900f87 100644 --- a/tests/openai_endpoints_tests/test_openai_batches_endpoint.py +++ b/tests/openai_endpoints_tests/test_openai_batches_endpoint.py @@ -5,13 +5,12 @@ import asyncio import aiohttp, openai from openai import OpenAI, AsyncOpenAI from typing import Optional, List, Union +from test_openai_files_endpoints import upload_file, delete_file import os import sys import time from unittest.mock import patch, MagicMock, AsyncMock -from litellm.proxy.batches_endpoints.endpoints import create_batch - BASE_URL = "http://localhost:4000" # Replace with your actual base URL API_KEY = "sk-1234" # Replace with your actual API key @@ -223,89 +222,6 @@ def test_vertex_batches_endpoint(): pass -@pytest.mark.asyncio -async def test_create_batch_respects_custom_llm_provider_header(): - request_body = { - "input_file_id": "file-batch-123", - "endpoint": "/v1/chat/completions", - "completion_window": "24h", - } - - response_object = MagicMock() - response_object._hidden_params = {} - response_object.id = "batch-id" - response_object.input_file_id = request_body["input_file_id"] - - mock_user_api_key_dict = MagicMock( - tpm_limit=0, - rpm_limit=0, - max_budget=0, - spend=0, - allowed_model_region="", - ) - mock_user_api_key_dict.metadata = {} - - request = MagicMock() - request.headers = {"custom-llm-provider": "vertex_ai"} - request.query_params = {} - request.method = "POST" - - fastapi_response = MagicMock() - fastapi_response.headers = {} - - logging_obj = MagicMock() - - mock_proxy_logging_obj = MagicMock() - mock_proxy_logging_obj.post_call_success_hook = AsyncMock( - return_value=response_object - ) - mock_proxy_logging_obj.post_call_failure_hook = AsyncMock() - mock_proxy_logging_obj.update_request_status = AsyncMock() - - with patch( - "litellm.proxy.batches_endpoints.endpoints._read_request_body", - new_callable=AsyncMock, - ) as mock_read_body, patch( - "litellm.proxy.batches_endpoints.endpoints.ProxyBaseLLMRequestProcessing.common_processing_pre_call_logic", - new_callable=AsyncMock, - ) as mock_common_processing, patch( - "litellm.proxy.batches_endpoints.endpoints.litellm.acreate_batch", - new_callable=AsyncMock, - ) as mock_acreate_batch, patch( - "litellm.proxy.proxy_server.general_settings", - new={}, - ), patch( - "litellm.proxy.proxy_server.proxy_config", - new={}, - ), patch( - "litellm.proxy.proxy_server.version", - new="test-version", - ), patch( - "litellm.proxy.proxy_server.llm_router", - new=None, - ), patch( - "litellm.proxy.proxy_server.proxy_logging_obj", - new=mock_proxy_logging_obj, - ), patch( - "litellm.proxy.batches_endpoints.endpoints.asyncio.create_task", - new=lambda coro: asyncio.ensure_future(coro), - ): - mock_read_body.return_value = dict(request_body) - mock_common_processing.return_value = (dict(request_body), logging_obj) - mock_acreate_batch.return_value = response_object - - response = await create_batch( - request=request, - fastapi_response=fastapi_response, - user_api_key_dict=mock_user_api_key_dict, - ) - - mock_acreate_batch.assert_awaited_once() - called_kwargs = mock_acreate_batch.call_args.kwargs - assert called_kwargs["custom_llm_provider"] == "vertex_ai" - assert response is response_object - - @pytest.mark.skip(reason="Local only test to verify if things work well") @pytest.mark.asyncio async def test_list_batches_with_target_model_names():