diff --git a/core/shared/src/main/scala/fs2/concurrent/Channel.scala b/core/shared/src/main/scala/fs2/concurrent/Channel.scala index df2f17b54a..8d97a29e35 100644 --- a/core/shared/src/main/scala/fs2/concurrent/Channel.scala +++ b/core/shared/src/main/scala/fs2/concurrent/Channel.scala @@ -24,6 +24,7 @@ package concurrent import cats.effect._ import cats.effect.implicits._ +import cats.effect.Resource.ExitCase import cats.syntax.all._ /** Stream aware, multiple producer, single consumer closeable channel. @@ -116,6 +117,8 @@ sealed trait Channel[F[_], A] { */ def closeWithElement(a: A): F[Either[Channel.Closed, Unit]] + private[concurrent] def closeWithExitCase(exitCase: ExitCase): F[Either[Channel.Closed, Unit]] + /** Returns true if this channel is closed */ def isClosed: F[Boolean] @@ -138,62 +141,62 @@ object Channel { size: Int, waiting: Option[Deferred[F, Unit]], producers: List[(A, Deferred[F, Unit])], - closed: Boolean + closed: Option[ExitCase] ) - val open = State(List.empty, 0, None, List.empty, closed = false) + val open = State(List.empty, 0, None, List.empty, closed = None) - def empty(isClosed: Boolean): State = - if (isClosed) State(List.empty, 0, None, List.empty, closed = true) + def empty(close: Option[ExitCase]): State = + if (close.nonEmpty) State(List.empty, 0, None, List.empty, closed = close) else open (F.ref(open), F.deferred[Unit]).mapN { (state, closedGate) => new Channel[F, A] { def sendAll: Pipe[F, A, Nothing] = { in => - (in ++ Stream.exec(close.void)) + in.onFinalizeCase(closeWithExitCase(_).void) .evalMap(send) .takeWhile(_.isRight) .drain } - def sendImpl(a: A, close: Boolean) = + def sendImpl(a: A, close: Option[ExitCase]) = F.deferred[Unit].flatMap { producer => state.flatModifyFull { case (poll, state) => state match { - case s @ State(_, _, _, _, closed @ true) => + case s @ State(_, _, _, _, Some(_)) => (s, Channel.closed[Unit].pure[F]) - case State(values, size, waiting, producers, closed @ false) => + case State(values, size, waiting, producers, None) => if (size < capacity) ( State(a :: values, size + 1, None, producers, close), - signalClosure.whenA(close) *> notifyStream(waiting).as(rightUnit) + signalClosure.whenA(close.nonEmpty) *> notifyStream(waiting).as(rightUnit) ) else ( State(values, size, None, (a, producer) :: producers, close), - signalClosure.whenA(close) *> + signalClosure.whenA(close.nonEmpty) *> notifyStream(waiting).as(rightUnit) <* - waitOnBound(producer, poll).unlessA(close) + waitOnBound(producer, poll).unlessA(close.nonEmpty) ) } } } - def send(a: A) = sendImpl(a, false) + def send(a: A) = sendImpl(a, None) - def closeWithElement(a: A) = sendImpl(a, true) + def closeWithElement(a: A) = sendImpl(a, Some(ExitCase.Succeeded)) def trySend(a: A) = state.flatModify { - case s @ State(_, _, _, _, closed @ true) => + case s @ State(_, _, _, _, Some(_)) => (s, Channel.closed[Boolean].pure[F]) - case s @ State(values, size, waiting, producers, closed @ false) => + case s @ State(values, size, waiting, producers, None) => if (size < capacity) ( - State(a :: values, size + 1, None, producers, false), + State(a :: values, size + 1, None, producers, None), notifyStream(waiting).as(rightTrue) ) else @@ -201,13 +204,16 @@ object Channel { } def close = + closeWithExitCase(ExitCase.Succeeded) + + def closeWithExitCase(exitCase: ExitCase): F[Either[Closed, Unit]] = state.flatModify { - case s @ State(_, _, _, _, closed @ true) => + case s @ State(_, _, _, _, Some(_)) => (s, Channel.closed[Unit].pure[F]) - case State(values, size, waiting, producers, closed @ false) => + case State(values, size, waiting, producers, None) => ( - State(values, size, None, producers, true), + State(values, size, None, producers, Some(exitCase)), notifyStream(waiting).as(rightUnit) <* signalClosure ) } @@ -250,8 +256,12 @@ object Channel { unblock.as(Pull.output(toEmit) >> consumeLoop) } else { F.pure( - if (closed) Pull.done - else Pull.eval(waiting.get) >> consumeLoop + closed match { + case Some(ExitCase.Succeeded) => Pull.done + case Some(ExitCase.Errored(e)) => Pull.raiseError(e) + case Some(ExitCase.Canceled) => Pull.eval(F.canceled) + case None => Pull.eval(waiting.get) >> consumeLoop + } ) } } diff --git a/core/shared/src/main/scala/fs2/concurrent/Topic.scala b/core/shared/src/main/scala/fs2/concurrent/Topic.scala index b069bcfe6f..6a175ff424 100644 --- a/core/shared/src/main/scala/fs2/concurrent/Topic.scala +++ b/core/shared/src/main/scala/fs2/concurrent/Topic.scala @@ -23,6 +23,7 @@ package fs2 package concurrent import cats.effect._ +import cats.effect.Resource.ExitCase import cats.syntax.all._ import scala.collection.immutable.LongMap @@ -220,7 +221,8 @@ object Topic { } def publish: Pipe[F, A, Nothing] = { in => - (in ++ Stream.exec(close.void)) + in + .onFinalizeCase(closeWithExitCase(_).void) .evalMap(publish1) .takeWhile(_.isRight) .drain @@ -235,9 +237,13 @@ object Topic { def subscribers: Stream[F, Int] = subscriberCount.discrete def close: F[Either[Topic.Closed, Unit]] = + closeWithExitCase(ExitCase.Succeeded) + + def closeWithExitCase(exitCase: ExitCase): F[Either[Closed, Unit]] = state.flatModify { case State.Active(subs, _) => - val action = foreach(subs)(_.close.void) *> signalClosure.complete(()) + val action = + foreach(subs)(_.closeWithExitCase(exitCase).void) *> signalClosure.complete(()) (State.Closed(), action.as(Topic.rightUnit)) case closed @ State.Closed() => (closed, Topic.closed.pure[F]) diff --git a/core/shared/src/test/scala/fs2/concurrent/ChannelSuite.scala b/core/shared/src/test/scala/fs2/concurrent/ChannelSuite.scala index 8da8487d6e..9235488253 100644 --- a/core/shared/src/test/scala/fs2/concurrent/ChannelSuite.scala +++ b/core/shared/src/test/scala/fs2/concurrent/ChannelSuite.scala @@ -29,6 +29,8 @@ import scala.concurrent.duration._ import org.scalacheck.effect.PropF.forAllF +import scala.concurrent.CancellationException + class ChannelSuite extends Fs2Suite { test("receives some simple elements above capacity and closes") { @@ -323,4 +325,21 @@ class ChannelSuite extends Fs2Suite { racingSendOperations(channel) } + test("stream should terminate when sendAll is interrupted") { + val program = + Channel + .bounded[IO, Unit](1) + .flatMap { ch => + val producer = + Stream + .eval(IO.canceled) + .through(ch.sendAll) + + ch.stream.concurrently(producer).compile.drain + } + + TestControl + .executeEmbed(program) + .intercept[CancellationException] + } } diff --git a/core/shared/src/test/scala/fs2/concurrent/TopicSuite.scala b/core/shared/src/test/scala/fs2/concurrent/TopicSuite.scala index f2d889b92d..c3bc812616 100644 --- a/core/shared/src/test/scala/fs2/concurrent/TopicSuite.scala +++ b/core/shared/src/test/scala/fs2/concurrent/TopicSuite.scala @@ -25,6 +25,8 @@ package concurrent import cats.syntax.all._ import cats.effect.IO import scala.concurrent.duration._ +import scala.concurrent.CancellationException + import cats.effect.testkit.TestControl class TopicSuite extends Fs2Suite { @@ -287,4 +289,27 @@ class TopicSuite extends Fs2Suite { check.replicateA_(10000) } + + test("publisher cancellation does not deadlock") { + val program = + Topic[IO, String] + .flatMap { topic => + val publisher = + Stream + .constant("1") + .covary[IO] + .evalTap(_ => IO.canceled) + .through(topic.publish) + + Stream + .resource(topic.subscribeAwait(1)) + .flatMap(subscriber => subscriber.concurrently(publisher)) + .compile + .drain + } + + TestControl + .executeEmbed(program) + .intercept[CancellationException] + } }