diff --git a/litellm/proxy/batches_endpoints/litellm_executed_batches.py b/litellm/proxy/batches_endpoints/litellm_executed_batches.py index b121ad1590e..5a7061d9ab1 100644 --- a/litellm/proxy/batches_endpoints/litellm_executed_batches.py +++ b/litellm/proxy/batches_endpoints/litellm_executed_batches.py @@ -42,7 +42,9 @@ if TYPE_CHECKING: from litellm.router import Router BatchEndpoint: TypeAlias = Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions", "/v1/responses"] -BatchStatus: TypeAlias = Literal["in_progress", "finalizing", "completed", "failed", "cancelling", "cancelled"] +BatchStatus: TypeAlias = Literal[ + "in_progress", "finalizing", "completed", "failed", "cancelling", "cancelled", "expired" +] TERMINAL_BATCH_STATUSES: Final[frozenset[str]] = frozenset({"completed", "failed", "cancelled", "expired"}) _STOP_STATUSES: Final[frozenset[str]] = TERMINAL_BATCH_STATUSES | frozenset({"cancelling"}) _BATCH_ENDPOINT_ADAPTER: Final[TypeAdapter[BatchEndpoint]] = TypeAdapter(BatchEndpoint) @@ -52,6 +54,7 @@ _STALE_AFTER_SECONDS: Final = 180.0 _FILES_API_PROBE_TIMEOUT_SECONDS: Final = 5.0 _COMPLETION_WINDOW_SECONDS: Final = 24 * 60 * 60 _RUNNER_LOST_MESSAGE: Final = "the proxy replica running this batch stopped before it finished; resubmit the batch" +_EXPIRED_MESSAGE: Final = "This request could not be executed before the completion window expired." _ROUTER_METHODS: Final[Mapping[BatchEndpoint, str]] = MappingProxyType( { "/v1/chat/completions": "acompletion", @@ -61,7 +64,7 @@ _ROUTER_METHODS: Final[Mapping[BatchEndpoint, str]] = MappingProxyType( } ) _CANCELLING_TRANSITIONS: Final[Mapping[BatchStatus, BatchStatus]] = MappingProxyType( - {"completed": "cancelled", "in_progress": "cancelling", "finalizing": "cancelling"} + {"completed": "cancelled", "expired": "cancelled", "in_progress": "cancelling", "finalizing": "cancelling"} ) LITELLM_EXECUTED_BATCH_UPLOAD_GUIDANCE: Final = ( "upload it through POST /v1/files with purpose=batch and either the x-litellm-model header or the " @@ -89,11 +92,16 @@ class _ResultResponse(TypedDict): body: ReadOnly[Mapping[str, object]] +class _LineError(TypedDict): + code: ReadOnly[str] + message: ReadOnly[str] + + class _ResultLine(TypedDict): id: ReadOnly[str] custom_id: ReadOnly[str] - response: ReadOnly[_ResultResponse] - error: ReadOnly[None] + response: ReadOnly[_ResultResponse | None] + error: ReadOnly[_LineError | None] class BatchInputLine(BaseModel): @@ -122,6 +130,11 @@ class RowOutcome: succeeded: bool +@dataclass(frozen=True, slots=True) +class ExpiredRow: + custom_id: str + + @dataclass(frozen=True, slots=True) class _BatchRun: unified_batch_id: str @@ -131,6 +144,7 @@ class _BatchRun: lines: tuple[BatchInputLine, ...] user_api_key_dict: UserAPIKeyAuth request_tags: tuple[str, ...] + deadline: float @runtime_checkable @@ -339,16 +353,30 @@ def _error_body(error: Exception) -> _ErrorBody: return body -def _result_line(outcome: RowOutcome) -> _ResultLine: +def _line_response(outcome: RowOutcome | ExpiredRow) -> _ResultResponse | None: + if isinstance(outcome, ExpiredRow): + return None + response: Final[_ResultResponse] = { + "status_code": outcome.status_code, + "request_id": f"req_{uuid_module.uuid4().hex[:24]}", + "body": outcome.body, + } + return response + + +def _line_error(outcome: RowOutcome | ExpiredRow) -> _LineError | None: + if isinstance(outcome, RowOutcome): + return None + error: Final[_LineError] = {"code": "batch_expired", "message": _EXPIRED_MESSAGE} + return error + + +def _result_line(outcome: RowOutcome | ExpiredRow) -> _ResultLine: line: Final[_ResultLine] = { "id": f"batch_req_{uuid_module.uuid4().hex[:24]}", "custom_id": outcome.custom_id, - "response": { - "status_code": outcome.status_code, - "request_id": f"req_{uuid_module.uuid4().hex[:24]}", - "body": outcome.body, - }, - "error": None, + "response": _line_response(outcome), + "error": _line_error(outcome), } return line @@ -399,6 +427,7 @@ class LiteLLMExecutedBatchRunner: general_settings: Mapping[str, object], concurrency: int = LITELLM_EXECUTED_BATCH_CONCURRENCY, heartbeat_seconds: float = _HEARTBEAT_SECONDS, + completion_window_seconds: float = _COMPLETION_WINDOW_SECONDS, storage_backend_factory: _StorageBackendFactory = get_storage_backend, upload_result_file: _ResultFileUploader = StorageBackendFileService.upload_file_to_storage_backend, ) -> None: @@ -409,6 +438,7 @@ class LiteLLMExecutedBatchRunner: self.general_settings = general_settings self.concurrency = concurrency self.heartbeat_seconds = heartbeat_seconds + self.completion_window_seconds = completion_window_seconds self.storage_backend_factory = storage_backend_factory self.upload_result_file = upload_result_file @@ -429,7 +459,8 @@ class LiteLLMExecutedBatchRunner: llm_batch_id: Final = f"{LITELLM_EXECUTED_BATCH_ID_PREFIX}{uuid_module.uuid4().hex}" model_id: Final = next(iter(self.llm_router.get_model_ids(model_name=model)), model) unified_batch_id: Final = self.managed_files.get_unified_batch_id(batch_id=llm_batch_id, model_id=model_id) - created_at: Final = int(time.time()) + now: Final = time.time() + created_at: Final = int(now) batch: Final = LiteLLMBatch( id=unified_batch_id, object="batch", @@ -438,7 +469,7 @@ class LiteLLMExecutedBatchRunner: completion_window="24h", status="validating", created_at=created_at, - expires_at=created_at + _COMPLETION_WINDOW_SECONDS, + expires_at=created_at + int(self.completion_window_seconds), metadata=create_request.get("metadata"), model=model, request_counts=BatchRequestCounts(completed=0, failed=0, total=len(parsed)), @@ -463,6 +494,7 @@ class LiteLLMExecutedBatchRunner: lines=parsed, user_api_key_dict=user_api_key_dict, request_tags=tuple(request_tags or ()), + deadline=now + self.completion_window_seconds, ) task: Final = asyncio.create_task(self._run(run)) _RUNNING_BATCHES.add(task) @@ -572,14 +604,21 @@ class LiteLLMExecutedBatchRunner: outcomes: Final = tuple(outcome for outcome in results if outcome is not None) if await self._advance(run, "finalizing") is None: return - succeeded: Final = tuple(outcome for outcome in outcomes if outcome.succeeded) - failed: Final = tuple(outcome for outcome in outcomes if not outcome.succeeded) + succeeded: Final = tuple( + outcome for outcome in outcomes if isinstance(outcome, RowOutcome) and outcome.succeeded + ) + failed: Final = tuple( + outcome for outcome in outcomes if isinstance(outcome, ExpiredRow) or not outcome.succeeded + ) output_file_id: Final = await self._upload_results(run, "output", succeeded) error_file_id: Final = await self._upload_results(run, "error", failed) request_counts: Final = BatchRequestCounts(completed=len(succeeded), failed=len(failed), total=len(run.lines)) + final_status: Final[BatchStatus] = ( + "expired" if any(isinstance(outcome, ExpiredRow) for outcome in outcomes) else "completed" + ) await self._advance( run, - "completed", + final_status, MappingProxyType( {"output_file_id": output_file_id, "error_file_id": error_file_id, "request_counts": request_counts} ), @@ -587,12 +626,17 @@ class LiteLLMExecutedBatchRunner: async def _run_row( self, run: _BatchRun, line: BatchInputLine, watch: _StopWatch, semaphore: asyncio.Semaphore - ) -> RowOutcome | None: + ) -> RowOutcome | ExpiredRow | None: async with semaphore: if await watch.stopped(): return None + remaining: Final = run.deadline - time.time() + if remaining <= 0: + return ExpiredRow(custom_id=line.custom_id) try: - body: Final = await self._dispatch(run, line) + body: Final = await asyncio.wait_for(self._dispatch(run, line), timeout=remaining) + except asyncio.TimeoutError: + return ExpiredRow(custom_id=line.custom_id) except Exception as e: # noqa: BLE001 # a provider error becomes the row's error line, never a crashed batch return RowOutcome( custom_id=line.custom_id, status_code=_status_code_of(e), body=_error_body(e), succeeded=False @@ -621,7 +665,7 @@ class LiteLLMExecutedBatchRunner: } async def _upload_results( - self, run: _BatchRun, kind: Literal["output", "error"], outcomes: Sequence[RowOutcome] + self, run: _BatchRun, kind: Literal["output", "error"], outcomes: Sequence[RowOutcome | ExpiredRow] ) -> str | None: if not outcomes: return None diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index 25365d34187..ccc3b5b7f5f 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -181,7 +181,8 @@ def get_batch_id_from_unified_batch_id(file_id: str) -> str: def is_litellm_executed_batch(decoded_unified_batch_id: str) -> bool: - return get_batch_id_from_unified_batch_id(decoded_unified_batch_id).startswith(LITELLM_EXECUTED_BATCH_ID_PREFIX) + _, marker, batch_id = decoded_unified_batch_id.partition("llm_batch_id:") + return bool(marker) and batch_id.startswith(LITELLM_EXECUTED_BATCH_ID_PREFIX) def encode_file_id_with_model(file_id: str, model: str, id_type: Literal["file", "batch"] = "file") -> str: 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 cb566d4dca7..860827e8fbd 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 @@ -53,6 +53,7 @@ ALL_STATUSES: Final[tuple[BatchStatus, ...]] = ( "failed", "cancelling", "cancelled", + "expired", ) @@ -375,6 +376,7 @@ def make_runner( store_factory: Callable[[Mapping[str, LiteLLM_ManagedFileTable]], FakeManagedBatchStore] = FakeManagedBatchStore, general_settings: Mapping[str, object] = MappingProxyType({}), heartbeat_seconds: float = 30.0, + completion_window_seconds: float = 24 * 60 * 60, ) -> Harness: store = store_factory({INPUT_FILE_ID: managed_input_file()} if files is None else files) router = FakeRouter() @@ -393,6 +395,7 @@ def make_runner( general_settings=general_settings, concurrency=concurrency, heartbeat_seconds=heartbeat_seconds, + completion_window_seconds=completion_window_seconds, storage_backend_factory=storage_factory, upload_result_file=uploads, ) @@ -464,6 +467,7 @@ def test_resolve_transition_keeps_the_requested_status_unless_cancelling(current ("requested", "expected"), [ ("completed", "cancelled"), + ("expired", "cancelled"), ("in_progress", "cancelling"), ("finalizing", "cancelling"), ("failed", "failed"), @@ -1014,6 +1018,46 @@ async def test_running_batch_skips_the_remaining_rows_after_an_operator_cancel( assert (finished.output_file_id, finished.error_file_id) == ("unified-output-1", None) +async def test_batch_expires_at_the_completion_window_and_keeps_what_finished() -> None: + rows = jsonl(chat_row("row-1", "hi 1"), chat_row("row-2", "hi 2"), chat_row("row-3", "hi 3")) + harness = make_runner(content=rows, concurrency=1, completion_window_seconds=0.2) + reply = chat_response("hi 1") + + async def dispatch(messages: Sequence[Mapping[str, str]], **_: object) -> ModelResponse: + if messages[0]["content"] == "hi 1": + return reply + await asyncio.Event().wait() + raise AssertionError("a row still running at the completion window must be cut off") + + harness.router.acompletion.side_effect = dispatch + created, finished = await harness.create_and_finish() + + assert created.expires_at == created.created_at + assert finished.status == "expired" + assert finished.expired_at is not None + assert finished.request_counts == BatchRequestCounts(completed=1, failed=2, total=3) + assert (finished.output_file_id, finished.error_file_id) == ("unified-output-1", "unified-output-2") + assert set(harness.uploads.calls[0].lines()) == {"row-1"} + error_lines = harness.uploads.calls[1].lines() + assert set(error_lines) == {"row-2", "row-3"} + for line in error_lines.values(): + assert line["response"] is None + error = line["error"] + assert isinstance(error, dict) + assert error["code"] == "batch_expired" + + +async def test_batch_created_past_its_window_dispatches_nothing() -> None: + harness = make_runner(completion_window_seconds=0) + _, finished = await harness.create_and_finish() + + assert harness.router.acompletion.await_count == 0 + assert finished.status == "expired" + assert finished.request_counts == BatchRequestCounts(completed=0, failed=2, total=2) + assert (finished.output_file_id, finished.error_file_id) == (None, "unified-output-1") + assert set(harness.uploads.calls[0].lines()) == {"row-1", "row-2"} + + async def test_upload_failure_marks_the_batch_failed() -> None: harness = make_runner(upload_error=RuntimeError("storage exploded")) _, finished = await harness.create_and_finish() diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py index f88d94d2c08..b7e3088ff01 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py @@ -487,6 +487,8 @@ class TestCompletedBatchSafeToRetire: ("litellm_proxy;model_id:my-vllm;llm_batch_id:litellm_batch_0123abcd", True), ("litellm_proxy;model_id:my-vllm;llm_batch_id:batch_0123abcd", False), ("litellm_proxy;model_id:my-vllm;generic_response_id:resp_0123abcd", False), + ("litellm_proxy;model_id:my-vllm;llm_output_file_id:file-0123abcd", False), + ("batch_0123abcd", False), ], ) def test_is_litellm_executed_batch_reads_the_llm_batch_id_prefix(decoded_unified_batch_id: str, executed: bool):