Revert "Respect custom llm provider in header" (#17211)

This commit is contained in:
Ishaan Jaff
2025-11-27 09:12:44 -08:00
committed by GitHub
parent 83c1138067
commit 31ecd4ce49
2 changed files with 3 additions and 89 deletions
+2 -4
View File
@@ -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)
@@ -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():