Merge pull request #19040 from Point72/ephrimstanley/batch-list

Fix /batches to return encoded ids (from managed objects table)
This commit is contained in:
Sameer Kankute 2026-01-27 13:02:05 +05:30 • committed by GitHub
commit f98eba24d4
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 349 additions and 5 deletions

View file

@ -244,6 +244,78 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
return managed_object.created_by == user_id
return True # don't raise error if managed object is not found
async def list_user_batches(
self,
user_api_key_dict: UserAPIKeyAuth,
limit: Optional[int] = None,
after: Optional[str] = None,
provider: Optional[str] = None,
target_model_names: Optional[str] = None,
llm_router: Optional[Router] = None,
) -> Dict[str, Any]:
# Provider filtering is not supported for managed batches
# This is because the encoded object ids stored in the managed objects table do not contain the provider information
# To support provider filtering, we would need to store the provider information in the encoded object ids
if provider:
raise Exception(
"Filtering by 'provider' is not supported when using managed batches."
)
# Model name filtering is not supported for managed batches
# This is because the encoded object ids stored in the managed objects table do not contain the model name
# A hash of the model name + litellm_params for the model name is encoded as the model id. This is not sufficient to reliably map the target model names to the model ids.
if target_model_names:
raise Exception(
"Filtering by 'target_model_names' is not supported when using managed batches."
)
where_clause: Dict[str, Any] = {"file_purpose": "batch"}
# Filter by user who created the batch
if user_api_key_dict.user_id:
where_clause["created_by"] = user_api_key_dict.user_id
if after:
where_clause["id"] = {"gt": after}
# Fetch more than needed to allow for post-fetch filtering
fetch_limit = limit or 20
if target_model_names:
# Fetch extra to account for filtering
fetch_limit = max(fetch_limit * 3, 100)
batches = await self.prisma_client.db.litellm_managedobjecttable.find_many(
where=where_clause,
take=fetch_limit,
order={"created_at": "desc"},
)
batch_objects: List[LiteLLMBatch] = []
for batch in batches:
try:
# Stop once we have enough after filtering
if len(batch_objects) >= (limit or 20):
break
batch_data = json.loads(batch.file_object) if isinstance(batch.file_object, str) else batch.file_object
batch_obj = LiteLLMBatch(**batch_data)
batch_obj.id = batch.unified_object_id
batch_objects.append(batch_obj)
except Exception as e:
verbose_logger.warning(
f"Failed to parse batch object {batch.unified_object_id}: {e}"
)
continue
return {
"object": "list",
"data": batch_objects,
"first_id": batch_objects[0].id if batch_objects else None,
"last_id": batch_objects[-1].id if batch_objects else None,
"has_more": len(batch_objects) == (limit or 20),
}
async def get_user_created_file_ids(
self, user_api_key_dict: UserAPIKeyAuth, model_object_ids: List[str]
) -> List[OpenAIFileObject]:
@ -673,6 +745,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
bytes=file_objects[0].bytes,
filename=file_objects[0].filename,
status="uploaded",
expires_at=file_objects[0].expires_at,
)
return response

View file

@ -542,14 +542,26 @@ async def list_batches(
route_type="alist_batches",
)
model_param = (
# Try to use managed objects table for listing batches (returns encoded IDs)
managed_files_obj = proxy_logging_obj.get_proxy_hook("managed_files")
if managed_files_obj is not None and hasattr(managed_files_obj, "list_user_batches"):
verbose_proxy_logger.debug(
"Using managed objects table for batch listing"
)
response = await managed_files_obj.list_user_batches(
user_api_key_dict=user_api_key_dict,
limit=limit,
after=after,
provider=provider,
target_model_names=target_model_names,
llm_router=llm_router,
)
elif (model_param := (
data.get("model")
or request.query_params.get("model")
or request.headers.get("x-litellm-model")
)
# SCENARIO 2: Use model-based routing from header/query/body
if model_param:
)):
# SCENARIO 2: Use model-based routing from header/query/body
credentials = get_credentials_for_model(
llm_router=llm_router,
model_id=model_param,

View file

@ -827,3 +827,262 @@ async def test_afile_retrieve_raises_error_for_non_managed_file():
)
assert "not found" in str(exc_info.value)
@pytest.mark.asyncio
async def test_list_batches_from_managed_objects_table():
from litellm.proxy._types import UserAPIKeyAuth
from openai.types.batch import BatchRequestCounts
prisma_client = AsyncMock()
batch_record_1 = MagicMock()
batch_record_1.unified_object_id = "unified-batch-id-1"
batch_record_1.file_object = json.dumps({
"id": "batch_abc123",
"object": "batch",
"endpoint": "/v1/chat/completions",
"completion_window": "24h",
"status": "completed",
"created_at": 1234567890,
"input_file_id": "file-input-1",
"request_counts": {"total": 1, "completed": 1, "failed": 0},
})
batch_record_2 = MagicMock()
batch_record_2.unified_object_id = "unified-batch-id-2"
batch_record_2.file_object = json.dumps({
"id": "batch_xyz789",
"object": "batch",
"endpoint": "/v1/chat/completions",
"completion_window": "24h",
"status": "in_progress",
"created_at": 1234567891,
"input_file_id": "file-input-2",
"request_counts": {"total": 5, "completed": 2, "failed": 0},
})
prisma_client.db.litellm_managedobjecttable.find_many.return_value = [
batch_record_1,
batch_record_2,
]
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
DualCache(), prisma_client=prisma_client
)
result = await proxy_managed_files.list_user_batches(
user_api_key_dict=UserAPIKeyAuth(user_id="test-user"),
limit=10,
)
assert result["object"] == "list"
assert len(result["data"]) == 2
assert result["data"][0].id == "unified-batch-id-1"
assert result["data"][1].id == "unified-batch-id-2"
assert result["first_id"] == "unified-batch-id-1"
assert result["last_id"] == "unified-batch-id-2"
# Should filter by user_id (created_by)
prisma_client.db.litellm_managedobjecttable.find_many.assert_called_once_with(
where={"file_purpose": "batch", "created_by": "test-user"},
take=10,
order={"created_at": "desc"},
)
@pytest.mark.asyncio
async def test_list_batches_from_managed_objects_table_empty_list():
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
prisma_client.db.litellm_managedobjecttable.find_many.return_value = []
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
DualCache(), prisma_client=prisma_client
)
result = await proxy_managed_files.list_user_batches(
user_api_key_dict=UserAPIKeyAuth(user_id="test-user"),
)
assert result["object"] == "list"
assert len(result["data"]) == 0
assert result["first_id"] is None
assert result["last_id"] is None
assert result["has_more"] is False
# Verify where clause includes created_by filter
# Default take is 20 when no limit is provided
prisma_client.db.litellm_managedobjecttable.find_many.assert_called_once_with(
where={"file_purpose": "batch", "created_by": "test-user"},
take=20,
order={"created_at": "desc"},
)
def _create_unified_batch_id(model_id: str, batch_id: str) -> str:
import base64
unified_str = f"litellm_proxy;model_id:{model_id};llm_batch_id:{batch_id}"
return base64.urlsafe_b64encode(unified_str.encode()).decode().rstrip("=")
@pytest.mark.asyncio
async def test_list_batches_from_managed_objects_table_provider_filter_raises_exception():
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
DualCache(), prisma_client=prisma_client
)
# Filtering by provider should raise Exception
with pytest.raises(Exception) as exc_info:
await proxy_managed_files.list_user_batches(
user_api_key_dict=UserAPIKeyAuth(user_id="test-user"),
limit=10,
provider="openai",
)
assert str(exc_info.value) == (
"Filtering by 'provider' is not supported when using managed batches."
)
# Verify find_many was NOT called since exception is raised before database query
prisma_client.db.litellm_managedobjecttable.find_many.assert_not_called()
@pytest.mark.asyncio
async def test_list_batches_from_managed_objects_table_target_model_name_filter_raises_exception():
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
DualCache(), prisma_client=prisma_client
)
# Filtering by provider should raise Exception
with pytest.raises(Exception) as exc_info:
await proxy_managed_files.list_user_batches(
user_api_key_dict=UserAPIKeyAuth(user_id="test-user"),
limit=10,
target_model_names="gpt-4o,gpt-3.5",
)
assert str(exc_info.value) == (
"Filtering by 'target_model_names' is not supported when using managed batches."
)
# Verify find_many was NOT called since exception is raised before database query
prisma_client.db.litellm_managedobjecttable.find_many.assert_not_called()
@pytest.mark.asyncio
async def test_list_batches_from_managed_objects_table_filters_by_created_by():
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
# Create batch for user1
batch_user1 = MagicMock()
batch_user1.unified_object_id = "unified-batch-user1"
batch_user1.file_object = json.dumps({
"id": "batch_user1_abc",
"object": "batch",
"endpoint": "/v1/chat/completions",
"completion_window": "24h",
"status": "completed",
"created_at": 1234567890,
"input_file_id": "file-input-user1",
"request_counts": {"total": 1, "completed": 1, "failed": 0},
})
# Create batch for user2
batch_user2 = MagicMock()
batch_user2.unified_object_id = "unified-batch-user2"
batch_user2.file_object = json.dumps({
"id": "batch_user2_xyz",
"object": "batch",
"endpoint": "/v1/chat/completions",
"completion_window": "24h",
"status": "completed",
"created_at": 1234567891,
"input_file_id": "file-input-user2",
"request_counts": {"total": 2, "completed": 2, "failed": 0},
})
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
DualCache(), prisma_client=prisma_client
)
# Query with user1's API key - should only return user1's batch
prisma_client.db.litellm_managedobjecttable.find_many.return_value = [batch_user1]
result_user1 = await proxy_managed_files.list_user_batches(
user_api_key_dict=UserAPIKeyAuth(user_id="user1"),
limit=10,
)
assert len(result_user1["data"]) == 1
assert result_user1["data"][0].id == "unified-batch-user1"
prisma_client.db.litellm_managedobjecttable.find_many.assert_called_with(
where={"file_purpose": "batch", "created_by": "user1"},
take=10,
order={"created_at": "desc"},
)
# Query with user2's API key - should only return user2's batch
prisma_client.db.litellm_managedobjecttable.find_many.return_value = [batch_user2]
result_user2 = await proxy_managed_files.list_user_batches(
user_api_key_dict=UserAPIKeyAuth(user_id="user2"),
limit=10,
)
assert len(result_user2["data"]) == 1
assert result_user2["data"][0].id == "unified-batch-user2"
prisma_client.db.litellm_managedobjecttable.find_many.assert_called_with(
where={"file_purpose": "batch", "created_by": "user2"},
take=10,
order={"created_at": "desc"},
)
@pytest.mark.asyncio
async def test_return_unified_file_id_includes_expires_at():
from litellm.types.llms.openai import OpenAIFileObject
# Create a mock file object with expires_at set
file_object = OpenAIFileObject(
id="file-abc123",
object="file",
bytes=1234,
created_at=1234567890,
filename="test.jsonl",
purpose="batch",
status="uploaded",
expires_at=1234657890,
)
file_object._hidden_params = {"model_id": "test-model-id"}
create_file_request = {
"file": ("test.jsonl", b"test content", "application/jsonl"),
"purpose": "batch",
}
internal_usage_cache = MagicMock()
result = await _PROXY_LiteLLMManagedFiles.return_unified_file_id(
file_objects=[file_object],
create_file_request=create_file_request,
internal_usage_cache=internal_usage_cache,
litellm_parent_otel_span=None,
target_model_names_list=["gpt-4o"],
)
# Verify expires_at is passed through
assert result.expires_at == 1234657890
assert result.purpose == "batch"
assert result.filename == "test.jsonl"
assert result.bytes == 1234
assert result.created_at == 1234567890
assert _is_base64_encoded_unified_file_id(result.id)