From 42271b282a7d195b0f8b0ac324b852ab8109b8fa Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 19 Sep 2026 04:44:26 -0700 Subject: [PATCH] fix(batches): guard batch status writes against stale reads and disable per-line fallbacks --- .../litellm_executed_batches.py | 65 +++--- .../test_litellm_executed_batches.py | 194 +++++++++++++++--- 2 files changed, 198 insertions(+), 61 deletions(-) diff --git a/litellm/proxy/batches_endpoints/litellm_executed_batches.py b/litellm/proxy/batches_endpoints/litellm_executed_batches.py index 7bd4c678183..b121ad1590e 100644 --- a/litellm/proxy/batches_endpoints/litellm_executed_batches.py +++ b/litellm/proxy/batches_endpoints/litellm_executed_batches.py @@ -3,7 +3,7 @@ import json import time from collections.abc import Awaitable, Callable, Mapping, Sequence from dataclasses import dataclass -from datetime import datetime, timezone +from datetime import datetime, timedelta, timezone from itertools import pairwise from types import MappingProxyType from typing import TYPE_CHECKING, Final, Literal, Protocol, TypeAlias, runtime_checkable @@ -28,11 +28,7 @@ from litellm.models.managed_files import LiteLLM_ManagedFileTable from litellm.proxy._types import ProxyErrorTypes, ProxyException, UserAPIKeyAuth from litellm.proxy.auth.auth_utils import is_request_body_safe from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup -from litellm.proxy.openai_files_endpoints.common_utils import ( - LITELLM_EXECUTED_BATCH_ID_PREFIX, - convert_b64_uid_to_unified_uid, - get_batch_id_from_unified_batch_id, -) +from litellm.proxy.openai_files_endpoints.common_utils import LITELLM_EXECUTED_BATCH_ID_PREFIX from litellm.proxy.openai_files_endpoints.storage_backend_service import StorageBackendFileService from litellm.proxy.utils import PrismaClient, ProxyLogging from litellm.repositories.table_repositories import ManagedObjectRepository @@ -41,6 +37,7 @@ from litellm.types.utils import LITELLM_EXECUTED_BATCH_PROVIDERS, ExtractedFileD if TYPE_CHECKING: from prisma import models as prisma_models + from prisma import types as prisma_types from litellm.router import Router @@ -154,7 +151,6 @@ class ManagedBatchStore(Protocol): user_api_key_dict: UserAPIKeyAuth, request_tags: Sequence[str] | None = None, persist_attribution: bool = False, - create_if_missing: bool = True, batch_processed: bool = False, ) -> None: ... @@ -375,10 +371,6 @@ def executed_batch_runner_lost(status: str, updated_at: datetime) -> bool: return (datetime.now(timezone.utc) - updated_at).total_seconds() > _STALE_AFTER_SECONDS -def _llm_batch_id_of(unified_batch_id: str) -> str: - return get_batch_id_from_unified_batch_id(convert_b64_uid_to_unified_uid(unified_batch_id)) - - class _StopWatch: def __init__(self, load_status: Callable[[], Awaitable[str | None]], interval_seconds: float) -> None: self._load_status = load_status @@ -488,8 +480,10 @@ class LiteLLMExecutedBatchRunner: cancelling: Final = current.model_copy( update=MappingProxyType({"status": "cancelling", "cancelling_at": int(time.time())}) ) - await self._store(cancelling, user_api_key_dict) - return cancelling + unchanged: Final[prisma_types.LiteLLM_ManagedObjectTableWhereInput] = {"status": current.status} + if await self._store_unless_changed(cancelling, unchanged, user_api_key_dict): + return cancelling + return await self.cancel(unified_batch_id, user_api_key_dict) async def fail_abandoned(self, batch: LiteLLMBatch, user_api_key_dict: UserAPIKeyAuth) -> LiteLLMBatch: error: Final = BatchError(message=_RUNNER_LOST_MESSAGE, code="runner_lost") @@ -497,8 +491,16 @@ class LiteLLMExecutedBatchRunner: failed: Final = batch.model_copy( update=MappingProxyType({"status": "failed", "failed_at": int(time.time()), "errors": errors}) ) - await self._store(failed, user_api_key_dict) - return failed + untouched: Final[prisma_types.DateTimeFilter] = { + "lt": datetime.now(timezone.utc) - timedelta(seconds=_STALE_AFTER_SECONDS) + } + still_abandoned: Final[prisma_types.LiteLLM_ManagedObjectTableWhereInput] = { + "status": batch.status, + "updated_at": untouched, + } + if await self._store_unless_changed(failed, still_abandoned, user_api_key_dict): + return failed + return await self._load_batch(batch.id) or batch def _body_rejection(self, model: str) -> BodyRejection: def reject(body: Mapping[str, object]) -> str | None: @@ -598,7 +600,9 @@ class LiteLLMExecutedBatchRunner: return RowOutcome(custom_id=line.custom_id, status_code=200, body=body, succeeded=True) async def _dispatch(self, run: _BatchRun, line: BatchInputLine) -> Mapping[str, object]: - params: Final = MappingProxyType({**line.body, "model": run.model, "metadata": self._row_metadata(run)}) + params: Final = MappingProxyType( + {**line.body, "model": run.model, "metadata": self._row_metadata(run), "disable_fallbacks": True} + ) return _dump(await self._router_call(run.endpoint)(**params)) def _router_call(self, endpoint: BatchEndpoint) -> _RouterCall: @@ -651,19 +655,26 @@ class LiteLLMExecutedBatchRunner: updated: Final = current.model_copy( update=MappingProxyType({**fields, "status": status, f"{status}_at": int(time.time())}) ) - await self._store(updated, run.user_api_key_dict) - return status + unchanged: Final[prisma_types.LiteLLM_ManagedObjectTableWhereInput] = {"status": current.status} + if await self._store_unless_changed(updated, unchanged, run.user_api_key_dict): + return status + return await self._advance(run, requested, fields) - async def _store(self, batch: LiteLLMBatch, user_api_key_dict: UserAPIKeyAuth) -> None: - await self.managed_files.store_unified_object_id( - unified_object_id=batch.id, - file_object=batch, - litellm_parent_otel_span=user_api_key_dict.parent_otel_span, - model_object_id=_llm_batch_id_of(batch.id), - file_purpose="batch", - user_api_key_dict=user_api_key_dict, - create_if_missing=False, + async def _store_unless_changed( + self, + batch: LiteLLMBatch, + guard: "prisma_types.LiteLLM_ManagedObjectTableWhereInput", + user_api_key_dict: UserAPIKeyAuth, + ) -> bool: + updated_rows: Final = await ManagedObjectRepository(self.prisma_client).table.update_many( + where={"unified_object_id": batch.id, **guard}, # mutable-ok: Prisma filter + data={ # mutable-ok: Prisma payload + "file_object": batch.model_dump_json(), + "status": batch.status, + "updated_by": user_api_key_dict.user_id, + }, ) + return updated_rows > 0 async def _find_row(self, unified_batch_id: str) -> "prisma_models.LiteLLM_ManagedObjectTable | None": return await ManagedObjectRepository(self.prisma_client).table.find_first( diff --git a/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py b/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py index e5ed873a29d..cb566d4dca7 100644 --- a/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py +++ b/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py @@ -105,6 +105,10 @@ class ProviderRateLimited(Exception): class StoredObject: file_object: str status: str + updated_at: datetime + + def batch(self) -> LiteLLMBatch: + return LiteLLMBatch.model_validate_json(self.file_object) @dataclass(frozen=True, slots=True) @@ -114,10 +118,20 @@ class StoreCall: status: str request_tags: tuple[str, ...] | None persist_attribution: bool - create_if_missing: bool batch_processed: bool +@dataclass(frozen=True, slots=True) +class StatusWrite: + unified_object_id: str + status: str + columns: frozenset[str] + + +STATUS_WRITE_COLUMNS: Final = frozenset({"file_object", "status", "updated_by"}) +STALE: Final = timedelta(seconds=litellm_executed_batches._STALE_AFTER_SECONDS + 20) + + class FakeManagedBatchStore: def __init__(self, files: Mapping[str, LiteLLM_ManagedFileTable]) -> None: self.files = files @@ -142,7 +156,6 @@ class FakeManagedBatchStore: user_api_key_dict: UserAPIKeyAuth, request_tags: Sequence[str] | None = None, persist_attribution: bool = False, - create_if_missing: bool = True, batch_processed: bool = False, ) -> None: self.calls.append( @@ -152,18 +165,18 @@ class FakeManagedBatchStore: status=file_object.status, request_tags=tuple(request_tags) if request_tags is not None else None, persist_attribution=persist_attribution, - create_if_missing=create_if_missing, batch_processed=batch_processed, ) ) - if create_if_missing or unified_object_id in self.objects: - self.write(file_object) + self.write(file_object) - def write(self, batch: LiteLLMBatch) -> None: - self.objects[batch.id] = StoredObject(file_object=batch.model_dump_json(), status=batch.status) + def write(self, batch: LiteLLMBatch, age: timedelta = timedelta(0)) -> None: + self.objects[batch.id] = StoredObject( + file_object=batch.model_dump_json(), status=batch.status, updated_at=datetime.now(timezone.utc) - age + ) def batch(self, unified_batch_id: str) -> LiteLLMBatch: - return LiteLLMBatch.model_validate_json(self.objects[unified_batch_id].file_object) + return self.objects[unified_batch_id].batch() REAL_HOOK: Final = _PROXY_LiteLLMManagedFiles(internal_usage_cache=MagicMock(), prisma_client=MagicMock()) @@ -174,26 +187,51 @@ class RealIdManagedBatchStore(FakeManagedBatchStore): return REAL_HOOK.get_unified_batch_id(batch_id=batch_id, model_id=model_id) +def row_matches(row: StoredObject, where: Mapping[str, object]) -> bool: + if "status" in where and row.status != where["status"]: + return False + match where.get("updated_at"): + case {"lt": datetime() as before}: + return row.updated_at < before + case _: + return True + + class FakeManagedObjectTable: - def __init__(self, objects: Mapping[str, StoredObject]) -> None: + def __init__(self, objects: dict[str, StoredObject]) -> None: self.objects = objects self.touches: list[tuple[str, str | None]] = [] + self.writes: list[StatusWrite] = [] + self.after_read: Callable[[StoredObject | None], None] | None = None async def find_first(self, where: Mapping[str, str]) -> StoredObject | None: - return self.objects.get(where["unified_object_id"]) + row = self.objects.get(where["unified_object_id"]) + if self.after_read is not None: + self.after_read(row) + return row - async def update_many(self, where: Mapping[str, str], data: Mapping[str, str | None]) -> int: - self.touches.append((where["unified_object_id"], data["updated_by"])) + async def update_many(self, where: Mapping[str, object], data: Mapping[str, str | None]) -> int: + unified_object_id = str(where["unified_object_id"]) + row = self.objects.get(unified_object_id) + if row is None or not row_matches(row, where): + return 0 + now = datetime.now(timezone.utc) + if "status" not in data: + self.touches.append((unified_object_id, data["updated_by"])) + self.objects[unified_object_id] = StoredObject(row.file_object, row.status, now) + return 1 + self.writes.append(StatusWrite(unified_object_id, str(data["status"]), frozenset(data))) + self.objects[unified_object_id] = StoredObject(str(data["file_object"]), str(data["status"]), now) return 1 class FakeDb: - def __init__(self, objects: Mapping[str, StoredObject]) -> None: + def __init__(self, objects: dict[str, StoredObject]) -> None: self.litellm_managedobjecttable = FakeManagedObjectTable(objects) class FakePrismaClient: - def __init__(self, objects: Mapping[str, StoredObject]) -> None: + def __init__(self, objects: dict[str, StoredObject]) -> None: self.db = FakeDb(objects) @@ -320,6 +358,13 @@ class Harness: await asyncio.gather(*list(litellm_executed_batches._RUNNING_BATCHES)) return created, self.store.batch(created.id) + @property + def table(self) -> FakeManagedObjectTable: + return self.prisma.db.litellm_managedobjecttable + + def written_statuses(self) -> list[str]: + return [write.status for write in self.table.writes] + def make_runner( content: bytes = TWO_CHAT_ROWS, @@ -354,7 +399,9 @@ def make_runner( return Harness(runner, store, router, uploads, storage, storage_factory, prisma, user) -def seeded_batch(store: FakeManagedBatchStore, status: Literal["in_progress", "completed"]) -> LiteLLMBatch: +def seeded_batch( + store: FakeManagedBatchStore, status: Literal["in_progress", "completed"], age: timedelta = timedelta(0) +) -> LiteLLMBatch: batch = LiteLLMBatch( id=store.get_unified_batch_id(batch_id="litellm_batch_seed", model_id=DEPLOYMENT_ID), object="batch", @@ -365,7 +412,7 @@ def seeded_batch(store: FakeManagedBatchStore, status: Literal["in_progress", "c created_at=1, model=BATCH_MODEL, ) - store.write(batch) + store.write(batch, age) return batch @@ -745,7 +792,9 @@ async def test_create_forwards_row_credentials_when_the_admin_opted_in() -> None assert finished.status == "completed" assert finished.request_counts == BatchRequestCounts(completed=2, failed=0, total=2) - by_content = {call.kwargs["messages"][0]["content"]: call.kwargs for call in harness.router.acompletion.await_args_list} + by_content = { + call.kwargs["messages"][0]["content"]: call.kwargs for call in harness.router.acompletion.await_args_list + } assert by_content["hi 2"]["api_base"] == "https://evil.example" assert "api_base" not in by_content["hi 1"] @@ -760,19 +809,20 @@ async def test_running_batch_touches_its_row_until_it_finishes() -> None: harness.router.acompletion.side_effect = slow_dispatch created, finished = await harness.create_and_finish() - touches = harness.prisma.db.litellm_managedobjecttable.touches + touches = harness.table.touches assert finished.status == "completed" assert touches assert set(touches) == {(created.id, "user-1")} - assert [call.status for call in harness.store.calls] == ["validating", "in_progress", "finalizing", "completed"] + assert [call.status for call in harness.store.calls] == ["validating"] + assert harness.written_statuses() == ["in_progress", "finalizing", "completed"] beats_at_finish = len(touches) await asyncio.sleep(0.05) assert len(touches) == beats_at_finish -async def test_fail_abandoned_marks_the_batch_failed_with_the_runner_lost_error() -> None: +async def test_fail_abandoned_marks_a_stale_batch_failed_with_the_runner_lost_error() -> None: harness = make_runner() - batch = seeded_batch(harness.store, "in_progress") + batch = seeded_batch(harness.store, "in_progress", age=STALE) failed = await harness.runner.fail_abandoned(batch, harness.user) @@ -783,7 +833,63 @@ async def test_fail_abandoned_marks_the_batch_failed_with_the_runner_lost_error( (litellm_executed_batches._RUNNER_LOST_MESSAGE, "runner_lost") ] assert harness.store.batch(batch.id).status == "failed" - assert [(call.status, call.create_if_missing) for call in harness.store.calls] == [("failed", False)] + assert harness.store.calls == [] + assert harness.table.writes == [StatusWrite(batch.id, "failed", STATUS_WRITE_COLUMNS)] + + +async def test_fail_abandoned_leaves_a_batch_that_finished_after_the_stale_read() -> None: + harness = make_runner() + stale_read = seeded_batch(harness.store, "in_progress", age=STALE) + harness.store.write(stale_read.model_copy(update={"status": "completed", "output_file_id": "out-1"}), age=STALE) + + current = await harness.runner.fail_abandoned(stale_read, harness.user) + + assert (current.status, current.output_file_id) == ("completed", "out-1") + assert harness.store.batch(stale_read.id).status == "completed" + assert harness.table.writes == [] + + +async def test_fail_abandoned_leaves_a_batch_its_runner_touched_since_the_read() -> None: + harness = make_runner() + batch = seeded_batch(harness.store, "in_progress", age=STALE) + harness.store.write(batch) + + current = await harness.runner.fail_abandoned(batch, harness.user) + + assert current.status == "in_progress" + assert harness.store.batch(batch.id).status == "in_progress" + assert harness.table.writes == [] + + +async def test_run_does_not_reverse_a_failure_written_between_its_read_and_its_completed_write() -> None: + harness = make_runner() + + def fail_once_finalizing_is_read(row: StoredObject | None) -> None: + if row is not None and row.status == "finalizing": + harness.store.write(row.batch().model_copy(update={"status": "failed"})) + + harness.table.after_read = fail_once_finalizing_is_read + _, finished = await harness.create_and_finish() + + assert finished.status == "failed" + assert finished.output_file_id is None + assert harness.written_statuses() == ["in_progress", "finalizing"] + + +async def test_run_honours_a_cancel_written_between_its_read_and_its_finalizing_write() -> None: + harness = make_runner(content=jsonl(chat_row("row-1", "hi 1"))) + + def cancel_once_the_row_is_dispatched(row: StoredObject | None) -> None: + if row is not None and row.status == "in_progress" and harness.router.acompletion.await_count == 1: + harness.store.write(row.batch().model_copy(update={"status": "cancelling"})) + + harness.table.after_read = cancel_once_the_row_is_dispatched + _, finished = await harness.create_and_finish() + + assert finished.status == "cancelled" + assert finished.request_counts == BatchRequestCounts(completed=1, failed=0, total=1) + assert finished.output_file_id == "unified-output-1" + assert harness.written_statuses() == ["in_progress", "cancelling", "cancelled"] async def test_running_batch_stops_and_writes_nothing_once_a_retriever_marked_it_failed( @@ -803,7 +909,8 @@ async def test_running_batch_stops_and_writes_nothing_once_a_retriever_marked_it assert harness.router.acompletion.await_count == 1 assert finished.status == "failed" - assert [call.status for call in harness.store.calls] == ["validating", "in_progress"] + assert [call.status for call in harness.store.calls] == ["validating"] + assert harness.written_statuses() == ["in_progress"] assert harness.uploads.calls == [] @@ -828,6 +935,7 @@ async def test_each_endpoint_awaits_only_its_router_method( assert awaited == {name: int(name == method) for name in ROUTER_METHODS} kwargs = getattr(harness.router, method).await_args.kwargs assert kwargs["model"] == BATCH_MODEL + assert kwargs["disable_fallbacks"] is True assert all(kwargs[key] == value for key, value in body.items()) @@ -845,7 +953,7 @@ async def test_cancel_terminal_batch_is_400() -> None: await harness.runner.cancel(batch.id, harness.user) assert raised.value.code == "400" assert "completed" in raised.value.message - assert harness.store.calls == [] + assert harness.table.writes == [] async def test_cancel_marks_a_running_batch_cancelling_once() -> None: @@ -857,12 +965,30 @@ async def test_cancel_marks_a_running_batch_cancelling_once() -> None: assert cancelled.status == "cancelling" assert cancelled.cancelling_at is not None assert harness.store.batch(batch.id).status == "cancelling" - assert [(call.status, call.create_if_missing) for call in harness.store.calls] == [("cancelling", False)] + assert harness.store.calls == [] + assert harness.table.writes == [StatusWrite(batch.id, "cancelling", STATUS_WRITE_COLUMNS)] again = await harness.runner.cancel(batch.id, harness.user) assert again.model_dump() == cancelled.model_dump() - assert len(harness.store.calls) == 1 + assert len(harness.table.writes) == 1 + + +async def test_cancel_racing_a_completion_is_400_and_leaves_the_batch_completed() -> None: + harness = make_runner() + batch = seeded_batch(harness.store, "in_progress") + + def complete_once_read(row: StoredObject | None) -> None: + if row is not None and row.status == "in_progress": + harness.store.write(row.batch().model_copy(update={"status": "completed"})) + + harness.table.after_read = complete_once_read + with pytest.raises(ProxyException) as raised: + await harness.runner.cancel(batch.id, harness.user) + + assert raised.value.code == "400" + assert harness.store.batch(batch.id).status == "completed" + assert harness.table.writes == [] async def test_running_batch_skips_the_remaining_rows_after_an_operator_cancel( @@ -905,10 +1031,10 @@ async def test_only_the_create_write_carries_attribution_and_billing_flags() -> harness = make_runner() await harness.create_and_finish() - assert [call.status for call in harness.store.calls] == ["validating", "in_progress", "finalizing", "completed"] - flags = [(call.persist_attribution, call.batch_processed, call.create_if_missing) for call in harness.store.calls] - assert flags[0] == (True, True, True) - assert flags[1:] == [(False, False, False)] * 3 + assert [(call.status, call.persist_attribution, call.batch_processed) for call in harness.store.calls] == [ + ("validating", True, True) + ] + assert [write.columns for write in harness.table.writes] == [STATUS_WRITE_COLUMNS] * 3 async def test_run_completes_under_the_real_hooks_base64_batch_id() -> None: @@ -917,9 +1043,8 @@ async def test_run_completes_under_the_real_hooks_base64_batch_id() -> None: assert _is_base64_encoded_unified_file_id(created.id) assert finished.status == "completed" - llm_batch_id = harness.store.calls[0].model_object_id - assert llm_batch_id.startswith("litellm_batch_") - assert [call.model_object_id for call in harness.store.calls] == [llm_batch_id] * 4 + assert [call.model_object_id.startswith("litellm_batch_") for call in harness.store.calls] == [True] + assert [write.unified_object_id for write in harness.table.writes] == [created.id] * 3 async def test_cancel_works_under_the_real_hooks_base64_batch_id() -> None: @@ -928,4 +1053,5 @@ async def test_cancel_works_under_the_real_hooks_base64_batch_id() -> None: cancelled = await harness.runner.cancel(batch.id, harness.user) assert cancelled.status == "cancelling" - assert [call.model_object_id for call in harness.store.calls] == ["litellm_batch_seed"] + assert harness.store.batch(batch.id).status == "cancelling" + assert [write.unified_object_id for write in harness.table.writes] == [batch.id]