Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
177 changes: 125 additions & 52 deletions agents/providers/openai.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,8 +5,10 @@
from collections.abc import Callable
from dataclasses import dataclass, replace
from datetime import UTC, datetime
from io import BytesIO, StringIO
from io import BytesIO
from tempfile import SpooledTemporaryFile
from typing import (
IO,
Literal,
TypedDict,
cast,
Expand Down Expand Up @@ -57,6 +59,24 @@
from ..observability import LLMUsage, Observable

DEFAULT_BATCH_SIZE = 1000

# Maximum size in bytes for the JSONL payload
# for batch tasks (any larger will raise errors from the API)
MAX_BATCH_FILE_SIZE = 200_000_000

# Offload batches larger than this size to disk
# to save memory
MIN_SIZE_FOR_OFFLOAD = 8_000_000

class RequestTooLargeError(ValueError):
def __init__(self, custom_id: str, size: int, limit: int) -> None:
self.custom_id = custom_id
self.size = size
self.limit = limit
super().__init__(
f"Batch request [{custom_id}] of size {size} bytes "
f"exceeds the maximum for the OpenAI API ({limit} bytes)."
)
logger = logging.getLogger(__name__)

OPENAI_BATCH_ACTIVE_STATUSES: frozenset[str] = frozenset(
Expand Down Expand Up @@ -308,39 +328,96 @@ async def _batcher(self):
It's started at init time if we select batch mode, and persists for the duration of the session.
"""
logger.info("OpenAIBatchHelper opening.")
holdout_request = None
while True:
# Define our batch and start the clock
batch_file = SpooledTemporaryFile( # noqa: SIM115
max_size=MIN_SIZE_FOR_OFFLOAD,
mode="w+b"
)
semaphore_acquired = False
sent = False
try:
# Define our batch and start the clock
batch = []
# Wait until first query comes in
req = await self._get_batch_request()
batch.append(req)

batch_ids = []
batch_file_size = 0
# Await new messages to load into the batch
# - Until we hit our max batch size, or
# - Until we've waited for the time indicated (default 2s)
while len(batch) < self.batch_size:
while len(batch_ids) < self.batch_size:
timeout = self.timeout if batch_ids else None
try:
req = await self._get_batch_request(timeout=self.timeout)
# NOTE: We would possibly end up with a holdout request if that
# request would have made the last batch too large
if holdout_request is not None:
req, serialized_req = holdout_request
holdout_request = None
else:
req = await self._get_batch_request(timeout=timeout)

# Try to compute size, or fail the request
try:
serialized_req = await asyncio.to_thread(
self.provider._serialize_request,
req
)
except (TypeError, ValueError, RecursionError, UnicodeError) as err:
self._fail_request(req, err)
continue
task_id = req["custom_id"]
req_length = len(serialized_req)

except TimeoutError:
break
batch.append(req)

if req_length > MAX_BATCH_FILE_SIZE:
# Case: Individual request would be too large, so we should fail that one and move on
self._fail_request(req, RequestTooLargeError(task_id, req_length, MAX_BATCH_FILE_SIZE))
continue
elif (req_length + batch_file_size) > MAX_BATCH_FILE_SIZE:
# Case: Batch size would be too big if we added the next request, so we should
# hold it for the next batch and break
holdout_request = req, serialized_req
break
else:
# Otherwise: add it to the batch
batch_file_size += req_length
await asyncio.to_thread(batch_file.write, serialized_req)
batch_ids.append(task_id)
del serialized_req

batch_file.seek(0)
# Wait for semaphore to send off batch task
await self.lock.acquire()
semaphore_acquired = True

batch_task = asyncio.create_task(self._batch_handler(batch))
# Send our batch
batch_task = asyncio.create_task(self._batch_handler(batch_file, batch_ids))
self.batch_tasks.add(batch_task)

sent = True
# Batch task should remove itself from the list once it's done
batch_task.add_done_callback(self._batch_handler_callback)

except (asyncio.CancelledError, GeneratorExit):
# If the task was cancelled, we should exit the loop
if not sent:
batch_file.close()
if semaphore_acquired:
self.lock.release()
logger.info("OpenAIBatchHelper closing.")

break

def _fail_request(self, req: BatchRequestInput, exc: Exception):
"""
Fail a single request future with an exception
"""
task_id = req["custom_id"]
result_future = self.provider.batch_out.get(task_id)

if result_future is not None and not result_future.done():
result_future.set_exception(exc)

async def _get_batch_request(
self, timeout: float | None = None
) -> BatchRequestInput:
Expand All @@ -362,7 +439,7 @@ async def _get_batch_request(
self.provider.batch_q.task_done()
return request

async def _batch_handler(self, batch: list[BatchRequestInput]) -> None:
async def _batch_handler(self, batch: IO[bytes], batch_ids: list[str]) -> None:
"""
A handler method that submits the batch of tasks to OpenAI and retrieves the results
when finished.
Expand All @@ -377,7 +454,9 @@ def status_callback(batch_status: Batch) -> None:

# Create batch file, send to OpenAI and execute
try:
batch_file = await self.provider.send_batch(batch)
# Close batch file after we're done with it
with batch:
batch_file = await self.provider.send_batch(batch)
batch_task = await self.provider.create_batch_task(
batch_file,
timeout=self.api_timeout,
Expand All @@ -398,7 +477,7 @@ def status_callback(batch_status: Batch) -> None:
results = await self.provider.get_batch_results(batch_task)

# Write out results to dict for agents to pick up
expected_ids = {batch_item["custom_id"] for batch_item in batch}
expected_ids = {batch_id for batch_id in batch_ids}
returned_ids = set()
for result in results:
if not isinstance(result, dict):
Expand Down Expand Up @@ -443,8 +522,8 @@ def status_callback(batch_status: Batch) -> None:
self._finish_interrupted_batch(
latest_snapshot, BATCH_TRACKING_CANCELLED_STATUS
)
for batch_item in batch:
future = self.provider.batch_out.get(batch_item["custom_id"])
for batch_id in batch_ids:
future = self.provider.batch_out.get(batch_id)
if future is not None and not future.done():
future.cancel()
raise
Expand All @@ -454,8 +533,8 @@ def status_callback(batch_status: Batch) -> None:
latest_snapshot, BATCH_TRACKING_FAILED_STATUS
)
# propagate the exception to the futures
for batch_item in batch:
fut = self.provider.batch_out.get(batch_item["custom_id"])
for batch_id in batch_ids:
fut = self.provider.batch_out.get(batch_id)
if fut is not None and not fut.done():
# If the future is not done, set it to an exception
fut.set_exception(e)
Expand Down Expand Up @@ -509,6 +588,7 @@ class _AzureProvider[AgentT: _Agent, ProviderModeT: Literal["chat", "batch"]](
model_name: str
interactive: bool
resource_endpoint: str
_credential: ClientSecretCredential | None = None

def __init__(
self,
Expand All @@ -531,16 +611,27 @@ def authenticate(self) -> None:
Retrieve Azure OpenAI API key via ClientSecret authentication and
"""

credential = ClientSecretCredential(
self._credential = ClientSecretCredential(
tenant_id=os.environ["AZURE_TENANT_ID"],
client_id=os.environ["AZURE_CLIENT_ID"],
client_secret=os.environ["AZURE_CLIENT_SECRET"],
)

self._bearer_token_generator = get_bearer_token_provider(
credential, self.resource_endpoint
self._credential, self.resource_endpoint
)

async def close(self) -> None:
"""Close the OpenAI client and Azure credential transports."""
try:
await self.llm.close()
finally:
if self._credential is not None:
await self._credential.close()

async def __aexit__(self, exc_type, exc_value, traceback) -> None:
await self.close()

@backoff.on_exception(backoff.expo, openai.APIError, max_tries=3)
async def prompt_agent(
self,
Expand Down Expand Up @@ -693,13 +784,16 @@ def batch_progress(self) -> BatchProgressState:

return self.batch_handler.batch_progress

async def __aexit__(self, exc_type, exc_value, traceback):
async def close(self) -> None:
"""
Clean up the batch handler and any ongoing tasks
Clean up the batch handler and provider transports.
"""
if hasattr(self, "batch_handler") and self.batch_handler is not None:
# Cancel the batch processing task
await self.batch_handler.close()
try:
if hasattr(self, "batch_handler") and self.batch_handler is not None:
# Cancel the batch processing task before closing its API client.
await self.batch_handler.close()
finally:
await super().close()

async def query_batch_mode(
self, messages: list[ChatCompletionMessageParam], model: str, **kwargs
Expand Down Expand Up @@ -735,48 +829,27 @@ async def query_batch_mode(
self.batch_out.pop(task_id, None)

@staticmethod
def _create_batch_file(
tasks: list[BatchRequestInput],
) -> tuple[str, bytes, str]:
"""
Create a batch file for the OpenAI Batch API

:param tasks: list of task dictionaries to be sent to OpenAI
:return: Tuple containing the file name, file content, and MIME type to send as an API payload
"""

with StringIO() as batch_file:
for task in tasks:
batch_file.write(json.dumps(task) + "\n")

return (
"batch_tasks.jsonl",
batch_file.getvalue().encode("utf-8"),
"application/json",
)
def _serialize_request(task: BatchRequestInput) -> bytes:
return (json.dumps(task) + "\n").encode("utf-8")

async def send_batch(
self,
tasks: list[BatchRequestInput],
file_content: IO[bytes],
**kwargs,
) -> FileObject:
"""
Send a batch file to OpenAI pending further processing.

:param tasks: list of task dictionaries to be sent to OpenAI
:param file_content: serialized JSON payload to send to OpenAI
:param kwargs: Additional keyword arguments for the file upload (see OpenAI API documentation)

:return: An OpenAI File object representing the uploaded batch file
"""
file_name, file_content, mime_type = await asyncio.to_thread(
self._create_batch_file, tasks
)

file = await self.llm.files.create(
file=(file_name, file_content, mime_type), purpose="batch", **kwargs
file=("batch_tasks.jsonl", file_content, "application/json"), purpose="batch", **kwargs
)

logger.info(f"Created file [{file.id}] with {len(tasks)} queries.")
logger.info(f"Created file [{file.id}]")

return file

Expand Down
3 changes: 3 additions & 0 deletions test/test_batch_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
import json
import time
from copy import copy
from unittest.mock import AsyncMock

import openai
import pytest
Expand Down Expand Up @@ -185,6 +186,7 @@ async def test_batch_api(mocker: MockFixture):
aoi = openai.AsyncAzureOpenAI()
aoi.files = AsyncFiles(aoi)
aoi.batches = AsyncBatches(aoi)
close_client = mocker.patch.object(aoi, "close", new_callable=AsyncMock)

mocker.patch.object(AzureOpenAIBatchProvider, "authenticate", return_value=None)
AzureOpenAIBatchProvider._bearer_token_generator = "lorem ipsum"
Expand All @@ -201,3 +203,4 @@ async def test_batch_api(mocker: MockFixture):
await ag()

assert ag.answer == "Lorem ipsum, dolor sit amet."
close_client.assert_awaited_once_with()
Loading
Loading