diff --git a/aioreactive/transform.py b/aioreactive/transform.py index 8a5eeab..ac6cf88 100644 --- a/aioreactive/transform.py +++ b/aioreactive/transform.py @@ -18,7 +18,6 @@ from expression.system import AsyncDisposable from .combine import merge_inner, zip_seq -from .create import fail from .msg import ( Key, Msg, @@ -433,18 +432,29 @@ def retry( retry_count: int, ) -> Callable[[AsyncObservable[_TSource]], AsyncObservable[_TSource]]: def _retry(source: AsyncObservable[_TSource]) -> AsyncObservable[_TSource]: - count = retry_count + async def subscribe_async(observer: AsyncObserver[_TSource]) -> AsyncDisposable: + # Each subscription gets its own retry count + async def action(remaining: int, src: AsyncObservable[_TSource]) -> AsyncDisposable: + async def asend(value: _TSource) -> None: + await observer.asend(value) + + async def athrow(error: Exception) -> None: + if not remaining: + await observer.athrow(error) + else: + # Retry with one less attempt + await action(remaining - 1, source) + + async def aclose() -> None: + await observer.aclose() - def factory(exn: Exception) -> AsyncObservable[_TSource]: - nonlocal count + _obv = AsyncAnonymousObserver(asend, athrow, aclose) + return await src.subscribe_async(_obv) - if not count: - return fail(exn) - else: - count -= count - return source + disposable = await action(retry_count, source) + return AsyncDisposable.create(disposable.dispose_async) - return pipe(source, catch(factory)) + return AsyncAnonymousObservable(subscribe_async) return _retry diff --git a/tests/test_retry.py b/tests/test_retry.py new file mode 100644 index 0000000..c39fc8e --- /dev/null +++ b/tests/test_retry.py @@ -0,0 +1,163 @@ +import asyncio + +import pytest +from expression.core import pipe + +import aioreactive as rx +from aioreactive import AsyncObserver +from aioreactive.create import fail +from aioreactive.testing import VirtualTimeEventLoop + + +class EventLoopPolicy(asyncio.DefaultEventLoopPolicy): + def get_event_loop(self) -> asyncio.AbstractEventLoop: + return VirtualTimeEventLoop() + + +@pytest.fixture(scope="module") +def event_loop_policy(): + return EventLoopPolicy() + + +@pytest.mark.asyncio(loop_scope="module") +async def test_retry_success(): + """Test that retry successfully retries after initial failure.""" + attempt_count = 0 + + def failing_source(): + nonlocal attempt_count + attempt_count += 1 + if attempt_count < 3: + return fail(Exception("Temporary failure")) + return rx.single("Success!") + + # Create an observable that fails twice then succeeds + xs = rx.defer(failing_source) + + # Retry up to 3 times + ys = pipe(xs, rx.retry(3)) + + # Create a test observer + values: list[str] = [] + + async def asend(value: str) -> None: + values.append(value) + + obv: AsyncObserver[str] = rx.AsyncAwaitableObserver(asend) + async with await ys.subscribe_async(obv): + result = await obv + assert result == "Success!" + assert values == ["Success!"] + assert attempt_count == 3 + + +@pytest.mark.asyncio(loop_scope="module") +async def test_retry_failure(): + """Test that retry propagates errors after exhausting retry attempts.""" + attempt_count = 0 + + def always_failing_source(): + nonlocal attempt_count + attempt_count += 1 + return fail(Exception("Permanent failure")) + + # Create an observable that always fails + xs = rx.defer(always_failing_source) + + # Retry only once (not enough) + ys = pipe(xs, rx.retry(1)) + + # Should fail after 2 attempts (initial + 1 retry) + exception = None + + async def athrow(ex: Exception): + nonlocal exception + exception = ex + + obv = rx.AsyncAwaitableObserver(athrow=athrow) + + await ys.subscribe_async(obv) + + try: + await obv + except Exception as ex: + assert str(ex) == "Permanent failure" + assert attempt_count == 2 + else: + assert False + + +@pytest.mark.asyncio(loop_scope="module") +async def test_retry_immediate_success(): + """Test that retry passes through successful completions without interference.""" + xs = rx.single("Success!") + ys = pipe(xs, rx.retry(3)) + + values: list[str] = [] + + async def asend(value: str) -> None: + values.append(value) + + obv: AsyncObserver[str] = rx.AsyncAwaitableObserver(asend) + async with await ys.subscribe_async(obv): + result = await obv + assert result == "Success!" + assert values == ["Success!"] + + +@pytest.mark.asyncio(loop_scope="module") +async def test_retry_zero_count(): + """Test that retry with zero count doesn't retry.""" + attempt_count = 0 + + def always_failing_source(): + nonlocal attempt_count + attempt_count += 1 + return fail(Exception("Failure")) + + # Create an observable that always fails + xs = rx.defer(always_failing_source) + + # Retry zero times (no retries) + ys = pipe(xs, rx.retry(0)) + + # Should fail immediately without retrying + exception = None + + async def athrow(ex: Exception): + nonlocal exception + exception = ex + + obv = rx.AsyncAwaitableObserver(athrow=athrow) + + await ys.subscribe_async(obv) + + try: + await obv + except Exception as ex: + assert str(ex) == "Failure" + assert attempt_count == 1 + else: + assert False + + +@pytest.mark.asyncio(loop_scope="module") +async def test_retry_subscription_cancel(): + """Test that retry properly handles subscription cancellation.""" + xs: rx.AsyncSubject[int] = rx.AsyncSubject() + result: list[int] = [] + + ys = pipe(xs, rx.retry(3)) + + async def asend(value: int) -> None: + result.append(value) + # Get the subscription from the context and dispose it + await asyncio.sleep(0) + + async with await ys.subscribe_async(rx.AsyncAnonymousObserver(asend)) as subscription: + await xs.asend(10) + await asyncio.sleep(0) + await subscription.dispose_async() + await xs.asend(20) + + assert result == [10]