Merge pull request #4914 from BerriAI/litellm_fix_batches

[Proxy-Fix + Test] - /batches endpoint
This commit is contained in:
Ishaan Jaff 2024-07-26 20:12:03 -07:00 • committed by GitHub
commit 3c463ccbe6
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 125 additions and 64 deletions

View file

@ -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)

View file

@ -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

View file

@ -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"

View file

@ -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,
)

View 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)

View file

@ -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"
)