diff --git a/miles/rollout/fully_async_rollout.py b/miles/rollout/fully_async_rollout.py index 4fcd41a0d42..3a3bd4fadac 100644 --- a/miles/rollout/fully_async_rollout.py +++ b/miles/rollout/fully_async_rollout.py @@ -20,11 +20,11 @@ import logging import random import time -from collections.abc import Callable, Iterator +from collections import deque +from collections.abc import Callable, Coroutine, Iterator from concurrent.futures import Future -from copy import deepcopy from dataclasses import dataclass -from typing import cast +from typing import TypeVar, cast import httpx @@ -39,8 +39,18 @@ TrainBatchLease, TrainBatchRollbackReason, ) +from miles.rollout.data_source import SourceReservation from miles.rollout.filter_hub.base_types import MetricGatherer, call_dynamic_filter +from miles.rollout.fully_async.execution import ( + FullyAsyncExecution, + FullyAsyncExecutionFailure, + FullyAsyncExecutionRetry, + FullyAsyncExecutionSuccess, + FullyAsyncRetryReason, + FullyAsyncTerminalPendingError, +) from miles.rollout.fully_async.ownership import ReservationOwnership, ReservationStageId, ReservationTerminalReceipt +from miles.rollout.inference_rollout.fully_async import InferenceFullyAsyncExecutor from miles.rollout.inference_rollout.inference_rollout_common import ( GenerateState, SubmissionScheduler, @@ -53,6 +63,8 @@ logger = logging.getLogger(__name__) +_T = TypeVar("_T") + OUTPUT_QUEUE_MAX_GROUPS = 1000 NO_PROGRESS_WARN_SECS = 30.0 WEIGHT_VERSION_QUERY_TIMEOUT_SECS = 2.0 @@ -60,6 +72,7 @@ # A finished group is list[Sample], or list[list[Sample]] when a generate function # returns multiple samples per trajectory (e.g. multi-agent). Group = list[Sample | list[Sample]] +LegacyBufferedGroup = tuple[list[Sample], Group] @dataclass(frozen=True) @@ -72,12 +85,28 @@ class _OwnedCompletedGroup: @dataclass(frozen=True) class _OwnedExecutionFailure: terminal_receipt: ReservationTerminalReceipt - error: Exception + error: BaseException + + +@dataclass(frozen=True) +class _OwnedExecutionRetry: + terminal_receipt: ReservationTerminalReceipt + reason: FullyAsyncRetryReason + + +_OwnedTerminalResult = _OwnedCompletedGroup | _OwnedExecutionRetry | _OwnedExecutionFailure +_OwnedTerminalObserver = Callable[[], Coroutine[object, object, _OwnedTerminalResult]] +_WorkerResult = LegacyBufferedGroup | _OwnedTerminalResult + + +@dataclass(frozen=True) +class _ActiveOwnedExecution: + execution: FullyAsyncExecution + observe_terminal: _OwnedTerminalObserver BufferSource = list[Sample] | _OwnedCompletedGroup BufferEntry = tuple[BufferSource, Group] -LegacyBufferedGroup = tuple[list[Sample], Group] class _OwnedTrainBatchLease(TrainBatchLease): @@ -216,7 +245,7 @@ def __init__( max_groups: int | None, max_staleness: int | None, on_evict: Callable[[BufferSource], None], - ): + ) -> None: assert order in ("fifo", "lifo"), f"unknown buffer order: {order}" self._order = order self._capacity = max_groups if max_groups is not None else OUTPUT_QUEUE_MAX_GROUPS @@ -244,8 +273,17 @@ async def put(self, entry: BufferEntry, *, current_version: int | None = None) - await self._cond.wait() self._entries.append(entry) self.entered_groups += 1 - if self._evict_on_overflow and len(self._entries) > self._capacity: - self._evict_overflow(current_version) + try: + if self._evict_on_overflow and len(self._entries) > self._capacity: + self._evict_overflow(current_version, incoming_entry=entry) + except BaseException: + for index, queued_entry in enumerate(self._entries): + if queued_entry is entry: + self._entries.pop(index) + self.entered_groups -= 1 + break + self._cond.notify_all() + raise self._cond.notify_all() async def get(self) -> BufferEntry: @@ -256,15 +294,26 @@ async def get(self) -> BufferEntry: self._cond.notify_all() return entry - def _evict_overflow(self, current_version: int | None) -> None: + def _evict_overflow( + self, + current_version: int | None, + *, + incoming_entry: BufferEntry, + ) -> None: if self._max_staleness is not None and current_version is not None: - index = 0 - while index < len(self._entries): - source, group = self._entries[index] + eviction_order = [entry for entry in self._entries if entry is not incoming_entry] + eviction_order.append(incoming_entry) + for entry in eviction_order: + source, group = entry oldest = group_oldest_weight_version(group) too_stale = oldest is not None and current_version - oldest > self._max_staleness if not too_stale: - index += 1 + continue + index = next( + (index for index, queued_entry in enumerate(self._entries) if queued_entry is entry), + None, + ) + if index is None: continue self._on_evict(source) self._entries.pop(index) @@ -296,16 +345,16 @@ def reset_counters(self) -> None: self.evicted_stale_groups = 0 self.evicted_overflow_groups = 0 - async def discard_all(self, on_discard: Callable[[BufferSource], None]) -> Exception | None: + async def discard_all(self, on_discard: Callable[[BufferSource], None]) -> BaseException | None: """Discard buffered entries after their ownership settlement succeeds.""" - first_error: Exception | None = None + first_error: BaseException | None = None async with self._cond: index = 0 while index < len(self._entries): source, _ = self._entries[index] try: on_discard(source) - except Exception as error: + except BaseException as error: if first_error is None: first_error = error index += 1 @@ -363,6 +412,16 @@ def __init__(self, input: RolloutFnConstructorInput) -> None: self._sample_filter = load_function(input.args.rollout_sample_filter_path) self._weight_version = _CachedWeightVersion() self._worker: asyncio.Task | None = None + self._closing = False + self._closed = False + self._close_task: asyncio.Task[None] | None = None + self._executor_closed = False + self._worker_error: BaseException | None = None + self._worker_failure_reported = False + self._legacy_executions: dict[asyncio.Task[LegacyBufferedGroup], list[Sample]] = {} + self._legacy_requeued_groups: deque[LegacyBufferedGroup] = deque() + self._legacy_close_pending_groups: deque[list[Sample]] = deque() + self._active_drains: set[asyncio.Task[RolloutFnTrainOutput]] = set() self._eval_prompt_dataset_cache: dict = {} self._producer_resumed = asyncio.Event() self._producer_resumed.set() @@ -381,6 +440,21 @@ def __init__(self, input: RolloutFnConstructorInput) -> None: ) self._completed_groups = cast(int, completed_groups) if self._uses_owned_capacity else None self._ownership = ReservationOwnership(self.data_source) if self._uses_owned_capacity else None + self._executor = ( + InferenceFullyAsyncExecutor( + self.state, + sample_done_callback=self._scheduler.sample_done_callback, + ) + if self._uses_owned_capacity + else None + ) + self._active_executions: dict[ + asyncio.Task[_OwnedTerminalResult], + _ActiveOwnedExecution, + ] = {} + self._pending_reserved_rollbacks: list[SourceReservation] = [] + self._pending_terminal_rollbacks: list[tuple[ReservationTerminalReceipt, bool]] = [] + self._pending_aborted_groups_recycled = 0 self._next_execution_id = 1 self._completed_slots: asyncio.Queue[object] | None = None self._completed_slot_available = asyncio.Event() @@ -391,6 +465,8 @@ def __init__(self, input: RolloutFnConstructorInput) -> None: self._completed_slot_available.set() async def __call__(self, input: RolloutFnInput) -> RolloutFnOutput: + if self._closing: + raise RuntimeError("Fully async rollout function is closed.") if input.evaluation: return await self._call_eval(input) if self._worker is None: @@ -402,7 +478,188 @@ async def __call__(self, input: RolloutFnInput) -> RolloutFnOutput: ) self._worker = asyncio.create_task(self._worker_loop()) logger.info("Started fully-async rollout worker") - return await self._drain(input.rollout_id) + drain_task = asyncio.create_task(self._drain(input.rollout_id)) + self._active_drains.add(drain_task) + try: + return await drain_task + finally: + self._active_drains.discard(drain_task) + + async def close(self) -> None: + """Stop rollout production and settle every retained reservation. + + Returns: + None after active executions are terminal and retained reservations + are requeued. + + Raises: + BaseException: The first cleanup or terminal execution failure. A + later call retries retained cleanup that did not settle. + """ + if self._closed: + return + close_task = self._close_task + if close_task is None: + close_task = asyncio.create_task(self._close()) + self._close_task = close_task + try: + await asyncio.shield(close_task) + except asyncio.CancelledError as cancellation: + try: + await _await_task_terminal(close_task) + except BaseException as terminal_error: + raise cancellation from terminal_error + raise + finally: + if self._close_task is close_task and close_task.done() and not self._closed: + self._close_task = None + + async def _close(self) -> None: + self._closing = True + cleanup_error: BaseException | None = None + shutdown_error: BaseException | None = None + + worker = self._worker + if worker is not None: + if not worker.done(): + worker.cancel() + worker_result: None | BaseException = (await asyncio.gather(worker, return_exceptions=True))[0] + if not self._worker_failure_reported: + recorded_worker_error = self._worker_error + if recorded_worker_error is not None and not isinstance(recorded_worker_error, asyncio.CancelledError): + shutdown_error = recorded_worker_error + elif isinstance(worker_result, BaseException) and not isinstance( + worker_result, asyncio.CancelledError + ): + shutdown_error = worker_result + + legacy_executions = tuple(self._legacy_executions.items()) + for task, _ in legacy_executions: + if not task.done(): + task.cancel() + if legacy_executions: + results = await asyncio.gather( + *(task for task, _ in legacy_executions), + return_exceptions=True, + ) + for (task, source_group), _ in zip(legacy_executions, results, strict=True): + self._legacy_close_pending_groups.append(source_group) + self._legacy_executions.pop(task, None) + + active_drains = tuple(self._active_drains) + if active_drains: + await asyncio.gather(*active_drains, return_exceptions=True) + + output = self._output + if not self._uses_owned_capacity and output is not None: + + def retain_legacy_source(source: BufferSource) -> None: + if isinstance(source, _OwnedCompletedGroup): + raise RuntimeError("Legacy fully async output buffer contained an owned group.") + self._legacy_close_pending_groups.append(source) + + buffered_retain_error = await output.discard_all(retain_legacy_source) + if buffered_retain_error is not None: + cleanup_error = buffered_retain_error + while self._legacy_requeued_groups: + source_group, _ = self._legacy_requeued_groups.popleft() + self._legacy_close_pending_groups.append(source_group) + + legacy_recycle_error = self._retry_pending_legacy_recycles() + if legacy_recycle_error is not None: + cleanup_error = legacy_recycle_error + + acquisition_rollback_error = self._retry_pending_acquisition_rollback() + if cleanup_error is None: + cleanup_error = acquisition_rollback_error + + reserved_rollback_error = self._retry_pending_reserved_rollbacks() + if cleanup_error is None: + cleanup_error = reserved_rollback_error + + pending_rollback_error = self._retry_pending_terminal_rollbacks() + if cleanup_error is None: + cleanup_error = pending_rollback_error + + active = list(self._active_executions.items()) + terminal_waits: list[ + tuple[ + asyncio.Task[_OwnedTerminalResult], + _ActiveOwnedExecution, + ] + ] = [] + for terminal_task, active_execution in active: + if _terminal_observation_needs_retry(terminal_task): + replacement: asyncio.Task[_OwnedTerminalResult] = asyncio.create_task( + active_execution.observe_terminal() + ) + del self._active_executions[terminal_task] + self._active_executions[replacement] = active_execution + terminal_task = replacement + if not terminal_task.done(): + try: + active_execution.execution.request_cancellation() + except BaseException as error: + if cleanup_error is None: + cleanup_error = error + continue + terminal_waits.append((terminal_task, active_execution)) + + for terminal_task, _ in terminal_waits: + try: + result = await asyncio.shield(terminal_task) + except BaseException as error: + if cleanup_error is None: + cleanup_error = error + continue + if not isinstance(result, (_OwnedCompletedGroup, _OwnedExecutionRetry, _OwnedExecutionFailure)): + if cleanup_error is None: + cleanup_error = RuntimeError(f"Fully async close observed unsupported {type(result).__name__}.") + continue + try: + self._rollback_owned_terminal(result.terminal_receipt, completed_slot_held=False) + except BaseException as error: + if cleanup_error is None: + cleanup_error = error + continue + self._active_executions.pop(terminal_task, None) + if ( + isinstance(result, _OwnedExecutionFailure) + and not isinstance(result.error, asyncio.CancelledError) + and shutdown_error is None + ): + shutdown_error = result.error + + if self._uses_owned_capacity and output is not None: + queued_error = await self._rollback_queued_owned_groups(output) + if cleanup_error is None: + cleanup_error = queued_error + + if not self._active_executions and not self._executor_closed and self._executor is not None: + try: + await self._executor.close() + except BaseException as error: + if cleanup_error is None: + cleanup_error = error + else: + self._executor_closed = True + + if cleanup_error is not None: + raise cleanup_error + if self._active_executions: + raise RuntimeError("Fully async close did not settle every active execution.") + if self._pending_reserved_rollbacks: + raise RuntimeError("Fully async close did not settle every retained reserved rollback.") + if self._pending_terminal_rollbacks: + raise RuntimeError("Fully async close did not settle every retained terminal rollback.") + if self._ownership is not None and self._ownership.has_pending_acquisition_rollback: + raise RuntimeError("Fully async close did not settle the retained acquisition rollback.") + if self._legacy_close_pending_groups: + raise RuntimeError("Fully async close did not recycle every retained legacy group.") + if shutdown_error is not None: + self._worker_failure_reported = True + raise shutdown_error + self._closed = True async def _call_eval(self, input: RolloutFnEvalInput) -> RolloutFnOutput: if input.generate_state is not None: @@ -459,7 +716,7 @@ async def _generate_group(self, prompt_group: list[Sample]) -> Group: def _submit_one_group( self, - ) -> asyncio.Task[LegacyBufferedGroup | _OwnedCompletedGroup | _OwnedExecutionFailure]: + ) -> asyncio.Task[_WorkerResult]: if not self._uses_owned_capacity: prompt_groups = self.data_source.get_samples(1) self._scheduler.on_submit(prompt_groups) @@ -468,7 +725,9 @@ def _submit_one_group( async def execute_legacy() -> LegacyBufferedGroup: return prompt_group, await self._generate_group(prompt_group) - return asyncio.create_task(execute_legacy()) + task = asyncio.create_task(execute_legacy()) + self._legacy_executions[task] = prompt_group + return task ownership = self._ownership retained_slots = self._retained_slots @@ -499,30 +758,66 @@ async def execute_legacy() -> LegacyBufferedGroup: f"Source reservation {reservation.reservation_id} has duplicate parent identities: " f"{list(expected_parent_identities)}." ) - group = deepcopy(list(reservation.samples)) stage_id = ReservationStageId(f"execution-{self._next_execution_id}") self._next_execution_id += 1 [executor_receipt] = ownership.begin_execution([reservation], stage_id=stage_id) - self._scheduler.on_submit([group]) - except Exception: - ownership.rollback_reserved([reservation]) + except Exception as validation_error: + try: + ownership.rollback_reserved([reservation]) + except BaseException as rollback_error: + self._pending_reserved_rollbacks.append(reservation) + raise validation_error from rollback_error retained_slots.release() raise - async def execute() -> Group | _OwnedCompletedGroup | _OwnedExecutionFailure: + executor = self._executor + if executor is None: + raise RuntimeError("Fully async executor is not initialized.") + try: + execution = executor.submit(reservation, executor_receipt) + except BaseException as submission_error: try: - samples = await self._generate_group(group) - except Exception as error: [terminal_receipt] = ownership.record_terminal([executor_receipt], stage_id=stage_id) - return _OwnedExecutionFailure(terminal_receipt=terminal_receipt, error=error) + except BaseException as terminal_error: + raise submission_error from terminal_error + try: + ownership.rollback_batch([terminal_receipt]) + except BaseException as settlement_error: + self._pending_terminal_rollbacks.append((terminal_receipt, False)) + raise submission_error from settlement_error + retained_slots.release() + raise + self._scheduler.on_submit([list(reservation.samples)]) + + async def observe_terminal() -> _OwnedCompletedGroup | _OwnedExecutionRetry | _OwnedExecutionFailure: + outcome = await execution.wait_terminal() + if outcome.executor_receipt is not executor_receipt: + # A foreign receipt cannot prove this reservation terminal; retain ownership fail-closed. + raise RuntimeError( + f"Execution receipt {executor_receipt.receipt_id} did not return its exact terminal receipt." + ) [terminal_receipt] = ownership.record_terminal([executor_receipt], stage_id=stage_id) + if isinstance(outcome, FullyAsyncExecutionFailure): + return _OwnedExecutionFailure(terminal_receipt=terminal_receipt, error=outcome.error) + if isinstance(outcome, FullyAsyncExecutionRetry): + return _OwnedExecutionRetry( + terminal_receipt=terminal_receipt, + reason=outcome.reason, + ) + if not isinstance(outcome, FullyAsyncExecutionSuccess): + raise RuntimeError(f"Fully async execution returned unsupported {type(outcome).__name__}.") return _OwnedCompletedGroup( terminal_receipt=terminal_receipt, - samples=samples, + samples=outcome.samples, expected_parent_identities=expected_parent_identities, ) - return asyncio.create_task(execute()) + terminal_task = asyncio.create_task(observe_terminal()) + self._active_executions[terminal_task] = _ActiveOwnedExecution( + execution=execution, + observe_terminal=observe_terminal, + ) + return terminal_task async def _acquire_retained_slot(self) -> bool: if self._retained_slots is None: @@ -535,7 +830,7 @@ async def _acquire_retained_slot(self) -> bool: async def _submit_active_group( self, - active: set[asyncio.Task[LegacyBufferedGroup | _OwnedCompletedGroup | _OwnedExecutionFailure]], + active: set[asyncio.Task[_WorkerResult]], ) -> bool: if self._uses_owned_capacity: retained_slot_acquired = await self._acquire_retained_slot() @@ -625,7 +920,7 @@ def _rollback_owned_terminals( async def _rollback_queued_owned_groups( self, output: GroupBuffer, - ) -> Exception | None: + ) -> BaseException | None: def rollback(source: BufferSource) -> None: if not isinstance(source, _OwnedCompletedGroup): raise RuntimeError(f"Owned fully async output buffer contained unsupported {type(source).__name__}.") @@ -633,13 +928,72 @@ def rollback(source: BufferSource) -> None: return await output.discard_all(rollback) + def _retry_pending_terminal_rollbacks(self) -> BaseException | None: + while self._pending_terminal_rollbacks: + terminal_receipt, completed_slot_held = self._pending_terminal_rollbacks[0] + try: + self._rollback_owned_terminal( + terminal_receipt, + completed_slot_held=completed_slot_held, + ) + except BaseException as error: + return error + del self._pending_terminal_rollbacks[0] + return None + + def _retry_pending_reserved_rollbacks(self) -> BaseException | None: + ownership = self._ownership + retained_slots = self._retained_slots + if ownership is None or retained_slots is None: + if self._pending_reserved_rollbacks: + return RuntimeError("Fully async ownership capacity is not initialized.") + return None + while self._pending_reserved_rollbacks: + reservation = self._pending_reserved_rollbacks[0] + try: + ownership.rollback_reserved([reservation]) + except BaseException as error: + return error + retained_slots.release() + del self._pending_reserved_rollbacks[0] + return None + + def _retry_pending_acquisition_rollback(self) -> BaseException | None: + ownership = self._ownership + retained_slots = self._retained_slots + if ownership is None or not ownership.has_pending_acquisition_rollback: + return None + if retained_slots is None: + return RuntimeError("Fully async ownership capacity is not initialized.") + try: + ownership.retry_failed_acquisition_rollback() + except BaseException as error: + return error + retained_slots.release() + return None + + def _retry_pending_legacy_recycles(self) -> BaseException | None: + while self._legacy_close_pending_groups: + try: + self._recycle(self._legacy_close_pending_groups[0]) + except BaseException as error: + return error + self._legacy_close_pending_groups.popleft() + return None + + def _record_worker_error(self, error: BaseException) -> BaseException: + if self._worker_error is None: + self._worker_error = error + return self._worker_error + async def _worker_loop(self) -> None: output = self._output if output is None: raise RuntimeError("Fully async output buffer is not initialized.") - active: set[asyncio.Task[LegacyBufferedGroup | _OwnedCompletedGroup | _OwnedExecutionFailure]] = set() - fatal_error: Exception | None = None - fatal_settlement_error: Exception | None = None + active: set[asyncio.Task[_WorkerResult]] = set() + cancellation_requested: set[asyncio.Task[_OwnedTerminalResult]] = set() + fatal_error: BaseException | None = None + fatal_settlement_error: BaseException | None = None while True: if fatal_error is None and self._producer_resumed.is_set(): self._scheduler.arm(pending_groups=len(active)) @@ -655,8 +1009,9 @@ async def _worker_loop(self) -> None: submitted = await self._submit_active_group(active) except Exception as submission_error: if not self._uses_owned_capacity: + self._record_worker_error(submission_error) raise - fatal_error = submission_error + fatal_error = self._record_worker_error(submission_error) break if not submitted: break @@ -670,6 +1025,28 @@ async def _worker_loop(self) -> None: queued_settlement_error = await self._rollback_queued_owned_groups(output) if fatal_settlement_error is None: fatal_settlement_error = queued_settlement_error + cancellation_error: BaseException | None = None + if self._uses_owned_capacity: + for task in active: + owned_task = cast(asyncio.Task[_OwnedTerminalResult], task) + if owned_task in cancellation_requested: + continue + active_execution = self._active_executions.get(owned_task) + if active_execution is None: + if cancellation_error is None: + cancellation_error = RuntimeError( + "Fully async worker lost an active execution record." + ) + continue + try: + active_execution.execution.request_cancellation() + except BaseException as error: + if cancellation_error is None: + cancellation_error = error + else: + cancellation_requested.add(owned_task) + if cancellation_error is not None: + raise fatal_error from cancellation_error if not active: if fatal_error is not None: if fatal_settlement_error is not None: @@ -686,58 +1063,107 @@ async def _worker_loop(self) -> None: done, active = await self._scheduler.wait_for_progress(active) else: done, active = await asyncio.wait(active, return_when=asyncio.FIRST_COMPLETED) + if not done: + continue + if not self._uses_owned_capacity: + completed_groups: list[LegacyBufferedGroup] = [] + done_source_groups: list[list[Sample]] = [] + task_error: BaseException | None = None + for task in done: + legacy_task = cast(asyncio.Task[LegacyBufferedGroup], task) + source_group = self._legacy_executions[legacy_task] + done_source_groups.append(source_group) + try: + completed_group = legacy_task.result() + except BaseException as error: + if task_error is None: + task_error = error + else: + completed_groups.append(completed_group) + if task_error is not None: + active_tasks = [cast(asyncio.Task[LegacyBufferedGroup], task) for task in active] + active_source_groups = [self._legacy_executions[task] for task in active_tasks] + for task in active_tasks: + task.cancel() + await asyncio.gather(*active_tasks, return_exceptions=True) + self._legacy_close_pending_groups.extend(done_source_groups) + self._legacy_close_pending_groups.extend(active_source_groups) + for task in done: + self._legacy_executions.pop(cast(asyncio.Task[LegacyBufferedGroup], task), None) + for task in active_tasks: + self._legacy_executions.pop(task, None) + self._record_worker_error(task_error) + raise task_error + for task in done: + self._legacy_executions.pop(cast(asyncio.Task[LegacyBufferedGroup], task), None) + for position, completed_group in enumerate(completed_groups): + try: + version = await self._weight_version.get(self.args) if output.wants_weight_version else None + await output.put(completed_group, current_version=version) + except BaseException as error: + self._legacy_requeued_groups.extend(completed_groups[position:]) + self._record_worker_error(error) + raise + continue for task in done: - if not self._uses_owned_capacity: - result = task.result() - if not isinstance(result, tuple): - raise RuntimeError( - f"Legacy fully async execution returned unsupported {type(result).__name__}." - ) - version = await self._weight_version.get(self.args) if output.wants_weight_version else None - await output.put(result, current_version=version) - continue + owned_task = cast(asyncio.Task[_OwnedTerminalResult], task) try: - result = task.result() - except Exception as task_error: + result = owned_task.result() + except BaseException as task_error: if fatal_error is None: - fatal_error = task_error + fatal_error = self._record_worker_error(task_error) + continue + if isinstance(result, _OwnedExecutionRetry): + try: + self._rollback_owned_terminal(result.terminal_receipt, completed_slot_held=False) + except BaseException as settlement_error: + self._pending_terminal_rollbacks.append((result.terminal_receipt, False)) + if fatal_error is None: + fatal_error = self._record_worker_error(settlement_error) + else: + if result.reason is FullyAsyncRetryReason.EXECUTION_ABORTED: + self._pending_aborted_groups_recycled += 1 + self._active_executions.pop(owned_task, None) continue if isinstance(result, _OwnedExecutionFailure): try: self._rollback_owned_terminal(result.terminal_receipt, completed_slot_held=False) - except Exception as settlement_error: + except BaseException as settlement_error: + self._pending_terminal_rollbacks.append((result.terminal_receipt, False)) if fatal_settlement_error is None: fatal_settlement_error = settlement_error if fatal_error is None: - fatal_error = result.error + fatal_error = self._record_worker_error(result.error) + self._active_executions.pop(owned_task, None) continue if not isinstance(result, _OwnedCompletedGroup): raise RuntimeError(f"Owned fully async execution returned unsupported {type(result).__name__}.") if fatal_error is not None or not self._try_acquire_completed_slot(): try: self._rollback_owned_terminal(result.terminal_receipt, completed_slot_held=False) - except Exception as settlement_error: + except BaseException as settlement_error: + self._pending_terminal_rollbacks.append((result.terminal_receipt, False)) if fatal_error is None: - fatal_error = settlement_error + fatal_error = self._record_worker_error(settlement_error) elif fatal_settlement_error is None: fatal_settlement_error = settlement_error + self._active_executions.pop(owned_task, None) continue try: version = await self._weight_version.get(self.args) if output.wants_weight_version else None - except Exception as version_error: + await output.put((result, result.samples), current_version=version) + except BaseException as buffer_error: try: self._rollback_owned_terminal(result.terminal_receipt, completed_slot_held=True) - except Exception as settlement_error: + except BaseException as settlement_error: + self._pending_terminal_rollbacks.append((result.terminal_receipt, True)) if fatal_settlement_error is None: fatal_settlement_error = settlement_error + self._active_executions.pop(owned_task, None) if fatal_error is None: - fatal_error = version_error - continue - try: - await output.put((result, result.samples), current_version=version) - except Exception as buffer_error: - if fatal_error is None: - fatal_error = buffer_error + fatal_error = self._record_worker_error(buffer_error) + else: + self._active_executions.pop(owned_task, None) # -------------------------- consumer -------------------------- @@ -746,6 +1172,8 @@ async def _next_group(self) -> BufferEntry: worker = self._worker if output is None or worker is None: raise RuntimeError("Fully async worker is not initialized.") + if self._legacy_requeued_groups: + return self._legacy_requeued_groups.popleft() queue_get = asyncio.create_task(output.get()) try: while True: @@ -761,9 +1189,34 @@ async def _next_group(self) -> BufferEntry: if queue_get in done: return queue_get.result() logger.warning(f"No completed rollout groups for {NO_PROGRESS_WARN_SECS}s (queued: {output.qsize()})") - finally: - if not queue_get.done(): - queue_get.cancel() + except BaseException as error: + try: + await self._settle_failed_queue_get(queue_get, output) + except BaseException as settlement_error: + raise error from settlement_error + raise + + async def _settle_failed_queue_get( + self, + queue_get: asyncio.Task[BufferEntry], + output: GroupBuffer, + ) -> None: + if not queue_get.done(): + queue_get.cancel() + await asyncio.gather(queue_get, return_exceptions=True) + if queue_get.cancelled(): + return + + completed = queue_get.result() + source, _ = completed + if isinstance(source, _OwnedCompletedGroup): + try: + self._rollback_owned_terminal(source.terminal_receipt, completed_slot_held=True) + except BaseException: + self._pending_terminal_rollbacks.append((source.terminal_receipt, True)) + raise + return + self._legacy_requeued_groups.append(completed) async def _drain(self, rollout_id: int) -> RolloutFnTrainOutput: args = self.args @@ -774,6 +1227,8 @@ async def _drain(self, rollout_id: int) -> RolloutFnTrainOutput: target_data_size = args.rollout_batch_size data: list[Group] = [] + accepted_legacy_groups: list[LegacyBufferedGroup] = [] + claimed_legacy_group: LegacyBufferedGroup | None = None terminal_receipts: list[ReservationTerminalReceipt] = [] aborted_groups_recycled = 0 stale_groups_recycled = 0 @@ -783,7 +1238,8 @@ async def _drain(self, rollout_id: int) -> RolloutFnTrainOutput: try: while len(data) < target_data_size: - source, group = await self._next_group() + buffered_group = await self._next_group() + source, group = buffered_group if isinstance(source, _OwnedCompletedGroup): prompt_group = None terminal_receipt = source.terminal_receipt @@ -791,6 +1247,7 @@ async def _drain(self, rollout_id: int) -> RolloutFnTrainOutput: else: prompt_group = source terminal_receipt = None + claimed_legacy_group = buffered_group if len(group) != args.n_samples_per_prompt: if terminal_receipt is None: @@ -812,6 +1269,7 @@ async def _drain(self, rollout_id: int) -> RolloutFnTrainOutput: if terminal_receipt is None: assert prompt_group is not None self._recycle(prompt_group) + claimed_legacy_group = None else: self._rollback_owned_terminal(terminal_receipt, completed_slot_held=True) assert terminal_receipts.pop() is terminal_receipt @@ -827,6 +1285,7 @@ async def _drain(self, rollout_id: int) -> RolloutFnTrainOutput: if terminal_receipt is None: assert prompt_group is not None self._recycle(prompt_group) + claimed_legacy_group = None else: self._rollback_owned_terminal(terminal_receipt, completed_slot_held=True) assert terminal_receipts.pop() is terminal_receipt @@ -846,6 +1305,8 @@ async def _drain(self, rollout_id: int) -> RolloutFnTrainOutput: completed_slot_held=True, ) assert terminal_receipts.pop() is terminal_receipt + else: + claimed_legacy_group = None # Filtered groups are consumed, not replayed: they have no usable gradient signal. metric_gatherer.on_dynamic_filter_drop(reason=filter_output.reason) continue @@ -859,6 +1320,9 @@ async def _drain(self, rollout_id: int) -> RolloutFnTrainOutput: do_print = False data.append(group) + if terminal_receipt is None: + accepted_legacy_groups.append(buffered_group) + claimed_legacy_group = None sample = _first_sample(data[-1]) logger.info( @@ -868,11 +1332,16 @@ async def _drain(self, rollout_id: int) -> RolloutFnTrainOutput: data.sort(key=lambda group: _first_sample(group).index) + if self._uses_owned_capacity and self._closing: + raise RuntimeError("Fully async rollout closed before the train batch lease was issued.") + if self._sample_filter is not None: self._sample_filter(args, data) - metrics = { - "rollout/fully_async/queue_size": output.qsize(), + aborted_groups_recycled += self._pending_aborted_groups_recycled + self._pending_aborted_groups_recycled = 0 + metrics: dict[str, int | float] = { + "rollout/fully_async/queue_size": output.qsize() + len(self._legacy_requeued_groups), "rollout/fully_async/aborted_groups_recycled": aborted_groups_recycled, "rollout/fully_async/stale_groups_recycled": stale_groups_recycled, "rollout/fully_async/evicted_stale_groups": output.evicted_stale_groups, @@ -911,6 +1380,15 @@ async def _drain(self, rollout_id: int) -> RolloutFnTrainOutput: ) return RolloutFnTrainOutput(samples=data, metrics=metrics) except BaseException as error: + if not self._uses_owned_capacity: + worker = self._worker + retained_groups = list(accepted_legacy_groups) + if claimed_legacy_group is not None: + retained_groups.append(claimed_legacy_group) + if self._closing or (worker is not None and worker.done()): + self._legacy_close_pending_groups.extend(source_group for source_group, _ in retained_groups) + else: + self._legacy_requeued_groups.extend(retained_groups) if terminal_receipts: try: self._rollback_owned_terminals( @@ -918,6 +1396,9 @@ async def _drain(self, rollout_id: int) -> RolloutFnTrainOutput: completed_slots=len(terminal_receipts), ) except BaseException as settlement_error: + self._pending_terminal_rollbacks.extend( + (terminal_receipt, True) for terminal_receipt in terminal_receipts + ) raise error from settlement_error raise @@ -931,3 +1412,20 @@ def _recycle(self, prompt_group: list[Sample]) -> None: for sample in prompt_group: sample.reset_for_retry() self.data_source.add_samples([prompt_group]) + + +async def _await_task_terminal(task: asyncio.Task[_T]) -> _T: + while True: + try: + return await asyncio.shield(task) + except asyncio.CancelledError: + if task.done(): + return task.result() + + +def _terminal_observation_needs_retry(task: asyncio.Task[_OwnedTerminalResult]) -> bool: + if not task.done(): + return False + if task.cancelled(): + return True + return isinstance(task.exception(), FullyAsyncTerminalPendingError) diff --git a/miles/rollout/inference_rollout/fully_async.py b/miles/rollout/inference_rollout/fully_async.py index 1d8ab3d767c..2165dadbf96 100644 --- a/miles/rollout/inference_rollout/fully_async.py +++ b/miles/rollout/inference_rollout/fully_async.py @@ -1,5 +1,5 @@ import asyncio -from collections.abc import Iterator +from collections.abc import Callable, Iterator from copy import deepcopy from typing import TypeVar, cast @@ -157,8 +157,14 @@ async def wait_terminal(self) -> FullyAsyncExecutionOutcome: class InferenceFullyAsyncExecutor(FullyAsyncExecutor): """Execute receipt-bound inference groups on the caller's event loop.""" - def __init__(self, state: GenerateState) -> None: + def __init__( + self, + state: GenerateState, + *, + sample_done_callback: Callable[[], None] | None, + ) -> None: self._state = state + self._sample_done_callback = sample_done_callback self._cancellation = _InferenceCancellationCoordinator(state) self._tasks: set[asyncio.Task[list[Sample | list[Sample]]]] = set() self._closed = False @@ -187,6 +193,7 @@ def submit( _execute_group( self._state, deepcopy(list(reservation.samples)), + sample_done_callback=self._sample_done_callback, ) ) self._tasks.add(task) @@ -225,7 +232,19 @@ async def close(self) -> None: async def _execute_group( state: GenerateState, samples: list[Sample], + *, + sample_done_callback: Callable[[], None] | None, ) -> list[Sample | list[Sample]]: + if sample_done_callback is None: + return cast( + list[Sample | list[Sample]], + await generate_and_rm_group( + state, + samples, + sampling_params=state.sampling_params.copy(), + evaluation=False, + ), + ) return cast( list[Sample | list[Sample]], await generate_and_rm_group( @@ -233,6 +252,7 @@ async def _execute_group( samples, sampling_params=state.sampling_params.copy(), evaluation=False, + sample_done_callback=sample_done_callback, ), ) diff --git a/miles/rollout/inference_rollout/inference_rollout_common.py b/miles/rollout/inference_rollout/inference_rollout_common.py index 0613a8de09e..5e0ed27eabb 100644 --- a/miles/rollout/inference_rollout/inference_rollout_common.py +++ b/miles/rollout/inference_rollout/inference_rollout_common.py @@ -4,7 +4,7 @@ from argparse import Namespace from collections.abc import Callable from copy import deepcopy -from typing import Any +from typing import Any, cast from miles.rollout.base_types import ( GenerateFnInput, @@ -199,6 +199,9 @@ async def generate_and_rm_group( args = state.args if state.aborted: + if sample_done_callback is not None: + for _ in group: + sample_done_callback() return group if policy_uses_routing_key(args): @@ -208,7 +211,7 @@ async def generate_and_rm_group( log_prefix = f"[group indices={[getattr(s, 'index', '?') for s in group]}]" logger.debug(f"{log_prefix} Starting group with {len(group)} samples") - tasks = [] + tasks: list[asyncio.Task[Sample | list[Sample]]] = [] for idx, sample in enumerate(group): current_sampling_params = sampling_params.copy() if getattr(args, "sglang_enable_deterministic_inference", False): @@ -221,7 +224,23 @@ async def generate_and_rm_group( task.add_done_callback(lambda _task: sample_done_callback()) tasks.append(task) - group = await asyncio.gather(*tasks) + terminal_wait = asyncio.gather(*tasks, return_exceptions=True) + cancellation: asyncio.CancelledError | None = None + while not terminal_wait.done(): + try: + await asyncio.shield(terminal_wait) + except asyncio.CancelledError as error: + if cancellation is None: + cancellation = error + results = terminal_wait.result() + errors = [result for result in results if isinstance(result, BaseException)] + if cancellation is not None: + if errors: + raise cancellation from errors[0] + raise cancellation + if errors: + raise errors[0] + group = cast(list[Sample], [task.result() for task in tasks]) logger.debug(f"{log_prefix} [group] All {len(group)} samples completed") if state.aborted: return group diff --git a/tests/fast/rollout/inference_rollout/test_fully_async.py b/tests/fast/rollout/inference_rollout/test_fully_async.py index 6220ede5d16..9508afcef2b 100644 --- a/tests/fast/rollout/inference_rollout/test_fully_async.py +++ b/tests/fast/rollout/inference_rollout/test_fully_async.py @@ -47,7 +47,7 @@ async def generate(input: GenerateFnInput) -> GenerateFnOutput: async def test_executor_returns_receipt_bound_success_without_mutating_reservation() -> None: - executor = InferenceFullyAsyncExecutor(make_generate_state()) + executor = InferenceFullyAsyncExecutor(make_generate_state(), sample_done_callback=None) reservation = SourceReservation( reservation_id=SourceReservationId("source-0"), samples=(Sample(group_index=0, index=0, prompt="prompt"),), @@ -87,7 +87,7 @@ async def generate(input: GenerateFnInput) -> GenerateFnOutput: return GenerateFnOutput(samples=sample) state.generate_function = generate - executor = InferenceFullyAsyncExecutor(state) + executor = InferenceFullyAsyncExecutor(state, sample_done_callback=None) reservation = SourceReservation( reservation_id=SourceReservationId("source-1"), samples=(Sample(group_index=1, index=10, prompt="prompt"),), @@ -119,7 +119,7 @@ async def generate(input: GenerateFnInput) -> GenerateFnOutput: return GenerateFnOutput(samples=sample) state.generate_function = generate - executor = InferenceFullyAsyncExecutor(state) + executor = InferenceFullyAsyncExecutor(state, sample_done_callback=None) reservation = SourceReservation( reservation_id=SourceReservationId("source-invalid"), samples=(Sample(group_index=2, index=20, prompt="prompt"),), @@ -153,7 +153,7 @@ async def generate_and_rm_group(state, samples, sampling_params, evaluation=Fals return generated_samples monkeypatch.setattr(fully_async_module, "generate_and_rm_group", generate_and_rm_group) - executor = InferenceFullyAsyncExecutor(make_generate_state()) + executor = InferenceFullyAsyncExecutor(make_generate_state(), sample_done_callback=None) reservation = SourceReservation( reservation_id=SourceReservationId("source-malformed"), samples=(Sample(group_index=3, index=30, prompt="prompt"),), @@ -171,6 +171,103 @@ async def generate_and_rm_group(state, samples, sampling_params, evaluation=Fals await executor.close() +async def test_cancellation_requests_abort_and_waits_for_terminal_generation( + monkeypatch: pytest.MonkeyPatch, +) -> None: + generation_started = asyncio.Event() + release_generation = asyncio.Event() + abort_requested = asyncio.Event() + release_abort = asyncio.Event() + + async def generate(input: GenerateFnInput) -> GenerateFnOutput: + generation_started.set() + await release_generation.wait() + sample = cast(Sample, input.sample) + sample.status = Sample.Status.COMPLETED + sample.reward = 1.0 + return GenerateFnOutput(samples=sample) + + async def request_abort(args: Namespace) -> None: + abort_requested.set() + await release_abort.wait() + + state = make_generate_state() + state.generate_function = generate + monkeypatch.setattr(fully_async_module, "request_abort", request_abort, raising=False) + executor = InferenceFullyAsyncExecutor(state, sample_done_callback=None) + reservation = SourceReservation( + reservation_id=SourceReservationId("source-2"), + samples=(Sample(group_index=2, index=20, prompt="prompt"),), + ) + executor_receipt = cast(ReservationExecutorReceipt, object()) + execution = executor.submit(reservation, executor_receipt) + terminal_wait = asyncio.create_task(execution.wait_terminal()) + + await generation_started.wait() + execution.request_cancellation() + + try: + await asyncio.wait_for(abort_requested.wait(), timeout=0.01) + release_abort.set() + await asyncio.sleep(0) + assert not terminal_wait.done() + finally: + release_abort.set() + release_generation.set() + outcome = await terminal_wait + await executor.close() + + assert outcome == FullyAsyncExecutionRetry( + executor_receipt=executor_receipt, + reason=FullyAsyncRetryReason.CANCELLATION_REQUESTED, + ) + assert state.aborted + + +async def test_cancellation_preserves_terminal_generation_failure( + monkeypatch: pytest.MonkeyPatch, +) -> None: + generation_started = asyncio.Event() + release_generation = asyncio.Event() + abort_requested = asyncio.Event() + generation_error = RuntimeError("generation failed after cancellation") + + async def generate(input: GenerateFnInput) -> GenerateFnOutput: + generation_started.set() + await release_generation.wait() + raise generation_error + + async def request_abort(args: Namespace) -> None: + abort_requested.set() + + state = make_generate_state() + state.generate_function = generate + monkeypatch.setattr(fully_async_module, "request_abort", request_abort) + executor = InferenceFullyAsyncExecutor(state, sample_done_callback=None) + reservation = SourceReservation( + reservation_id=SourceReservationId("source-3"), + samples=(Sample(group_index=3, index=30, prompt="prompt"),), + ) + executor_receipt = cast(ReservationExecutorReceipt, object()) + execution = executor.submit(reservation, executor_receipt) + terminal_wait = asyncio.create_task(execution.wait_terminal()) + + await generation_started.wait() + execution.request_cancellation() + await abort_requested.wait() + assert not terminal_wait.done() + + release_generation.set() + outcome = await terminal_wait + + assert outcome == FullyAsyncExecutionFailure( + executor_receipt=executor_receipt, + error=generation_error, + ) + + await executor.close() + + async def test_executor_close_settles_siblings_before_raising_failure( monkeypatch: pytest.MonkeyPatch, ) -> None: @@ -196,7 +293,7 @@ async def generate_and_rm_group(state, samples, sampling_params, evaluation=Fals return samples monkeypatch.setattr(fully_async_module, "generate_and_rm_group", generate_and_rm_group) - executor = InferenceFullyAsyncExecutor(make_generate_state()) + executor = InferenceFullyAsyncExecutor(make_generate_state(), sample_done_callback=None) first_execution = executor.submit( SourceReservation( reservation_id=SourceReservationId("source-4"), diff --git a/tests/fast/rollout/inference_rollout/test_inference_rollout_common.py b/tests/fast/rollout/inference_rollout/test_inference_rollout_common.py new file mode 100644 index 00000000000..40649d0c58d --- /dev/null +++ b/tests/fast/rollout/inference_rollout/test_inference_rollout_common.py @@ -0,0 +1,174 @@ +from tests.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=5, suite="stage-a-cpu", labels=[]) + +import asyncio +from argparse import Namespace + +import pytest + +import miles.rollout.inference_rollout.inference_rollout_common as inference_rollout_common +from miles.utils.types import Sample + + +class FakeGenerateState(inference_rollout_common.GenerateState): + def __init__(self) -> None: + self.args = Namespace(group_rm=False) + self.aborted = False + + +async def test_aborted_group_releases_every_sample_completion_credit() -> None: + state = FakeGenerateState() + state.aborted = True + group = [Sample(index=0), Sample(index=1)] + completed_samples = 0 + + def on_sample_done() -> None: + nonlocal completed_samples + completed_samples += 1 + + result = await inference_rollout_common.generate_and_rm_group( + state, + group, + sampling_params={}, + sample_done_callback=on_sample_done, + ) + + assert result == group + assert completed_samples == len(group) + + +async def test_group_failure_waits_for_sibling_generation_to_finish(monkeypatch) -> None: + failure = RuntimeError("first parent failed") + sibling_started = asyncio.Event() + first_parent_failed = asyncio.Event() + release_sibling = asyncio.Event() + sibling_finished = asyncio.Event() + + async def generate_and_rm(state, sample, sampling_params, evaluation=False): + if sample.index == 0: + await sibling_started.wait() + first_parent_failed.set() + raise failure + sibling_started.set() + await release_sibling.wait() + sibling_finished.set() + return sample + + monkeypatch.setattr(inference_rollout_common, "generate_and_rm", generate_and_rm) + monkeypatch.setattr(inference_rollout_common, "policy_uses_routing_key", lambda args: False) + group_task = asyncio.create_task( + inference_rollout_common.generate_and_rm_group( + FakeGenerateState(), + [Sample(index=0), Sample(index=1)], + sampling_params={}, + ) + ) + + await sibling_started.wait() + await first_parent_failed.wait() + + try: + with pytest.raises(asyncio.TimeoutError): + await asyncio.wait_for(asyncio.shield(group_task), timeout=0.01) + finally: + release_sibling.set() + await sibling_finished.wait() + + with pytest.raises(RuntimeError) as error: + await group_task + + assert error.value is failure + assert sibling_finished.is_set() + + +async def test_group_cancellation_waits_for_child_generation_to_finish(monkeypatch) -> None: + children_started = 0 + children_finished = 0 + all_children_started = asyncio.Event() + release_children = asyncio.Event() + all_children_finished = asyncio.Event() + child_cancelled = asyncio.Event() + + async def generate_and_rm(state, sample, sampling_params, evaluation=False): + nonlocal children_finished, children_started + children_started += 1 + if children_started == 2: + all_children_started.set() + try: + await release_children.wait() + except asyncio.CancelledError: + child_cancelled.set() + await release_children.wait() + children_finished += 1 + if children_finished == 2: + all_children_finished.set() + return sample + + monkeypatch.setattr(inference_rollout_common, "generate_and_rm", generate_and_rm) + monkeypatch.setattr(inference_rollout_common, "policy_uses_routing_key", lambda args: False) + group_task = asyncio.create_task( + inference_rollout_common.generate_and_rm_group( + FakeGenerateState(), + [Sample(index=0), Sample(index=1)], + sampling_params={}, + ) + ) + + await all_children_started.wait() + group_task.cancel() + + try: + await asyncio.sleep(0) + assert not group_task.done() + assert not child_cancelled.is_set() + finally: + release_children.set() + await all_children_finished.wait() + if not group_task.done(): + group_task.cancel() + with pytest.raises(asyncio.CancelledError): + await group_task + + with pytest.raises(asyncio.CancelledError): + await group_task + assert children_finished == 2 + + +async def test_group_cancellation_chains_later_child_failure_after_drain(monkeypatch) -> None: + failure = RuntimeError("child failed after cancellation") + children_started = 0 + all_children_started = asyncio.Event() + release_children = asyncio.Event() + sibling_finished = asyncio.Event() + + async def generate_and_rm(state, sample, sampling_params, evaluation=False): + nonlocal children_started + children_started += 1 + if children_started == 2: + all_children_started.set() + await release_children.wait() + if sample.index == 0: + raise failure + sibling_finished.set() + return sample + + monkeypatch.setattr(inference_rollout_common, "generate_and_rm", generate_and_rm) + monkeypatch.setattr(inference_rollout_common, "policy_uses_routing_key", lambda args: False) + group_task = asyncio.create_task( + inference_rollout_common.generate_and_rm_group( + FakeGenerateState(), + [Sample(index=0), Sample(index=1)], + sampling_params={}, + ) + ) + + await all_children_started.wait() + group_task.cancel() + release_children.set() + + with pytest.raises(asyncio.CancelledError) as cancellation: + await group_task + + assert cancellation.value.__cause__ is failure + assert sibling_finished.is_set() diff --git a/tests/fast/rollout/inference_rollout/test_inference_rollout_train.py b/tests/fast/rollout/inference_rollout/test_inference_rollout_train.py new file mode 100644 index 00000000000..40c1d37f3a8 --- /dev/null +++ b/tests/fast/rollout/inference_rollout/test_inference_rollout_train.py @@ -0,0 +1,175 @@ +from tests.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=5, suite="stage-a-cpu", labels=[]) + +import asyncio +from argparse import Namespace + +import pytest + +import miles.rollout.inference_rollout.inference_rollout_train as inference_rollout_train + + +async def test_request_abort_settles_every_worker_before_propagating_failure(monkeypatch) -> None: + failure = RuntimeError("first worker abort failed") + first_finished = asyncio.Event() + second_started = asyncio.Event() + release_second = asyncio.Event() + second_finished = asyncio.Event() + agent_abort_called = asyncio.Event() + + async def get_worker_urls(args: Namespace) -> list[str]: + return ["http://worker-0", "http://worker-1"] + + async def post(url: str, payload: dict[str, bool]) -> None: + assert payload == {"abort_all": True} + if url == "http://worker-0/abort_request": + first_finished.set() + raise failure + assert url == "http://worker-1/abort_request" + second_started.set() + await release_second.wait() + second_finished.set() + + async def call_agent_abort_hook(args: Namespace) -> None: + agent_abort_called.set() + + monkeypatch.setattr(inference_rollout_train, "get_worker_urls", get_worker_urls) + monkeypatch.setattr(inference_rollout_train, "post", post) + monkeypatch.setattr(inference_rollout_train, "call_agent_abort_hook", call_agent_abort_hook) + abort_task = asyncio.create_task(inference_rollout_train.request_abort(Namespace())) + await first_finished.wait() + await second_started.wait() + + try: + await asyncio.wait_for(agent_abort_called.wait(), timeout=0.01) + with pytest.raises(asyncio.TimeoutError): + await asyncio.wait_for(asyncio.shield(abort_task), timeout=0.01) + assert not second_finished.is_set() + finally: + release_second.set() + await asyncio.gather(abort_task, return_exceptions=True) + + with pytest.raises(RuntimeError) as error: + await abort_task + + assert error.value is failure + assert second_finished.is_set() + assert agent_abort_called.is_set() + + +async def test_request_abort_settles_agent_hook_when_worker_requests_time_out(monkeypatch) -> None: + worker_started = asyncio.Event() + worker_cancelled = asyncio.Event() + agent_abort_finished = asyncio.Event() + + async def get_worker_urls(args: Namespace) -> list[str]: + return ["http://worker-0"] + + async def post(url: str, payload: dict[str, bool]) -> None: + worker_started.set() + try: + await asyncio.Future() + except asyncio.CancelledError: + worker_cancelled.set() + raise + + async def call_agent_abort_hook(args: Namespace) -> None: + await asyncio.sleep(0) + agent_abort_finished.set() + + monkeypatch.setattr(inference_rollout_train, "get_worker_urls", get_worker_urls) + monkeypatch.setattr(inference_rollout_train, "post", post) + monkeypatch.setattr(inference_rollout_train, "call_agent_abort_hook", call_agent_abort_hook) + abort_task = asyncio.create_task( + asyncio.wait_for( + inference_rollout_train.request_abort(Namespace()), + timeout=0.01, + ) + ) + await worker_started.wait() + + with pytest.raises(asyncio.TimeoutError): + await abort_task + + assert worker_cancelled.is_set() + assert agent_abort_finished.is_set() + + +async def test_request_abort_cancellation_terminates_the_agent_hook(monkeypatch) -> None: + worker_started = asyncio.Event() + agent_abort_started = asyncio.Event() + agent_abort_cancelled = asyncio.Event() + release_agent_abort = asyncio.Event() + + async def get_worker_urls(args: Namespace) -> list[str]: + return ["http://worker-0"] + + async def post(url: str, payload: dict[str, bool]) -> None: + worker_started.set() + await asyncio.Future() + + async def call_agent_abort_hook(args: Namespace) -> None: + agent_abort_started.set() + try: + await release_agent_abort.wait() + except asyncio.CancelledError: + agent_abort_cancelled.set() + raise + + monkeypatch.setattr(inference_rollout_train, "get_worker_urls", get_worker_urls) + monkeypatch.setattr(inference_rollout_train, "post", post) + monkeypatch.setattr(inference_rollout_train, "call_agent_abort_hook", call_agent_abort_hook) + abort_task = asyncio.create_task(inference_rollout_train.request_abort(Namespace())) + await worker_started.wait() + await agent_abort_started.wait() + + abort_task.cancel() + done, pending = await asyncio.wait({abort_task}, timeout=0.1) + try: + assert (done, pending) == ({abort_task}, set()) + assert agent_abort_cancelled.is_set() + finally: + release_agent_abort.set() + await asyncio.gather(abort_task, return_exceptions=True) + + with pytest.raises(asyncio.CancelledError): + await abort_task + + +async def test_request_abort_preserves_worker_discovery_failure_after_agent_hook_settles(monkeypatch) -> None: + discovery_error = RuntimeError("worker discovery failed") + agent_abort_finished = asyncio.Event() + + async def get_worker_urls(args: Namespace) -> list[str]: + raise discovery_error + + async def call_agent_abort_hook(args: Namespace) -> None: + agent_abort_finished.set() + + monkeypatch.setattr(inference_rollout_train, "get_worker_urls", get_worker_urls) + monkeypatch.setattr(inference_rollout_train, "call_agent_abort_hook", call_agent_abort_hook) + + with pytest.raises(RuntimeError) as error: + await inference_rollout_train.request_abort(Namespace()) + + assert error.value is discovery_error + assert agent_abort_finished.is_set() + + +async def test_request_abort_propagates_agent_hook_cancellation(monkeypatch) -> None: + cancellation = asyncio.CancelledError("agent abort cancelled") + + async def get_worker_urls(args: Namespace) -> list[str]: + return [] + + async def call_agent_abort_hook(args: Namespace) -> None: + raise cancellation + + monkeypatch.setattr(inference_rollout_train, "get_worker_urls", get_worker_urls) + monkeypatch.setattr(inference_rollout_train, "call_agent_abort_hook", call_agent_abort_hook) + + with pytest.raises(asyncio.CancelledError) as error: + await inference_rollout_train.request_abort(Namespace()) + + assert error.value is cancellation diff --git a/tests/fast/rollout/test_fully_async_rollout.py b/tests/fast/rollout/test_fully_async_rollout.py index 5504ec73f68..3b893ddd61c 100644 --- a/tests/fast/rollout/test_fully_async_rollout.py +++ b/tests/fast/rollout/test_fully_async_rollout.py @@ -3,6 +3,7 @@ register_cpu_ci(est_time=60, suite="stage-a-cpu", labels=[]) import asyncio +import gc from argparse import Namespace from collections import deque from collections.abc import Sequence @@ -13,6 +14,7 @@ import pytest import miles.rollout.fully_async_rollout as fully_async +import miles.rollout.inference_rollout.fully_async as inference_fully_async from miles.rollout.base_types import ( LeasedRolloutFnTrainOutput, RolloutFnConstructorInput, @@ -22,6 +24,13 @@ ) from miles.rollout.data_source import SourceReservation, SourceReservationId from miles.rollout.filter_hub.base_types import DynamicFilterOutput +from miles.rollout.fully_async.execution import ( + FullyAsyncExecutionRetry, + FullyAsyncExecutionSuccess, + FullyAsyncRetryReason, + FullyAsyncTerminalPendingError, +) +from miles.rollout.fully_async.ownership import ReservationExecutorReceipt, ReservationStageId from miles.utils.async_utils import AsyncLoopThread from miles.utils.types import Sample @@ -149,6 +158,10 @@ def make_args(**overrides) -> Namespace: async_buffer_order="fifo", dynamic_sampling_filter_path=None, rollout_sample_filter_path=None, + sglang_server_concurrency=8, + rollout_num_gpus=4, + rollout_num_gpus_per_engine=1, + rollout_health_check_timeout=0.1, sglang_router_ip="127.0.0.1", sglang_router_port=30000, eval_num_gpus=0, @@ -172,6 +185,7 @@ async def default_generate(state, group, sampling_params, evaluation=False, samp monkeypatch.setattr(fully_async, "GenerateState", FakeGenerateState) monkeypatch.setattr(fully_async, "generate_and_rm_group", generate or default_generate) + monkeypatch.setattr(inference_fully_async, "generate_and_rm_group", generate or default_generate) fn = fully_async.FullyAsyncRolloutFn(RolloutFnConstructorInput(args=args, data_source=data_source)) # Staleness accounting queries the router on every drain; fake it out. fn._weight_version = FakeWeightVersion() @@ -215,7 +229,7 @@ async def test_train_call_leases_one_to_many_output_by_parent_group(monkeypatch) generated_group = [[first_child, second_child], completed_second_parent] data_source = FakeReservationDataSource([reservation]) - async def generate(state, group, sampling_params, evaluation=False): + async def generate(state, group, sampling_params, evaluation=False, sample_done_callback=None): assert [(sample.group_index, sample.index) for sample in group] == [(1, 10), (1, 11)] return generated_group @@ -249,7 +263,7 @@ async def test_owned_admission_rejects_missing_parent_identity(monkeypatch): reservation.samples[0].index = None data_source = FakeReservationDataSource([reservation]) - async def generate(state, group, sampling_params, evaluation=False): + async def generate(state, group, sampling_params, evaluation=False, sample_done_callback=None): raise AssertionError("generation must not start for an invalid source reservation") fn = make_owned_fn(monkeypatch, data_source, generate) @@ -269,7 +283,7 @@ async def test_owned_admission_rejects_duplicate_parent_identities(monkeypatch): reservation.samples[1].index = reservation.samples[0].index data_source = FakeReservationDataSource([reservation]) - async def generate(state, group, sampling_params, evaluation=False): + async def generate(state, group, sampling_params, evaluation=False, sample_done_callback=None): raise AssertionError("generation must not start for an invalid source reservation") fn = make_owned_fn(monkeypatch, data_source, generate) @@ -291,7 +305,7 @@ async def test_terminal_failure_requeues_exact_source_reservation(monkeypatch): data_source = FakeReservationDataSource([reservation]) failure = RuntimeError("generation failed") - async def generate(state, group, sampling_params, evaluation=False): + async def generate(state, group, sampling_params, evaluation=False, sample_done_callback=None): raise failure fn = make_owned_fn(monkeypatch, data_source, generate) @@ -304,15 +318,67 @@ async def generate(state, group, sampling_params, evaluation=False): assert data_source.requeued == [[reservation]] +async def test_mismatched_execution_receipt_retains_ownership_fail_closed(monkeypatch): + reservation = make_reservation(50) + data_source = FakeReservationDataSource([reservation]) + + class MismatchedReceiptExecution: + def __init__(self, executor_receipt: ReservationExecutorReceipt) -> None: + self.executor_receipt = executor_receipt + self.cancellation_requests = 0 + + def request_cancellation(self) -> None: + self.cancellation_requests += 1 + + async def wait_terminal(self) -> FullyAsyncExecutionSuccess: + return FullyAsyncExecutionSuccess( + executor_receipt=replace(self.executor_receipt), + samples=[deepcopy(sample) for sample in reservation.samples], + ) + + fn = make_owned_fn(monkeypatch, data_source) + assert fn._executor is not None + executions: list[MismatchedReceiptExecution] = [] + + def submit( + source_reservation: SourceReservation, + executor_receipt: ReservationExecutorReceipt, + ) -> MismatchedReceiptExecution: + assert source_reservation is reservation + execution = MismatchedReceiptExecution(executor_receipt) + executions.append(execution) + return execution + + monkeypatch.setattr(fn._executor, "submit", submit) + + with pytest.raises(RuntimeError) as train_error: + await fn(RolloutFnTrainInput(rollout_id=50)) + + assert str(train_error.value) == "Execution receipt 0 did not return its exact terminal receipt." + assert data_source.reserved == [reservation] + assert data_source.acknowledged == [] + assert data_source.requeued == [] + assert len(fn._active_executions) == 1 + + with pytest.raises(RuntimeError) as close_error: + await fn.close() + + assert close_error.value is train_error.value + assert executions[0].cancellation_requests == 0 + assert data_source.acknowledged == [] + assert data_source.requeued == [] + + async def test_terminal_failure_drains_and_requeues_active_siblings(monkeypatch): reservations = [make_reservation(index) for index in range(26, 29)] data_source = FakeReservationDataSource(reservations) all_started = asyncio.Event() release_successes = asyncio.Event() + abort_requested = asyncio.Event() started: list[int] = [] failure = RuntimeError("generation failed") - async def generate(state, group, sampling_params, evaluation=False): + async def generate(state, group, sampling_params, evaluation=False, sample_done_callback=None): group_index = group[0].group_index started.append(group_index) if len(started) == 3: @@ -323,14 +389,22 @@ async def generate(state, group, sampling_params, evaluation=False): await release_successes.wait() return group + async def request_abort(args) -> None: + abort_requested.set() + + monkeypatch.setattr(inference_fully_async, "request_abort", request_abort) fn = make_owned_fn(monkeypatch, data_source, generate, execution_samples=6, retained_groups=3) drain = asyncio.create_task(fn(RolloutFnTrainInput(rollout_id=26))) await all_started.wait() - for _ in range(10): - await asyncio.sleep(0) - finished_before_siblings = drain.done() - release_successes.set() + try: + await asyncio.wait_for(abort_requested.wait(), timeout=0.01) + for _ in range(10): + await asyncio.sleep(0) + finished_before_siblings = drain.done() + finally: + release_successes.set() + with pytest.raises(RuntimeError) as error: await drain for _ in range(10): @@ -355,6 +429,61 @@ async def generate(state, group, sampling_params, evaluation=False): ) +async def test_close_resurfaces_recorded_worker_failure_after_cancelling_worker(monkeypatch): + reservations = [make_reservation(index) for index in range(80, 82)] + data_source = FakeReservationDataSource(reservations) + all_started = asyncio.Event() + release_sibling = asyncio.Event() + abort_requested = asyncio.Event() + started: list[int] = [] + generation_error = RuntimeError("generation failed before close") + + async def generate(state, group, sampling_params, evaluation=False, sample_done_callback=None): + group_index = group[0].group_index + started.append(group_index) + if len(started) == 2: + all_started.set() + await all_started.wait() + if group_index == 80: + raise generation_error + await release_sibling.wait() + return group + + async def request_abort(args) -> None: + abort_requested.set() + + monkeypatch.setattr(inference_fully_async, "request_abort", request_abort) + fn = make_owned_fn(monkeypatch, data_source, generate, execution_samples=4, retained_groups=2) + train = asyncio.create_task(fn(RolloutFnTrainInput(rollout_id=80))) + await all_started.wait() + await abort_requested.wait() + close = asyncio.create_task(fn.close()) + + try: + await asyncio.sleep(0) + assert not close.done() + finally: + release_sibling.set() + + with pytest.raises(RuntimeError) as close_error: + await close + + assert close_error.value is generation_error + with pytest.raises(asyncio.CancelledError): + await train + assert data_source.acknowledged == [] + assert ( + sorted( + (reservation for batch in data_source.requeued for reservation in batch), + key=lambda reservation: reservation.reservation_id, + ) + == reservations + ) + + await fn.close() + assert fn._closed + + async def test_submission_failure_drains_and_requeues_active_sibling(monkeypatch): reservation = make_reservation(31) data_source = FakeReservationDataSource([reservation]) @@ -373,7 +502,7 @@ def reserve_samples(num_groups: int) -> list[SourceReservation]: data_source.reserve_samples = reserve_samples - async def generate(state, group, sampling_params, evaluation=False): + async def generate(state, group, sampling_params, evaluation=False, sample_done_callback=None): generation_started.set() await release_generation.wait() return group @@ -415,7 +544,7 @@ def reserve_samples(num_groups: int) -> list[SourceReservation]: data_source.reserve_samples = reserve_samples - async def generate(state, group, sampling_params, evaluation=False): + async def generate(state, group, sampling_params, evaluation=False, sample_done_callback=None): if group[0].group_index == 36: await release_prefetch.wait() return group @@ -445,11 +574,221 @@ async def generate(state, group, sampling_params, evaluation=False): assert data_source.requeued == [[prefetched_reservation], [leased_reservation]] -async def test_local_execution_cancellation_does_not_requeue_without_terminal_receipt(monkeypatch): +async def test_close_retries_failed_requeue_after_executor_rejects_submission(monkeypatch): + reservation = make_reservation(45) + data_source = FakeReservationDataSource([reservation]) + submission_error = RuntimeError("submission rejected") + requeue_error = RuntimeError("submission requeue failed") + original_requeue = data_source.requeue_reservations + requeue_attempts = 0 + + def requeue_reservations(reservations: Sequence[SourceReservation]) -> None: + nonlocal requeue_attempts + requeue_attempts += 1 + if requeue_attempts == 1: + raise requeue_error + original_requeue(reservations) + + data_source.requeue_reservations = requeue_reservations + fn = make_owned_fn(monkeypatch, data_source) + assert fn._executor is not None + + def submit(reservation, executor_receipt): + raise submission_error + + monkeypatch.setattr(fn._executor, "submit", submit) + + with pytest.raises(RuntimeError) as train_error: + await fn(RolloutFnTrainInput(rollout_id=45)) + + assert train_error.value is submission_error + assert train_error.value.__cause__ is requeue_error + assert data_source.requeued == [] + + with pytest.raises(RuntimeError) as close_error: + await fn.close() + + assert close_error.value is submission_error + assert requeue_attempts == 2 + assert data_source.requeued == [[reservation]] + + await fn.close() + + +async def test_executor_rejection_does_not_charge_sample_backfill(monkeypatch): + reservation = make_reservation(55) + data_source = FakeReservationDataSource([reservation]) + submission_error = RuntimeError("submission rejected") + fn = make_owned_fn( + monkeypatch, + data_source, + rollout_submission_granularity="sample", + ) + assert fn._executor is not None + + def submit(reservation, executor_receipt): + raise submission_error + + monkeypatch.setattr(fn._executor, "submit", submit) + + with pytest.raises(RuntimeError) as train_error: + await fn(RolloutFnTrainInput(rollout_id=55)) + + assert train_error.value is submission_error + assert fn._scheduler.samples_in_flight == 0 + assert data_source.requeued == [[reservation]] + + with pytest.raises(RuntimeError) as close_error: + await fn.close() + + assert close_error.value is submission_error + + await fn.close() + + +async def test_failed_owned_buffer_handoff_restores_completed_capacity(monkeypatch): + reservation = make_reservation(59) + data_source = FakeReservationDataSource([reservation]) + buffer_error = RuntimeError("buffer handoff failed") + fn = make_owned_fn(monkeypatch, data_source) + + async def reject_put(self, buffered_group, *, current_version=None): + raise buffer_error + + monkeypatch.setattr(fully_async.GroupBuffer, "put", reject_put) + + with pytest.raises(RuntimeError) as train_error: + await fn(RolloutFnTrainInput(rollout_id=59)) + + assert train_error.value is buffer_error + assert data_source.requeued == [[reservation]] + assert fn._active_executions == {} + assert fn._pending_terminal_rollbacks == [] + assert fn._retained_slots is not None + assert not fn._retained_slots.locked() + assert fn._completed_slots is not None + assert fn._completed_slots.qsize() == 1 + + with pytest.raises(RuntimeError) as close_error: + await fn.close() + + assert close_error.value is buffer_error + + await fn.close() + + +async def test_close_settles_pending_acquisition_rollback_before_releasing_capacity(monkeypatch): + reservations = [make_reservation(56), make_reservation(57)] + data_source = FakeReservationDataSource([]) + acquisition_requeue_error = RuntimeError("acquisition requeue failed") + close_requeue_error = RuntimeError("close acquisition requeue failed") + requeue_attempts = 0 + + def reserve_samples(num_groups: int) -> list[SourceReservation]: + assert num_groups == 1 + return reservations + + def requeue_reservations(requeued: Sequence[SourceReservation]) -> None: + nonlocal requeue_attempts + requeue_attempts += 1 + if requeue_attempts == 1: + raise acquisition_requeue_error + if requeue_attempts == 2: + raise close_requeue_error + data_source.requeued.append(list(requeued)) + + data_source.reserve_samples = reserve_samples + data_source.requeue_reservations = requeue_reservations + fn = make_owned_fn(monkeypatch, data_source) + + with pytest.raises(RuntimeError) as train_error: + await fn(RolloutFnTrainInput(rollout_id=56)) + + assert train_error.value.__cause__ is acquisition_requeue_error + assert fn._retained_slots is not None + assert fn._retained_slots.locked() + + with pytest.raises(RuntimeError) as first_close_error: + await fn.close() + + assert first_close_error.value is close_requeue_error + assert requeue_attempts == 2 + assert fn._retained_slots.locked() + assert not fn._closed + + with pytest.raises(RuntimeError) as second_close_error: + await fn.close() + + assert second_close_error.value is train_error.value + assert requeue_attempts == 3 + assert data_source.requeued == [reservations] + assert not fn._retained_slots.locked() + assert not fn._closed + + await fn.close() + + assert fn._closed + + +async def test_close_retries_failed_reserved_validation_rollback_before_releasing_capacity(monkeypatch): + group = make_group(58) + reservation = SourceReservation( + reservation_id=SourceReservationId("source-58"), + samples=(group[0],), + ) + data_source = FakeReservationDataSource([reservation]) + validation_requeue_error = RuntimeError("validation requeue failed") + close_requeue_error = RuntimeError("close validation requeue failed") + requeue_attempts = 0 + + def requeue_reservations(requeued: Sequence[SourceReservation]) -> None: + nonlocal requeue_attempts + requeue_attempts += 1 + if requeue_attempts == 1: + raise validation_requeue_error + if requeue_attempts == 2: + raise close_requeue_error + data_source.requeued.append(list(requeued)) + + data_source.requeue_reservations = requeue_reservations + fn = make_owned_fn(monkeypatch, data_source) + + with pytest.raises(ValueError) as train_error: + await fn(RolloutFnTrainInput(rollout_id=58)) + + assert str(train_error.value) == "Source reservation source-58 contains 1 parent slots; expected 2." + assert train_error.value.__cause__ is validation_requeue_error + assert fn._retained_slots is not None + assert fn._retained_slots.locked() + + with pytest.raises(RuntimeError) as first_close_error: + await fn.close() + + assert first_close_error.value is close_requeue_error + assert requeue_attempts == 2 + assert data_source.requeued == [] + assert fn._retained_slots.locked() + assert not fn._closed + + with pytest.raises(ValueError) as second_close_error: + await fn.close() + + assert second_close_error.value is train_error.value + assert requeue_attempts == 3 + assert data_source.requeued == [[reservation]] + assert not fn._retained_slots.locked() + assert not fn._closed + + await fn.close() + + assert fn._closed + + +async def test_terminal_local_execution_cancellation_requeues_exact_source_reservation(monkeypatch): reservation = make_reservation(29) data_source = FakeReservationDataSource([reservation]) - async def generate(state, group, sampling_params, evaluation=False): + async def generate(state, group, sampling_params, evaluation=False, sample_done_callback=None): raise asyncio.CancelledError() fn = make_owned_fn(monkeypatch, data_source, generate) @@ -459,15 +798,767 @@ async def generate(state, group, sampling_params, evaluation=False): assert data_source.reserved == [reservation] assert data_source.acknowledged == [] + assert data_source.requeued == [[reservation]] + assert data_source.requeued[0][0] is reservation + + +async def test_cancelled_train_waiter_does_not_consume_next_completed_group(monkeypatch): + reservation = make_reservation(39) + data_source = FakeReservationDataSource([reservation]) + generation_started = asyncio.Event() + release_generation = asyncio.Event() + + async def generate(state, group, sampling_params, evaluation=False, sample_done_callback=None): + generation_started.set() + await release_generation.wait() + return group + + fn = make_owned_fn(monkeypatch, data_source, generate) + cancelled_train = asyncio.create_task(fn(RolloutFnTrainInput(rollout_id=39))) + await generation_started.wait() + + cancelled_train.cancel() + with pytest.raises(asyncio.CancelledError): + await cancelled_train + + release_generation.set() + output = await asyncio.wait_for(fn(RolloutFnTrainInput(rollout_id=40)), timeout=1) + + assert output.samples == [list(reservation.samples)] + assert data_source.acknowledged == [] + assert data_source.requeued == [] + + output.lease.rollback(TrainBatchRollbackReason.HANDOFF_FAILED) + + assert data_source.requeued == [[reservation]] + + +async def test_legacy_claimed_group_survives_output_queue_refill(monkeypatch): + first_group = make_group(52) + second_group = make_group(53) + first_buffered_group = (first_group, first_group) + second_buffered_group = (second_group, second_group) + fn = make_fn(monkeypatch, make_args(), FakeDataSource()) + output = fully_async.GroupBuffer( + order="fifo", + max_groups=None, + max_staleness=None, + on_evict=fn._recycle_buffer_source, + ) + await output.put(first_buffered_group) + fn._output = output + hold_worker = asyncio.Event() + worker = asyncio.create_task(hold_worker.wait()) + fn._worker = worker + queue_get = asyncio.create_task(output.get()) + await queue_get + await output.put(second_buffered_group) + + try: + await fn._settle_failed_queue_get(queue_get, output) + + assert await fn._next_group() is first_buffered_group + assert await fn._next_group() is second_buffered_group + finally: + worker.cancel() + await asyncio.gather(worker, return_exceptions=True) + + +async def test_worker_failure_wins_when_output_queue_completes_in_the_same_wait(monkeypatch): + reservation = make_reservation(59) + data_source = FakeReservationDataSource([reservation]) + fn = make_owned_fn(monkeypatch, data_source) + ownership = fn._ownership + retained_slots = fn._retained_slots + completed_slots = fn._completed_slots + assert ownership is not None + assert retained_slots is not None + assert completed_slots is not None + + await retained_slots.acquire() + [reserved] = ownership.reserve_samples(1) + stage_id = ReservationStageId("execution-tie") + [executor_receipt] = ownership.begin_execution([reserved], stage_id=stage_id) + [terminal_receipt] = ownership.record_terminal([executor_receipt], stage_id=stage_id) + completed_slots.get_nowait() + output = fully_async.GroupBuffer( + order="fifo", + max_groups=None, + max_staleness=None, + on_evict=fn._recycle_buffer_source, + ) + completed = fully_async._OwnedCompletedGroup( + terminal_receipt=terminal_receipt, + samples=list(reservation.samples), + expected_parent_identities=tuple((sample.group_index, sample.index) for sample in reservation.samples), + ) + await output.put((completed, completed.samples)) + fn._output = output + failure = RuntimeError("worker failed with output ready") + + async def fail_worker() -> None: + raise failure + + worker = asyncio.create_task(fail_worker()) + await asyncio.sleep(0) + fn._worker = worker + + async def complete_both(fs, **kwargs): + await asyncio.gather(*fs, return_exceptions=True) + return set(fs), set() + + monkeypatch.setattr(fully_async.asyncio, "wait", complete_both) + + with pytest.raises(RuntimeError) as next_group_error: + await fn._next_group() + + assert next_group_error.value is failure + assert data_source.acknowledged == [] + assert data_source.requeued == [[reservation]] + assert not retained_slots.locked() + assert completed_slots.qsize() == 1 + + fn._worker_failure_reported = True + await fn.close() + + +async def test_close_waits_for_terminal_late_success_before_requeue(monkeypatch): + reservation = make_reservation(40) + data_source = FakeReservationDataSource([reservation]) + generation_started = asyncio.Event() + release_generation = asyncio.Event() + abort_requested = asyncio.Event() + + async def generate(state, group, sampling_params, evaluation=False, sample_done_callback=None): + generation_started.set() + await release_generation.wait() + return group + + async def request_abort(args) -> None: + abort_requested.set() + + monkeypatch.setattr(inference_fully_async, "request_abort", request_abort) + fn = make_owned_fn(monkeypatch, data_source, generate) + train = asyncio.create_task(fn(RolloutFnTrainInput(rollout_id=40))) + await generation_started.wait() + close = None + + try: + close = asyncio.create_task(fn.close()) + await abort_requested.wait() + await asyncio.sleep(0) + + assert not close.done() + assert data_source.acknowledged == [] + assert data_source.requeued == [] + + close.cancel() + await asyncio.sleep(0) + assert not close.done() + assert data_source.requeued == [] + + release_generation.set() + with pytest.raises(asyncio.CancelledError): + await close + finally: + release_generation.set() + if close is not None: + await asyncio.gather(close, return_exceptions=True) + if fn._worker is not None and not fn._worker.done(): + fn._worker.cancel() + train.cancel() + await asyncio.gather(train, fn._worker, return_exceptions=True) + + with pytest.raises(asyncio.CancelledError): + await train + assert data_source.acknowledged == [] + assert data_source.requeued == [[reservation]] + assert data_source.requeued[0][0] is reservation + + await fn.close() + + +async def test_close_waits_for_legacy_in_flight_generation(monkeypatch): + generation_started = asyncio.Event() + cancellation_received = asyncio.Event() + release_generation = asyncio.Event() + generation_tasks: list[asyncio.Task[list[Sample]]] = [] + + async def generate(state, group, sampling_params, evaluation=False, sample_done_callback=None): + generation_task = asyncio.current_task() + assert generation_task is not None + generation_tasks.append(generation_task) + generation_started.set() + try: + await release_generation.wait() + except asyncio.CancelledError: + cancellation_received.set() + await release_generation.wait() + return group + + fn = make_fn( + monkeypatch, + make_args(rollout_batch_size=1), + FakeDataSource(), + generate=generate, + ) + train = asyncio.create_task(fn(RolloutFnTrainInput(rollout_id=48))) + await generation_started.wait() + close = asyncio.create_task(fn.close()) + + try: + await asyncio.wait_for(cancellation_received.wait(), timeout=0.01) + await asyncio.sleep(0) + assert not close.done() + finally: + release_generation.set() + await asyncio.gather(close, train, *generation_tasks, return_exceptions=True) + + await close + with pytest.raises(asyncio.CancelledError): + await train + assert fn._closed + + +async def test_close_recycles_unconsumed_legacy_active_and_prefetched_groups(monkeypatch): + active_group = make_group(54) + cancelled_group = make_group(56) + prefetched_group = make_group(55) + claimed_group = make_group(58) + data_source = FakeDataSource(scripted=[active_group, cancelled_group]) + cancellation_received = asyncio.Event() + all_started = asyncio.Event() + started: list[int] = [] + + async def finish_after_cancellation( + state, + group, + sampling_params, + evaluation=False, + sample_done_callback=None, + ): + started.append(group[0].group_index) + if len(started) == 2: + all_started.set() + try: + await asyncio.Future() + except asyncio.CancelledError: + if group is active_group: + cancellation_received.set() + return group + raise + + fn = make_fn( + monkeypatch, + make_args(rollout_batch_size=1), + data_source, + generate=finish_after_cancellation, + ) + output = fully_async.GroupBuffer( + order="fifo", + max_groups=None, + max_staleness=None, + on_evict=fn._recycle_buffer_source, + ) + await output.put((prefetched_group, prefetched_group)) + fn._output = output + fn._legacy_requeued_groups.append((claimed_group, claimed_group)) + fn._submit_one_group() + fn._submit_one_group() + await all_started.wait() + + await fn.close() + + assert cancellation_received.is_set() + assert data_source.recycled == [active_group, cancelled_group, prefetched_group, claimed_group] + assert fn._closed + + +async def test_close_recycles_legacy_group_blocked_on_output_put(monkeypatch): + prefetched_group = make_group(59) + blocked_group = make_group(60) + data_source = FakeDataSource(scripted=[blocked_group]) + fn = make_fn(monkeypatch, make_args(rollout_batch_size=1), data_source) + output = fully_async.GroupBuffer( + order="fifo", + max_groups=None, + max_staleness=None, + on_evict=fn._recycle_buffer_source, + ) + output._capacity = 1 + await output.put((prefetched_group, prefetched_group)) + fn._output = output + fn._worker = asyncio.create_task(fn._worker_loop()) + await wait_until(lambda: data_source.num_get_calls == 1 and not fn._legacy_executions) + + await fn.close() + + assert data_source.recycled == [prefetched_group, blocked_group] + assert fn._closed + + +async def test_close_recycles_legacy_partial_drain(monkeypatch): + first_group = make_group(61) + second_group = make_group(62) + data_source = FakeDataSource(scripted=[first_group, second_group]) + release_second = asyncio.Event() + first_claimed = asyncio.Event() + + async def generate(state, group, sampling_params, evaluation=False, sample_done_callback=None): + if group[0].group_index == 62: + await release_second.wait() + return group + + fn = make_fn( + monkeypatch, + make_args(rollout_batch_size=2, async_max_concurrent_samples=2), + data_source, + generate=generate, + ) + original_next_group = fn._next_group + + async def next_group(): + buffered_group = await original_next_group() + if buffered_group[1] is first_group: + first_claimed.set() + return buffered_group + + monkeypatch.setattr(fn, "_next_group", next_group) + train = asyncio.create_task(fn(RolloutFnTrainInput(rollout_id=61))) + await first_claimed.wait() + await asyncio.sleep(0) + + await fn.close() + + with pytest.raises(asyncio.CancelledError): + await train + assert sorted(group[0].group_index for group in data_source.recycled) == [61, 62] + assert fn._closed + + +async def test_close_retries_failed_legacy_recycle(monkeypatch): + prefetched_group = make_group(63) + data_source = FakeDataSource() + recycle_error = RuntimeError("legacy recycle failed") + recycle_attempts = 0 + + def add_samples(groups): + nonlocal recycle_attempts + recycle_attempts += 1 + if recycle_attempts == 1: + raise recycle_error + data_source.recycled.extend(groups) + + data_source.add_samples = add_samples + fn = make_fn(monkeypatch, make_args(rollout_batch_size=1), data_source) + output = fully_async.GroupBuffer( + order="fifo", + max_groups=None, + max_staleness=None, + on_evict=fn._recycle_buffer_source, + ) + await output.put((prefetched_group, prefetched_group)) + fn._output = output + + with pytest.raises(RuntimeError) as first_close_error: + await fn.close() + + assert first_close_error.value is recycle_error + assert data_source.recycled == [] + assert not fn._closed + + await fn.close() + + assert recycle_attempts == 2 + assert data_source.recycled == [prefetched_group] + assert fn._closed + + +async def test_close_does_not_abort_completed_terminal_observation(monkeypatch): + reservation = make_reservation(48) + data_source = FakeReservationDataSource([reservation]) + abort_requests = 0 + + async def request_abort(args) -> None: + nonlocal abort_requests + abort_requests += 1 + + monkeypatch.setattr(inference_fully_async, "request_abort", request_abort) + fn = make_owned_fn(monkeypatch, data_source) + retained_slots = fn._retained_slots + assert retained_slots is not None + await retained_slots.acquire() + + terminal_task = fn._submit_one_group() + await terminal_task + await fn.close() + + assert abort_requests == 0 + assert data_source.acknowledged == [] + assert data_source.requeued == [[reservation]] + assert data_source.requeued[0][0] is reservation + + +async def test_close_reobserves_cancelled_terminal_observer(monkeypatch): + reservation = make_reservation(49) + data_source = FakeReservationDataSource([reservation]) + first_observation_started = asyncio.Event() + + class CancelOnceExecution: + def __init__(self) -> None: + self.executor_receipt: ReservationExecutorReceipt | None = None + self.cancellation_requests = 0 + self.observation_attempts = 0 + + def request_cancellation(self) -> None: + self.cancellation_requests += 1 + + async def wait_terminal(self) -> FullyAsyncExecutionRetry: + self.observation_attempts += 1 + if self.observation_attempts == 1: + first_observation_started.set() + await asyncio.Future() + raise AssertionError("cancelled terminal observation resumed") + assert self.executor_receipt is not None + return FullyAsyncExecutionRetry( + executor_receipt=self.executor_receipt, + reason=FullyAsyncRetryReason.CANCELLATION_REQUESTED, + ) + + execution = CancelOnceExecution() + fn = make_owned_fn(monkeypatch, data_source) + assert fn._executor is not None + + def submit( + reservation: SourceReservation, + executor_receipt: ReservationExecutorReceipt, + ) -> CancelOnceExecution: + execution.executor_receipt = executor_receipt + return execution + + monkeypatch.setattr(fn._executor, "submit", submit) + retained_slots = fn._retained_slots + assert retained_slots is not None + await retained_slots.acquire() + terminal_task = fn._submit_one_group() + await first_observation_started.wait() + terminal_task.cancel() + with pytest.raises(asyncio.CancelledError): + await terminal_task + + await fn.close() + + assert execution.cancellation_requests == 1 + assert execution.observation_attempts == 2 + assert data_source.acknowledged == [] + assert data_source.requeued == [[reservation]] + assert data_source.requeued[0][0] is reservation + + +async def test_close_retries_terminal_observation_without_releasing_ownership(monkeypatch): + reservation = make_reservation(41) + data_source = FakeReservationDataSource([reservation]) + generation_started = asyncio.Event() + release_generation = asyncio.Event() + abort_requests = 0 + + async def generate(state, group, sampling_params, evaluation=False, sample_done_callback=None): + generation_started.set() + await release_generation.wait() + return group + + async def request_abort(args) -> None: + nonlocal abort_requests + abort_requests += 1 + + monkeypatch.setattr(inference_fully_async, "request_abort", request_abort) + fn = make_owned_fn( + monkeypatch, + data_source, + generate, + rollout_health_check_timeout=0.01, + ) + train = asyncio.create_task(fn(RolloutFnTrainInput(rollout_id=41))) + await generation_started.wait() + + try: + with pytest.raises(FullyAsyncTerminalPendingError): + await fn.close() + + assert abort_requests == 1 + assert data_source.acknowledged == [] + assert data_source.requeued == [] + + release_generation.set() + await fn.close() + finally: + release_generation.set() + if fn._worker is not None and not fn._worker.done(): + fn._worker.cancel() + train.cancel() + await asyncio.gather(train, fn._worker, return_exceptions=True) + + assert abort_requests == 2 + assert data_source.acknowledged == [] + assert data_source.requeued == [[reservation]] + assert data_source.requeued[0][0] is reservation + + +async def test_close_retries_failed_prefetched_terminal_requeue(monkeypatch): + reservation = make_reservation(42) + data_source = FakeReservationDataSource([reservation]) + release_generation = asyncio.Event() + + async def generate(state, group, sampling_params, evaluation=False, sample_done_callback=None): + await release_generation.wait() + return group + + fn = make_owned_fn(monkeypatch, data_source, generate) + train = asyncio.create_task(fn(RolloutFnTrainInput(rollout_id=42))) + await wait_until(lambda: len(data_source.reserved) == 1) + train.cancel() + with pytest.raises(asyncio.CancelledError): + await train + + release_generation.set() + await wait_until(lambda: fn._output is not None and fn._output.qsize() == 1) + requeue_error = RuntimeError("prefetched requeue failed") + original_requeue = data_source.requeue_reservations + requeue_attempts = 0 + + def requeue_reservations(reservations: Sequence[SourceReservation]) -> None: + nonlocal requeue_attempts + requeue_attempts += 1 + if requeue_attempts == 1: + raise requeue_error + original_requeue(reservations) + + data_source.requeue_reservations = requeue_reservations + + with pytest.raises(RuntimeError) as first_close_error: + await fn.close() + + assert first_close_error.value is requeue_error + assert data_source.acknowledged == [] + assert data_source.requeued == [] + + await fn.close() + + assert requeue_attempts == 2 + assert data_source.acknowledged == [] + assert data_source.requeued == [[reservation]] + assert data_source.requeued[0][0] is reservation + + +async def test_close_waits_for_claimed_drain_and_retries_rollback(monkeypatch): + reservation = make_reservation(51) + data_source = FakeReservationDataSource([reservation]) + requeue_error = RuntimeError("claimed drain requeue failed") + original_requeue = data_source.requeue_reservations + requeue_attempts = 0 + + def requeue_reservations(reservations: Sequence[SourceReservation]) -> None: + nonlocal requeue_attempts + requeue_attempts += 1 + if requeue_attempts == 1: + raise requeue_error + original_requeue(reservations) + + data_source.requeue_reservations = requeue_reservations + fn = make_owned_fn(monkeypatch, data_source) + claimed = asyncio.Event() + release_claimed = asyncio.Event() + original_next_group = fn._next_group + + async def next_group(): + completed = await original_next_group() + claimed.set() + await release_claimed.wait() + return completed + + monkeypatch.setattr(fn, "_next_group", next_group) + train = asyncio.create_task(fn(RolloutFnTrainInput(rollout_id=51))) + await claimed.wait() + close = asyncio.create_task(fn.close()) + + try: + await asyncio.sleep(0) + assert not close.done() + finally: + release_claimed.set() + await asyncio.gather(close, train, return_exceptions=True) + + with pytest.raises(RuntimeError) as train_error: + await train + await close + + assert str(train_error.value) == "Fully async rollout closed before the train batch lease was issued." + assert train_error.value.__cause__ is requeue_error + assert requeue_attempts == 2 + assert data_source.acknowledged == [] + assert data_source.requeued == [[reservation]] + assert data_source.requeued[0][0] is reservation + assert fn._pending_terminal_rollbacks == [] + assert fn._closed + + +async def test_cancelled_claimed_waiter_retains_failed_rollback_for_close(monkeypatch): + reservation = make_reservation(43) + data_source = FakeReservationDataSource([reservation]) + fn = make_owned_fn(monkeypatch, data_source) + original_wait = asyncio.wait + queue_get_completed = asyncio.Event() + hold_wait_result = asyncio.Event() + + async def gated_wait(fs, **kwargs): + done, pending = await original_wait(fs, **kwargs) + if fn._worker is not None and fn._worker in fs and any(task is not fn._worker for task in done): + queue_get_completed.set() + await hold_wait_result.wait() + return done, pending + + monkeypatch.setattr(fully_async.asyncio, "wait", gated_wait) + requeue_error = RuntimeError("claimed rollback failed") + original_requeue = data_source.requeue_reservations + requeue_attempts = 0 + + def requeue_reservations(reservations: Sequence[SourceReservation]) -> None: + nonlocal requeue_attempts + requeue_attempts += 1 + if requeue_attempts == 1: + raise requeue_error + original_requeue(reservations) + + data_source.requeue_reservations = requeue_reservations + train = asyncio.create_task(fn(RolloutFnTrainInput(rollout_id=43))) + await queue_get_completed.wait() + train.cancel() + + with pytest.raises(asyncio.CancelledError) as cancellation: + await train + + assert cancellation.value.__cause__ is requeue_error + assert requeue_attempts == 1 + assert data_source.requeued == [] + + monkeypatch.setattr(fully_async.asyncio, "wait", original_wait) + await fn.close() + + assert requeue_attempts == 2 + assert data_source.requeued == [[reservation]] + assert data_source.requeued[0][0] is reservation + + +async def test_cancelled_partial_drain_retains_failed_batch_rollback_for_close(monkeypatch): + reservations = [make_reservation(46), make_reservation(47)] + data_source = FakeReservationDataSource(reservations) + release_second = asyncio.Event() + + async def generate(state, group, sampling_params, evaluation=False, sample_done_callback=None): + if group[0].group_index == 47: + await release_second.wait() + return group + + async def request_abort(args) -> None: + pass + + monkeypatch.setattr(inference_fully_async, "request_abort", request_abort) + fn = make_owned_fn( + monkeypatch, + data_source, + generate, + batch_size=2, + execution_samples=4, + retained_groups=2, + completed_groups=2, + ) + first_group_claimed = asyncio.Event() + original_next_group = fn._next_group + + async def next_group(): + completed = await original_next_group() + if not first_group_claimed.is_set(): + first_group_claimed.set() + return completed + + monkeypatch.setattr(fn, "_next_group", next_group) + requeue_error = RuntimeError("partial drain rollback failed") + original_requeue = data_source.requeue_reservations + requeue_attempts = 0 + + def requeue_reservations(requeued: Sequence[SourceReservation]) -> None: + nonlocal requeue_attempts + requeue_attempts += 1 + if requeue_attempts == 1: + raise requeue_error + original_requeue(requeued) + + data_source.requeue_reservations = requeue_reservations + train = asyncio.create_task(fn(RolloutFnTrainInput(rollout_id=46))) + await first_group_claimed.wait() + train.cancel() + + with pytest.raises(asyncio.CancelledError) as cancellation: + await train + + assert cancellation.value.__cause__ is requeue_error assert data_source.requeued == [] + release_second.set() + await fn.close() + + assert data_source.requeued == [[reservations[0]], [reservations[1]]] + assert data_source.requeued[0][0] is reservations[0] + assert data_source.requeued[1][0] is reservations[1] + + +async def test_close_retries_worker_terminal_requeue_before_resurfacing_failure(monkeypatch): + reservation = make_reservation(44) + data_source = FakeReservationDataSource([reservation]) + generation_error = RuntimeError("generation failed") + requeue_error = RuntimeError("worker terminal requeue failed") + original_requeue = data_source.requeue_reservations + requeue_attempts = 0 + + def requeue_reservations(reservations: Sequence[SourceReservation]) -> None: + nonlocal requeue_attempts + requeue_attempts += 1 + if requeue_attempts == 1: + raise requeue_error + original_requeue(reservations) + + data_source.requeue_reservations = requeue_reservations + + async def generate(state, group, sampling_params, evaluation=False, sample_done_callback=None): + raise generation_error + + fn = make_owned_fn(monkeypatch, data_source, generate) + + with pytest.raises(RuntimeError) as train_error: + await fn(RolloutFnTrainInput(rollout_id=44)) + + assert train_error.value is generation_error + assert train_error.value.__cause__ is requeue_error + assert requeue_attempts == 1 + assert data_source.requeued == [] + + with pytest.raises(RuntimeError) as close_error: + await fn.close() + + assert close_error.value is generation_error + assert requeue_attempts == 2 + assert data_source.requeued == [[reservation]] + assert data_source.requeued[0][0] is reservation + + await fn.close() + async def test_aborted_owned_group_requeues_pristine_reservation(monkeypatch): aborted_reservation = make_reservation(3) completed_reservation = make_reservation(4) data_source = FakeReservationDataSource([aborted_reservation, completed_reservation]) - async def generate(state, group, sampling_params, evaluation=False): + async def generate(state, group, sampling_params, evaluation=False, sample_done_callback=None): if group[0].group_index == 3: for sample in group: sample.response = "aborted output" @@ -501,7 +1592,7 @@ async def test_owned_group_rejects_missing_parent_slot_without_losing_reservatio reservation = make_reservation(5) data_source = FakeReservationDataSource([reservation]) - async def generate(state, group, sampling_params, evaluation=False): + async def generate(state, group, sampling_params, evaluation=False, sample_done_callback=None): return group[:1] fn = make_owned_fn(monkeypatch, data_source, generate) @@ -518,7 +1609,7 @@ async def test_owned_group_rejects_reordered_parent_identity_without_losing_rese reservation = make_reservation(24) data_source = FakeReservationDataSource([reservation]) - async def generate(state, group, sampling_params, evaluation=False): + async def generate(state, group, sampling_params, evaluation=False, sample_done_callback=None): return [deepcopy(group[1]), deepcopy(group[0])] fn = make_owned_fn(monkeypatch, data_source, generate) @@ -549,7 +1640,7 @@ async def test_owned_group_rejects_foreign_one_to_many_child_identity(monkeypatc ] data_source = FakeReservationDataSource([reservation]) - async def generate(state, group, sampling_params, evaluation=False): + async def generate(state, group, sampling_params, evaluation=False, sample_done_callback=None): return generated_group fn = make_owned_fn(monkeypatch, data_source, generate) @@ -595,7 +1686,7 @@ async def test_owned_execution_capacity_bounds_started_source_samples(monkeypatc started: list[int] = [] data_source = FakeReservationDataSource([make_reservation(index) for index in range(6, 10)]) - async def generate(state, group, sampling_params, evaluation=False): + async def generate(state, group, sampling_params, evaluation=False, sample_done_callback=None): started.append(group[0].group_index) await release.wait() return group @@ -628,7 +1719,7 @@ async def test_owned_retained_limit_does_not_block_active_completion(monkeypatch generation_started = asyncio.Event() release_generation = asyncio.Event() - async def generate(state, group, sampling_params, evaluation=False): + async def generate(state, group, sampling_params, evaluation=False, sample_done_callback=None): generation_started.set() await release_generation.wait() return group @@ -703,28 +1794,7 @@ async def test_owned_lease_settlement_wakes_capacity_waiter_on_lifecycle_loop( assert data_source.acknowledged == [] assert data_source.requeued == [[reservations[0]]] - async def teardown_on_lifecycle_loop() -> None: - output_buffer = fn._output - worker = fn._worker - ownership = fn._ownership - assert output_buffer is not None - assert worker is not None - assert ownership is not None - - await wait_until(lambda: output_buffer.qsize() == 1) - worker.cancel() - [worker_result] = await asyncio.gather(worker, return_exceptions=True) - - assert isinstance(worker_result, asyncio.CancelledError) - assert await fn._rollback_queued_owned_groups(output_buffer) is None - assert output_buffer.qsize() == 0 - assert ownership._records == {} - - await asyncio.to_thread(lifecycle_loop.run, teardown_on_lifecycle_loop()) - - expected_requeued = [[reservations[1]]] if rollback_reason is None else [[reservations[0]], [reservations[1]]] - assert data_source.requeued == expected_requeued - assert data_source.requeued[-1][0] is reservations[1] + await asyncio.to_thread(lifecycle_loop.run, fn.close()) async def test_owned_batch_rolls_back_valid_and_identity_invalid_reservations_together(monkeypatch): @@ -739,7 +1809,7 @@ def identity_error(completed): first_identity_validated.set() return error - async def generate(state, group, sampling_params, evaluation=False): + async def generate(state, group, sampling_params, evaluation=False, sample_done_callback=None): if group[0].group_index == 33: return group await first_identity_validated.wait() @@ -774,7 +1844,7 @@ async def test_owned_completed_prefetch_overflow_requeues_and_blocks_admission(m reservations = [make_reservation(index) for index in range(12, 20)] data_source = FakeReservationDataSource(reservations) - async def generate(state, group, sampling_params, evaluation=False): + async def generate(state, group, sampling_params, evaluation=False, sample_done_callback=None): started.append(group[0].group_index) await release.wait() return group @@ -816,9 +1886,9 @@ async def generate(state, group, sampling_params, evaluation=False): ) output.lease.commit() - await wait_until(lambda: len(data_source.reserved) == 4) + await wait_until(lambda: len(data_source.reserved) >= 4) - assert data_source.reserved == reservations[:4] + assert data_source.reserved[:4] == reservations[:4] assert data_source.acknowledged == [([leased_reservation], 23)] @@ -1032,14 +2102,82 @@ async def get(self, args): assert output.metrics["rollout/fully_async/max_staleness"] == 5 -async def test_worker_error_propagates(monkeypatch): - async def failing_generate(state, group, sampling_params, evaluation=False, sample_done_callback=None): - raise RuntimeError("generation exploded") +def test_worker_error_propagates_without_leaking_sibling_failure(monkeypatch): + unhandled_messages: list[str] = [] - fn = make_fn(monkeypatch, make_args(), FakeDataSource(), generate=failing_generate) + async def run_failure() -> None: + async def failing_generate(state, group, sampling_params, evaluation=False, sample_done_callback=None): + raise RuntimeError("generation exploded") - with pytest.raises(RuntimeError, match="generation exploded"): - await fn(RolloutFnTrainInput(rollout_id=0)) + loop = asyncio.get_running_loop() + loop.set_exception_handler(lambda _, context: unhandled_messages.append(str(context["message"]))) + fn = make_fn(monkeypatch, make_args(), FakeDataSource(), generate=failing_generate) + + with pytest.raises(RuntimeError, match="generation exploded"): + await fn(RolloutFnTrainInput(rollout_id=0)) + + worker = fn._worker + fn._worker = None + del worker + del fn + gc.collect() + await asyncio.sleep(0) + + asyncio.run(run_failure()) + + assert unhandled_messages == [] + + +async def test_legacy_worker_failure_retains_done_and_cancelled_siblings_for_close(monkeypatch): + groups = [make_group(index) for index in range(77, 80)] + data_source = FakeDataSource(scripted=groups) + all_started = asyncio.Event() + release_terminals = asyncio.Event() + cancelled_sibling = asyncio.Event() + started: list[int] = [] + generation_error = RuntimeError("legacy generation failed") + + async def generate(state, group, sampling_params, evaluation=False, sample_done_callback=None): + group_index = group[0].group_index + started.append(group_index) + if len(started) == 3: + all_started.set() + await all_started.wait() + if group_index == 79: + try: + await asyncio.Future() + except asyncio.CancelledError: + cancelled_sibling.set() + return group + await release_terminals.wait() + if group_index == 78: + raise generation_error + return group + + fn = make_fn( + monkeypatch, + make_args(rollout_batch_size=3), + data_source, + generate=generate, + ) + train = asyncio.create_task(fn(RolloutFnTrainInput(rollout_id=77))) + await all_started.wait() + release_terminals.set() + + with pytest.raises(RuntimeError) as train_error: + await train + + assert train_error.value is generation_error + await cancelled_sibling.wait() + + with pytest.raises(RuntimeError) as close_error: + await fn.close() + + assert close_error.value is generation_error + await fn.close() + + assert sorted(group[0].group_index for group in data_source.recycled) == [77, 78, 79] + assert len(data_source.recycled) == 3 async def test_worker_bounds_in_flight_groups(monkeypatch): @@ -1215,6 +2353,40 @@ async def blocking_generate(state, group, sampling_params, evaluation=False, sam assert len(output.samples) == 1 +async def test_owned_backfill_submits_replacement_before_the_group_returns(monkeypatch): + callbacks = [] + release = asyncio.Event() + reservations = [make_reservation(90), make_reservation(91)] + data_source = FakeReservationDataSource(reservations) + + async def blocking_generate(state, group, sampling_params, evaluation=False, sample_done_callback=None): + callbacks.append(sample_done_callback) + await release.wait() + return group + + fn = make_owned_fn( + monkeypatch, + data_source, + blocking_generate, + retained_groups=2, + completed_groups=2, + rollout_submission_granularity="sample", + ) + drain = asyncio.create_task(fn(RolloutFnTrainInput(rollout_id=90))) + await wait_until(lambda: len(data_source.reserved) == 1 and len(callbacks) == 1) + + for _ in range(N_SAMPLES_PER_PROMPT): + callbacks[0]() + await wait_until(lambda: len(data_source.reserved) == 2) + + release.set() + output = await drain + output.lease.rollback(TrainBatchRollbackReason.HANDOFF_FAILED) + await fn.close() + + assert data_source.reserved[:2] == reservations + + async def test_group_granularity_opts_the_worker_out_of_backfill(monkeypatch): callbacks = [] release = asyncio.Event() @@ -1319,6 +2491,36 @@ async def test_buffer_threshold_evicts_all_over_staleness_first(): assert buffer.qsize() == 2 +async def test_buffer_failed_eviction_returns_incoming_group_to_caller(): + first_group = make_group(1, weight_versions=["5"]) + incoming_group = make_group(2, weight_versions=["6"]) + first_entry = (first_group, first_group) + incoming_entry = (incoming_group, incoming_group) + settled = [] + failure = RuntimeError("incoming settlement failed") + + def settle(source): + if source is incoming_group: + raise failure + settled.append(source) + + buffer = fully_async.GroupBuffer( + order="fifo", + max_groups=1, + max_staleness=2, + on_evict=settle, + ) + await buffer.put(first_entry, current_version=10) + + with pytest.raises(RuntimeError) as error: + await buffer.put(incoming_entry, current_version=10) + + assert error.value is failure + assert settled == [first_group] + assert buffer.qsize() == 0 + assert buffer.entered_groups == 1 + + async def test_buffer_lifo_serves_freshest_first(): buffer, _ = make_buffer(order="lifo") await put_group(buffer, make_group(1))