diff --git a/core/shared/src/main/scala/fs2/Stream.scala b/core/shared/src/main/scala/fs2/Stream.scala index c95539b649..faf312c9bb 100644 --- a/core/shared/src/main/scala/fs2/Stream.scala +++ b/core/shared/src/main/scala/fs2/Stream.scala @@ -1152,6 +1152,22 @@ final class Stream[+F[_], +O] private[fs2] (private[fs2] val underlying: Pull[F, underlying.flatMapOutput(tapOut).streamNoScope } + /** Executes an effect for certain errors, then rethrows the original error. + * Voids any error thrown by the given effect. Any non-matching error is rethrown as well. + */ + def onErrorEvalTap[F2[x] >: F[x], O2]( + pf: PartialFunction[Throwable, F2[O2]] + )(implicit F: ApplicativeError[F2, Throwable]): Stream[F2, O] = + handleErrorWith { t => + import fs2.Stream.NotApplied + val rethrow = new Stream[F2, O](Pull.fail(t)) + val eff = pf.applyOrElse(t, NotApplied) + + if (eff.asInstanceOf[AnyRef] ne NotApplied) + Stream.exec(eff.asInstanceOf[F2[O2]].void.voidError) ++ rethrow + else rethrow + } + @deprecated("Use overload without functor", "3.7.0") private[fs2] def evalTap[F2[x] >: F[x], O2](f: O => F2[O2], F: Functor[F2]): Stream[F2, O] = evalTap(f) @@ -3315,6 +3331,9 @@ object Stream extends StreamLowPriority { */ val unit: Stream[Pure, Unit] = Pull.outUnit.streamNoScope + // A special value that is used to indicate that whether a PartialFunction is not applied + private final val NotApplied: Any => Any = _ => Stream.NotApplied + /** Creates a single element stream that gets its value by evaluating the supplied effect. If the effect fails, a `Left` * is emitted. Otherwise, a `Right` is emitted. * diff --git a/core/shared/src/test/scala/fs2/StreamCombinatorsSuite.scala b/core/shared/src/test/scala/fs2/StreamCombinatorsSuite.scala index 48ed6e1011..b88deb2ca9 100644 --- a/core/shared/src/test/scala/fs2/StreamCombinatorsSuite.scala +++ b/core/shared/src/test/scala/fs2/StreamCombinatorsSuite.scala @@ -510,6 +510,50 @@ class StreamCombinatorsSuite extends Fs2Suite { } } + group("onErrorEvalTap") { + test("rethrows the original error") { + Counter[SyncIO].flatMap { counter => + Stream + .range(0, 3) + .append(Stream.raiseError[SyncIO](new Err)) + .onErrorEvalTap { + case _: IllegalStateException => counter.decrement + case _: Err => counter.increment + } + .compile + .drain + .intercept[Err] >> counter.get.assertEquals(1L) + } + } + + test("fires an effect for matching errors only") { + Counter[SyncIO].flatMap { counter => + Stream + .range(0, 3) + .append(Stream.raiseError[SyncIO](new Err)) + .onErrorEvalTap { case _: IllegalStateException => + counter.increment + } + .compile + .drain + .intercept[Err] >> counter.get.assertEquals(0L) + } + } + + test("voids the underlying effect errors") { + Counter[SyncIO].flatMap { counter => + Stream + .raiseError[SyncIO](new Err) + .onErrorEvalTap { case _: Err => + counter.increment *> SyncIO.raiseError(new IllegalStateException("Oops!")) + } + .compile + .drain + .intercept[Err] >> counter.get.assertEquals(1L) + } + } + } + test("evalScan") { forAllF { (s: Stream[Pure, Int], n: String) => val f: (String, Int) => IO[String] = (a: String, b: Int) => IO.pure(a + b)