mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
Merge pull request #4914 from BerriAI/litellm_fix_batches
[Proxy-Fix + Test] - /batches endpoint
This commit is contained in:
commit
3c463ccbe6
6 changed files with 125 additions and 64 deletions
|
|
@ -1,15 +1,13 @@
|
|||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
|
||||
# Batches API
|
||||
# [BETA] Batches API
|
||||
|
||||
Covers Batches, Files
|
||||
|
||||
|
||||
## Quick Start
|
||||
|
||||
Call an existing Assistant.
|
||||
|
||||
- Create File for Batch Completion
|
||||
|
||||
- Create Batch Request
|
||||
|
|
@ -18,6 +16,47 @@ Call an existing Assistant.
|
|||
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="proxy" label="LiteLLM PROXY Server">
|
||||
|
||||
```bash
|
||||
$ export OPENAI_API_KEY="sk-..."
|
||||
|
||||
$ litellm
|
||||
|
||||
# RUNNING on http://0.0.0.0:4000
|
||||
```
|
||||
|
||||
**Create File for Batch Completion**
|
||||
|
||||
```shell
|
||||
curl http://localhost:4000/v1/files \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-F purpose="batch" \
|
||||
-F file="@mydata.jsonl"
|
||||
```
|
||||
|
||||
**Create Batch Request**
|
||||
|
||||
```bash
|
||||
curl http://localhost:4000/v1/batches \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"input_file_id": "file-abc123",
|
||||
"endpoint": "/v1/chat/completions",
|
||||
"completion_window": "24h"
|
||||
}'
|
||||
```
|
||||
|
||||
**Retrieve the Specific Batch**
|
||||
|
||||
```bash
|
||||
curl http://localhost:4000/v1/batches/batch_abc123 \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-H "Content-Type: application/json" \
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
||||
**Create File for Batch Completion**
|
||||
|
|
@ -78,47 +117,7 @@ print("file content = ", file_content)
|
|||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="proxy" label="PROXY">
|
||||
|
||||
```bash
|
||||
$ export OPENAI_API_KEY="sk-..."
|
||||
|
||||
$ litellm
|
||||
|
||||
# RUNNING on http://0.0.0.0:4000
|
||||
```
|
||||
|
||||
**Create File for Batch Completion**
|
||||
|
||||
```shell
|
||||
curl https://api.openai.com/v1/files \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-F purpose="batch" \
|
||||
-F file="@mydata.jsonl"
|
||||
```
|
||||
|
||||
**Create Batch Request**
|
||||
|
||||
```bash
|
||||
curl http://localhost:4000/v1/batches \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"input_file_id": "file-abc123",
|
||||
"endpoint": "/v1/chat/completions",
|
||||
"completion_window": "24h"
|
||||
}'
|
||||
```
|
||||
|
||||
**Retrieve the Specific Batch**
|
||||
|
||||
```bash
|
||||
curl http://localhost:4000/v1/batches/batch_abc123 \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-H "Content-Type: application/json" \
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## [👉 Proxy API Reference](https://litellm-api.up.railway.app/#/batch)
|
||||
|
|
|
|||
|
|
@ -24,6 +24,7 @@ from pydantic import BaseModel
|
|||
from typing_extensions import overload, override
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.types.utils import ProviderField
|
||||
from litellm.utils import (
|
||||
|
|
@ -2534,6 +2535,7 @@ class OpenAIBatchesAPI(BaseLLM):
|
|||
retrieve_batch_data: RetrieveBatchRequest,
|
||||
openai_client: AsyncOpenAI,
|
||||
) -> Batch:
|
||||
verbose_logger.debug("retrieving batch, args= %s", retrieve_batch_data)
|
||||
response = await openai_client.batches.retrieve(**retrieve_batch_data)
|
||||
return response
|
||||
|
||||
|
|
|
|||
|
|
@ -39,6 +39,8 @@ def _get_metadata_variable_name(request: Request) -> str:
|
|||
"""
|
||||
if "thread" in request.url.path or "assistant" in request.url.path:
|
||||
return "litellm_metadata"
|
||||
if "batches" in request.url.path:
|
||||
return "litellm_metadata"
|
||||
if "/v1/messages" in request.url.path:
|
||||
# anthropic API has a field called metadata
|
||||
return "litellm_metadata"
|
||||
|
|
|
|||
|
|
@ -4808,10 +4808,18 @@ async def create_batch(
|
|||
"""
|
||||
global proxy_logging_obj
|
||||
data: Dict = {}
|
||||
|
||||
try:
|
||||
# Use orjson to parse JSON data, orjson speeds up requests significantly
|
||||
form_data = await request.form()
|
||||
data = {key: value for key, value in form_data.items() if key != "file"}
|
||||
body = await request.body()
|
||||
body_str = body.decode()
|
||||
try:
|
||||
data = ast.literal_eval(body_str)
|
||||
except:
|
||||
data = json.loads(body_str)
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
"Request received by LiteLLM:\n{}".format(json.dumps(data, indent=4)),
|
||||
)
|
||||
|
||||
# Include original request and headers in the data
|
||||
data = await add_litellm_data_to_request(
|
||||
|
|
@ -4883,12 +4891,12 @@ async def create_batch(
|
|||
|
||||
|
||||
@router.get(
|
||||
"/v1/batches{batch_id}",
|
||||
"/v1/batches{batch_id:path}",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
tags=["batch"],
|
||||
)
|
||||
@router.get(
|
||||
"/batches{batch_id}",
|
||||
"/batches{batch_id:path}",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
tags=["batch"],
|
||||
)
|
||||
|
|
@ -4916,20 +4924,6 @@ async def retrieve_batch(
|
|||
global proxy_logging_obj
|
||||
data: Dict = {}
|
||||
try:
|
||||
# Use orjson to parse JSON data, orjson speeds up requests significantly
|
||||
form_data = await request.form()
|
||||
data = {key: value for key, value in form_data.items() if key != "file"}
|
||||
|
||||
# Include original request and headers in the data
|
||||
data = await add_litellm_data_to_request(
|
||||
data=data,
|
||||
request=request,
|
||||
general_settings=general_settings,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
version=version,
|
||||
proxy_config=proxy_config,
|
||||
)
|
||||
|
||||
_retrieve_batch_request = RetrieveBatchRequest(
|
||||
batch_id=batch_id,
|
||||
)
|
||||
|
|
|
|||
64
tests/test_openai_batches_endpoint.py
Normal file
64
tests/test_openai_batches_endpoint.py
Normal file
|
|
@ -0,0 +1,64 @@
|
|||
# What this tests ?
|
||||
## Tests /batches endpoints
|
||||
import pytest
|
||||
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
|
||||
|
||||
|
||||
BASE_URL = "http://localhost:4000" # Replace with your actual base URL
|
||||
API_KEY = "sk-1234" # Replace with your actual API key
|
||||
|
||||
|
||||
async def create_batch(session, input_file_id, endpoint, completion_window):
|
||||
url = f"{BASE_URL}/v1/batches"
|
||||
headers = {"Authorization": f"Bearer {API_KEY}", "Content-Type": "application/json"}
|
||||
payload = {
|
||||
"input_file_id": input_file_id,
|
||||
"endpoint": endpoint,
|
||||
"completion_window": completion_window,
|
||||
}
|
||||
|
||||
async with session.post(url, headers=headers, json=payload) as response:
|
||||
assert response.status == 200, f"Expected status 200, got {response.status}"
|
||||
result = await response.json()
|
||||
print(f"Batch creation successful. Batch ID: {result.get('id', 'N/A')}")
|
||||
return result
|
||||
|
||||
|
||||
async def get_batch_by_id(session, batch_id):
|
||||
url = f"{BASE_URL}/v1/batches/{batch_id}"
|
||||
headers = {"Authorization": f"Bearer {API_KEY}"}
|
||||
|
||||
async with session.get(url, headers=headers) as response:
|
||||
if response.status == 200:
|
||||
result = await response.json()
|
||||
return result
|
||||
else:
|
||||
print(f"Error: Failed to get batch. Status code: {response.status}")
|
||||
return None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_batches_operations():
|
||||
async with aiohttp.ClientSession() as session:
|
||||
# Test file upload and get file_id
|
||||
file_id = await upload_file(session, purpose="batch")
|
||||
|
||||
create_batch_response = await create_batch(
|
||||
session, file_id, "/v1/chat/completions", "24h"
|
||||
)
|
||||
batch_id = create_batch_response.get("id")
|
||||
assert batch_id is not None
|
||||
|
||||
# Test get batch
|
||||
get_batch_response = await get_batch_by_id(session, batch_id)
|
||||
print("response from get batch", get_batch_response)
|
||||
|
||||
assert get_batch_response["id"] == batch_id
|
||||
assert get_batch_response["input_file_id"] == file_id
|
||||
|
||||
# Test delete file
|
||||
await delete_file(session, file_id)
|
||||
|
|
@ -30,11 +30,11 @@ async def test_file_operations():
|
|||
await delete_file(session, file_id)
|
||||
|
||||
|
||||
async def upload_file(session):
|
||||
async def upload_file(session, purpose="fine-tune"):
|
||||
url = f"{BASE_URL}/v1/files"
|
||||
headers = {"Authorization": f"Bearer {API_KEY}"}
|
||||
data = aiohttp.FormData()
|
||||
data.add_field("purpose", "fine-tune")
|
||||
data.add_field("purpose", purpose)
|
||||
data.add_field(
|
||||
"file", b'{"prompt": "Hello", "completion": "Hi"}', filename="mydata.jsonl"
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue