mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-09 16:24:38 +00:00
Revert "Respect custom llm provider in header" (#17211)
This commit is contained in:
@@ -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():
|
||||
|
||||
Reference in New Issue
Block a user