fix(batches): reject unresolvable list cursors and derive has_more from row count

An `after` that does not resolve to a batch the caller can list now returns
400 instead of an empty page. An empty page is indistinguishable from the end
of the list, so a stale or malformed cursor silently truncated a client's batch
list. The lookup is scoped to the caller's own rows, so a Prisma cursor can no
longer be anchored to another user's batch.

`has_more` now comes from whether an extra row exists rather than from whether
the page came back full. Reporting fullness made every client fetch one extra
empty page when the batch count was an exact multiple of `limit`, and made a
page shortened by an unparseable row look like the end of the list, hiding the
older batches behind it.

Also drops the unreachable `target_model_names` oversampling branch; that
argument raises a few lines above it.
This commit is contained in:
mateo-berri 2026-07-24 19:33:33 -07:00
parent 93f27641ae
commit c1ea54a1d0
2 changed files with 261 additions and 25 deletions

View file

@ -315,29 +315,37 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
where_clause: Dict[str, Any] = {"file_purpose": "batch", **owner_filter}
fetch_limit = limit or 20
if target_model_names:
# Oversample so post-fetch model-name filtering still has enough rows.
fetch_limit = max(fetch_limit * 3, 100)
if after is not None:
cursor_row = (
await self.prisma_client.db.litellm_managedobjecttable.find_first(
where={**where_clause, "unified_object_id": after}
)
)
if cursor_row is None:
raise HTTPException(
status_code=400,
detail=f"Invalid 'after' cursor: no batch found with id '{after}'.",
)
page_size = limit or 20
cursor_args: Dict[str, Any] = (
{"cursor": {"unified_object_id": after}, "skip": 1} if after else {}
{"cursor": {"unified_object_id": after}, "skip": 1}
if after is not None
else {}
)
batches = await self.prisma_client.db.litellm_managedobjecttable.find_many(
where=where_clause,
take=fetch_limit,
take=page_size + 1,
order=[{"created_at": "desc"}, {"unified_object_id": "desc"}],
**cursor_args,
)
batch_objects: List[LiteLLMBatch] = []
for batch in batches:
try:
# Stop once we have enough after filtering
if len(batch_objects) >= (limit or 20):
break
has_more = len(batches) > page_size
batch_objects: List[LiteLLMBatch] = []
for batch in batches[:page_size]:
try:
batch_data = (
json.loads(batch.file_object)
if isinstance(batch.file_object, str)
@ -353,9 +361,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
)
continue
return build_list_page(
batch_objects, has_more=len(batch_objects) == (limit or 20)
)
return build_list_page(batch_objects, has_more=has_more)
async def get_user_created_file_ids(
self, user_api_key_dict: UserAPIKeyAuth, model_object_ids: List[str]

View file

@ -1771,7 +1771,7 @@ async def test_list_batches_from_managed_objects_table():
# 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,
take=11,
order=[{"created_at": "desc"}, {"unified_object_id": "desc"}],
)
@ -1801,7 +1801,7 @@ async def test_list_batches_from_managed_objects_table_empty_list():
# 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,
take=21,
order=[{"created_at": "desc"}, {"unified_object_id": "desc"}],
)
@ -1918,7 +1918,7 @@ async def test_list_batches_from_managed_objects_table_filters_by_created_by():
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,
take=11,
order=[{"created_at": "desc"}, {"unified_object_id": "desc"}],
)
@ -1933,7 +1933,7 @@ async def test_list_batches_from_managed_objects_table_filters_by_created_by():
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,
take=11,
order=[{"created_at": "desc"}, {"unified_object_id": "desc"}],
)
@ -1950,6 +1950,7 @@ async def test_list_batches_pagination_uses_unified_object_id_cursor():
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
prisma_client.db.litellm_managedobjecttable.find_first.return_value = MagicMock()
prisma_client.db.litellm_managedobjecttable.find_many.return_value = []
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
@ -1964,7 +1965,7 @@ async def test_list_batches_pagination_uses_unified_object_id_cursor():
prisma_client.db.litellm_managedobjecttable.find_many.assert_called_once_with(
where={"file_purpose": "batch", "created_by": "test-user"},
take=5,
take=6,
order=[{"created_at": "desc"}, {"unified_object_id": "desc"}],
cursor={"unified_object_id": "unified-batch-id-7"},
skip=1,
@ -2035,10 +2036,19 @@ async def test_list_batches_pagination_walks_all_pages_without_loops_or_gaps():
result = result[idx + skip:]
return result[:take]
async def fake_find_first(where):
return next(
(r for r in rows if r.unified_object_id == where.get("unified_object_id")),
None,
)
prisma_client = AsyncMock()
prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
side_effect=fake_find_many
)
prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock(
side_effect=fake_find_first
)
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
DualCache(), prisma_client=prisma_client
@ -2052,9 +2062,10 @@ async def test_list_batches_pagination_walks_all_pages_without_loops_or_gaps():
user_api_key_dict=user, limit=3, after=after
)
page_ids = [b.id for b in resp["data"]]
if not page_ids:
break
seen.extend(page_ids)
if not resp["has_more"]:
break
assert page_ids, "has_more was true but the page was empty"
assert resp["last_id"] != after, "cursor did not advance (pagination loop)"
after = resp["last_id"]
@ -2146,10 +2157,19 @@ async def test_list_batches_pagination_stable_when_created_at_ties():
result = result[idx + skip:]
return result[:take]
async def fake_find_first(where):
return next(
(r for r in rows if r.unified_object_id == where.get("unified_object_id")),
None,
)
prisma_client = AsyncMock()
prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
side_effect=fake_find_many
)
prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock(
side_effect=fake_find_first
)
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
DualCache(), prisma_client=prisma_client
@ -2163,9 +2183,10 @@ async def test_list_batches_pagination_stable_when_created_at_ties():
user_api_key_dict=user, limit=2, after=after
)
page_ids = [b.id for b in resp["data"]]
if not page_ids:
break
seen.extend(page_ids)
if not resp["has_more"]:
break
assert page_ids, "has_more was true but the page was empty"
assert resp["last_id"] != after, "cursor did not advance (pagination loop)"
after = resp["last_id"]
@ -2173,6 +2194,215 @@ async def test_list_batches_pagination_stable_when_created_at_ties():
assert len(seen) == len(set(seen)), "a tied batch was returned more than once"
def _managed_batch_row(index, file_object=None):
row = MagicMock()
row.id = f"pk-{index:03d}"
raw = f"litellm_proxy;model_id:gpt-4o-batch;llm_batch_id:batch_{index:03d}"
row.unified_object_id = base64.urlsafe_b64encode(raw.encode()).decode().rstrip("=")
row.created_at = 1_000_000 + index
row.file_object = (
file_object
if file_object is not None
else json.dumps(
{
"id": f"batch_provider_{index:03d}",
"object": "batch",
"endpoint": "/v1/chat/completions",
"completion_window": "24h",
"status": "completed",
"created_at": 1_000_000 + index,
"input_file_id": f"file-input-{index:03d}",
"request_counts": {"total": 1, "completed": 1, "failed": 0},
}
)
)
return row
def _fake_managed_object_table(rows):
async def find_many(where, take, order, cursor=None, skip=0):
result = sorted(
rows, key=lambda r: (r.created_at, r.unified_object_id), reverse=True
)
if cursor is not None:
(cur_field, cur_val), = cursor.items()
idx = next(
(i for i, r in enumerate(result) if getattr(r, cur_field) == cur_val),
None,
)
if idx is None:
return []
result = result[idx + skip:]
return result[:take]
async def find_first(where):
return next(
(r for r in rows if r.unified_object_id == where.get("unified_object_id")),
None,
)
prisma_client = AsyncMock()
prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
side_effect=find_many
)
prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock(
side_effect=find_first
)
return prisma_client
async def _walk_batch_pages(proxy_managed_files, user, limit, max_pages=20):
pages = []
after = None
for _ in range(max_pages):
resp = await proxy_managed_files.list_user_batches(
user_api_key_dict=user, limit=limit, after=after
)
pages.append(resp)
if not resp["has_more"]:
break
assert resp["last_id"] is not None, "has_more was true but there is no cursor"
after = resp["last_id"]
return pages
@pytest.mark.asyncio
async def test_list_batches_rejects_unknown_after_cursor():
"""An ``after`` that does not resolve to a batch the caller can see is a
client error, not an empty page.
Returning ``[]`` for an unresolvable cursor is indistinguishable from
"you have reached the end of the list", so a client walking pages with a
stale or malformed cursor silently sees a truncated batch list instead of
an error it can act on.
"""
from fastapi import HTTPException
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock(
return_value=None
)
prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[])
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
DualCache(), prisma_client=prisma_client
)
with pytest.raises(HTTPException) as exc_info:
await proxy_managed_files.list_user_batches(
user_api_key_dict=UserAPIKeyAuth(user_id="test-user"),
limit=3,
after="does-not-exist-xyz",
)
assert exc_info.value.status_code == 400
assert "does-not-exist-xyz" in str(exc_info.value.detail)
prisma_client.db.litellm_managedobjecttable.find_many.assert_not_called()
prisma_client.db.litellm_managedobjecttable.find_first.assert_called_once_with(
where={
"file_purpose": "batch",
"created_by": "test-user",
"unified_object_id": "does-not-exist-xyz",
}
)
@pytest.mark.asyncio
async def test_list_batches_rejects_after_cursor_owned_by_another_user():
"""The cursor lookup must be scoped to the rows the caller can list.
A Prisma cursor resolves by unique column regardless of the ``where``
filter, so an unscoped cursor would let one user anchor their page window
to another user's batch and learn when it was created.
"""
from fastapi import HTTPException
from litellm.proxy._types import UserAPIKeyAuth
other_users_batch = _managed_batch_row(0)
async def find_first(where):
if where.get("created_by") != "user-b":
return None
return other_users_batch
prisma_client = AsyncMock()
prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock(
side_effect=find_first
)
prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[])
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
DualCache(), prisma_client=prisma_client
)
with pytest.raises(HTTPException) as exc_info:
await proxy_managed_files.list_user_batches(
user_api_key_dict=UserAPIKeyAuth(user_id="user-a"),
limit=3,
after=other_users_batch.unified_object_id,
)
assert exc_info.value.status_code == 400
prisma_client.db.litellm_managedobjecttable.find_many.assert_not_called()
@pytest.mark.asyncio
async def test_list_batches_has_more_false_on_exactly_full_final_page():
"""``has_more`` must mean "another row exists", not "this page is full".
With a batch count that is an exact multiple of ``limit``, reporting
``has_more`` off page fullness makes every client fetch one extra empty
page before it can stop.
"""
from litellm.proxy._types import UserAPIKeyAuth
rows = [_managed_batch_row(i) for i in range(4)]
prisma_client = _fake_managed_object_table(rows)
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
DualCache(), prisma_client=prisma_client
)
pages = await _walk_batch_pages(
proxy_managed_files, UserAPIKeyAuth(user_id="test-user"), limit=2
)
seen = [batch.id for page in pages for batch in page["data"]]
assert seen == [r.unified_object_id for r in reversed(rows)]
assert [page["has_more"] for page in pages] == [True, False]
assert prisma_client.db.litellm_managedobjecttable.find_many.call_count == 2
@pytest.mark.asyncio
async def test_list_batches_unparseable_row_does_not_truncate_pagination():
"""A row that fails to parse must not end pagination early.
Skipping a corrupt row shortens the page, so deriving ``has_more`` from
the number of returned batches reports "no more results" while older
batches are still unread, silently hiding them from the caller.
"""
from litellm.proxy._types import UserAPIKeyAuth
rows = [_managed_batch_row(i) for i in range(4)]
rows[2].file_object = "{ not valid json"
prisma_client = _fake_managed_object_table(rows)
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
DualCache(), prisma_client=prisma_client
)
pages = await _walk_batch_pages(
proxy_managed_files, UserAPIKeyAuth(user_id="test-user"), limit=2
)
seen = [batch.id for page in pages for batch in page["data"]]
assert seen == [rows[3].unified_object_id, rows[1].unified_object_id, rows[0].unified_object_id]
assert len(seen) == len(set(seen))
@pytest.mark.asyncio
async def test_return_unified_file_id_includes_expires_at():
from litellm.types.llms.openai import OpenAIFileObject
@ -2547,7 +2777,7 @@ async def test_list_batches_only_returns_user_own_batches():
# Verify the database query filtered by user_id
prisma_client.db.litellm_managedobjecttable.find_many.assert_called_once_with(
where={"file_purpose": "batch", "created_by": "user_a_id"},
take=10,
take=11,
order=[{"created_at": "desc"}, {"unified_object_id": "desc"}],
)