diff --git a/osprey_async_worker/src/osprey/async_worker/executor.py b/osprey_async_worker/src/osprey/async_worker/executor.py index ee55193..19d6c92 100644 --- a/osprey_async_worker/src/osprey/async_worker/executor.py +++ b/osprey_async_worker/src/osprey/async_worker/executor.py @@ -40,6 +40,7 @@ from osprey.engine.executor.execution_graph import ExecutionGraph from osprey.engine.executor.node_executor.call_executor import CallExecutor from osprey.engine.executor.udf_execution_helpers import UDFHelpers from osprey.engine.stdlib.udfs.json_utils import MissingJsonPath +from osprey.engine.udf.arguments import ArgumentsBase from osprey.engine.udf.base import BatchableUDFBase from osprey.worker.lib.instruments import metrics from osprey.worker.lib.osprey_shared.logging import get_logger @@ -51,6 +52,7 @@ logger = get_logger(__name__) _UNSET_RESULT: NodeResult = Err(None) _DEFAULT_MAX_ASYNC_PER_EXECUTION = 12 +_OWNED_TASK_CLEANUP_SECONDS = 10.0 def _get_ready_sync_and_async( @@ -230,7 +232,7 @@ async def _execute_async_udf( chain: DependencyChain, context: ExecutionContext, error_info_: list[NodeErrorInfo], - pre_resolved_arguments: Any | None = None, + pre_resolved_arguments: ArgumentsBase | None = None, ) -> NodeResult: """Execute a native async UDF. Awaited directly on the event loop. @@ -261,7 +263,7 @@ async def _execute_async_udf( caught_exception = e finally: _record_udf_metric(metric_tags, execution_result, caught_exception) - return execution_result + return execution_result async def _execute_async_batch( @@ -334,14 +336,14 @@ async def _enqueue_batches( error_infos: list[NodeErrorInfo], ready_async: Sequence[DependencyChain], ) -> tuple[ - Sequence[tuple[DependencyChain, Any | None]], + Sequence[tuple[DependencyChain, ArgumentsBase | None]], dict[asyncio.Task[Sequence[NodeResult]], Sequence[DependencyChain]], ]: """Launch batches and return the chains that remain. A native async chain can include arguments that this function resolved. Legacy chains include `None`. """ - batch_chains: dict[tuple[type, str], list[tuple[DependencyChain, Any, Any]]] = defaultdict(list) + batch_chains: dict[tuple[type, str], list[tuple[DependencyChain, ArgumentsBase, Any]]] = defaultdict(list) chains_to_remove: list[DependencyChain] = [] for async_chain in ready_async: @@ -365,7 +367,7 @@ async def _enqueue_batches( context.set_resolved_value(async_chain, Err(None)) new_batch_tasks: dict[asyncio.Task[Sequence[NodeResult]], Sequence[DependencyChain]] = {} - pre_resolved_by_chain: dict[DependencyChain, Any] = {} + pre_resolved_by_chain: dict[DependencyChain, ArgumentsBase] = {} for _, chains_and_args in batch_chains.items(): if len(chains_and_args) < 2: @@ -403,6 +405,10 @@ async def _enqueue_batches( # --- Main executor --- +async def _drain_owned_tasks(owned_tasks: Sequence[asyncio.Task[Any]]) -> None: + await asyncio.gather(*owned_tasks, return_exceptions=True) + + async def execute( execution_graph: ExecutionGraph, udf_helpers: UDFHelpers, @@ -419,6 +425,11 @@ async def execute( - Legacy UDFBase with execute_async=True: run in thread pool via run_in_executor (may fail on gevent calls, errors captured gracefully) """ + execution_task = asyncio.current_task() + if execution_task is None: + raise RuntimeError('async executor requires a running task') + entry_cancelling_count = execution_task.cancelling() + if parent_tracer_span: parent_tracer_span.set_tag('action-name', action.action_name) @@ -433,71 +444,129 @@ async def execute( ready_sync, ready_async = _get_ready_sync_and_async(allow_async, context) - while ready_sync or ready_async or in_progress_singlets or in_progress_batches: - # Check for already-finished tasks (non-blocking) - finished_singlets = [t for t in in_progress_singlets if t.done()] - finished_batches = [t for t in in_progress_batches if t.done()] - - if not ready_sync and not ready_async and not finished_singlets and not finished_batches: - # Block until at least one async task finishes - all_pending: set[asyncio.Task[Any]] = set(in_progress_singlets.keys()) | set(in_progress_batches.keys()) - if all_pending: - with tracer.start_span('osprey.rules.async_wait_nodes', child_of=parent_tracer_span): - done, _ = await asyncio.wait(all_pending, return_when=asyncio.FIRST_COMPLETED) - finished_singlets = [t for t in done if t in in_progress_singlets] - finished_batches = [t for t in done if t in in_progress_batches] - - # Process finished singlets - for task in finished_singlets: - chain = in_progress_singlets.pop(task) - context.set_resolved_value(chain, task.result()) - - # Process finished batches - # Distinct loop variable from the singlet `task` above: the two dicts hold - # tasks with different result types (Task[NodeResult] vs Task[Sequence[NodeResult]]), - # and reusing one name would unify them to the singlet type. - for batch_task in finished_batches: - chains = in_progress_batches.pop(batch_task) - results = batch_task.result() - for i, chain in enumerate(chains): - context.set_resolved_value(chain, results[i]) - - # Enqueue async tasks - if allow_async and ready_async: - with tracer.start_span('osprey.rules.try_enqueue_batches', child_of=parent_tracer_span): - remaining_ready_async, new_batch_tasks = await _enqueue_batches( - loop, semaphore, context, error_infos, ready_async - ) - in_progress_batches.update(new_batch_tasks) - - for async_chain, pre_resolved_arguments in remaining_ready_async: - # Native async UDF → await on event loop - if isinstance(async_chain.executor, CallExecutor) and isinstance( - async_chain.executor._udf, (AsyncUDFBase, AsyncBatchableUDFBase) - ): - task = asyncio.create_task( - _execute_async_udf(semaphore, async_chain, context, error_infos, pre_resolved_arguments) - ) - else: - # Legacy sync UDF with execute_async=True → thread pool fallback - task = asyncio.create_task( - _execute_legacy_in_executor(loop, semaphore, async_chain, context, error_infos) + try: + while ready_sync or ready_async or in_progress_singlets or in_progress_batches: + # Check for already-finished tasks (non-blocking) + finished_singlets = [t for t in in_progress_singlets if t.done()] + finished_batches = [t for t in in_progress_batches if t.done()] + + if not ready_sync and not ready_async and not finished_singlets and not finished_batches: + # Block until at least one async task finishes + all_pending: set[asyncio.Task[Any]] = set(in_progress_singlets.keys()) | set(in_progress_batches.keys()) + if all_pending: + with tracer.start_span('osprey.rules.async_wait_nodes', child_of=parent_tracer_span): + done, _ = await asyncio.wait(all_pending, return_when=asyncio.FIRST_COMPLETED) + finished_singlets = [t for t in done if t in in_progress_singlets] + finished_batches = [t for t in done if t in in_progress_batches] + + # Process finished singlets + for task in finished_singlets: + chain = in_progress_singlets.pop(task) + context.set_resolved_value(chain, task.result()) + + # Process finished batches + # Distinct loop variable from the singlet `task` above: the two dicts hold + # tasks with different result types (Task[NodeResult] vs Task[Sequence[NodeResult]]), + # and reusing one name would unify them to the singlet type. + for batch_task in finished_batches: + chains = in_progress_batches.pop(batch_task) + results = batch_task.result() + for i, chain in enumerate(chains): + context.set_resolved_value(chain, results[i]) + + # Enqueue async tasks + if allow_async and ready_async: + with tracer.start_span('osprey.rules.try_enqueue_batches', child_of=parent_tracer_span): + remaining_ready_async, new_batch_tasks = await _enqueue_batches( + loop, semaphore, context, error_infos, ready_async ) - in_progress_singlets[task] = async_chain - - # Execute sync chains inline (pure computation, fast, no I/O). - # Only yield deep into a long sync round when async tasks are in flight. - # Each sleep(0) triggers a full event loop cycle including gRPC C-core polling, - # so unnecessary yields cause significant context-switch overhead. - # Short rounds (<100 chains) skip yielding entirely — the asyncio.wait() at the - # top of the loop provides natural yield points between rounds. - for i, sync_chain in enumerate(ready_sync): - if (in_progress_singlets or in_progress_batches) and i > 0 and i % 100 == 0: - await asyncio.sleep(0) - result = _execute_sync(sync_chain, context, error_infos) - context.set_resolved_value(sync_chain, result) - - ready_sync, ready_async = _get_ready_sync_and_async(allow_async, context) + in_progress_batches.update(new_batch_tasks) + + for async_chain, pre_resolved_arguments in remaining_ready_async: + # Native async UDF → await on event loop + if isinstance(async_chain.executor, CallExecutor) and isinstance( + async_chain.executor._udf, (AsyncUDFBase, AsyncBatchableUDFBase) + ): + task = asyncio.create_task( + _execute_async_udf(semaphore, async_chain, context, error_infos, pre_resolved_arguments) + ) + else: + # Legacy sync UDF with execute_async=True → thread pool fallback + task = asyncio.create_task( + _execute_legacy_in_executor(loop, semaphore, async_chain, context, error_infos) + ) + in_progress_singlets[task] = async_chain + + # Execute sync chains inline (pure computation, fast, no I/O). + # Only yield deep into a long sync round when async tasks are in flight. + # Each sleep(0) triggers a full event loop cycle including gRPC C-core polling, + # so unnecessary yields cause significant context-switch overhead. + # Short rounds (<100 chains) skip yielding entirely — the asyncio.wait() at the + # top of the loop provides natural yield points between rounds. + for i, sync_chain in enumerate(ready_sync): + if (in_progress_singlets or in_progress_batches) and i > 0 and i % 100 == 0: + await asyncio.sleep(0) + result = _execute_sync(sync_chain, context, error_infos) + context.set_resolved_value(sync_chain, result) + + ready_sync, ready_async = _get_ready_sync_and_async(allow_async, context) + except GeneratorExit: + for owned_task in [*in_progress_singlets, *in_progress_batches]: + if owned_task.done(): + if not owned_task.cancelled(): + owned_task.exception() + elif not loop.is_closed(): + owned_task.cancel() + raise + except BaseException as execution_error: + owned_tasks = [*in_progress_singlets, *in_progress_batches] + for owned_task in owned_tasks: + owned_task.cancel() + cleanup_waiter = execution_task + initial_cancelling_count = cleanup_waiter.cancelling() + self_cancellation = ( + isinstance(execution_error, asyncio.CancelledError) and initial_cancelling_count > entry_cancelling_count + ) + cancellation_during_cleanup: asyncio.CancelledError | None = None + if owned_tasks: + cleanup = asyncio.create_task(_drain_owned_tasks(owned_tasks)) + cleanup_deadline = loop.time() + _OWNED_TASK_CLEANUP_SECONDS + while not cleanup.done(): + remaining = cleanup_deadline - loop.time() + if remaining <= 0: + break + try: + done, _ = await asyncio.wait({cleanup}, timeout=remaining) + if not done: + break + except asyncio.CancelledError as cleanup_cancellation: + if self_cancellation: + while cleanup_waiter.cancelling() > initial_cancelling_count: + cleanup_waiter.uncancel() + else: + cancellation_during_cleanup = cleanup_cancellation + continue + if not cleanup.done(): + unfinished_count = sum(not owned_task.done() for owned_task in owned_tasks) + logger.warning( + 'Owned task cleanup exceeded %ss; cancelling %d unfinished tasks', + _OWNED_TASK_CLEANUP_SECONDS, + unfinished_count, + ) + cleanup.cancel() + if self_cancellation: + while cleanup_waiter.cancelling() > initial_cancelling_count: + cleanup_waiter.uncancel() + else: + if cancellation_during_cleanup is None and cleanup_waiter.cancelling() > initial_cancelling_count: + cancellation_during_cleanup = asyncio.CancelledError() + if cancellation_during_cleanup is not None: + logger.warning( + 'Cancellation requested while cleaning up execution failure', + exc_info=(type(execution_error), execution_error, execution_error.__traceback__), + ) + raise cancellation_during_cleanup from execution_error + raise # --- Build result --- diff --git a/osprey_async_worker/src/osprey/async_worker/tests/test_async_executor.py b/osprey_async_worker/src/osprey/async_worker/tests/test_async_executor.py index 7cf1327..c28f3af 100644 --- a/osprey_async_worker/src/osprey/async_worker/tests/test_async_executor.py +++ b/osprey_async_worker/src/osprey/async_worker/tests/test_async_executor.py @@ -9,9 +9,10 @@ from collections import defaultdict from dataclasses import dataclass from datetime import datetime, timezone from textwrap import dedent -from typing import ClassVar, Sequence +from typing import Any, ClassVar, Sequence import pytest +from osprey.async_worker import executor as async_executor from osprey.async_worker.adaptor.interfaces import AsyncBatchableUDFBase, AsyncUDFBase from osprey.async_worker.executor import execute from osprey.engine.ast.grammar import Source @@ -24,8 +25,9 @@ from osprey.engine.executor.execution_plan import ExecutionPlanState from osprey.engine.executor.udf_execution_helpers import UDFHelpers from osprey.engine.stdlib import get_config_registry from osprey.engine.udf.arguments import ArgumentsBase +from osprey.engine.udf.base import UDFBase from osprey.engine.udf.registry import UDFRegistry -from result import Ok +from result import Ok, Result class CountingBatchableArguments(ArgumentsBase): @@ -83,6 +85,15 @@ class GatedArguments(ArgumentsBase): value: str +class CancellationArguments(ArgumentsBase): + value: str + + +class BatchCancellationArguments(ArgumentsBase): + id: str + routing_key: str + + class GatedAsyncUdf(AsyncUDFBase[GatedArguments, str]): entered = 0 both_entered: asyncio.Event @@ -473,3 +484,351 @@ async def test_batch_of_one_resolve_failure_surfaces_once( assert isinstance(result.error_infos[0].error, RuntimeError) assert str(result.error_infos[0].error) == 'resolve boom' assert counting_batchable_udf.resolve_call_count == 1 + + +@pytest.mark.asyncio +async def test_cancelled_udf_task_cancels_execution(async_execute_with_result): + started = asyncio.Event() + cancelled = asyncio.Event() + + class SlowUDF(AsyncUDFBase[CancellationArguments, str]): + @classmethod + def _get_udf_base_args(cls): + return (CancellationArguments, str) + + async def async_execute(self, execution_context: ExecutionContext, arguments: CancellationArguments) -> str: + started.set() + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + cancelled.set() + raise + raise AssertionError('unreachable') + + execution = asyncio.create_task( + async_execute_with_result( + 'Result = SlowUDF(value="test")', + udf_registry=UDFRegistry.with_udfs(SlowUDF), + ) + ) + await asyncio.wait_for(started.wait(), timeout=1) + + udf_tasks = [ + task + for task in asyncio.all_tasks() + if task is not asyncio.current_task() and getattr(task.get_coro(), '__name__', '') == '_execute_async_udf' + ] + assert len(udf_tasks) == 1 + udf_tasks[0].cancel() + + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(execution, timeout=1) + assert cancelled.is_set() + + +@pytest.mark.asyncio +async def test_cancelled_execution_cancels_owned_udf_tasks(async_execute_with_result): + singlet_started = asyncio.Event() + singlet_cancelled = asyncio.Event() + singlet_release = asyncio.Event() + batch_started = asyncio.Event() + batch_cancelled = asyncio.Event() + batch_release = asyncio.Event() + + class SlowSingletUDF(AsyncUDFBase[CancellationArguments, str]): + @classmethod + def _get_udf_base_args(cls): + return (CancellationArguments, str) + + async def async_execute(self, execution_context: ExecutionContext, arguments: CancellationArguments) -> str: + singlet_started.set() + try: + await singlet_release.wait() + except asyncio.CancelledError: + singlet_cancelled.set() + raise + return arguments.value + + class SlowBatchUDF(AsyncBatchableUDFBase[BatchCancellationArguments, str, BatchCancellationArguments]): + @classmethod + def _get_udf_base_args(cls): + return (BatchCancellationArguments, str, BatchCancellationArguments) + + def get_batchable_arguments(self, arguments: BatchCancellationArguments) -> BatchCancellationArguments: + return arguments + + async def async_execute( + self, execution_context: ExecutionContext, arguments: BatchCancellationArguments + ) -> str: + return arguments.id + + async def async_execute_batch( + self, + execution_context: ExecutionContext, + udfs: Sequence[UDFBase[Any, Any]], + arguments: Sequence[BatchCancellationArguments], + ) -> Sequence[Result[str, Exception]]: + batch_started.set() + try: + await batch_release.wait() + except asyncio.CancelledError: + batch_cancelled.set() + raise + return [Ok(argument.id) for argument in arguments] + + execution = asyncio.create_task( + async_execute_with_result( + """ + Single = SlowSingletUDF(value="single") + Batch1 = SlowBatchUDF(id="one", routing_key="shared") + Batch2 = SlowBatchUDF(id="two", routing_key="shared") + """, + udf_registry=UDFRegistry.with_udfs(SlowSingletUDF, SlowBatchUDF), + ) + ) + + await asyncio.wait_for(asyncio.gather(singlet_started.wait(), batch_started.wait()), timeout=1) + execution.cancel() + try: + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(execution, timeout=1) + assert singlet_cancelled.is_set() + assert batch_cancelled.is_set() + finally: + singlet_release.set() + batch_release.set() + + +@pytest.mark.asyncio +async def test_repeated_cancellation_waits_for_owned_task_cleanup(async_execute_with_result): + started = asyncio.Event() + work_release = asyncio.Event() + cleanup_started = asyncio.Event() + cleanup_release = asyncio.Event() + + class SlowUDF(AsyncUDFBase[CancellationArguments, str]): + @classmethod + def _get_udf_base_args(cls): + return (CancellationArguments, str) + + async def async_execute(self, execution_context: ExecutionContext, arguments: CancellationArguments) -> str: + started.set() + try: + await work_release.wait() + except asyncio.CancelledError: + cleanup_started.set() + await cleanup_release.wait() + raise + raise AssertionError('unreachable') + + execution = asyncio.create_task( + async_execute_with_result( + 'Result = SlowUDF(value="test")', + udf_registry=UDFRegistry.with_udfs(SlowUDF), + ) + ) + await asyncio.wait_for(started.wait(), timeout=1) + + try: + execution.cancel() + await asyncio.wait_for(cleanup_started.wait(), timeout=1) + for _ in range(3): + execution.cancel() + for _ in range(10): + await asyncio.sleep(0) + + assert not execution.done() + cleanup_release.set() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(execution, timeout=1) + assert execution.cancelling() == 1 + finally: + work_release.set() + cleanup_release.set() + await asyncio.wait_for(asyncio.gather(execution, return_exceptions=True), timeout=1) + + +@pytest.mark.asyncio +@pytest.mark.parametrize(('stale_cancelling_count', 'expected_cancelling_count'), [(False, 1), (True, 2)]) +async def test_cancellation_during_child_cancellation_cleanup_preserves_request( + async_execute_with_result, stale_cancelling_count: bool, expected_cancelling_count: int +): + cancelled_started = asyncio.Event() + cancelled_release = asyncio.Event() + slow_started = asyncio.Event() + work_release = asyncio.Event() + cleanup_started = asyncio.Event() + cleanup_release = asyncio.Event() + + class CancelledUDF(AsyncUDFBase[CancellationArguments, str]): + @classmethod + def _get_udf_base_args(cls): + return (CancellationArguments, str) + + async def async_execute(self, execution_context: ExecutionContext, arguments: CancellationArguments) -> str: + cancelled_started.set() + await cancelled_release.wait() + raise asyncio.CancelledError + + class SlowUDF(AsyncUDFBase[CancellationArguments, str]): + @classmethod + def _get_udf_base_args(cls): + return (CancellationArguments, str) + + async def async_execute(self, execution_context: ExecutionContext, arguments: CancellationArguments) -> str: + slow_started.set() + try: + await work_release.wait() + except asyncio.CancelledError: + cleanup_started.set() + await cleanup_release.wait() + raise + raise AssertionError('unreachable') + + async def run_execution(): + if stale_cancelling_count: + current_task = asyncio.current_task() + assert current_task is not None + current_task.cancel() + try: + await asyncio.sleep(0) + except asyncio.CancelledError: + # Preserve the stale cancellation count for this test + pass + return await async_execute_with_result( + 'Cancelled = CancelledUDF(value="cancelled")\nSlow = SlowUDF(value="slow")', + udf_registry=UDFRegistry.with_udfs(CancelledUDF, SlowUDF), + ) + + execution = asyncio.create_task(run_execution()) + await asyncio.wait_for(asyncio.gather(cancelled_started.wait(), slow_started.wait()), timeout=1) + + try: + cancelled_release.set() + await asyncio.wait_for(cleanup_started.wait(), timeout=1) + execution.cancel() + cleanup_release.set() + + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(execution, timeout=1) + assert execution.cancelling() == expected_cancelling_count + finally: + cancelled_release.set() + work_release.set() + cleanup_release.set() + await asyncio.wait_for(asyncio.gather(execution, return_exceptions=True), timeout=1) + + +@pytest.mark.asyncio +async def test_owned_task_cleanup_timeout_does_not_block_cancellation( + async_execute_with_result, monkeypatch: pytest.MonkeyPatch +): + started = asyncio.Event() + work_release = asyncio.Event() + cleanup_started = asyncio.Event() + cleanup_release = asyncio.Event() + cleanup_finished = asyncio.Event() + monkeypatch.setattr(async_executor, '_OWNED_TASK_CLEANUP_SECONDS', 0.01) + + class SlowUDF(AsyncUDFBase[CancellationArguments, str]): + @classmethod + def _get_udf_base_args(cls): + return (CancellationArguments, str) + + async def async_execute(self, execution_context: ExecutionContext, arguments: CancellationArguments) -> str: + started.set() + try: + await work_release.wait() + except asyncio.CancelledError: + cleanup_started.set() + try: + await cleanup_release.wait() + finally: + cleanup_finished.set() + raise + raise AssertionError('unreachable') + + execution = asyncio.create_task( + async_execute_with_result( + 'Result = SlowUDF(value="test")', + udf_registry=UDFRegistry.with_udfs(SlowUDF), + ) + ) + await asyncio.wait_for(started.wait(), timeout=1) + + execution.cancel() + try: + await asyncio.wait_for(cleanup_started.wait(), timeout=1) + done, _ = await asyncio.wait({execution}, timeout=1) + + assert execution in done + with pytest.raises(asyncio.CancelledError): + execution.result() + await asyncio.wait_for(cleanup_finished.wait(), timeout=1) + finally: + work_release.set() + cleanup_release.set() + await asyncio.wait_for(cleanup_finished.wait(), timeout=1) + await asyncio.wait_for(asyncio.gather(execution, return_exceptions=True), timeout=1) + + +@pytest.mark.asyncio +async def test_cancellation_during_error_cleanup_supersedes_original_error(async_execute_with_result): + class FatalError(BaseException): + pass + + fatal_started = asyncio.Event() + fatal_release = asyncio.Event() + slow_started = asyncio.Event() + work_release = asyncio.Event() + cleanup_started = asyncio.Event() + cleanup_release = asyncio.Event() + + class FatalUDF(AsyncUDFBase[CancellationArguments, str]): + @classmethod + def _get_udf_base_args(cls): + return (CancellationArguments, str) + + async def async_execute(self, execution_context: ExecutionContext, arguments: CancellationArguments) -> str: + fatal_started.set() + await fatal_release.wait() + raise FatalError + + class SlowUDF(AsyncUDFBase[CancellationArguments, str]): + @classmethod + def _get_udf_base_args(cls): + return (CancellationArguments, str) + + async def async_execute(self, execution_context: ExecutionContext, arguments: CancellationArguments) -> str: + slow_started.set() + try: + await work_release.wait() + except asyncio.CancelledError: + cleanup_started.set() + await cleanup_release.wait() + raise + raise AssertionError('unreachable') + + execution = asyncio.create_task( + async_execute_with_result( + 'Fatal = FatalUDF(value="fatal")\nSlow = SlowUDF(value="slow")', + udf_registry=UDFRegistry.with_udfs(FatalUDF, SlowUDF), + ) + ) + await asyncio.wait_for(asyncio.gather(fatal_started.wait(), slow_started.wait()), timeout=1) + + try: + fatal_release.set() + await asyncio.wait_for(cleanup_started.wait(), timeout=1) + execution.cancel() + cleanup_release.set() + + with pytest.raises(asyncio.CancelledError) as exc_info: + await asyncio.wait_for(execution, timeout=1) + assert execution.cancelling() == 1 + assert isinstance(exc_info.value.__cause__, FatalError) + finally: + fatal_release.set() + work_release.set() + cleanup_release.set() + await asyncio.wait_for(asyncio.gather(execution, return_exceptions=True), timeout=1) diff --git a/osprey_async_worker/src/osprey/async_worker/tests/test_external_service.py b/osprey_async_worker/src/osprey/async_worker/tests/test_external_service.py index 789c240..da3354d 100644 --- a/osprey_async_worker/src/osprey/async_worker/tests/test_external_service.py +++ b/osprey_async_worker/src/osprey/async_worker/tests/test_external_service.py @@ -204,7 +204,7 @@ async def test_cancelling_get_without_cache_does_not_cancel_shared_get(): service = CancelOnceService() accessor = ExternalServiceAccessor(service) owner = asyncio.create_task(accessor.get_without_cache('foo')) - await service.started.wait() + await asyncio.wait_for(service.started.wait(), timeout=1) owner.cancel() with pytest.raises(asyncio.CancelledError): @@ -229,7 +229,7 @@ async def test_cancelling_waiter_does_not_cancel_shared_get(): service = CancelOnceService() accessor = ExternalServiceAccessor(service) owner = asyncio.create_task(accessor.get('foo')) - await service.started.wait() + await asyncio.wait_for(service.started.wait(), timeout=1) waiter = asyncio.create_task(accessor.get('foo')) waiter.cancel() @@ -246,7 +246,7 @@ async def test_cancelling_owner_does_not_cancel_shared_get(): service = CancelOnceService() accessor = ExternalServiceAccessor(service) owner = asyncio.create_task(accessor.get('foo')) - await service.started.wait() + await asyncio.wait_for(service.started.wait(), timeout=1) survivor = asyncio.create_task(accessor.get('foo')) owner.cancel() @@ -283,9 +283,9 @@ async def test_failed_task_does_not_evict_replacement(): service = ReplacementService() accessor = ExternalServiceAccessor(service) first = asyncio.create_task(accessor.get('foo')) - await service.first_started.wait() + await asyncio.wait_for(service.first_started.wait(), timeout=1) replacement = asyncio.create_task(accessor.get_without_cache('foo')) - await service.second_started.wait() + await asyncio.wait_for(service.second_started.wait(), timeout=1) service.first_release.set() with pytest.raises(ValueError): @@ -302,9 +302,16 @@ async def test_count_error_once_with_concurrent_waiter(): service = CountErrorOnceGatedService() accessor = ExternalServiceAccessor(service) creator = asyncio.create_task(accessor.get('foo')) - await service.started.wait() - waiter = asyncio.create_task(accessor.get('foo')) - await asyncio.sleep(0) + await asyncio.wait_for(service.started.wait(), timeout=1) + waiter_started = asyncio.Event() + + async def wait_for_cached_get(): + # Do not suspend before the cached future is attached below + waiter_started.set() + return await accessor.get('foo') + + waiter = asyncio.create_task(wait_for_cached_get()) + await asyncio.wait_for(waiter_started.wait(), timeout=1) service.release.set() with pytest.raises(ValueError): @@ -319,7 +326,7 @@ async def test_count_error_once_does_not_apply_to_get_without_cache(): service = CountErrorOnceGatedService() accessor = ExternalServiceAccessor(service) creator = asyncio.create_task(accessor.get_without_cache('foo')) - await service.started.wait() + await asyncio.wait_for(service.started.wait(), timeout=1) service.release.set() with pytest.raises(ValueError, match='service fails'): @@ -366,12 +373,26 @@ async def test_batch_get_deduplicates_duplicate_keys(): assert service.batch_call_count == 1 +@pytest.mark.asyncio +async def test_batch_get_resolves_its_owned_future_after_cache_replacement(): + service = GatedBatchService() + accessor = ExternalServiceAccessor(service) + batch_get = asyncio.create_task(accessor.batch_get(['a'])) + await asyncio.wait_for(service.started.wait(), timeout=1) + + assert await accessor.get_without_cache('a') == 'value_a' + service.release.set() + + assert await asyncio.wait_for(batch_get, timeout=1) == [Ok('batch_a')] + assert await accessor.get('a') == 'value_a' + + @pytest.mark.asyncio async def test_cancelled_batch_loader_evicts_its_cache_entries(): service = GatedBatchService() accessor = ExternalServiceAccessor(service) batch = asyncio.create_task(accessor.batch_get(['a'])) - await service.started.wait() + await asyncio.wait_for(service.started.wait(), timeout=1) loader = next(iter(accessor._active_batch_loaders)) loader.cancel() @@ -410,7 +431,7 @@ async def test_cancelling_batch_owner_does_not_cancel_shared_get(): service = GatedBatchService() accessor = ExternalServiceAccessor(service) batch = asyncio.create_task(accessor.batch_get(['a'])) - await service.started.wait() + await asyncio.wait_for(service.started.wait(), timeout=1) survivor = asyncio.create_task(accessor.get('a')) batch.cancel() @@ -428,7 +449,7 @@ async def test_cancelled_batch_owner_keeps_loader_alive_through_garbage_collecti service = GatedBatchService() accessor = ExternalServiceAccessor(service) batch = asyncio.create_task(accessor.batch_get(['a'])) - await service.started.wait() + await asyncio.wait_for(service.started.wait(), timeout=1) survivor = asyncio.create_task(accessor.get('a')) batch.cancel() @@ -451,14 +472,14 @@ async def test_cancelled_failed_batch_consumes_future_exceptions(): loop.set_exception_handler(lambda _loop, context: contexts.append(context)) try: batch = asyncio.create_task(accessor.batch_get(['a'])) - await service.started.wait() + await asyncio.wait_for(service.started.wait(), timeout=1) + loader = next(iter(accessor._active_batch_loaders)) batch.cancel() with pytest.raises(asyncio.CancelledError): _ = await batch service.release.set() - await asyncio.sleep(0) - await asyncio.sleep(0) + await asyncio.wait_for(asyncio.shield(loader), timeout=1) accessor._cache.clear() gc.collect() await asyncio.sleep(0) @@ -473,9 +494,16 @@ async def test_cancelling_batch_waiter_does_not_cancel_shared_get(): service = CancelOnceService() accessor = ExternalServiceAccessor(service) owner = asyncio.create_task(accessor.get('a')) - await service.started.wait() - batch = asyncio.create_task(accessor.batch_get(['a'])) - await asyncio.sleep(0) + await asyncio.wait_for(service.started.wait(), timeout=1) + batch_started = asyncio.Event() + + async def wait_for_cached_batch(): + # Do not suspend before the cached future is attached below + batch_started.set() + return await accessor.batch_get(['a']) + + batch = asyncio.create_task(wait_for_cached_batch()) + await asyncio.wait_for(batch_started.wait(), timeout=1) batch.cancel() with pytest.raises(asyncio.CancelledError):