diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index cf5c2b0905d..56e7f7ba633 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -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] diff --git a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py index 347a8fcd023..bbcdbac2709 100644 --- a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py +++ b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py @@ -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"}], )