From 14c8489acff05335daea6a3937aabdca0f586108 Mon Sep 17 00:00:00 2001 From: Daniel Spiewak Date: Wed, 19 Jun 2024 18:58:14 -0500 Subject: [PATCH 01/10] Implemented asynchronous cancelation in `PureConc` --- .../cats/effect/kernel/testkit/pure.scala | 141 +++++++++++------- .../cats/effect/laws/PureConcSuite.scala | 89 +++++++++++ 2 files changed, 178 insertions(+), 52 deletions(-) diff --git a/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala b/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala index 52578252d7..bbb9c310fd 100644 --- a/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala +++ b/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala @@ -17,7 +17,7 @@ package cats.effect.kernel package testkit -import cats.{~>, Defer, Eq, Functor, Id, Monad, MonadError, Order, Show} +import cats.{~>, Applicative, Defer, Eq, Functor, Id, Monad, MonadError, Order, Show} import cats.data.{Kleisli, State, WriterT} import cats.effect.kernel._ import cats.free.FreeT @@ -97,45 +97,48 @@ object pure { type Main[X] = MVarR[ResolvedPC[E, *], X] MVar.empty[Main, Outcome[PureConc[E, *], E, A]].flatMap { state0 => - val state = state0[Main] - - val fiber = new PureFiber[E, A](state0) + MVar.empty[Main, Unit] flatMap { canceled0 => + val state = state0[Main] + val fiber = new PureFiber[E, A](state0, canceled0) + + val identified = canceled mapF { ta => + val fk = new (FiberR[E, *] ~> IdOC[E, *]) { + def apply[a](ke: FiberR[E, a]) = + ke.run(FiberCtx(fiber)) + } - val identified = canceled mapF { ta => - val fk = new (FiberR[E, *] ~> IdOC[E, *]) { - def apply[a](ke: FiberR[E, a]) = - ke.run(FiberCtx(fiber)) + ta.mapK(fk) } - ta.mapK(fk) - } + import Outcome._ - import Outcome._ + val body = identified flatMap { a => + state.tryPut(Succeeded(a.pure[PureConc[E, *]])) + } handleErrorWith { e => state.tryPut(Errored(e)) } - val body = identified flatMap { a => - state.tryPut(Succeeded(a.pure[PureConc[E, *]])) - } handleErrorWith { e => state.tryPut(Errored(e)) } + val results = state.read.flatMap { + case Canceled() => (Outcome.Canceled(): IdOC[E, A]).pure[Main] + case Errored(e) => (Outcome.Errored(e): IdOC[E, A]).pure[Main] - val results = state.read.flatMap { - case Canceled() => (Outcome.Canceled(): IdOC[E, A]).pure[Main] - case Errored(e) => (Outcome.Errored(e): IdOC[E, A]).pure[Main] + case Succeeded(fa) => + val identifiedCompletion = fa.mapF { ta => + val fk = new (FiberR[E, *] ~> IdOC[E, *]) { + def apply[a](ke: FiberR[E, a]) = + ke.run(FiberCtx(fiber)) + } - case Succeeded(fa) => - val identifiedCompletion = fa.mapF { ta => - val fk = new (FiberR[E, *] ~> IdOC[E, *]) { - def apply[a](ke: FiberR[E, a]) = - ke.run(FiberCtx(fiber)) + ta.mapK(fk) } - ta.mapK(fk) - } + identifiedCompletion.map(a => Succeeded[Id, E, A](a): IdOC[E, A]) handleError { + e => Errored(e) + } + } - identifiedCompletion.map(a => Succeeded[Id, E, A](a): IdOC[E, A]) handleError { e => - Errored(e) - } + Kleisli.ask[ResolvedPC[E, *], MVar.Universe].map { u => + ApplicativeThread[ResolvedPC[E, *]].start(body.run(u)) >> results.run(u) + } } - - Kleisli.ask[ResolvedPC[E, *], MVar.Universe].map { u => body.run(u) >> results.run(u) } } } @@ -164,8 +167,9 @@ object pure { case (List(results), _) => results.mapK(optLift) case (_, false) => Outcome.Succeeded(None) - // we could make a writer that only receives one object, but that seems meh. just pretend we deadlocked - case _ => Outcome.Succeeded(None) + // in the case of never and such, we are awaiting the async cancel monitor + // this scenario only arises if the main fiber self cancels + case _ => Outcome.Canceled() } } @@ -197,20 +201,16 @@ object pure { } def canceled: PureConc[E, Unit] = - Thread.annotate("canceled") { - withCtx { ctx => - if (ctx.masks.isEmpty) - uncancelable(_ => ctx.self.cancel >> ctx.finalizers.sequence_ >> Thread.done) - else - ctx.self.cancel - } - } + Thread.annotate("canceled")(withCtx(_.self.cancelAndRealize.ifM(Thread.done, unit))) def cede: PureConc[E, Unit] = Thread.cede def never[A]: PureConc[E, A] = - Thread.annotate("never")(Thread.done[A]) + withCtx[E, A] { ctx => + // we monitor for asynchronous cancelation. if we're masked, this won't cancel and we hang + Thread.annotate("never")(ctx.self.awaitCancelation *> Thread.done) + } def ref[A](a: A): PureConc[E, Ref[PureConc[E, *], A]] = MVar[PureConc[E, *], A](a).flatMap(mVar => Kleisli.pure(unsafeRef(mVar))) @@ -273,7 +273,19 @@ object pure { private def unsafeDeferred[A](mVar: MVar[A]): Deferred[PureConc[E, *], A] = new Deferred[PureConc[E, *], A] { - override def get: PureConc[E, A] = mVar.read[PureConc[E, *]] + override def get: PureConc[E, A] = + withCtx { ctx => + // we need to race cancelation against reading the mvar + // if cancelation wins, we shut down the thread + MVar.empty[PureConc[E, *], Option[A]] flatMap { signal => + val left = Thread.start(ctx.self.awaitCancelation.ifM(signal.tryPut[PureConc[E, *]](None).void, unit)) + val right = Thread.start(mVar.read[PureConc[E, *]].flatMap(a => signal.tryPut[PureConc[E, *]](Some(a)))) + + left *> + right *> + signal.read[PureConc[E, *]].flatMap(_.map(_.pure[PureConc[E, *]]).getOrElse(Thread.done)) + } + } override def complete(a: A): PureConc[E, Boolean] = mVar.tryPut[PureConc[E, *]](a) @@ -283,12 +295,14 @@ object pure { def start[A](fa: PureConc[E, A]): PureConc[E, Fiber[PureConc[E, *], E, A]] = Thread.annotate("start", true) { MVar.empty[PureConc[E, *], Outcome[PureConc[E, *], E, A]].flatMap { state => - val fiber = new PureFiber[E, A](state) + MVar.empty[PureConc[E, *], Unit] flatMap { canceled => + val fiber = new PureFiber[E, A](state, canceled) - // the tryPut here is interesting: it encodes first-wins semantics on cancelation/completion - val body = guaranteeCase(fa)(state.tryPut[PureConc[E, *]](_).void) - val identified = localCtx(FiberCtx(fiber), body) - Thread.start(identified.attempt.void).as(fiber) + // the tryPut here is interesting: it encodes first-wins semantics on cancelation/completion + val body = guaranteeCase(fa)(state.tryPut[PureConc[E, *]](_).void) + val identified = localCtx(FiberCtx(fiber), body) + Thread.start(identified.attempt.void).as(fiber) + } } } @@ -363,14 +377,16 @@ object pure { } // todo: MVar is not Serializable, release then update here - final class PureFiber[E, A](val state0: MVar[Outcome[PureConc[E, *], E, A]]) + final class PureFiber[E, A]( + val state0: MVar[Outcome[PureConc[E, *], E, A]], + val canceled0: MVar[Unit]) extends Fiber[PureConc[E, *], E, A] with Serializable { private[this] val state = state0[PureConc[E, *]] private[pure] val canceled: PureConc[E, Boolean] = - state.tryRead.map(_.map(_.fold(true, _ => false, _ => false)).getOrElse(false)) + canceled0.tryRead[PureConc[E, *]].map(_.as(true).getOrElse(false)) private[pure] val realizeCancelation: PureConc[E, Boolean] = withCtx { ctx => @@ -379,7 +395,10 @@ object pure { checkM.ifM( canceled.ifM( // if unmasked and canceled, finalize - allocateForPureConc[E].uncancelable(_ => ctx.finalizers.sequence_.as(true)), + allocateForPureConc[E] uncancelable { _ => + ctx.finalizers.sequence_.as(true) <* state0.tryPut[PureConc[E, *]]( + Outcome.Canceled()) + }, // if unmasked but not canceled, ignore false.pure[PureConc[E, *]] ), @@ -388,9 +407,27 @@ object pure { ) } - val cancel: PureConc[E, Unit] = state.tryPut(Outcome.Canceled()).void + private[pure] val awaitCancelation: PureConc[E, Boolean] = + canceled0.read[PureConc[E, *]] *> realizeCancelation + + private[pure] val cancelAndRealize: PureConc[E, Boolean] = + canceled0.tryPut[PureConc[E, *]](()) *> realizeCancelation + + val join: PureConc[E, Outcome[PureConc[E, *], E, A]] = { + val Thread = ApplicativeThread[PureConc[E, *]] + + withCtx { ctx => + // this is exactly like Deferred#get + MVar.empty[PureConc[E, *], Option[Outcome[PureConc[E, *], E, A]]] flatMap { signal => + // note we must read our *own* canceled, not the target fiber's + val left = Thread.start(ctx.self.awaitCancelation.ifM(signal.tryPut[PureConc[E, *]](None).void, Applicative[PureConc[E, *]].unit)) + val right = Thread.start(state.read.flatMap(oc => signal.tryPut[PureConc[E, *]](Some(oc)))) + + left *> right *> signal.read[PureConc[E, *]].flatMap(_.map(_.pure[PureConc[E, *]]).getOrElse(Thread.done)) + } + } + } - val join: PureConc[E, Outcome[PureConc[E, *], E, A]] = - state.read + val cancel: PureConc[E, Unit] = canceled0.tryPut[PureConc[E, *]](()) *> join.void } } diff --git a/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala b/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala index 5da3183f1e..0ca35b6c56 100644 --- a/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala +++ b/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala @@ -52,6 +52,10 @@ class PureConcSuite extends DisciplineSuite with BaseSuite { } test("short-circuit on canceled") { + assertEquals(pure.run(F.canceled), Outcome.Canceled[Option, Int, Unit]()) + assertEquals( + pure.run((F.never[Unit], F.canceled).parTupled), + Outcome.Canceled[Option, Int, (Unit, Unit)]()) assert( pure.run((F.never[Unit], F.canceled).parTupled.start.flatMap(_.join)) === Outcome .Succeeded(Some(Outcome.canceled[F, Int, (Unit, Unit)]))) @@ -78,6 +82,91 @@ class PureConcSuite extends DisciplineSuite with BaseSuite { } } + { + import cats.effect.kernel.{GenConcurrent, Outcome} + import cats.effect.kernel.implicits._ + import cats.syntax.all._ + + type F[A] = PureConc[Int, A] + val F = GenConcurrent[F] + + test("run finalizers when canceling never") { + val t = for { + c <- F.ref(0) + latch <- F.deferred[Unit] + fib <- F.start((latch.complete(()) *> F.never[Unit]).onCancel(c.update(_ + 1))) + _ <- latch.get + _ <- fib.cancel + v <- c.get + } yield v + + assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, Int](Some(1))) + } + + test("run finalizers when canceling Deferred#get") { + val t = for { + c <- F.ref(0) + latch <- F.deferred[Unit] + hang <- F.deferred[Unit] + fib <- F.start((latch.complete(()) *> hang.get).onCancel(c.update(_ + 1))) + _ <- latch.get + _ <- fib.cancel + v <- c.get + } yield v + + assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, Int](Some(1))) + } + + test("run finalizers when canceling Fiber#join") { + val t = for { + c <- F.ref(0) + latch <- F.deferred[Unit] + hang <- F.start(F.never[Unit]) + fib <- F.start((latch.complete(()) *> hang.join).onCancel(c.update(_ + 1))) + _ <- latch.get + _ <- fib.cancel + v <- c.get + } yield v + + assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, Int](Some(1))) + } + + test("hang when canceling uncancelable never") { + val t = for { + latch <- F.deferred[Unit] + f <- F.start((latch.complete(()) *> F.never[Unit]).uncancelable) + _ <- latch.get + _ <- f.cancel + } yield () + + assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, Unit](None)) + } + + test("hang when canceling uncancelable Deferred#get") { + val t = for { + latch <- F.deferred[Unit] + hang <- F.deferred[Unit] + f <- F.start((latch.complete(()) *> hang.get).uncancelable) + _ <- latch.get + _ <- f.cancel + } yield () + + assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, Unit](None)) + } + + test("hang when canceling uncancelable Fiber#join") { + val t = for { + latch <- F.deferred[Unit] + hang <- F.start(F.never[Unit]) + f <- F.start((latch.complete(()) *> hang.join).uncancelable) + _ <- latch.get + _ <- f.cancel + } yield () + + assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, Unit](None)) + } + } + checkAll( "TimeT[PureConc]", GenTemporalTests[TimeT[PureConc[Int, *], *], Int].temporal[Int, Int, Int](10.millis) From be06e67be4af143207b584384419af2dfcf64943 Mon Sep 17 00:00:00 2001 From: Daniel Spiewak Date: Fri, 25 Oct 2024 12:09:55 -0500 Subject: [PATCH 02/10] Beefed up PC tests --- .../cats/effect/laws/PureConcSuite.scala | 61 +++++++++++++++++++ 1 file changed, 61 insertions(+) diff --git a/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala b/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala index 0ca35b6c56..45db4b45db 100644 --- a/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala +++ b/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala @@ -32,6 +32,9 @@ class PureConcSuite extends DisciplineSuite with BaseSuite { import PureConcGenerators._ import OutcomeGenerators._ + override def scalaCheckInitialSeed = + "ogn64yom4GXCEX0mXdqSfsqSeJxI2RbPUFC5YkvDtzD=" + implicit def exec(fb: TimeT[PureConc[Int, *], Boolean]): Prop = Prop(pure.run(TimeT.run(fb)).fold(false, _ => false, _.getOrElse(false))) @@ -165,6 +168,64 @@ class PureConcSuite extends DisciplineSuite with BaseSuite { assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, Unit](None)) } + + test("run finalizers in order") { + val t = for { + results <- F.ref[String]("") + f <- F start { + F.canceled.onCancel(results.update(_ + "A")).onCancel(results.update(_ + "B")) + } + _ <- f.join + back <- results.get + } yield back + + assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, String](Some("AB"))) + } + + test("correctly interpret uncancelable cancelation followed by suspension") { + val t = F.uncancelable(_ => F.canceled *> F.never[Unit]) + assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, Unit](None)) + + val forked = pure.run(F.start(t).flatMap(_.joinWith(F.canceled *> F.never[Unit]))) + assertEquals(forked, Outcome.Succeeded[Option, Int, Unit](None)) + } + + test("implement locals via Kleisli and FreeT") { + import cats.{~>, Eval, Id} + import cats.data.Kleisli + import cats.free.FreeT + import cats.syntax.all._ + + type F[A] = FreeT[Id, Kleisli[Eval, Int, *], A] + + def read[A](f: Int => F[A]): F[A] = + FreeT.liftT(Kleisli.ask[Eval, Int]).flatMap(f) + + def withLocal[A](i: Int)(fa: F[A]): F[A] = + fa.mapK(new (Kleisli[Eval, Int, *] ~> Kleisli[Eval, Int, *]) { + def apply[a](kea: Kleisli[Eval, Int, a]) = + Kleisli((_: Int) => kea(i)) + }) + + def run[A](i: Int)(fa: F[A]): A = + fa.runM(fta => Kleisli.liftF(Eval.now(fta))).apply(i).value + + val _ = run(1) { + withLocal(42) { + read { i => + FreeT + .liftT[Id, Kleisli[Eval, Int, *], Unit]( + Kleisli.liftF[Eval, Int, Unit](Eval.later(assertEquals(i, 42)))) + .flatMap(_ => + read { i2 => + FreeT.liftT(Kleisli.liftF(Eval.later(assertEquals(i2, 42)))) + }) + } + } *> read { i => + FreeT.liftT(Kleisli.liftF(Eval.later(assertEquals(i, 1)))) + } + } + } } checkAll( From 89fb2c5bc46c7afdfb9836e208d0258f20907d2c Mon Sep 17 00:00:00 2001 From: Daniel Spiewak Date: Sun, 21 Jun 2026 18:40:54 -0500 Subject: [PATCH 03/10] Fixed PureConc poll cancelation --- .../cats/effect/kernel/testkit/pure.scala | 573 +++++++++++++++--- .../cats/effect/laws/PureConcSuite.scala | 56 ++ 2 files changed, 528 insertions(+), 101 deletions(-) diff --git a/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala b/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala index bbb9c310fd..925b0b236f 100644 --- a/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala +++ b/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala @@ -17,7 +17,7 @@ package cats.effect.kernel package testkit -import cats.{~>, Applicative, Defer, Eq, Functor, Id, Monad, MonadError, Order, Show} +import cats.{~>, Defer, Eq, Functor, Id, Monad, MonadError, Order, Show} import cats.data.{Kleisli, State, WriterT} import cats.effect.kernel._ import cats.free.FreeT @@ -41,10 +41,40 @@ object pure { implicit val eq: Eq[MaskId] = Eq.fromUniversalEquals[MaskId] } + private[pure] final case class MaskFrame(id: MaskId) + + private[pure] sealed trait CancelationSignal[E] + + private[pure] object CancelationSignal { + final case class External[E]() extends CancelationSignal[E] + final case class Self[E](finalizers: List[PureConc[E, Unit]]) extends CancelationSignal[E] + } + + private[pure] final class CancelationListenerId + + private[pure] object CancelationListenerId { + implicit val eq: Eq[CancelationListenerId] = + Eq.fromUniversalEquals[CancelationListenerId] + } + + private[pure] final case class CancelationListener[E]( + id: CancelationListenerId, + action: PureConc[E, Unit]) + + private sealed trait MaskUpdate + + private object MaskUpdate { + case object Removed extends MaskUpdate + case object Shadowed extends MaskUpdate + case object Absent extends MaskUpdate + } + final case class FiberCtx[E]( self: PureFiber[E, ?], masks: List[MaskId] = Nil, - finalizers: List[PureConc[E, Unit]] = Nil) + finalizers: List[PureConc[E, Unit]] = Nil, + selfCancelationBoundary: Option[Int] = None, + finalizing: Boolean = false) type ResolvedPC[E, A] = ThreadT[IdOC[E, *], A] @@ -69,8 +99,20 @@ object pure { val back = Kleisli.ask[IdOC[E, *], FiberCtx[E]] map { ctx => val checker = ctx .self - .realizeCancelation - .ifM(ApplicativeThread[PureConc[E, *]].done, ().pure[PureConc[E, *]]) + .isFinalizing + .ifM( + ().pure[PureConc[E, *]], + ctx + .self + .hasActivePoll + .ifM( + ().pure[PureConc[E, *]], + ctx + .self + .realizeCancelationWith(ctx) + .ifM( + ApplicativeThread[PureConc[E, *]].done[Unit], + ().pure[PureConc[E, *]]))) checker >> mvarLiftF(ThreadT.liftF(ka)) } @@ -97,46 +139,61 @@ object pure { type Main[X] = MVarR[ResolvedPC[E, *], X] MVar.empty[Main, Outcome[PureConc[E, *], E, A]].flatMap { state0 => - MVar.empty[Main, Unit] flatMap { canceled0 => - val state = state0[Main] - val fiber = new PureFiber[E, A](state0, canceled0) - - val identified = canceled mapF { ta => - val fk = new (FiberR[E, *] ~> IdOC[E, *]) { - def apply[a](ke: FiberR[E, a]) = - ke.run(FiberCtx(fiber)) - } - - ta.mapK(fk) - } - - import Outcome._ - - val body = identified flatMap { a => - state.tryPut(Succeeded(a.pure[PureConc[E, *]])) - } handleErrorWith { e => state.tryPut(Errored(e)) } - - val results = state.read.flatMap { - case Canceled() => (Outcome.Canceled(): IdOC[E, A]).pure[Main] - case Errored(e) => (Outcome.Errored(e): IdOC[E, A]).pure[Main] - - case Succeeded(fa) => - val identifiedCompletion = fa.mapF { ta => - val fk = new (FiberR[E, *] ~> IdOC[E, *]) { - def apply[a](ke: FiberR[E, a]) = - ke.run(FiberCtx(fiber)) + MVar.empty[Main, CancelationSignal[E]] flatMap { canceled0 => + MVar[Main, List[MaskFrame]](Nil) flatMap { masks => + MVar[Main, List[CancelationListener[E]]](Nil) flatMap { cancelationListeners => + MVar[Main, Boolean](false) flatMap { finalizing => + MVar[Main, Int](0) flatMap { activePolls => + val state = state0[Main] + val fiber = + new PureFiber[E, A]( + state0, + canceled0, + masks, + cancelationListeners, + finalizing, + activePolls) + + val identified = canceled mapF { ta => + val fk = new (FiberR[E, *] ~> IdOC[E, *]) { + def apply[a](ke: FiberR[E, a]) = + ke.run(FiberCtx(fiber)) + } + + ta.mapK(fk) + } + + import Outcome._ + + val body = identified flatMap { a => + state.tryPut(Succeeded(a.pure[PureConc[E, *]])) + } handleErrorWith { e => state.tryPut(Errored(e)) } + + val results = state.read.flatMap { + case Canceled() => (Outcome.Canceled(): IdOC[E, A]).pure[Main] + case Errored(e) => (Outcome.Errored(e): IdOC[E, A]).pure[Main] + + case Succeeded(fa) => + val identifiedCompletion = fa.mapF { ta => + val fk = new (FiberR[E, *] ~> IdOC[E, *]) { + def apply[a](ke: FiberR[E, a]) = + ke.run(FiberCtx(fiber)) + } + + ta.mapK(fk) + } + + identifiedCompletion.map(a => Succeeded[Id, E, A](a): IdOC[E, A]) handleError { + e => Errored(e) + } + } + + Kleisli.ask[ResolvedPC[E, *], MVar.Universe].map { u => + ApplicativeThread[ResolvedPC[E, *]].start(body.run(u)) >> results.run(u) + } } - - ta.mapK(fk) - } - - identifiedCompletion.map(a => Succeeded[Id, E, A](a): IdOC[E, A]) handleError { - e => Errored(e) } - } - - Kleisli.ask[ResolvedPC[E, *], MVar.Universe].map { u => - ApplicativeThread[ResolvedPC[E, *]].start(body.run(u)) >> results.run(u) + } } } } @@ -201,7 +258,9 @@ object pure { } def canceled: PureConc[E, Unit] = - Thread.annotate("canceled")(withCtx(_.self.cancelAndRealize.ifM(Thread.done, unit))) + Thread.annotate("canceled")(withCtx { ctx => + ctx.self.cancelAndRealizeWith(ctx).ifM(Thread.done, unit) + }) def cede: PureConc[E, Unit] = Thread.cede @@ -209,7 +268,7 @@ object pure { def never[A]: PureConc[E, A] = withCtx[E, A] { ctx => // we monitor for asynchronous cancelation. if we're masked, this won't cancel and we hang - Thread.annotate("never")(ctx.self.awaitCancelation *> Thread.done) + Thread.annotate("never")(ctx.self.awaitCancelationWith(ctx) *> Thread.done) } def ref[A](a: A): PureConc[E, Ref[PureConc[E, *], A]] = @@ -218,6 +277,11 @@ object pure { def deferred[A]: PureConc[E, Deferred[PureConc[E, *], A]] = MVar.empty[PureConc[E, *], A].flatMap(mVar => Kleisli.pure(unsafeDeferred(mVar))) + private[this] def interruptible[A]( + ctx: FiberCtx[E], + fa: PureConc[E, A]): PureConc[E, A] = + ctx.self.interruptible(ctx)(fa) + private def unsafeRef[A](mVar: MVar[A]): Ref[PureConc[E, *], A] = new Ref[PureConc[E, *], A] { override def get: PureConc[E, A] = mVar.read[PureConc[E, *]] @@ -275,16 +339,7 @@ object pure { new Deferred[PureConc[E, *], A] { override def get: PureConc[E, A] = withCtx { ctx => - // we need to race cancelation against reading the mvar - // if cancelation wins, we shut down the thread - MVar.empty[PureConc[E, *], Option[A]] flatMap { signal => - val left = Thread.start(ctx.self.awaitCancelation.ifM(signal.tryPut[PureConc[E, *]](None).void, unit)) - val right = Thread.start(mVar.read[PureConc[E, *]].flatMap(a => signal.tryPut[PureConc[E, *]](Some(a)))) - - left *> - right *> - signal.read[PureConc[E, *]].flatMap(_.map(_.pure[PureConc[E, *]]).getOrElse(Thread.done)) - } + interruptible(ctx, mVar.read[PureConc[E, *]]) } override def complete(a: A): PureConc[E, Boolean] = mVar.tryPut[PureConc[E, *]](a) @@ -295,32 +350,157 @@ object pure { def start[A](fa: PureConc[E, A]): PureConc[E, Fiber[PureConc[E, *], E, A]] = Thread.annotate("start", true) { MVar.empty[PureConc[E, *], Outcome[PureConc[E, *], E, A]].flatMap { state => - MVar.empty[PureConc[E, *], Unit] flatMap { canceled => - val fiber = new PureFiber[E, A](state, canceled) - - // the tryPut here is interesting: it encodes first-wins semantics on cancelation/completion - val body = guaranteeCase(fa)(state.tryPut[PureConc[E, *]](_).void) - val identified = localCtx(FiberCtx(fiber), body) - Thread.start(identified.attempt.void).as(fiber) + MVar.empty[PureConc[E, *], CancelationSignal[E]] flatMap { canceled => + MVar[PureConc[E, *], List[MaskFrame]](Nil) flatMap { masks => + MVar[PureConc[E, *], List[CancelationListener[E]]](Nil) flatMap { cancelationListeners => + MVar[PureConc[E, *], Boolean](false) flatMap { finalizing => + MVar[PureConc[E, *], Int](0) flatMap { activePolls => + val fiber = + new PureFiber[E, A]( + state, + canceled, + masks, + cancelationListeners, + finalizing, + activePolls) + + // the tryPut here is interesting: it encodes first-wins semantics on cancelation/completion + val body = guaranteeCase(fa)(state.tryPut[PureConc[E, *]](_).void) + val identified = localCtx(FiberCtx(fiber), body) + Thread.start(identified.attempt.void).as(fiber) + } + } + } + } } } } + override def racePair[A, B](fa: PureConc[E, A], fb: PureConc[E, B]): PureConc[ + E, + Either[ + (Outcome[PureConc[E, *], E, A], Fiber[PureConc[E, *], E, B]), + (Fiber[PureConc[E, *], E, A], Outcome[PureConc[E, *], E, B])]] = + uncancelable { poll => + for { + result <- deferred[Either[Outcome[PureConc[E, *], E, A], Outcome[PureConc[E, *], E, B]]] + + fibA <- start(fa) + fibB <- start(fb) + + _ <- Thread.start( + fibA.join.flatMap(oc => + result.complete( + Left(oc): Either[ + Outcome[PureConc[E, *], E, A], + Outcome[PureConc[E, *], E, B]]).void)) + _ <- Thread.start( + fibB.join.flatMap(oc => + result.complete( + Right(oc): Either[ + Outcome[PureConc[E, *], E, A], + Outcome[PureConc[E, *], E, B]]).void)) + + back <- onCancel( + poll(result.get), + for { + canA <- start(fibA.cancel) + canB <- start(fibB.cancel) + + _ <- canA.join + _ <- canB.join + } yield ()) + } yield back match { + case Left(oc) => Left((oc, fibB)) + case Right(oc) => Right((fibA, oc)) + } + } + def uncancelable[A](body: Poll[PureConc[E, *]] => PureConc[E, A]): PureConc[E, A] = Thread.annotate("uncancelable", true) { - val mask = new MaskId + withCtx { ctx => + val mask = new MaskId + val selfCancelationBoundary = + ctx.selfCancelationBoundary.getOrElse(ctx.finalizers.length) + + val self = ctx.self + def updateMasks[B](f: List[MaskFrame] => (List[MaskFrame], B)): PureConc[E, B] = + self.masks.read[PureConc[E, *]].flatMap { ms => + val (updated, b) = f(ms) + self.masks.swap[PureConc[E, *]](updated).as(b) + } - val poll = new Poll[PureConc[E, *]] { - def apply[a](fa: PureConc[E, a]) = - withCtx { ctx => - val ctx2 = ctx.copy(masks = ctx.masks.dropWhile(mask === _)) - localCtx(ctx2, fa.attempt <* ctx.self.realizeCancelation).rethrow + val addF = updateMasks(ms => (MaskFrame(mask) :: ms, ())) + val removeF = updateMasks { + case MaskFrame(`mask`) :: ms => (ms, MaskUpdate.Removed) + case ms if ms.exists(_.id === mask) => (ms, MaskUpdate.Shadowed) + case ms => (ms, MaskUpdate.Absent) + } + + def restore(update: MaskUpdate) = + update match { + case MaskUpdate.Removed => self.exitPoll *> addF + case MaskUpdate.Shadowed | MaskUpdate.Absent => unit } - } - withCtx { ctx => - val ctx2 = ctx.copy(masks = mask :: ctx.masks) - localCtx(ctx2, body(poll)) + val poll = new Poll[PureConc[E, *]] { + def apply[a](fa: PureConc[E, a]) = + withCtx { callCtx => + if (callCtx.self eq self) + removeF.flatMap { update => + val restoreF = restore(update) + val pollCtx = update match { + case MaskUpdate.Removed => + callCtx.copy( + selfCancelationBoundary = + callCtx.selfCancelationBoundary.orElse(Some(selfCancelationBoundary))) + + case MaskUpdate.Shadowed | MaskUpdate.Absent => + callCtx + } + + val enterF = update match { + case MaskUpdate.Removed => self.enterPoll + case MaskUpdate.Shadowed | MaskUpdate.Absent => unit + } + + enterF *> + localCtx( + pollCtx, + onCancel( + self + .realizeCancelationWith(pollCtx) + .ifM( + Thread.done, + fa.attempt.flatMap { result => + self + .realizeCancelationWith(pollCtx) + .ifM( + Thread.done, + restoreF *> result.pure[PureConc[E, *]].rethrow) + }), + restoreF)) + } + else fa + } + } + + val runBody = + addF *> body(poll).attempt.flatMap { result => + removeF.flatMap { + case MaskUpdate.Removed => + val back = result.pure[PureConc[E, *]].rethrow + + self.hasActivePoll.ifM( + back, + self.realizeCancelationWith(ctx).ifM(Thread.done, back)) + + case MaskUpdate.Shadowed | MaskUpdate.Absent => + result.pure[PureConc[E, *]].rethrow + } + } + + onCancel(runBody, removeF.void) } } @@ -329,7 +509,7 @@ object pure { Defer[PureConc[E, *]].defer(pure(new Unique.Token())) def forceR[A, B](fa: PureConc[E, A])(fb: PureConc[E, B]): PureConc[E, B] = - Thread.annotate("forceR")(productR(attempt(fa))(fb)) + Thread.annotate("forceR")(productR(handleError(fa.void)(_ => ()))(fb)) def flatMap[A, B](fa: PureConc[E, A])(f: A => PureConc[E, B]): PureConc[E, B] = M.flatMap(fa)(f) @@ -379,55 +559,246 @@ object pure { // todo: MVar is not Serializable, release then update here final class PureFiber[E, A]( val state0: MVar[Outcome[PureConc[E, *], E, A]], - val canceled0: MVar[Unit]) + val canceled0: MVar[CancelationSignal[E]], + val masks: MVar[List[MaskFrame]], + val cancelationListeners: MVar[List[CancelationListener[E]]], + val finalizing: MVar[Boolean], + val activePolls: MVar[Int]) extends Fiber[PureConc[E, *], E, A] with Serializable { + def this( + state0: MVar[Outcome[PureConc[E, *], E, A]], + canceled0: MVar[CancelationSignal[E]], + masks: MVar[List[MaskFrame]]) = + this(state0, canceled0, masks, null, null, null) + + def this( + state0: MVar[Outcome[PureConc[E, *], E, A]], + canceled0: MVar[CancelationSignal[E]]) = + this(state0, canceled0, null, null, null, null) + private[this] val state = state0[PureConc[E, *]] + private[pure] val hasActivePoll: PureConc[E, Boolean] = + if (activePolls eq null) false.pure[PureConc[E, *]] + else activePolls.read[PureConc[E, *]].map(_ > 0) + + private[pure] val enterPoll: PureConc[E, Unit] = + if (activePolls eq null) ().pure[PureConc[E, *]] + else + activePolls.read[PureConc[E, *]].flatMap { n => + activePolls.swap[PureConc[E, *]](n + 1).void + } + + private[pure] val exitPoll: PureConc[E, Unit] = + if (activePolls eq null) ().pure[PureConc[E, *]] + else + activePolls.read[PureConc[E, *]].flatMap { n => + activePolls.swap[PureConc[E, *]]((n - 1) max 0).void + } + private[pure] val canceled: PureConc[E, Boolean] = canceled0.tryRead[PureConc[E, *]].map(_.as(true).getOrElse(false)) - private[pure] val realizeCancelation: PureConc[E, Boolean] = - withCtx { ctx => - val checkM = ctx.masks.isEmpty.pure[PureConc[E, *]] - - checkM.ifM( - canceled.ifM( - // if unmasked and canceled, finalize - allocateForPureConc[E] uncancelable { _ => - ctx.finalizers.sequence_.as(true) <* state0.tryPut[PureConc[E, *]]( - Outcome.Canceled()) - }, - // if unmasked but not canceled, ignore - false.pure[PureConc[E, *]] - ), + private[pure] def registerCancelationListener( + notify: PureConc[E, Unit]): PureConc[E, CancelationListenerId] = { + val id = new CancelationListenerId + + if (cancelationListeners eq null) id.pure[PureConc[E, *]] + else + cancelationListeners.read[PureConc[E, *]].flatMap { listeners => + cancelationListeners + .swap[PureConc[E, *]](CancelationListener(id, notify) :: listeners) + .as(id) + } + } + + private[pure] def removeCancelationListener(id: CancelationListenerId): PureConc[E, Unit] = + if (cancelationListeners eq null) ().pure[PureConc[E, *]] + else + cancelationListeners.read[PureConc[E, *]].flatMap { listeners => + cancelationListeners + .swap[PureConc[E, *]](listeners.filterNot(_.id === id)) + .void + } + + private[this] def notifyCancelationListeners: PureConc[E, Unit] = + if (cancelationListeners eq null) ().pure[PureConc[E, *]] + else cancelationListeners.swap[PureConc[E, *]](Nil).flatMap(_.traverse_(_.action)) + + private[pure] def interruptible[B](ctx: FiberCtx[E])( + fb: PureConc[E, B]): PureConc[E, B] = { + val Thread = ApplicativeThread[PureConc[E, *]] + + ctx.self.masks.read[PureConc[E, *]].flatMap { + case Nil => + MVar.empty[PureConc[E, *], Option[B]].flatMap { signal => + val notifyCancelation = signal.tryPut[PureConc[E, *]](None).void + + ctx.self.registerCancelationListener(notifyCancelation).flatMap { listener => + val awaitCompletion = + Thread.start(fb.flatMap(b => signal.tryPut[PureConc[E, *]](Some(b)).void)) + + val checkCancelation = + signal.tryRead[PureConc[E, *]].flatMap { + case Some(_) => ().pure[PureConc[E, *]] + case None => + ctx.self + .realizeCancelationWith(ctx) + .ifM(notifyCancelation, ().pure[PureConc[E, *]]) + } + + awaitCompletion *> + checkCancelation *> + signal.read[PureConc[E, *]].flatMap { + case Some(b) => + ctx.self.removeCancelationListener(listener).as(b) + + case None => + ctx.self.removeCancelationListener(listener) *> + ctx.self.realizeCancelationWith(ctx) *> + Thread.done + } + } + } + + case _ => + fb + } + } + + private[pure] val isFinalizing: PureConc[E, Boolean] = + if (finalizing eq null) false.pure[PureConc[E, *]] + else finalizing.read[PureConc[E, *]] + + private[this] def setFinalizing(value: Boolean): PureConc[E, Unit] = + if (finalizing eq null) ().pure[PureConc[E, *]] + else finalizing.swap[PureConc[E, *]](value).void + + private[this] def finalizeWith( + ctx: FiberCtx[E], + finalizers: List[PureConc[E, Unit]]): PureConc[E, Boolean] = + localCtx( + ctx.copy(finalizers = Nil, finalizing = true), + allocateForPureConc[E].uncancelable(_ => finalizers.sequence_) *> + (state0.tryPut[PureConc[E, *]](Outcome.Canceled()).flatMap { + case true => true.pure[PureConc[E, *]] + case false => + state.read.map { + case Outcome.Canceled() => true + case _ => false + } + } <* setFinalizing(false))) + + private[this] def cancelationFinalizers( + signal: CancelationSignal[E], + ctx: FiberCtx[E]): List[PureConc[E, Unit]] = { + signal match { + case CancelationSignal.External() => ctx.finalizers + case CancelationSignal.Self(Nil) if ctx.selfCancelationBoundary.nonEmpty => + selfCancelationFinalizers(ctx) + case CancelationSignal.Self(finalizers) => finalizers + } + } + + private[this] def selfCancelationFinalizers(ctx: FiberCtx[E]): List[PureConc[E, Unit]] = + ctx.selfCancelationBoundary match { + case Some(boundary) => ctx.finalizers.take((ctx.finalizers.length - boundary) max 0) + case None => ctx.finalizers + } + + private[this] def whileFinalizing[B](ctx: FiberCtx[E])(fb: PureConc[E, B]): PureConc[E, B] = + localCtx(ctx.copy(finalizers = Nil, finalizing = true), setFinalizing(true) *> fb) + + private[this] def finalizationOutcome: PureConc[E, Boolean] = + state.read.map { + case Outcome.Canceled() => true + case _ => false + } + + private[this] def realizeCancelationWithSignal( + ctx: FiberCtx[E], + signal: CancelationSignal[E]): PureConc[E, Boolean] = + whileFinalizing(ctx)(finalizeWith(ctx, cancelationFinalizers(signal, ctx))) + + private[pure] def realizeCancelationWith(ctx: FiberCtx[E]): PureConc[E, Boolean] = + if (ctx.finalizing) false.pure[PureConc[E, *]] + else isFinalizing.ifM( + finalizationOutcome, + ctx.self.masks.read[PureConc[E, *]].map(_.isEmpty).ifM( + canceled0.tryRead[PureConc[E, *]].flatMap { + case Some(signal) => + // if unmasked and canceled, finalize + realizeCancelationWithSignal(ctx, signal) + + case None => + // if unmasked but not canceled, ignore + false.pure[PureConc[E, *]] + }, // if masked, ignore cancelation state but retain until unmasked false.pure[PureConc[E, *]] - ) + )) + + private[pure] val realizeCancelation: PureConc[E, Boolean] = + withCtx(realizeCancelationWith) + + private[pure] def awaitCancelationWith(ctx: FiberCtx[E]): PureConc[E, Boolean] = { + def blocked = + MVar.empty[PureConc[E, *], Unit].flatMap(_.read[PureConc[E, *]]).as(false) + + if (ctx.finalizing) blocked + else ctx.self.masks.read[PureConc[E, *]].flatMap { + case Nil => + isFinalizing.ifM( + canceled0.tryRead[PureConc[E, *]].map(_.isEmpty), + canceled0.read[PureConc[E, *]].flatMap(realizeCancelationWithSignal(ctx, _))) + + case _ => + isFinalizing.ifM( + canceled0.tryRead[PureConc[E, *]].map(_.isEmpty), + blocked) } + } private[pure] val awaitCancelation: PureConc[E, Boolean] = - canceled0.read[PureConc[E, *]] *> realizeCancelation + withCtx(awaitCancelationWith) + + private[pure] def cancelAndRealizeWith(ctx: FiberCtx[E]): PureConc[E, Boolean] = + if (ctx.finalizing) false.pure[PureConc[E, *]] + else isFinalizing.ifM( + ctx.self.masks.read[PureConc[E, *]].map(_.isEmpty), + ctx.self.masks.read[PureConc[E, *]].map(_.isEmpty).ifM( + whileFinalizing(ctx) { + val finalizers = selfCancelationFinalizers(ctx) + + canceled0.tryPut[PureConc[E, *]](CancelationSignal.Self(finalizers)).flatMap { + case true => notifyCancelationListeners *> finalizeWith(ctx, finalizers) + case false => + canceled0 + .tryRead[PureConc[E, *]] + .flatMap( + _.fold(finalizeWith(ctx, finalizers))(signal => + finalizeWith(ctx, cancelationFinalizers(signal, ctx)))) + } + }, + canceled0 + .tryPut[PureConc[E, *]](CancelationSignal.Self(Nil)) + .flatMap(inserted => if (inserted) notifyCancelationListeners.as(false) else false.pure[PureConc[E, *]]))) private[pure] val cancelAndRealize: PureConc[E, Boolean] = - canceled0.tryPut[PureConc[E, *]](()) *> realizeCancelation + withCtx(cancelAndRealizeWith) val join: PureConc[E, Outcome[PureConc[E, *], E, A]] = { - val Thread = ApplicativeThread[PureConc[E, *]] - withCtx { ctx => - // this is exactly like Deferred#get - MVar.empty[PureConc[E, *], Option[Outcome[PureConc[E, *], E, A]]] flatMap { signal => - // note we must read our *own* canceled, not the target fiber's - val left = Thread.start(ctx.self.awaitCancelation.ifM(signal.tryPut[PureConc[E, *]](None).void, Applicative[PureConc[E, *]].unit)) - val right = Thread.start(state.read.flatMap(oc => signal.tryPut[PureConc[E, *]](Some(oc)))) - - left *> right *> signal.read[PureConc[E, *]].flatMap(_.map(_.pure[PureConc[E, *]]).getOrElse(Thread.done)) - } + ctx.self.interruptible(ctx)(state.read) } } - val cancel: PureConc[E, Unit] = canceled0.tryPut[PureConc[E, *]](()) *> join.void + val cancel: PureConc[E, Unit] = + canceled0.tryPut[PureConc[E, *]](CancelationSignal.External()).flatMap { + case true => notifyCancelationListeners + case false => ().pure[PureConc[E, *]] + } *> join.void } } diff --git a/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala b/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala index 45db4b45db..67db137caf 100644 --- a/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala +++ b/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala @@ -190,6 +190,62 @@ class PureConcSuite extends DisciplineSuite with BaseSuite { assertEquals(forked, Outcome.Succeeded[Option, Int, Unit](None)) } + test("ignore poll from another fiber") { + val t = for { + started <- F.deferred[Unit] + polled <- F.deferred[Unit] + + parent <- F.start { + F.uncancelable { poll => + started.complete(()) *> + F.start(poll(polled.complete(()) *> F.never[Unit])).void *> + polled.get *> + F.never[Unit] + } + } + + _ <- started.get + _ <- polled.get + _ <- parent.cancel + } yield () + + assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, Unit](None)) + } + + test("run finalizers around a self-canceling polled region") { + val t = for { + finalized <- F.ref(0) + fiber <- F.start { + F.uncancelable { poll => + F.onCancel(poll(F.canceled), finalized.update(_ + 1)) + } + } + _ <- fiber.join + back <- finalized.get + } yield back + + assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, Int](Some(1))) + } + + test("observe pending self-cancel before running a polled region") { + val t = for { + finalized <- F.ref(0) + ran <- F.ref(false) + fiber <- F.start { + F.uncancelable { poll => + F.canceled *> F.onCancel(poll(ran.set(true)), finalized.update(_ + 1)) + } + } + _ <- fiber.join + fin <- finalized.get + body <- ran.get + } yield (fin, body) + + assertEquals( + pure.run(t), + Outcome.Succeeded[Option, Int, (Int, Boolean)](Some((1, false)))) + } + test("implement locals via Kleisli and FreeT") { import cats.{~>, Eval, Id} import cats.data.Kleisli From 0f8e6969b9d0a2f8738946c846582064cae6ad1f Mon Sep 17 00:00:00 2001 From: Daniel Spiewak Date: Mon, 22 Jun 2026 00:13:25 -0500 Subject: [PATCH 04/10] Fix PureConc cancellation laws --- .../kernel/testkit/PureConcGenerators.scala | 4 +- .../cats/effect/kernel/testkit/TimeT.scala | 87 ++++++++- .../cats/effect/kernel/testkit/pure.scala | 170 +++++++++++------ .../cats/effect/laws/PureConcSuite.scala | 178 +++++++++++++++++- 4 files changed, 362 insertions(+), 77 deletions(-) diff --git a/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/PureConcGenerators.scala b/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/PureConcGenerators.scala index e35a7e5e68..9941847536 100644 --- a/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/PureConcGenerators.scala +++ b/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/PureConcGenerators.scala @@ -40,8 +40,8 @@ object PureConcGenerators { super .recursiveGen[B](deeper) .filterNot( - _._1 == "racePair" - ) // remove the racePair generator since it reifies nondeterminism, which cannot be law-tested + gen => gen._1 == "racePair" || gen._1 == "join" + ) // remove generators which reify nondeterminism and cannot be law-tested } implicit def arbitraryPureConc[E: Arbitrary: Cogen, A: Arbitrary: Cogen] diff --git a/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/TimeT.scala b/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/TimeT.scala index baa42a033d..0702e442e9 100644 --- a/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/TimeT.scala +++ b/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/TimeT.scala @@ -18,7 +18,7 @@ package cats.effect package kernel package testkit -import cats.{~>, Group, Monad, Monoid, Order} +import cats.{~>, Eq, Group, Monad, Monoid, Order} import cats.data.Kleisli import cats.syntax.all._ @@ -90,6 +90,9 @@ private[effect] object TimeT { a.map(_.inverse()) } + implicit def eqTimeT[F[_], A](implicit FA: Eq[F[A]]): Eq[TimeT[F, A]] = + Eq.by(TimeT.run) + implicit def orderTimeT[F[_], A](implicit FA: Order[F[A]]): Order[TimeT[F, A]] = Order.by(TimeT.run) @@ -111,15 +114,70 @@ private[effect] object TimeT { val forkA = time.fork() val forkB = time.fork() - // TODO this doesn't work (yet) because we need to force the "faster" effect to win the race, which right now isn't happening - F.racePair(fa.run(forkA), fb.run(forkB)).map { + def liftOutcome[C](oc: Outcome[F, E, C]): Outcome[TimeT[F, *], E, C] = + oc.mapK(TimeT.liftK[F]) + + F.racePair(fa.run(forkA), fb.run(forkB)).flatMap { case Left((oca, delegate)) => - time.now = forkA.now - Left((oca.mapK(TimeT.liftK[F]), fiberize(forkB, delegate))) + F.onCancel(F.race(delegate.join, F.cede), delegate.cancel).map { + case Left(ocb) if forkB.now < forkA.now => + time.now = forkB.now + Right((completedFiber(forkA, liftOutcome(oca)), liftOutcome(ocb))) + + case _ => + time.now = forkA.now + Left((liftOutcome(oca), fiberize(forkB, delegate))) + } case Right((delegate, ocb)) => - time.now = forkB.now - Right((fiberize(forkA, delegate), ocb.mapK(TimeT.liftK[F]))) + F.onCancel(F.race(delegate.join, F.cede), delegate.cancel).map { + case Left(oca) if forkA.now < forkB.now => + time.now = forkA.now + Left((liftOutcome(oca), completedFiber(forkB, liftOutcome(ocb)))) + + case _ => + time.now = forkB.now + Right((fiberize(forkA, delegate), liftOutcome(ocb))) + } + } + } + + override def race[A, B](fa: TimeT[F, A], fb: TimeT[F, B]): TimeT[F, Either[A, B]] = + uncancelable { poll => + poll(racePair(fa, fb)).flatMap { + case Left((oc, f)) => + oc match { + case Outcome.Succeeded(fa) => f.cancel *> fa.map(Left(_)) + case Outcome.Errored(ea) => f.cancel *> raiseError(ea) + case Outcome.Canceled() => + f.cancel *> poll(f.join) flatMap { + case Outcome.Succeeded(fb) => fb.map(Right(_)) + case Outcome.Errored(eb) => raiseError(eb) + case Outcome.Canceled() => poll(canceled) *> never + } + } + + case Right((f, oc)) => + oc match { + case Outcome.Succeeded(fb) => f.cancel *> fb.map(Right(_)) + case Outcome.Errored(eb) => f.cancel *> raiseError(eb) + case Outcome.Canceled() => + f.cancel *> poll(f.join) flatMap { + case Outcome.Succeeded(fa) => fa.map(Left(_)) + case Outcome.Errored(ea) => raiseError(ea) + case Outcome.Canceled() => poll(canceled) *> never + } + } + } + } + + override def raceOutcome[A, B](fa: TimeT[F, A], fb: TimeT[F, B]): TimeT[ + F, + Either[Outcome[TimeT[F, *], E, A], Outcome[TimeT[F, *], E, B]]] = + uncancelable { poll => + poll(racePair(fa, fb)).flatMap { + case Left((oc, f)) => f.cancel.as(Left(oc)) + case Right((f, oc)) => f.cancel.as(Right(oc)) } } @@ -156,5 +214,20 @@ private[effect] object TimeT { } } } + + private[this] def completedFiber[A]( + forked: Time, + outcome: Outcome[TimeT[F, *], E, A]): Fiber[TimeT[F, *], E, A] = + new Fiber[TimeT[F, *], E, A] { + + val cancel = + unit + + val join = + Kleisli { outerTime => + outerTime.now = outerTime.now.max(forked.now) + F.pure(outcome) + } + } } } diff --git a/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala b/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala index 925b0b236f..0341458c8d 100644 --- a/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala +++ b/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala @@ -69,12 +69,42 @@ object pure { case object Absent extends MaskUpdate } - final case class FiberCtx[E]( + sealed case class FiberCtx[E]( self: PureFiber[E, ?], masks: List[MaskId] = Nil, - finalizers: List[PureConc[E, Unit]] = Nil, - selfCancelationBoundary: Option[Int] = None, - finalizing: Boolean = false) + finalizers: List[PureConc[E, Unit]] = Nil) { + + private[pure] def selfCancelationBoundary: Option[Int] = None + + private[pure] def finalizing: Boolean = false + + private[pure] def withFinalizers(value: List[PureConc[E, Unit]]): FiberCtx[E] = + FiberCtx.internal(self, masks, value, selfCancelationBoundary, finalizing) + + private[pure] def withSelfCancelationBoundary(value: Option[Int]): FiberCtx[E] = + FiberCtx.internal(self, masks, finalizers, value, finalizing) + + private[pure] def withFinalizing(value: Boolean): FiberCtx[E] = + FiberCtx.internal(self, masks, finalizers, selfCancelationBoundary, value) + } + + object FiberCtx { + private final class Internal[E]( + self: PureFiber[E, ?], + masks: List[MaskId], + finalizers: List[PureConc[E, Unit]], + override private[pure] val selfCancelationBoundary: Option[Int], + override private[pure] val finalizing: Boolean) + extends FiberCtx[E](self, masks, finalizers) + + private[pure] def internal[E]( + self: PureFiber[E, ?], + masks: List[MaskId], + finalizers: List[PureConc[E, Unit]], + selfCancelationBoundary: Option[Int], + finalizing: Boolean): FiberCtx[E] = + new Internal(self, masks, finalizers, selfCancelationBoundary, finalizing) + } type ResolvedPC[E, A] = ThreadT[IdOC[E, *], A] @@ -106,7 +136,12 @@ object pure { .self .hasActivePoll .ifM( - ().pure[PureConc[E, *]], + ctx + .self + .realizeExternalCancelationWith(ctx) + .ifM( + ApplicativeThread[PureConc[E, *]].done[Unit], + ().pure[PureConc[E, *]]), ctx .self .realizeCancelationWith(ctx) @@ -252,7 +287,7 @@ object pure { def onCancel[A](fa: PureConc[E, A], fin: PureConc[E, Unit]): PureConc[E, A] = Thread.annotate("onCancel", true) { withCtx[E, A] { ctx => - val ctx2 = ctx.copy(finalizers = fin.attempt.void :: ctx.finalizers) + val ctx2 = ctx.withFinalizers(fin.attempt.void :: ctx.finalizers) localCtx(ctx2, fa) } } @@ -263,7 +298,10 @@ object pure { }) def cede: PureConc[E, Unit] = - Thread.cede + withCtx { ctx => + Thread.cede *> + ctx.self.realizeExternalCancelationWith(ctx).ifM(Thread.done, unit) + } def never[A]: PureConc[E, A] = withCtx[E, A] { ctx => @@ -388,13 +426,13 @@ object pure { fibA <- start(fa) fibB <- start(fb) - _ <- Thread.start( + _ <- start( fibA.join.flatMap(oc => result.complete( Left(oc): Either[ Outcome[PureConc[E, *], E, A], Outcome[PureConc[E, *], E, B]]).void)) - _ <- Thread.start( + _ <- start( fibB.join.flatMap(oc => result.complete( Right(oc): Either[ @@ -451,9 +489,8 @@ object pure { val restoreF = restore(update) val pollCtx = update match { case MaskUpdate.Removed => - callCtx.copy( - selfCancelationBoundary = - callCtx.selfCancelationBoundary.orElse(Some(selfCancelationBoundary))) + callCtx.withSelfCancelationBoundary( + callCtx.selfCancelationBoundary.orElse(Some(selfCancelationBoundary))) case MaskUpdate.Shadowed | MaskUpdate.Absent => callCtx @@ -492,7 +529,7 @@ object pure { val back = result.pure[PureConc[E, *]].rethrow self.hasActivePoll.ifM( - back, + self.realizeSelfCancelationWith(ctx).ifM(Thread.done, back), self.realizeCancelationWith(ctx).ifM(Thread.done, back)) case MaskUpdate.Shadowed | MaskUpdate.Absent => @@ -559,27 +596,23 @@ object pure { // todo: MVar is not Serializable, release then update here final class PureFiber[E, A]( val state0: MVar[Outcome[PureConc[E, *], E, A]], - val canceled0: MVar[CancelationSignal[E]], - val masks: MVar[List[MaskFrame]], - val cancelationListeners: MVar[List[CancelationListener[E]]], - val finalizing: MVar[Boolean], - val activePolls: MVar[Int]) + private[this] val canceled0: MVar[CancelationSignal[E]], + private[pure] val masks: MVar[List[MaskFrame]], + private[this] val cancelationListeners: MVar[List[CancelationListener[E]]], + private[this] val finalizing: MVar[Boolean], + private[this] val activePolls: MVar[Int]) extends Fiber[PureConc[E, *], E, A] with Serializable { - def this( - state0: MVar[Outcome[PureConc[E, *], E, A]], - canceled0: MVar[CancelationSignal[E]], - masks: MVar[List[MaskFrame]]) = - this(state0, canceled0, masks, null, null, null) - - def this( - state0: MVar[Outcome[PureConc[E, *], E, A]], - canceled0: MVar[CancelationSignal[E]]) = - this(state0, canceled0, null, null, null, null) + def this(state0: MVar[Outcome[PureConc[E, *], E, A]]) = + this(state0, null, null, null, null, null) private[this] val state = state0[PureConc[E, *]] + private[pure] val currentMasks: PureConc[E, List[MaskFrame]] = + if (masks eq null) List.empty[MaskFrame].pure[PureConc[E, *]] + else masks.read[PureConc[E, *]] + private[pure] val hasActivePoll: PureConc[E, Boolean] = if (activePolls eq null) false.pure[PureConc[E, *]] else activePolls.read[PureConc[E, *]].map(_ > 0) @@ -598,9 +631,6 @@ object pure { activePolls.swap[PureConc[E, *]]((n - 1) max 0).void } - private[pure] val canceled: PureConc[E, Boolean] = - canceled0.tryRead[PureConc[E, *]].map(_.as(true).getOrElse(false)) - private[pure] def registerCancelationListener( notify: PureConc[E, Unit]): PureConc[E, CancelationListenerId] = { val id = new CancelationListenerId @@ -631,7 +661,7 @@ object pure { fb: PureConc[E, B]): PureConc[E, B] = { val Thread = ApplicativeThread[PureConc[E, *]] - ctx.self.masks.read[PureConc[E, *]].flatMap { + ctx.self.currentMasks.flatMap { case Nil => MVar.empty[PureConc[E, *], Option[B]].flatMap { signal => val notifyCancelation = signal.tryPut[PureConc[E, *]](None).void @@ -680,7 +710,7 @@ object pure { ctx: FiberCtx[E], finalizers: List[PureConc[E, Unit]]): PureConc[E, Boolean] = localCtx( - ctx.copy(finalizers = Nil, finalizing = true), + ctx.withFinalizers(Nil).withFinalizing(true), allocateForPureConc[E].uncancelable(_ => finalizers.sequence_) *> (state0.tryPut[PureConc[E, *]](Outcome.Canceled()).flatMap { case true => true.pure[PureConc[E, *]] @@ -709,7 +739,7 @@ object pure { } private[this] def whileFinalizing[B](ctx: FiberCtx[E])(fb: PureConc[E, B]): PureConc[E, B] = - localCtx(ctx.copy(finalizers = Nil, finalizing = true), setFinalizing(true) *> fb) + localCtx(ctx.withFinalizers(Nil).withFinalizing(true), setFinalizing(true) *> fb) private[this] def finalizationOutcome: PureConc[E, Boolean] = state.read.map { @@ -722,33 +752,43 @@ object pure { signal: CancelationSignal[E]): PureConc[E, Boolean] = whileFinalizing(ctx)(finalizeWith(ctx, cancelationFinalizers(signal, ctx))) - private[pure] def realizeCancelationWith(ctx: FiberCtx[E]): PureConc[E, Boolean] = + private[this] def realizeCancelationIf(ctx: FiberCtx[E])( + accepts: CancelationSignal[E] => Boolean): PureConc[E, Boolean] = if (ctx.finalizing) false.pure[PureConc[E, *]] else isFinalizing.ifM( finalizationOutcome, - ctx.self.masks.read[PureConc[E, *]].map(_.isEmpty).ifM( + ctx.self.currentMasks.map(_.isEmpty).ifM( canceled0.tryRead[PureConc[E, *]].flatMap { - case Some(signal) => - // if unmasked and canceled, finalize + case Some(signal) if accepts(signal) => realizeCancelationWithSignal(ctx, signal) - case None => - // if unmasked but not canceled, ignore + case Some(_) | None => false.pure[PureConc[E, *]] }, - // if masked, ignore cancelation state but retain until unmasked false.pure[PureConc[E, *]] )) - private[pure] val realizeCancelation: PureConc[E, Boolean] = - withCtx(realizeCancelationWith) + private[pure] def realizeExternalCancelationWith(ctx: FiberCtx[E]): PureConc[E, Boolean] = + realizeCancelationIf(ctx) { + case CancelationSignal.External() => true + case _ => false + } + + private[pure] def realizeSelfCancelationWith(ctx: FiberCtx[E]): PureConc[E, Boolean] = + realizeCancelationIf(ctx) { + case CancelationSignal.Self(_) => true + case _ => false + } + + private[pure] def realizeCancelationWith(ctx: FiberCtx[E]): PureConc[E, Boolean] = + realizeCancelationIf(ctx)((_: CancelationSignal[E]) => true) private[pure] def awaitCancelationWith(ctx: FiberCtx[E]): PureConc[E, Boolean] = { def blocked = MVar.empty[PureConc[E, *], Unit].flatMap(_.read[PureConc[E, *]]).as(false) if (ctx.finalizing) blocked - else ctx.self.masks.read[PureConc[E, *]].flatMap { + else ctx.self.currentMasks.flatMap { case Nil => isFinalizing.ifM( canceled0.tryRead[PureConc[E, *]].map(_.isEmpty), @@ -761,18 +801,17 @@ object pure { } } - private[pure] val awaitCancelation: PureConc[E, Boolean] = - withCtx(awaitCancelationWith) - private[pure] def cancelAndRealizeWith(ctx: FiberCtx[E]): PureConc[E, Boolean] = if (ctx.finalizing) false.pure[PureConc[E, *]] else isFinalizing.ifM( - ctx.self.masks.read[PureConc[E, *]].map(_.isEmpty), - ctx.self.masks.read[PureConc[E, *]].map(_.isEmpty).ifM( + ctx.self.currentMasks.map(_.isEmpty), + ctx.self.currentMasks.map(_.isEmpty).ifM( whileFinalizing(ctx) { - val finalizers = selfCancelationFinalizers(ctx) + val finalizers = + selfCancelationFinalizers(ctx) + val selfCancelation = CancelationSignal.Self(finalizers) - canceled0.tryPut[PureConc[E, *]](CancelationSignal.Self(finalizers)).flatMap { + canceled0.tryPut[PureConc[E, *]](selfCancelation).flatMap { case true => notifyCancelationListeners *> finalizeWith(ctx, finalizers) case false => canceled0 @@ -786,19 +825,26 @@ object pure { .tryPut[PureConc[E, *]](CancelationSignal.Self(Nil)) .flatMap(inserted => if (inserted) notifyCancelationListeners.as(false) else false.pure[PureConc[E, *]]))) - private[pure] val cancelAndRealize: PureConc[E, Boolean] = - withCtx(cancelAndRealizeWith) - - val join: PureConc[E, Outcome[PureConc[E, *], E, A]] = { - withCtx { ctx => - ctx.self.interruptible(ctx)(state.read) + val join: PureConc[E, Outcome[PureConc[E, *], E, A]] = + if (canceled0 eq null) state.read + else { + withCtx { ctx => + ctx.self.interruptible(ctx)(state.read) + } } - } val cancel: PureConc[E, Unit] = - canceled0.tryPut[PureConc[E, *]](CancelationSignal.External()).flatMap { - case true => notifyCancelationListeners - case false => ().pure[PureConc[E, *]] - } *> join.void + if (canceled0 eq null) state.tryPut(Outcome.Canceled()).void + else + allocateForPureConc[E].uncancelable { _ => + state.tryRead.flatMap { + case Some(_) => ().pure[PureConc[E, *]] + case None => + canceled0.tryPut[PureConc[E, *]](CancelationSignal.External()).flatMap { + case true => notifyCancelationListeners + case false => ().pure[PureConc[E, *]] + } *> state.read.void + } + } } } diff --git a/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala b/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala index 67db137caf..ed4cd520d2 100644 --- a/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala +++ b/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala @@ -17,8 +17,9 @@ package cats.effect package laws +import cats.{Eq, Order} import cats.effect.kernel.testkit.{pure, OutcomeGenerators, PureConcGenerators, TimeT} -import cats.effect.kernel.testkit.TimeT._ +import cats.effect.kernel.testkit.TimeT.{eqTimeT => _, orderTimeT => _, _} import cats.effect.kernel.testkit.pure._ import cats.laws.discipline.arbitrary._ @@ -28,16 +29,27 @@ import scala.concurrent.duration._ import munit.DisciplineSuite -class PureConcSuite extends DisciplineSuite with BaseSuite { +private[laws] trait PureConcSuiteLowPriorityTimeTInstances { + implicit def orderTimeTPureConcFiniteDuration( + implicit FA: Order[PureConc[Int, FiniteDuration]]): Order[ + TimeT[PureConc[Int, *], FiniteDuration]] = + TimeT.orderTimeT +} + +class PureConcSuite + extends DisciplineSuite + with BaseSuite + with PureConcSuiteLowPriorityTimeTInstances { import PureConcGenerators._ import OutcomeGenerators._ - override def scalaCheckInitialSeed = - "ogn64yom4GXCEX0mXdqSfsqSeJxI2RbPUFC5YkvDtzD=" - implicit def exec(fb: TimeT[PureConc[Int, *], Boolean]): Prop = Prop(pure.run(TimeT.run(fb)).fold(false, _ => false, _.getOrElse(false))) + implicit def eqTimeTPureConc[A](implicit FA: Eq[PureConc[Int, A]]): Eq[ + TimeT[PureConc[Int, *], A]] = + TimeT.eqTimeT + { import cats.effect.kernel.{GenConcurrent, Outcome} import cats.effect.kernel.implicits._ @@ -86,7 +98,7 @@ class PureConcSuite extends DisciplineSuite with BaseSuite { } { - import cats.effect.kernel.{GenConcurrent, Outcome} + import cats.effect.kernel.{GenConcurrent, GenTemporal, Outcome} import cats.effect.kernel.implicits._ import cats.syntax.all._ @@ -169,6 +181,22 @@ class PureConcSuite extends DisciplineSuite with BaseSuite { assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, Unit](None)) } + test("hang when canceling fiber blocked on cancel finalization") { + val t = for { + targetStarted <- F.deferred[Unit] + finalizerStarted <- F.deferred[Unit] + target <- F.start( + (targetStarted.complete(()) *> F.never[Unit]) + .onCancel(finalizerStarted.complete(()) *> F.never[Unit])) + _ <- targetStarted.get + canceler <- F.start(target.cancel) + _ <- finalizerStarted.get + _ <- canceler.cancel + } yield () + + assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, Unit](None)) + } + test("run finalizers in order") { val t = for { results <- F.ref[String]("") @@ -182,6 +210,36 @@ class PureConcSuite extends DisciplineSuite with BaseSuite { assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, String](Some("AB"))) } + test("ignore cancelation of a fiber after racePair has completed") { + val t = for { + finalized <- F.ref(0) + fiber <- F.start { + F.racePair(F.unit, F.never[Unit]).void.onCancel(finalized.update(_ + 1)) + } + _ <- fiber.join + _ <- fiber.cancel + _ <- F.cede + back <- finalized.get + } yield back + + assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, Int](Some(0))) + } + + test("ignore cancelation of a fiber after race has completed") { + val t = for { + finalized <- F.ref(0) + fiber <- F.start { + F.race(F.unit, F.never[Unit]).void.onCancel(finalized.update(_ + 1)) + } + _ <- fiber.join + _ <- fiber.cancel + _ <- F.cede + back <- finalized.get + } yield back + + assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, Int](Some(0))) + } + test("correctly interpret uncancelable cancelation followed by suspension") { val t = F.uncancelable(_ => F.canceled *> F.never[Unit]) assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, Unit](None)) @@ -212,6 +270,27 @@ class PureConcSuite extends DisciplineSuite with BaseSuite { assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, Unit](None)) } + test("observe external cancelation while blocked inside poll") { + val t = for { + started <- F.deferred[Unit] + polled <- F.deferred[Unit] + gate <- F.deferred[Unit] + ran <- F.ref(false) + fiber <- F.start { + F.uncancelable { poll => + started.complete(()) *> + poll(polled.complete(()) *> gate.get *> ran.set(true)) + } + } + canceler <- F.start(polled.get *> fiber.cancel) + _ <- started.get + _ <- canceler.join + back <- ran.get + } yield back + + assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, Boolean](Some(false))) + } + test("run finalizers around a self-canceling polled region") { val t = for { finalized <- F.ref(0) @@ -227,6 +306,23 @@ class PureConcSuite extends DisciplineSuite with BaseSuite { assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, Int](Some(1))) } + test("run outer finalizers around a self-canceling polled region") { + val t = for { + finalized <- F.ref(0) + fiber <- F.start { + F.onCancel( + F.uncancelable { poll => + poll(F.canceled) + }, + finalized.update(_ + 1)) + } + _ <- fiber.join + back <- finalized.get + } yield back + + assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, Int](Some(1))) + } + test("observe pending self-cancel before running a polled region") { val t = for { finalized <- F.ref(0) @@ -246,6 +342,23 @@ class PureConcSuite extends DisciplineSuite with BaseSuite { Outcome.Succeeded[Option, Int, (Int, Boolean)](Some((1, false)))) } + test("observe nested self-cancel inside a polled region before continuing") { + val t = for { + ran <- F.ref(false) + fiber <- F.start { + F.uncancelable { poll => + poll { + F.uncancelable(_ => F.canceled) *> ran.set(true) + } + } + } + _ <- fiber.join + back <- ran.get + } yield back + + assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, Boolean](Some(false))) + } + test("implement locals via Kleisli and FreeT") { import cats.{~>, Eval, Id} import cats.data.Kleisli @@ -282,6 +395,59 @@ class PureConcSuite extends DisciplineSuite with BaseSuite { } } } + + test("race TimeT values against never") { + type T[A] = TimeT[F, A] + val T = GenTemporal[T, Int] + + assertEquals( + pure.run(TimeT.run(T.race(T.pure(1), T.never[Unit]))), + Outcome.Succeeded[Option, Int, Either[Int, Unit]](Some(Left(1)))) + assertEquals( + pure.run(TimeT.run(T.race(T.never[Unit], T.pure(1)))), + Outcome.Succeeded[Option, Int, Either[Unit, Int]](Some(Right(1)))) + assertEquals( + pure.run(TimeT.run(T.race(T.sleep(1.second).as(1), T.never[Unit]))), + Outcome.Succeeded[Option, Int, Either[Int, Unit]](Some(Left(1)))) + assertEquals( + pure.run(TimeT.run(T.race(T.never[Unit], T.sleep(1.second).as(1)))), + Outcome.Succeeded[Option, Int, Either[Unit, Int]](Some(Right(1)))) + assertEquals( + pure.run( + TimeT.run(T.race(T.sleep(2.seconds).as("slow"), T.sleep(1.second).as("fast")))), + Outcome.Succeeded[Option, Int, Either[String, String]](Some(Right("fast")))) + assertEquals( + pure.run( + TimeT.run(T.race(T.sleep(1.second).as("fast"), T.sleep(2.seconds).as("slow")))), + Outcome.Succeeded[Option, Int, Either[String, String]](Some(Left("fast")))) + assertEquals( + pure.run(TimeT.run(T.race(T.canceled, T.never[Unit]).void)), + Outcome.Canceled[Option, Int, Unit]()) + assertEquals( + pure.run(TimeT.run(T.race(T.never[Unit], T.canceled).void)), + Outcome.Canceled[Option, Int, Unit]()) + assertEquals( + pure.run( + TimeT.run( + T.race(TimeT.liftF(F.uncancelable(_ => F.canceled.as(1))), T.never[Unit]))), + Outcome.Canceled[Option, Int, Either[Int, Unit]]()) + assertEquals( + pure.run( + TimeT.run( + T.race(T.never[Unit], TimeT.liftF(F.uncancelable(_ => F.canceled.as(1)))))), + Outcome.Canceled[Option, Int, Either[Unit, Int]]()) + assertEquals( + pure.run( + TimeT.run( + T.race(TimeT.liftF(F.start(F.unit).flatMap(_.join).as(1)), T.never[Unit]))), + Outcome.Succeeded[Option, Int, Either[Int, Unit]](Some(Left(1)))) + assertEquals( + pure.run( + TimeT.run( + T.race(T.never[Unit], TimeT.liftF(F.start(F.unit).flatMap(_.join).as(1))))), + Outcome.Succeeded[Option, Int, Either[Unit, Int]](Some(Right(1)))) + } + } checkAll( From 408010296b304e5ca0a3429fa335dc7f7340598c Mon Sep 17 00:00:00 2001 From: Daniel Spiewak Date: Mon, 22 Jun 2026 10:36:37 -0500 Subject: [PATCH 05/10] Re-unified external and self-cancelation --- .../cats/effect/kernel/testkit/pure.scala | 133 +++++++----------- .../cats/effect/laws/PureConcSuite.scala | 17 +++ 2 files changed, 69 insertions(+), 81 deletions(-) diff --git a/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala b/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala index 0341458c8d..dda87ec54c 100644 --- a/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala +++ b/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala @@ -43,12 +43,9 @@ object pure { private[pure] final case class MaskFrame(id: MaskId) - private[pure] sealed trait CancelationSignal[E] - - private[pure] object CancelationSignal { - final case class External[E]() extends CancelationSignal[E] - final case class Self[E](finalizers: List[PureConc[E, Unit]]) extends CancelationSignal[E] - } + // None defers finalizer selection until observation; Some scopes an in-fiber request. + private[pure] final case class CancelationSignal[E]( + finalizers: Option[List[PureConc[E, Unit]]]) private[pure] final class CancelationListenerId @@ -129,19 +126,14 @@ object pure { val back = Kleisli.ask[IdOC[E, *], FiberCtx[E]] map { ctx => val checker = ctx .self - .isFinalizing + .hasActivePoll .ifM( ().pure[PureConc[E, *]], ctx .self - .hasActivePoll + .isFinalizing .ifM( - ctx - .self - .realizeExternalCancelationWith(ctx) - .ifM( - ApplicativeThread[PureConc[E, *]].done[Unit], - ().pure[PureConc[E, *]]), + ().pure[PureConc[E, *]], ctx .self .realizeCancelationWith(ctx) @@ -221,7 +213,7 @@ object pure { identifiedCompletion.map(a => Succeeded[Id, E, A](a): IdOC[E, A]) handleError { e => Errored(e) } - } + } Kleisli.ask[ResolvedPC[E, *], MVar.Universe].map { u => ApplicativeThread[ResolvedPC[E, *]].start(body.run(u)) >> results.run(u) @@ -300,7 +292,7 @@ object pure { def cede: PureConc[E, Unit] = withCtx { ctx => Thread.cede *> - ctx.self.realizeExternalCancelationWith(ctx).ifM(Thread.done, unit) + ctx.self.realizeCancelationWith(ctx).ifM(Thread.done, unit) } def never[A]: PureConc[E, A] = @@ -528,9 +520,7 @@ object pure { case MaskUpdate.Removed => val back = result.pure[PureConc[E, *]].rethrow - self.hasActivePoll.ifM( - self.realizeSelfCancelationWith(ctx).ifM(Thread.done, back), - self.realizeCancelationWith(ctx).ifM(Thread.done, back)) + self.realizeCancelationWith(ctx).ifM(Thread.done, back) case MaskUpdate.Shadowed | MaskUpdate.Absent => result.pure[PureConc[E, *]].rethrow @@ -721,23 +711,6 @@ object pure { } } <* setFinalizing(false))) - private[this] def cancelationFinalizers( - signal: CancelationSignal[E], - ctx: FiberCtx[E]): List[PureConc[E, Unit]] = { - signal match { - case CancelationSignal.External() => ctx.finalizers - case CancelationSignal.Self(Nil) if ctx.selfCancelationBoundary.nonEmpty => - selfCancelationFinalizers(ctx) - case CancelationSignal.Self(finalizers) => finalizers - } - } - - private[this] def selfCancelationFinalizers(ctx: FiberCtx[E]): List[PureConc[E, Unit]] = - ctx.selfCancelationBoundary match { - case Some(boundary) => ctx.finalizers.take((ctx.finalizers.length - boundary) max 0) - case None => ctx.finalizers - } - private[this] def whileFinalizing[B](ctx: FiberCtx[E])(fb: PureConc[E, B]): PureConc[E, B] = localCtx(ctx.withFinalizers(Nil).withFinalizing(true), setFinalizing(true) *> fb) @@ -747,42 +720,42 @@ object pure { case _ => false } + private[this] def cancelationFinalizers( + signal: CancelationSignal[E], + ctx: FiberCtx[E]): List[PureConc[E, Unit]] = + signal.finalizers match { + case None => ctx.finalizers + case Some(Nil) if ctx.selfCancelationBoundary.nonEmpty => + cancelationBoundaryFinalizers(ctx) + case Some(finalizers) => finalizers + } + private[this] def realizeCancelationWithSignal( ctx: FiberCtx[E], signal: CancelationSignal[E]): PureConc[E, Boolean] = - whileFinalizing(ctx)(finalizeWith(ctx, cancelationFinalizers(signal, ctx))) - - private[this] def realizeCancelationIf(ctx: FiberCtx[E])( - accepts: CancelationSignal[E] => Boolean): PureConc[E, Boolean] = if (ctx.finalizing) false.pure[PureConc[E, *]] else isFinalizing.ifM( finalizationOutcome, ctx.self.currentMasks.map(_.isEmpty).ifM( - canceled0.tryRead[PureConc[E, *]].flatMap { - case Some(signal) if accepts(signal) => - realizeCancelationWithSignal(ctx, signal) - - case Some(_) | None => - false.pure[PureConc[E, *]] - }, + whileFinalizing(ctx)(finalizeWith(ctx, cancelationFinalizers(signal, ctx))), false.pure[PureConc[E, *]] )) - private[pure] def realizeExternalCancelationWith(ctx: FiberCtx[E]): PureConc[E, Boolean] = - realizeCancelationIf(ctx) { - case CancelationSignal.External() => true - case _ => false - } + private[pure] def realizeCancelationWith(ctx: FiberCtx[E]): PureConc[E, Boolean] = + if (ctx.finalizing) false.pure[PureConc[E, *]] + else isFinalizing.ifM( + finalizationOutcome, + canceled0.tryRead[PureConc[E, *]].flatMap { + case Some(signal) => realizeCancelationWithSignal(ctx, signal) + case None => false.pure[PureConc[E, *]] + }) - private[pure] def realizeSelfCancelationWith(ctx: FiberCtx[E]): PureConc[E, Boolean] = - realizeCancelationIf(ctx) { - case CancelationSignal.Self(_) => true - case _ => false + private[this] def cancelationBoundaryFinalizers(ctx: FiberCtx[E]): List[PureConc[E, Unit]] = + ctx.selfCancelationBoundary match { + case Some(boundary) => ctx.finalizers.take((ctx.finalizers.length - boundary) max 0) + case None => Nil } - private[pure] def realizeCancelationWith(ctx: FiberCtx[E]): PureConc[E, Boolean] = - realizeCancelationIf(ctx)((_: CancelationSignal[E]) => true) - private[pure] def awaitCancelationWith(ctx: FiberCtx[E]): PureConc[E, Boolean] = { def blocked = MVar.empty[PureConc[E, *], Unit].flatMap(_.read[PureConc[E, *]]).as(false) @@ -805,25 +778,26 @@ object pure { if (ctx.finalizing) false.pure[PureConc[E, *]] else isFinalizing.ifM( ctx.self.currentMasks.map(_.isEmpty), - ctx.self.currentMasks.map(_.isEmpty).ifM( - whileFinalizing(ctx) { - val finalizers = - selfCancelationFinalizers(ctx) - val selfCancelation = CancelationSignal.Self(finalizers) - - canceled0.tryPut[PureConc[E, *]](selfCancelation).flatMap { - case true => notifyCancelationListeners *> finalizeWith(ctx, finalizers) - case false => - canceled0 - .tryRead[PureConc[E, *]] - .flatMap( - _.fold(finalizeWith(ctx, finalizers))(signal => - finalizeWith(ctx, cancelationFinalizers(signal, ctx)))) + ctx.self.currentMasks.flatMap { + case Nil => + whileFinalizing(ctx) { + canceled0 + .tryPut[PureConc[E, *]](CancelationSignal[E](Some(ctx.finalizers))) + .flatMap { + case true => notifyCancelationListeners *> finalizeWith(ctx, ctx.finalizers) + case false => finalizeWith(ctx, ctx.finalizers) + } } - }, - canceled0 - .tryPut[PureConc[E, *]](CancelationSignal.Self(Nil)) - .flatMap(inserted => if (inserted) notifyCancelationListeners.as(false) else false.pure[PureConc[E, *]]))) + + case _ => + requestCancelation(Some(Nil)).as(false) + }) + + private[this] def requestCancelation( + finalizers: Option[List[PureConc[E, Unit]]]): PureConc[E, Unit] = + canceled0 + .tryPut[PureConc[E, *]](CancelationSignal[E](finalizers)) + .flatMap(inserted => if (inserted) notifyCancelationListeners else ().pure[PureConc[E, *]]) val join: PureConc[E, Outcome[PureConc[E, *], E, A]] = if (canceled0 eq null) state.read @@ -840,10 +814,7 @@ object pure { state.tryRead.flatMap { case Some(_) => ().pure[PureConc[E, *]] case None => - canceled0.tryPut[PureConc[E, *]](CancelationSignal.External()).flatMap { - case true => notifyCancelationListeners - case false => ().pure[PureConc[E, *]] - } *> state.read.void + requestCancelation(None) *> state.read.void } } } diff --git a/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala b/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala index ed4cd520d2..cd63a4aed0 100644 --- a/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala +++ b/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala @@ -323,6 +323,23 @@ class PureConcSuite assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, Int](Some(1))) } + test("run outer finalizers when a masked self-cancel is observed inside poll") { + val t = for { + finalized <- F.ref(0) + fiber <- F.start { + F.onCancel( + F.uncancelable { poll => + poll(F.uncancelable(_ => F.canceled)) + }, + finalized.update(_ + 1)) + } + _ <- fiber.join + back <- finalized.get + } yield back + + assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, Int](Some(1))) + } + test("observe pending self-cancel before running a polled region") { val t = for { finalized <- F.ref(0) From b2dc913dc6e6447536e14bed59a6812c9bb16b23 Mon Sep 17 00:00:00 2001 From: Daniel Spiewak Date: Mon, 22 Jun 2026 10:49:14 -0500 Subject: [PATCH 06/10] Scalafmt --- .../kernel/testkit/PureConcGenerators.scala | 5 +- .../cats/effect/kernel/testkit/TimeT.scala | 5 +- .../cats/effect/kernel/testkit/pure.scala | 194 ++++++++++-------- .../cats/effect/laws/PureConcSuite.scala | 54 ++--- 4 files changed, 129 insertions(+), 129 deletions(-) diff --git a/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/PureConcGenerators.scala b/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/PureConcGenerators.scala index 9941847536..64e44512dd 100644 --- a/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/PureConcGenerators.scala +++ b/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/PureConcGenerators.scala @@ -39,9 +39,8 @@ object PureConcGenerators { override def recursiveGen[B: Arbitrary: Cogen](deeper: GenK[PureConc[E, *]]) = super .recursiveGen[B](deeper) - .filterNot( - gen => gen._1 == "racePair" || gen._1 == "join" - ) // remove generators which reify nondeterminism and cannot be law-tested + .filterNot(gen => + gen._1 == "racePair" || gen._1 == "join") // remove generators which reify nondeterminism and cannot be law-tested } implicit def arbitraryPureConc[E: Arbitrary: Cogen, A: Arbitrary: Cogen] diff --git a/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/TimeT.scala b/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/TimeT.scala index 0702e442e9..8068a639b4 100644 --- a/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/TimeT.scala +++ b/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/TimeT.scala @@ -171,9 +171,8 @@ private[effect] object TimeT { } } - override def raceOutcome[A, B](fa: TimeT[F, A], fb: TimeT[F, B]): TimeT[ - F, - Either[Outcome[TimeT[F, *], E, A], Outcome[TimeT[F, *], E, B]]] = + override def raceOutcome[A, B](fa: TimeT[F, A], fb: TimeT[F, B]) + : TimeT[F, Either[Outcome[TimeT[F, *], E, A], Outcome[TimeT[F, *], E, B]]] = uncancelable { poll => poll(racePair(fa, fb)).flatMap { case Left((oc, f)) => f.cancel.as(Left(oc)) diff --git a/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala b/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala index dda87ec54c..ec6a3ce8a9 100644 --- a/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala +++ b/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala @@ -137,9 +137,8 @@ object pure { ctx .self .realizeCancelationWith(ctx) - .ifM( - ApplicativeThread[PureConc[E, *]].done[Unit], - ().pure[PureConc[E, *]]))) + .ifM(ApplicativeThread[PureConc[E, *]].done[Unit], ().pure[PureConc[E, *]])) + ) checker >> mvarLiftF(ThreadT.liftF(ka)) } @@ -210,9 +209,8 @@ object pure { ta.mapK(fk) } - identifiedCompletion.map(a => Succeeded[Id, E, A](a): IdOC[E, A]) handleError { - e => Errored(e) - } + identifiedCompletion.map(a => + Succeeded[Id, E, A](a): IdOC[E, A]) handleError { e => Errored(e) } } Kleisli.ask[ResolvedPC[E, *], MVar.Universe].map { u => @@ -307,9 +305,7 @@ object pure { def deferred[A]: PureConc[E, Deferred[PureConc[E, *], A]] = MVar.empty[PureConc[E, *], A].flatMap(mVar => Kleisli.pure(unsafeDeferred(mVar))) - private[this] def interruptible[A]( - ctx: FiberCtx[E], - fa: PureConc[E, A]): PureConc[E, A] = + private[this] def interruptible[A](ctx: FiberCtx[E], fa: PureConc[E, A]): PureConc[E, A] = ctx.self.interruptible(ctx)(fa) private def unsafeRef[A](mVar: MVar[A]): Ref[PureConc[E, *], A] = @@ -368,9 +364,7 @@ object pure { private def unsafeDeferred[A](mVar: MVar[A]): Deferred[PureConc[E, *], A] = new Deferred[PureConc[E, *], A] { override def get: PureConc[E, A] = - withCtx { ctx => - interruptible(ctx, mVar.read[PureConc[E, *]]) - } + withCtx { ctx => interruptible(ctx, mVar.read[PureConc[E, *]]) } override def complete(a: A): PureConc[E, Boolean] = mVar.tryPut[PureConc[E, *]](a) @@ -382,24 +376,25 @@ object pure { MVar.empty[PureConc[E, *], Outcome[PureConc[E, *], E, A]].flatMap { state => MVar.empty[PureConc[E, *], CancelationSignal[E]] flatMap { canceled => MVar[PureConc[E, *], List[MaskFrame]](Nil) flatMap { masks => - MVar[PureConc[E, *], List[CancelationListener[E]]](Nil) flatMap { cancelationListeners => - MVar[PureConc[E, *], Boolean](false) flatMap { finalizing => - MVar[PureConc[E, *], Int](0) flatMap { activePolls => - val fiber = - new PureFiber[E, A]( - state, - canceled, - masks, - cancelationListeners, - finalizing, - activePolls) - - // the tryPut here is interesting: it encodes first-wins semantics on cancelation/completion - val body = guaranteeCase(fa)(state.tryPut[PureConc[E, *]](_).void) - val identified = localCtx(FiberCtx(fiber), body) - Thread.start(identified.attempt.void).as(fiber) + MVar[PureConc[E, *], List[CancelationListener[E]]](Nil) flatMap { + cancelationListeners => + MVar[PureConc[E, *], Boolean](false) flatMap { finalizing => + MVar[PureConc[E, *], Int](0) flatMap { activePolls => + val fiber = + new PureFiber[E, A]( + state, + canceled, + masks, + cancelationListeners, + finalizing, + activePolls) + + // the tryPut here is interesting: it encodes first-wins semantics on cancelation/completion + val body = guaranteeCase(fa)(state.tryPut[PureConc[E, *]](_).void) + val identified = localCtx(FiberCtx(fiber), body) + Thread.start(identified.attempt.void).as(fiber) + } } - } } } } @@ -413,23 +408,30 @@ object pure { (Fiber[PureConc[E, *], E, A], Outcome[PureConc[E, *], E, B])]] = uncancelable { poll => for { - result <- deferred[Either[Outcome[PureConc[E, *], E, A], Outcome[PureConc[E, *], E, B]]] + result <- deferred[ + Either[Outcome[PureConc[E, *], E, A], Outcome[PureConc[E, *], E, B]]] fibA <- start(fa) fibB <- start(fb) _ <- start( - fibA.join.flatMap(oc => - result.complete( - Left(oc): Either[ - Outcome[PureConc[E, *], E, A], - Outcome[PureConc[E, *], E, B]]).void)) + fibA + .join + .flatMap(oc => + result + .complete(Left(oc): Either[ + Outcome[PureConc[E, *], E, A], + Outcome[PureConc[E, *], E, B]]) + .void)) _ <- start( - fibB.join.flatMap(oc => - result.complete( - Right(oc): Either[ - Outcome[PureConc[E, *], E, A], - Outcome[PureConc[E, *], E, B]]).void)) + fibB + .join + .flatMap(oc => + result + .complete(Right(oc): Either[ + Outcome[PureConc[E, *], E, A], + Outcome[PureConc[E, *], E, B]]) + .void)) back <- onCancel( poll(result.get), @@ -482,7 +484,9 @@ object pure { val pollCtx = update match { case MaskUpdate.Removed => callCtx.withSelfCancelationBoundary( - callCtx.selfCancelationBoundary.orElse(Some(selfCancelationBoundary))) + callCtx + .selfCancelationBoundary + .orElse(Some(selfCancelationBoundary))) case MaskUpdate.Shadowed | MaskUpdate.Absent => callCtx @@ -508,7 +512,9 @@ object pure { Thread.done, restoreF *> result.pure[PureConc[E, *]].rethrow) }), - restoreF)) + restoreF + ) + ) } else fa } @@ -638,17 +644,14 @@ object pure { if (cancelationListeners eq null) ().pure[PureConc[E, *]] else cancelationListeners.read[PureConc[E, *]].flatMap { listeners => - cancelationListeners - .swap[PureConc[E, *]](listeners.filterNot(_.id === id)) - .void + cancelationListeners.swap[PureConc[E, *]](listeners.filterNot(_.id === id)).void } private[this] def notifyCancelationListeners: PureConc[E, Unit] = if (cancelationListeners eq null) ().pure[PureConc[E, *]] else cancelationListeners.swap[PureConc[E, *]](Nil).flatMap(_.traverse_(_.action)) - private[pure] def interruptible[B](ctx: FiberCtx[E])( - fb: PureConc[E, B]): PureConc[E, B] = { + private[pure] def interruptible[B](ctx: FiberCtx[E])(fb: PureConc[E, B]): PureConc[E, B] = { val Thread = ApplicativeThread[PureConc[E, *]] ctx.self.currentMasks.flatMap { @@ -664,7 +667,8 @@ object pure { signal.tryRead[PureConc[E, *]].flatMap { case Some(_) => ().pure[PureConc[E, *]] case None => - ctx.self + ctx + .self .realizeCancelationWith(ctx) .ifM(notifyCancelation, ().pure[PureConc[E, *]]) } @@ -709,7 +713,8 @@ object pure { case Outcome.Canceled() => true case _ => false } - } <* setFinalizing(false))) + } <* setFinalizing(false)) + ) private[this] def whileFinalizing[B](ctx: FiberCtx[E])(fb: PureConc[E, B]): PureConc[E, B] = localCtx(ctx.withFinalizers(Nil).withFinalizing(true), setFinalizing(true) *> fb) @@ -734,21 +739,29 @@ object pure { ctx: FiberCtx[E], signal: CancelationSignal[E]): PureConc[E, Boolean] = if (ctx.finalizing) false.pure[PureConc[E, *]] - else isFinalizing.ifM( - finalizationOutcome, - ctx.self.currentMasks.map(_.isEmpty).ifM( - whileFinalizing(ctx)(finalizeWith(ctx, cancelationFinalizers(signal, ctx))), - false.pure[PureConc[E, *]] - )) + else + isFinalizing.ifM( + finalizationOutcome, + ctx + .self + .currentMasks + .map(_.isEmpty) + .ifM( + whileFinalizing(ctx)(finalizeWith(ctx, cancelationFinalizers(signal, ctx))), + false.pure[PureConc[E, *]] + ) + ) private[pure] def realizeCancelationWith(ctx: FiberCtx[E]): PureConc[E, Boolean] = if (ctx.finalizing) false.pure[PureConc[E, *]] - else isFinalizing.ifM( - finalizationOutcome, - canceled0.tryRead[PureConc[E, *]].flatMap { - case Some(signal) => realizeCancelationWithSignal(ctx, signal) - case None => false.pure[PureConc[E, *]] - }) + else + isFinalizing.ifM( + finalizationOutcome, + canceled0.tryRead[PureConc[E, *]].flatMap { + case Some(signal) => realizeCancelationWithSignal(ctx, signal) + case None => false.pure[PureConc[E, *]] + } + ) private[this] def cancelationBoundaryFinalizers(ctx: FiberCtx[E]): List[PureConc[E, Unit]] = ctx.selfCancelationBoundary match { @@ -761,50 +774,51 @@ object pure { MVar.empty[PureConc[E, *], Unit].flatMap(_.read[PureConc[E, *]]).as(false) if (ctx.finalizing) blocked - else ctx.self.currentMasks.flatMap { - case Nil => - isFinalizing.ifM( - canceled0.tryRead[PureConc[E, *]].map(_.isEmpty), - canceled0.read[PureConc[E, *]].flatMap(realizeCancelationWithSignal(ctx, _))) + else + ctx.self.currentMasks.flatMap { + case Nil => + isFinalizing.ifM( + canceled0.tryRead[PureConc[E, *]].map(_.isEmpty), + canceled0.read[PureConc[E, *]].flatMap(realizeCancelationWithSignal(ctx, _))) - case _ => - isFinalizing.ifM( - canceled0.tryRead[PureConc[E, *]].map(_.isEmpty), - blocked) - } + case _ => + isFinalizing.ifM(canceled0.tryRead[PureConc[E, *]].map(_.isEmpty), blocked) + } } private[pure] def cancelAndRealizeWith(ctx: FiberCtx[E]): PureConc[E, Boolean] = if (ctx.finalizing) false.pure[PureConc[E, *]] - else isFinalizing.ifM( - ctx.self.currentMasks.map(_.isEmpty), - ctx.self.currentMasks.flatMap { - case Nil => - whileFinalizing(ctx) { - canceled0 - .tryPut[PureConc[E, *]](CancelationSignal[E](Some(ctx.finalizers))) - .flatMap { - case true => notifyCancelationListeners *> finalizeWith(ctx, ctx.finalizers) - case false => finalizeWith(ctx, ctx.finalizers) - } - } + else + isFinalizing.ifM( + ctx.self.currentMasks.map(_.isEmpty), + ctx.self.currentMasks.flatMap { + case Nil => + whileFinalizing(ctx) { + canceled0 + .tryPut[PureConc[E, *]](CancelationSignal[E](Some(ctx.finalizers))) + .flatMap { + case true => + notifyCancelationListeners *> finalizeWith(ctx, ctx.finalizers) + case false => finalizeWith(ctx, ctx.finalizers) + } + } - case _ => - requestCancelation(Some(Nil)).as(false) - }) + case _ => + requestCancelation(Some(Nil)).as(false) + } + ) private[this] def requestCancelation( finalizers: Option[List[PureConc[E, Unit]]]): PureConc[E, Unit] = canceled0 .tryPut[PureConc[E, *]](CancelationSignal[E](finalizers)) - .flatMap(inserted => if (inserted) notifyCancelationListeners else ().pure[PureConc[E, *]]) + .flatMap(inserted => + if (inserted) notifyCancelationListeners else ().pure[PureConc[E, *]]) val join: PureConc[E, Outcome[PureConc[E, *], E, A]] = if (canceled0 eq null) state.read else { - withCtx { ctx => - ctx.self.interruptible(ctx)(state.read) - } + withCtx { ctx => ctx.self.interruptible(ctx)(state.read) } } val cancel: PureConc[E, Unit] = diff --git a/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala b/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala index cd63a4aed0..f5544cc1a1 100644 --- a/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala +++ b/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala @@ -31,8 +31,8 @@ import munit.DisciplineSuite private[laws] trait PureConcSuiteLowPriorityTimeTInstances { implicit def orderTimeTPureConcFiniteDuration( - implicit FA: Order[PureConc[Int, FiniteDuration]]): Order[ - TimeT[PureConc[Int, *], FiniteDuration]] = + implicit FA: Order[PureConc[Int, FiniteDuration]]) + : Order[TimeT[PureConc[Int, *], FiniteDuration]] = TimeT.orderTimeT } @@ -46,8 +46,8 @@ class PureConcSuite implicit def exec(fb: TimeT[PureConc[Int, *], Boolean]): Prop = Prop(pure.run(TimeT.run(fb)).fold(false, _ => false, _.getOrElse(false))) - implicit def eqTimeTPureConc[A](implicit FA: Eq[PureConc[Int, A]]): Eq[ - TimeT[PureConc[Int, *], A]] = + implicit def eqTimeTPureConc[A]( + implicit FA: Eq[PureConc[Int, A]]): Eq[TimeT[PureConc[Int, *], A]] = TimeT.eqTimeT { @@ -295,9 +295,7 @@ class PureConcSuite val t = for { finalized <- F.ref(0) fiber <- F.start { - F.uncancelable { poll => - F.onCancel(poll(F.canceled), finalized.update(_ + 1)) - } + F.uncancelable { poll => F.onCancel(poll(F.canceled), finalized.update(_ + 1)) } } _ <- fiber.join back <- finalized.get @@ -310,11 +308,7 @@ class PureConcSuite val t = for { finalized <- F.ref(0) fiber <- F.start { - F.onCancel( - F.uncancelable { poll => - poll(F.canceled) - }, - finalized.update(_ + 1)) + F.onCancel(F.uncancelable { poll => poll(F.canceled) }, finalized.update(_ + 1)) } _ <- fiber.join back <- finalized.get @@ -328,9 +322,7 @@ class PureConcSuite finalized <- F.ref(0) fiber <- F.start { F.onCancel( - F.uncancelable { poll => - poll(F.uncancelable(_ => F.canceled)) - }, + F.uncancelable { poll => poll(F.uncancelable(_ => F.canceled)) }, finalized.update(_ + 1)) } _ <- fiber.join @@ -403,13 +395,9 @@ class PureConcSuite .liftT[Id, Kleisli[Eval, Int, *], Unit]( Kleisli.liftF[Eval, Int, Unit](Eval.later(assertEquals(i, 42)))) .flatMap(_ => - read { i2 => - FreeT.liftT(Kleisli.liftF(Eval.later(assertEquals(i2, 42)))) - }) + read { i2 => FreeT.liftT(Kleisli.liftF(Eval.later(assertEquals(i2, 42)))) }) } - } *> read { i => - FreeT.liftT(Kleisli.liftF(Eval.later(assertEquals(i, 1)))) - } + } *> read { i => FreeT.liftT(Kleisli.liftF(Eval.later(assertEquals(i, 1)))) } } } @@ -432,11 +420,13 @@ class PureConcSuite assertEquals( pure.run( TimeT.run(T.race(T.sleep(2.seconds).as("slow"), T.sleep(1.second).as("fast")))), - Outcome.Succeeded[Option, Int, Either[String, String]](Some(Right("fast")))) + Outcome.Succeeded[Option, Int, Either[String, String]](Some(Right("fast"))) + ) assertEquals( pure.run( TimeT.run(T.race(T.sleep(1.second).as("fast"), T.sleep(2.seconds).as("slow")))), - Outcome.Succeeded[Option, Int, Either[String, String]](Some(Left("fast")))) + Outcome.Succeeded[Option, Int, Either[String, String]](Some(Left("fast"))) + ) assertEquals( pure.run(TimeT.run(T.race(T.canceled, T.never[Unit]).void)), Outcome.Canceled[Option, Int, Unit]()) @@ -445,24 +435,22 @@ class PureConcSuite Outcome.Canceled[Option, Int, Unit]()) assertEquals( pure.run( - TimeT.run( - T.race(TimeT.liftF(F.uncancelable(_ => F.canceled.as(1))), T.never[Unit]))), + TimeT.run(T.race(TimeT.liftF(F.uncancelable(_ => F.canceled.as(1))), T.never[Unit]))), Outcome.Canceled[Option, Int, Either[Int, Unit]]()) assertEquals( pure.run( - TimeT.run( - T.race(T.never[Unit], TimeT.liftF(F.uncancelable(_ => F.canceled.as(1)))))), + TimeT.run(T.race(T.never[Unit], TimeT.liftF(F.uncancelable(_ => F.canceled.as(1)))))), Outcome.Canceled[Option, Int, Either[Unit, Int]]()) assertEquals( pure.run( - TimeT.run( - T.race(TimeT.liftF(F.start(F.unit).flatMap(_.join).as(1)), T.never[Unit]))), - Outcome.Succeeded[Option, Int, Either[Int, Unit]](Some(Left(1)))) + TimeT.run(T.race(TimeT.liftF(F.start(F.unit).flatMap(_.join).as(1)), T.never[Unit]))), + Outcome.Succeeded[Option, Int, Either[Int, Unit]](Some(Left(1))) + ) assertEquals( pure.run( - TimeT.run( - T.race(T.never[Unit], TimeT.liftF(F.start(F.unit).flatMap(_.join).as(1))))), - Outcome.Succeeded[Option, Int, Either[Unit, Int]](Some(Right(1)))) + TimeT.run(T.race(T.never[Unit], TimeT.liftF(F.start(F.unit).flatMap(_.join).as(1))))), + Outcome.Succeeded[Option, Int, Either[Unit, Int]](Some(Right(1))) + ) } } From 3cfb3b2a1c0bef8dfc9d94914aaec03f10a2d6ea Mon Sep 17 00:00:00 2001 From: Daniel Spiewak Date: Tue, 28 Jul 2026 17:37:43 +0200 Subject: [PATCH 07/10] Fix PureConc deferred cancellation finalizers --- .../cats/effect/kernel/testkit/pure.scala | 142 ++++++++---------- .../cats/effect/laws/PureConcSuite.scala | 97 ++++++++++++ 2 files changed, 158 insertions(+), 81 deletions(-) diff --git a/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala b/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala index ec6a3ce8a9..3288587d7b 100644 --- a/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala +++ b/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala @@ -43,9 +43,12 @@ object pure { private[pure] final case class MaskFrame(id: MaskId) - // None defers finalizer selection until observation; Some scopes an in-fiber request. - private[pure] final case class CancelationSignal[E]( - finalizers: Option[List[PureConc[E, Unit]]]) + private[pure] sealed trait CancelationSignal + + private[pure] object CancelationSignal { + case object Self extends CancelationSignal + case object External extends CancelationSignal + } private[pure] final class CancelationListenerId @@ -71,18 +74,13 @@ object pure { masks: List[MaskId] = Nil, finalizers: List[PureConc[E, Unit]] = Nil) { - private[pure] def selfCancelationBoundary: Option[Int] = None - private[pure] def finalizing: Boolean = false private[pure] def withFinalizers(value: List[PureConc[E, Unit]]): FiberCtx[E] = - FiberCtx.internal(self, masks, value, selfCancelationBoundary, finalizing) - - private[pure] def withSelfCancelationBoundary(value: Option[Int]): FiberCtx[E] = - FiberCtx.internal(self, masks, finalizers, value, finalizing) + FiberCtx.internal(self, masks, value, finalizing) private[pure] def withFinalizing(value: Boolean): FiberCtx[E] = - FiberCtx.internal(self, masks, finalizers, selfCancelationBoundary, value) + FiberCtx.internal(self, masks, finalizers, value) } object FiberCtx { @@ -90,7 +88,6 @@ object pure { self: PureFiber[E, ?], masks: List[MaskId], finalizers: List[PureConc[E, Unit]], - override private[pure] val selfCancelationBoundary: Option[Int], override private[pure] val finalizing: Boolean) extends FiberCtx[E](self, masks, finalizers) @@ -98,9 +95,8 @@ object pure { self: PureFiber[E, ?], masks: List[MaskId], finalizers: List[PureConc[E, Unit]], - selfCancelationBoundary: Option[Int], finalizing: Boolean): FiberCtx[E] = - new Internal(self, masks, finalizers, selfCancelationBoundary, finalizing) + new Internal(self, masks, finalizers, finalizing) } type ResolvedPC[E, A] = ThreadT[IdOC[E, *], A] @@ -165,11 +161,11 @@ object pure { type Main[X] = MVarR[ResolvedPC[E, *], X] MVar.empty[Main, Outcome[PureConc[E, *], E, A]].flatMap { state0 => - MVar.empty[Main, CancelationSignal[E]] flatMap { canceled0 => + MVar.empty[Main, CancelationSignal] flatMap { canceled0 => MVar[Main, List[MaskFrame]](Nil) flatMap { masks => MVar[Main, List[CancelationListener[E]]](Nil) flatMap { cancelationListeners => MVar[Main, Boolean](false) flatMap { finalizing => - MVar[Main, Int](0) flatMap { activePolls => + MVar[Main, List[List[Finalizer[E]]]](Nil) flatMap { activePolls => val state = state0[Main] val fiber = new PureFiber[E, A]( @@ -374,25 +370,26 @@ object pure { def start[A](fa: PureConc[E, A]): PureConc[E, Fiber[PureConc[E, *], E, A]] = Thread.annotate("start", true) { MVar.empty[PureConc[E, *], Outcome[PureConc[E, *], E, A]].flatMap { state => - MVar.empty[PureConc[E, *], CancelationSignal[E]] flatMap { canceled => + MVar.empty[PureConc[E, *], CancelationSignal] flatMap { canceled => MVar[PureConc[E, *], List[MaskFrame]](Nil) flatMap { masks => MVar[PureConc[E, *], List[CancelationListener[E]]](Nil) flatMap { cancelationListeners => MVar[PureConc[E, *], Boolean](false) flatMap { finalizing => - MVar[PureConc[E, *], Int](0) flatMap { activePolls => - val fiber = - new PureFiber[E, A]( - state, - canceled, - masks, - cancelationListeners, - finalizing, - activePolls) - - // the tryPut here is interesting: it encodes first-wins semantics on cancelation/completion - val body = guaranteeCase(fa)(state.tryPut[PureConc[E, *]](_).void) - val identified = localCtx(FiberCtx(fiber), body) - Thread.start(identified.attempt.void).as(fiber) + MVar[PureConc[E, *], List[List[Finalizer[E]]]](Nil) flatMap { + activePolls => + val fiber = + new PureFiber[E, A]( + state, + canceled, + masks, + cancelationListeners, + finalizing, + activePolls) + + // the tryPut here is interesting: it encodes first-wins semantics on cancelation/completion + val body = guaranteeCase(fa)(state.tryPut[PureConc[E, *]](_).void) + val identified = localCtx(FiberCtx(fiber), body) + Thread.start(identified.attempt.void).as(fiber) } } } @@ -452,8 +449,6 @@ object pure { Thread.annotate("uncancelable", true) { withCtx { ctx => val mask = new MaskId - val selfCancelationBoundary = - ctx.selfCancelationBoundary.getOrElse(ctx.finalizers.length) val self = ctx.self def updateMasks[B](f: List[MaskFrame] => (List[MaskFrame], B)): PureConc[E, B] = @@ -481,33 +476,23 @@ object pure { if (callCtx.self eq self) removeF.flatMap { update => val restoreF = restore(update) - val pollCtx = update match { - case MaskUpdate.Removed => - callCtx.withSelfCancelationBoundary( - callCtx - .selfCancelationBoundary - .orElse(Some(selfCancelationBoundary))) - - case MaskUpdate.Shadowed | MaskUpdate.Absent => - callCtx - } val enterF = update match { - case MaskUpdate.Removed => self.enterPoll + case MaskUpdate.Removed => self.enterPoll(callCtx.finalizers) case MaskUpdate.Shadowed | MaskUpdate.Absent => unit } enterF *> localCtx( - pollCtx, + callCtx, onCancel( self - .realizeCancelationWith(pollCtx) + .realizeCancelationWith(callCtx) .ifM( Thread.done, fa.attempt.flatMap { result => self - .realizeCancelationWith(pollCtx) + .realizeCancelationWith(callCtx) .ifM( Thread.done, restoreF *> result.pure[PureConc[E, *]].rethrow) @@ -592,11 +577,11 @@ object pure { // todo: MVar is not Serializable, release then update here final class PureFiber[E, A]( val state0: MVar[Outcome[PureConc[E, *], E, A]], - private[this] val canceled0: MVar[CancelationSignal[E]], + private[this] val canceled0: MVar[CancelationSignal], private[pure] val masks: MVar[List[MaskFrame]], private[this] val cancelationListeners: MVar[List[CancelationListener[E]]], private[this] val finalizing: MVar[Boolean], - private[this] val activePolls: MVar[Int]) + private[this] val activePolls: MVar[List[List[Finalizer[E]]]]) extends Fiber[PureConc[E, *], E, A] with Serializable { @@ -611,22 +596,27 @@ object pure { private[pure] val hasActivePoll: PureConc[E, Boolean] = if (activePolls eq null) false.pure[PureConc[E, *]] - else activePolls.read[PureConc[E, *]].map(_ > 0) + else activePolls.read[PureConc[E, *]].map(_.nonEmpty) - private[pure] val enterPoll: PureConc[E, Unit] = + private[pure] def enterPoll(finalizers: List[Finalizer[E]]): PureConc[E, Unit] = if (activePolls eq null) ().pure[PureConc[E, *]] else - activePolls.read[PureConc[E, *]].flatMap { n => - activePolls.swap[PureConc[E, *]](n + 1).void + activePolls.read[PureConc[E, *]].flatMap { polls => + activePolls.swap[PureConc[E, *]](finalizers :: polls).void } private[pure] val exitPoll: PureConc[E, Unit] = if (activePolls eq null) ().pure[PureConc[E, *]] else - activePolls.read[PureConc[E, *]].flatMap { n => - activePolls.swap[PureConc[E, *]]((n - 1) max 0).void + activePolls.read[PureConc[E, *]].flatMap { + case _ :: polls => activePolls.swap[PureConc[E, *]](polls).void + case Nil => ().pure[PureConc[E, *]] } + private[pure] val currentPollFinalizers: PureConc[E, List[Finalizer[E]]] = + if (activePolls eq null) List.empty[Finalizer[E]].pure[PureConc[E, *]] + else activePolls.read[PureConc[E, *]].map(_.headOption.getOrElse(Nil)) + private[pure] def registerCancelationListener( notify: PureConc[E, Unit]): PureConc[E, CancelationListenerId] = { val id = new CancelationListenerId @@ -726,18 +716,16 @@ object pure { } private[this] def cancelationFinalizers( - signal: CancelationSignal[E], - ctx: FiberCtx[E]): List[PureConc[E, Unit]] = - signal.finalizers match { - case None => ctx.finalizers - case Some(Nil) if ctx.selfCancelationBoundary.nonEmpty => - cancelationBoundaryFinalizers(ctx) - case Some(finalizers) => finalizers + signal: CancelationSignal, + ctx: FiberCtx[E]): PureConc[E, List[Finalizer[E]]] = + signal match { + case CancelationSignal.Self => ctx.self.currentPollFinalizers + case CancelationSignal.External => ctx.finalizers.pure[PureConc[E, *]] } private[this] def realizeCancelationWithSignal( ctx: FiberCtx[E], - signal: CancelationSignal[E]): PureConc[E, Boolean] = + signal: CancelationSignal): PureConc[E, Boolean] = if (ctx.finalizing) false.pure[PureConc[E, *]] else isFinalizing.ifM( @@ -747,7 +735,8 @@ object pure { .currentMasks .map(_.isEmpty) .ifM( - whileFinalizing(ctx)(finalizeWith(ctx, cancelationFinalizers(signal, ctx))), + cancelationFinalizers(signal, ctx) + .flatMap(finalizers => whileFinalizing(ctx)(finalizeWith(ctx, finalizers))), false.pure[PureConc[E, *]] ) ) @@ -763,12 +752,6 @@ object pure { } ) - private[this] def cancelationBoundaryFinalizers(ctx: FiberCtx[E]): List[PureConc[E, Unit]] = - ctx.selfCancelationBoundary match { - case Some(boundary) => ctx.finalizers.take((ctx.finalizers.length - boundary) max 0) - case None => Nil - } - private[pure] def awaitCancelationWith(ctx: FiberCtx[E]): PureConc[E, Boolean] = { def blocked = MVar.empty[PureConc[E, *], Unit].flatMap(_.read[PureConc[E, *]]).as(false) @@ -794,24 +777,21 @@ object pure { ctx.self.currentMasks.flatMap { case Nil => whileFinalizing(ctx) { - canceled0 - .tryPut[PureConc[E, *]](CancelationSignal[E](Some(ctx.finalizers))) - .flatMap { - case true => - notifyCancelationListeners *> finalizeWith(ctx, ctx.finalizers) - case false => finalizeWith(ctx, ctx.finalizers) - } + canceled0.tryPut[PureConc[E, *]](CancelationSignal.Self).flatMap { + case true => + notifyCancelationListeners *> finalizeWith(ctx, ctx.finalizers) + case false => finalizeWith(ctx, ctx.finalizers) + } } case _ => - requestCancelation(Some(Nil)).as(false) + requestCancelation(CancelationSignal.Self).as(false) } ) - private[this] def requestCancelation( - finalizers: Option[List[PureConc[E, Unit]]]): PureConc[E, Unit] = + private[this] def requestCancelation(signal: CancelationSignal): PureConc[E, Unit] = canceled0 - .tryPut[PureConc[E, *]](CancelationSignal[E](finalizers)) + .tryPut[PureConc[E, *]](signal) .flatMap(inserted => if (inserted) notifyCancelationListeners else ().pure[PureConc[E, *]]) @@ -828,7 +808,7 @@ object pure { state.tryRead.flatMap { case Some(_) => ().pure[PureConc[E, *]] case None => - requestCancelation(None) *> state.read.void + requestCancelation(CancelationSignal.External) *> state.read.void } } } diff --git a/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala b/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala index f5544cc1a1..7bb49d236b 100644 --- a/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala +++ b/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala @@ -351,6 +351,75 @@ class PureConcSuite Outcome.Succeeded[Option, Int, (Int, Boolean)](Some((1, false)))) } + test("run only finalizers installed after a masked self-cancel") { + val t = for { + before <- F.ref(0) + after <- F.ref(0) + ran <- F.ref(false) + fiber <- F.start { + F.uncancelable { poll => + F.onCancel(F.canceled, before.update(_ + 1)) *> + F.onCancel(poll(ran.set(true)), after.update(_ + 1)) + } + } + _ <- fiber.join + beforeCount <- before.get + afterCount <- after.get + body <- ran.get + } yield (beforeCount, afterCount, body) + + assertEquals( + pure.run(t), + Outcome.Succeeded[Option, Int, (Int, Int, Boolean)](Some((0, 1, false)))) + } + + test("select the innermost active poll finalizers") { + val t = for { + finalized <- F.ref("") + fiber <- F.start { + F.uncancelable { outerPoll => + F.onCancel( + outerPoll { + F.uncancelable { innerPoll => + F.onCancel( + innerPoll(F.uncancelable(_ => F.canceled)), + finalized.update(_ + "B")) + } + }, + finalized.update(_ + "A")) + } + } + _ <- fiber.join + back <- finalized.get + } yield back + + assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, String](Some("BA"))) + } + + test("restore outer poll finalizers after an inner poll completes") { + val t = for { + outerFinalized <- F.ref(0) + innerFinalized <- F.ref(0) + fiber <- F.start { + F.uncancelable { outerPoll => + F.onCancel( + outerPoll { + F.uncancelable { innerPoll => + F.onCancel(innerPoll(F.unit), innerFinalized.update(_ + 1)) + } *> F.uncancelable(_ => F.canceled) + }, + outerFinalized.update(_ + 1) + ) + } + } + _ <- fiber.join + outer <- outerFinalized.get + inner <- innerFinalized.get + } yield (outer, inner) + + assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, (Int, Int)](Some((1, 0)))) + } + test("observe nested self-cancel inside a polled region before continuing") { val t = for { ran <- F.ref(false) @@ -368,6 +437,34 @@ class PureConcSuite assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, Boolean](Some(false))) } + test("preserve masked self-cancel through poll") { + val maskedCancel = F.uncancelable(_ => F.canceled) + val fa = F.onCancel(maskedCancel, F.never[Unit]) + + assertEquals(pure.run(maskedCancel), Outcome.Canceled[Option, Int, Unit]()) + assertEquals( + pure.run(F.uncancelable(poll => poll(maskedCancel))), + Outcome.Canceled[Option, Int, Unit]()) + assertEquals(pure.run(fa), Outcome.Canceled[Option, Int, Unit]()) + assertEquals( + pure.run(F.uncancelable(poll => poll(fa))), + Outcome.Canceled[Option, Int, Unit]()) + } + + test("run a guarantee finalizer around a masked self-cancel") { + val fa = F.guarantee(F.uncancelable(_ => F.canceled), F.never[Unit]) + + assertEquals(pure.run(fa), Outcome.Succeeded[Option, Int, Unit](None)) + } + + test("associate finalizers across an uncancelable boundary") { + val left = F.uncancelable(_ => F.onCancel(F.canceled, F.never[Unit])) + val right = F.onCancel(F.uncancelable(_ => F.canceled), F.never[Unit]) + + assertEquals(pure.run(left), Outcome.Canceled[Option, Int, Unit]()) + assertEquals(pure.run(right), Outcome.Canceled[Option, Int, Unit]()) + } + test("implement locals via Kleisli and FreeT") { import cats.{~>, Eval, Id} import cats.data.Kleisli From 13bddf3d49b0233eb2ce51a59bb53acfe9b5926e Mon Sep 17 00:00:00 2001 From: Daniel Spiewak Date: Wed, 29 Jul 2026 08:53:35 +0200 Subject: [PATCH 08/10] Make PureConc cancellation origin-neutral --- .../cats/effect/kernel/testkit/pure.scala | 90 ++++++++----------- .../cats/effect/laws/PureConcSuite.scala | 39 ++++++++ 2 files changed, 77 insertions(+), 52 deletions(-) diff --git a/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala b/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala index 3288587d7b..bce5877de4 100644 --- a/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala +++ b/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala @@ -43,13 +43,6 @@ object pure { private[pure] final case class MaskFrame(id: MaskId) - private[pure] sealed trait CancelationSignal - - private[pure] object CancelationSignal { - case object Self extends CancelationSignal - case object External extends CancelationSignal - } - private[pure] final class CancelationListenerId private[pure] object CancelationListenerId { @@ -161,7 +154,7 @@ object pure { type Main[X] = MVarR[ResolvedPC[E, *], X] MVar.empty[Main, Outcome[PureConc[E, *], E, A]].flatMap { state0 => - MVar.empty[Main, CancelationSignal] flatMap { canceled0 => + MVar.empty[Main, Unit] flatMap { canceled0 => MVar[Main, List[MaskFrame]](Nil) flatMap { masks => MVar[Main, List[CancelationListener[E]]](Nil) flatMap { cancelationListeners => MVar[Main, Boolean](false) flatMap { finalizing => @@ -176,7 +169,16 @@ object pure { finalizing, activePolls) - val identified = canceled mapF { ta => + val completed = canceled.flatMap { a => + withCtx { ctx => + ctx + .self + .realizeCancelationWith(ctx) + .ifM(ApplicativeThread[PureConc[E, *]].done[A], a.pure[PureConc[E, *]]) + } + } + + val identified = completed mapF { ta => val fk = new (FiberR[E, *] ~> IdOC[E, *]) { def apply[a](ke: FiberR[E, a]) = ke.run(FiberCtx(fiber)) @@ -370,7 +372,7 @@ object pure { def start[A](fa: PureConc[E, A]): PureConc[E, Fiber[PureConc[E, *], E, A]] = Thread.annotate("start", true) { MVar.empty[PureConc[E, *], Outcome[PureConc[E, *], E, A]].flatMap { state => - MVar.empty[PureConc[E, *], CancelationSignal] flatMap { canceled => + MVar.empty[PureConc[E, *], Unit] flatMap { canceled => MVar[PureConc[E, *], List[MaskFrame]](Nil) flatMap { masks => MVar[PureConc[E, *], List[CancelationListener[E]]](Nil) flatMap { cancelationListeners => @@ -511,7 +513,13 @@ object pure { case MaskUpdate.Removed => val back = result.pure[PureConc[E, *]].rethrow - self.realizeCancelationWith(ctx).ifM(Thread.done, back) + self.currentPollFinalizers.flatMap { + case Some(finalizers) => + self + .realizeCancelationWith(ctx.withFinalizers(finalizers)) + .ifM(Thread.done, back) + case None => back + } case MaskUpdate.Shadowed | MaskUpdate.Absent => result.pure[PureConc[E, *]].rethrow @@ -577,7 +585,7 @@ object pure { // todo: MVar is not Serializable, release then update here final class PureFiber[E, A]( val state0: MVar[Outcome[PureConc[E, *], E, A]], - private[this] val canceled0: MVar[CancelationSignal], + private[this] val canceled0: MVar[Unit], private[pure] val masks: MVar[List[MaskFrame]], private[this] val cancelationListeners: MVar[List[CancelationListener[E]]], private[this] val finalizing: MVar[Boolean], @@ -613,9 +621,9 @@ object pure { case Nil => ().pure[PureConc[E, *]] } - private[pure] val currentPollFinalizers: PureConc[E, List[Finalizer[E]]] = - if (activePolls eq null) List.empty[Finalizer[E]].pure[PureConc[E, *]] - else activePolls.read[PureConc[E, *]].map(_.headOption.getOrElse(Nil)) + private[pure] val currentPollFinalizers: PureConc[E, Option[List[Finalizer[E]]]] = + if (activePolls eq null) none[List[Finalizer[E]]].pure[PureConc[E, *]] + else activePolls.read[PureConc[E, *]].map(_.headOption) private[pure] def registerCancelationListener( notify: PureConc[E, Unit]): PureConc[E, CancelationListenerId] = { @@ -715,39 +723,21 @@ object pure { case _ => false } - private[this] def cancelationFinalizers( - signal: CancelationSignal, - ctx: FiberCtx[E]): PureConc[E, List[Finalizer[E]]] = - signal match { - case CancelationSignal.Self => ctx.self.currentPollFinalizers - case CancelationSignal.External => ctx.finalizers.pure[PureConc[E, *]] - } - - private[this] def realizeCancelationWithSignal( - ctx: FiberCtx[E], - signal: CancelationSignal): PureConc[E, Boolean] = - if (ctx.finalizing) false.pure[PureConc[E, *]] - else - isFinalizing.ifM( - finalizationOutcome, - ctx - .self - .currentMasks - .map(_.isEmpty) - .ifM( - cancelationFinalizers(signal, ctx) - .flatMap(finalizers => whileFinalizing(ctx)(finalizeWith(ctx, finalizers))), - false.pure[PureConc[E, *]] - ) - ) - private[pure] def realizeCancelationWith(ctx: FiberCtx[E]): PureConc[E, Boolean] = if (ctx.finalizing) false.pure[PureConc[E, *]] else isFinalizing.ifM( finalizationOutcome, canceled0.tryRead[PureConc[E, *]].flatMap { - case Some(signal) => realizeCancelationWithSignal(ctx, signal) + case Some(_) => + ctx + .self + .currentMasks + .map(_.isEmpty) + .ifM( + whileFinalizing(ctx)(finalizeWith(ctx, ctx.finalizers)), + false.pure[PureConc[E, *]] + ) case None => false.pure[PureConc[E, *]] } ) @@ -762,7 +752,7 @@ object pure { case Nil => isFinalizing.ifM( canceled0.tryRead[PureConc[E, *]].map(_.isEmpty), - canceled0.read[PureConc[E, *]].flatMap(realizeCancelationWithSignal(ctx, _))) + canceled0.read[PureConc[E, *]] *> realizeCancelationWith(ctx)) case _ => isFinalizing.ifM(canceled0.tryRead[PureConc[E, *]].map(_.isEmpty), blocked) @@ -777,21 +767,17 @@ object pure { ctx.self.currentMasks.flatMap { case Nil => whileFinalizing(ctx) { - canceled0.tryPut[PureConc[E, *]](CancelationSignal.Self).flatMap { - case true => - notifyCancelationListeners *> finalizeWith(ctx, ctx.finalizers) - case false => finalizeWith(ctx, ctx.finalizers) - } + requestCancelation *> finalizeWith(ctx, ctx.finalizers) } case _ => - requestCancelation(CancelationSignal.Self).as(false) + requestCancelation.as(false) } ) - private[this] def requestCancelation(signal: CancelationSignal): PureConc[E, Unit] = + private[this] def requestCancelation: PureConc[E, Unit] = canceled0 - .tryPut[PureConc[E, *]](signal) + .tryPut[PureConc[E, *]](()) .flatMap(inserted => if (inserted) notifyCancelationListeners else ().pure[PureConc[E, *]]) @@ -808,7 +794,7 @@ object pure { state.tryRead.flatMap { case Some(_) => ().pure[PureConc[E, *]] case None => - requestCancelation(CancelationSignal.External) *> state.read.void + requestCancelation *> state.read.void } } } diff --git a/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala b/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala index 7bb49d236b..487f5612d7 100644 --- a/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala +++ b/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala @@ -291,6 +291,45 @@ class PureConcSuite assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, Boolean](Some(false))) } + test("select current finalizers for external cancelation inside poll") { + val t = for { + polled <- F.deferred[Unit] + gate <- F.deferred[Unit] + finalized <- F.ref(0) + fiber <- F.start { + F.uncancelable { poll => + poll(F.onCancel(polled.complete(()) *> gate.get, finalized.update(_ + 1))) + } + } + _ <- polled.get + _ <- fiber.cancel + back <- finalized.get + } yield back + + assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, Int](Some(1))) + } + + test("unregister finalizers before observing masked external cancelation") { + val t = for { + masked <- F.deferred[Unit] + gate <- F.deferred[Unit] + finalized <- F.ref(0) + fiber <- F.start { + F.onCancel( + F.uncancelable(_ => masked.complete(()) *> gate.get), + finalized.update(_ + 1)) + } + _ <- masked.get + releaser <- F.start(F.cede *> gate.complete(())) + _ <- fiber.cancel + _ <- releaser.join + outcome <- fiber.join + back <- finalized.get + } yield (outcome === Outcome.canceled[F, Int, Unit], back) + + assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, (Boolean, Int)](Some((true, 0)))) + } + test("run finalizers around a self-canceling polled region") { val t = for { finalized <- F.ref(0) From bc990fd901461b4a09985b4fead7c3624fc2c495 Mon Sep 17 00:00:00 2001 From: Daniel Spiewak Date: Sat, 1 Aug 2026 22:25:03 +0200 Subject: [PATCH 09/10] Fix PureConc cancellation boundaries --- .../cats/effect/kernel/testkit/pure.scala | 717 ++++++++++-------- .../cats/effect/laws/PureConcSuite.scala | 89 ++- .../effect/laws/ResourcePureConcSuite.scala | 56 +- 3 files changed, 549 insertions(+), 313 deletions(-) diff --git a/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala b/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala index bce5877de4..a6b8d195f6 100644 --- a/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala +++ b/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala @@ -43,6 +43,16 @@ object pure { private[pure] final case class MaskFrame(id: MaskId) + private[pure] final class FinalizerId + + private[pure] final case class RegisteredFinalizer[E](id: FinalizerId, action: Finalizer[E]) + + // These are the purely functional analogue of IOFiber's mask and finalizer stacks. + private[pure] final case class FiberState[E]( + frames: List[MaskFrame], + activePolls: Int, + finalizers: List[RegisteredFinalizer[E]]) + private[pure] final class CancelationListenerId private[pure] object CancelationListenerId { @@ -99,13 +109,13 @@ object pure { Outcome.monadError[Id, E] def resolveMain[E, A](pc: PureConc[E, A]): ResolvedPC[E, IdOC[E, A]] = { + val M = rawMonad[E] + /* - * The cancelation implementation is here. The failures of type inference make this look - * HORRIBLE but the general idea is fairly simple: mapK over the FreeT into a new monad - * which sequences a cancelation check within each flatten. Thus, we go from Kleisli[FreeT[Kleisli[Outcome[Id, ...]]]] + * The failures of type inference make this look HORRIBLE but the general idea is fairly + * simple: mapK over the FreeT into a new monad. Thus, we go from Kleisli[FreeT[Kleisli[Outcome[Id, ...]]]] * to Kleisli[FreeT[Kleisli[FreeT[Kleisli[Outcome[Id, ...]]]]]]]], which we then need to go - * through and flatten. The cancelation check *itself* is in `cancelationCheck`, while the flattening - * process is in the definition of `val canceled`. + * through and flatten. The flattening process is in the definition of `val canceled`. * * FlatMapK and TraverseK typeclasses would make this a one-liner. */ @@ -113,26 +123,25 @@ object pure { val cancelationCheck = new (FiberR[E, *] ~> PureConc[E, *]) { def apply[α](ka: FiberR[E, α]): PureConc[E, α] = { val back = Kleisli.ask[IdOC[E, *], FiberCtx[E]] map { ctx => - val checker = ctx - .self - .hasActivePoll - .ifM( - ().pure[PureConc[E, *]], - ctx - .self - .isFinalizing - .ifM( - ().pure[PureConc[E, *]], - ctx - .self - .realizeCancelationWith(ctx) - .ifM(ApplicativeThread[PureConc[E, *]].done[Unit], ().pure[PureConc[E, *]])) - ) - - checker >> mvarLiftF(ThreadT.liftF(ka)) + val checker = + M.flatMap(ctx.self.hasActivePoll) { active => + if (active) M.unit + else + M.flatMap(ctx.self.isFinalizing) { finalizing => + if (finalizing) M.unit + else + M.flatMap(ctx.self.realizeCancelationAtEvaluatorBoundaryWith(ctx)) { + canceled => + if (canceled) ApplicativeThread[PureConc[E, *]].done[Unit] + else M.unit + } + } + } + + M.productR(checker)(mvarLiftF(ThreadT.liftF(ka))) } - mvarLiftF(ThreadT.liftF(back)).flatten + M.flatten(mvarLiftF(ThreadT.liftF(back))) } } @@ -155,65 +164,62 @@ object pure { MVar.empty[Main, Outcome[PureConc[E, *], E, A]].flatMap { state0 => MVar.empty[Main, Unit] flatMap { canceled0 => - MVar[Main, List[MaskFrame]](Nil) flatMap { masks => + MVar[Main, FiberState[E]](FiberState(Nil, 0, Nil)) flatMap { fiberState => MVar[Main, List[CancelationListener[E]]](Nil) flatMap { cancelationListeners => MVar[Main, Boolean](false) flatMap { finalizing => - MVar[Main, List[List[Finalizer[E]]]](Nil) flatMap { activePolls => - val state = state0[Main] - val fiber = - new PureFiber[E, A]( - state0, - canceled0, - masks, - cancelationListeners, - finalizing, - activePolls) - - val completed = canceled.flatMap { a => - withCtx { ctx => - ctx - .self - .realizeCancelationWith(ctx) - .ifM(ApplicativeThread[PureConc[E, *]].done[A], a.pure[PureConc[E, *]]) + val state = state0[Main] + val fiber = + new PureFiber[E, A]( + state0, + canceled0, + fiberState, + cancelationListeners, + finalizing) + + val completed = M.flatMap(canceled) { a => + withCtx { ctx => + M.flatMap(ctx.self.realizeCancelationWith(ctx)) { canceled => + if (canceled) ApplicativeThread[PureConc[E, *]].done[A] + else M.pure(a) } } + } - val identified = completed mapF { ta => - val fk = new (FiberR[E, *] ~> IdOC[E, *]) { - def apply[a](ke: FiberR[E, a]) = - ke.run(FiberCtx(fiber)) - } - - ta.mapK(fk) + val identified = completed mapF { ta => + val fk = new (FiberR[E, *] ~> IdOC[E, *]) { + def apply[a](ke: FiberR[E, a]) = + ke.run(FiberCtx(fiber)) } - import Outcome._ + ta.mapK(fk) + } - val body = identified flatMap { a => - state.tryPut(Succeeded(a.pure[PureConc[E, *]])) - } handleErrorWith { e => state.tryPut(Errored(e)) } + import Outcome._ - val results = state.read.flatMap { - case Canceled() => (Outcome.Canceled(): IdOC[E, A]).pure[Main] - case Errored(e) => (Outcome.Errored(e): IdOC[E, A]).pure[Main] + val body = identified flatMap { a => + state.tryPut(Succeeded(M.pure(a))) + } handleErrorWith { e => state.tryPut(Errored(e)) } - case Succeeded(fa) => - val identifiedCompletion = fa.mapF { ta => - val fk = new (FiberR[E, *] ~> IdOC[E, *]) { - def apply[a](ke: FiberR[E, a]) = - ke.run(FiberCtx(fiber)) - } + val results = state.read.flatMap { + case Canceled() => (Outcome.Canceled(): IdOC[E, A]).pure[Main] + case Errored(e) => (Outcome.Errored(e): IdOC[E, A]).pure[Main] - ta.mapK(fk) + case Succeeded(fa) => + val identifiedCompletion = fa.mapF { ta => + val fk = new (FiberR[E, *] ~> IdOC[E, *]) { + def apply[a](ke: FiberR[E, a]) = + ke.run(FiberCtx(fiber)) } - identifiedCompletion.map(a => - Succeeded[Id, E, A](a): IdOC[E, A]) handleError { e => Errored(e) } - } + ta.mapK(fk) + } - Kleisli.ask[ResolvedPC[E, *], MVar.Universe].map { u => - ApplicativeThread[ResolvedPC[E, *]].start(body.run(u)) >> results.run(u) - } + identifiedCompletion.map(a => + Succeeded[Id, E, A](a): IdOC[E, A]) handleError { e => Errored(e) } + } + + Kleisli.ask[ResolvedPC[E, *], MVar.Universe].map { u => + ApplicativeThread[ResolvedPC[E, *]].start(body.run(u)) >> results.run(u) } } } @@ -258,93 +264,137 @@ object pure { implicit def allocateForPureConc[E]: GenConcurrent[PureConc[E, *], E] = new GenConcurrent[PureConc[E, *], E] { - private[this] val M: MonadError[PureConc[E, *], E] = - Kleisli.catsDataMonadErrorForKleisli + private[this] val M: MonadError[PureConc[E, *], E] = rawMonad[E] private[this] val Thread = ApplicativeThread[PureConc[E, *]] + private[this] val Ask = implicitly[MVar.Ask[PureConc[E, *]]] + + private[this] def emptyMVar[A]: PureConc[E, MVar[A]] = + MVar.empty[PureConc[E, *], A](M, Thread) + + private[this] def mvar[A](a: A): PureConc[E, MVar[A]] = + MVar[PureConc[E, *], A](a)(M, Thread, Ask) + + private[this] def readMVar[A](mvar: MVar[A]): PureConc[E, A] = + mvar.read[PureConc[E, *]](M, Thread, Ask) + + private[this] def tryReadMVar[A](mvar: MVar[A]): PureConc[E, Option[A]] = + mvar.tryRead[PureConc[E, *]](M, Ask) + + private[this] def tryPutMVar[A](mvar: MVar[A], a: A): PureConc[E, Boolean] = + mvar.tryPut[PureConc[E, *]](a)(M, Thread, Ask) + + private[this] def putMVar[A](mvar: MVar[A], a: A): PureConc[E, Unit] = + mvar.put[PureConc[E, *]](a)(M, Thread, Ask) + + private[this] def takeMVar[A](mvar: MVar[A]): PureConc[E, A] = + mvar.take[PureConc[E, *]](M, Thread, Ask) + + private[this] def swapMVar[A](mvar: MVar[A], a: A): PureConc[E, A] = + mvar.swap[PureConc[E, *]](a)(M, Thread, Ask) + + private[this] def cancelationBoundaryWith[A]( + ctx: FiberCtx[E], + fa: => PureConc[E, A]): PureConc[E, A] = + M.flatMap(ctx.self.realizeCancelationWith(ctx)) { canceled => + if (canceled) Thread.done[A] else fa + } + + private[this] def cancelationBoundary[A](fa: => PureConc[E, A]): PureConc[E, A] = + withCtx(ctx => cancelationBoundaryWith(ctx, fa)) + + private[this] def onCancelWith[A]( + fa: PureConc[E, A], + fin: PureConc[E, Unit]): PureConc[E, A] = + withCtx[E, A] { ctx => + // Frame unwinding is structural, so it must not introduce user flatMap boundaries. + M.flatMap(ctx.self.registerFinalizer(M.void(M.attempt(fin)))) { id => + M.flatMap(M.attempt(fa)) { result => + M.flatMap(ctx.self.removeFinalizer(id))(_ => M.rethrow(M.pure(result))) + } + } + } def pure[A](x: A): PureConc[E, A] = M.pure(x) def handleErrorWith[A](fa: PureConc[E, A])(f: E => PureConc[E, A]): PureConc[E, A] = - Thread.annotate("handleErrorWith", true)(M.handleErrorWith(fa)(f)) + Thread.annotate("handleErrorWith", true)( + M.handleErrorWith(fa)(e => cancelationBoundary(f(e)))) def raiseError[A](e: E): PureConc[E, A] = Thread.annotate("raiseError")(M.raiseError(e)) def onCancel[A](fa: PureConc[E, A], fin: PureConc[E, Unit]): PureConc[E, A] = - Thread.annotate("onCancel", true) { - withCtx[E, A] { ctx => - val ctx2 = ctx.withFinalizers(fin.attempt.void :: ctx.finalizers) - localCtx(ctx2, fa) - } - } + Thread.annotate("onCancel", true)(onCancelWith(fa, fin)) def canceled: PureConc[E, Unit] = Thread.annotate("canceled")(withCtx { ctx => - ctx.self.cancelAndRealizeWith(ctx).ifM(Thread.done, unit) + M.flatMap(ctx.self.cancelAndRealizeWith(ctx)) { canceled => + if (canceled) Thread.done else M.unit + } }) def cede: PureConc[E, Unit] = withCtx { ctx => - Thread.cede *> - ctx.self.realizeCancelationWith(ctx).ifM(Thread.done, unit) + M.productR(Thread.cede)(M.flatMap(ctx.self.realizeCancelationWith(ctx)) { canceled => + if (canceled) Thread.done else M.unit + }) } def never[A]: PureConc[E, A] = withCtx[E, A] { ctx => // we monitor for asynchronous cancelation. if we're masked, this won't cancel and we hang - Thread.annotate("never")(ctx.self.awaitCancelationWith(ctx) *> Thread.done) + Thread.annotate("never")(M.productR(ctx.self.awaitCancelationWith(ctx))(Thread.done)) } def ref[A](a: A): PureConc[E, Ref[PureConc[E, *], A]] = - MVar[PureConc[E, *], A](a).flatMap(mVar => Kleisli.pure(unsafeRef(mVar))) + M.map(mvar(a))(unsafeRef(_)) def deferred[A]: PureConc[E, Deferred[PureConc[E, *], A]] = - MVar.empty[PureConc[E, *], A].flatMap(mVar => Kleisli.pure(unsafeDeferred(mVar))) + M.map(emptyMVar[A])(unsafeDeferred(_)) private[this] def interruptible[A](ctx: FiberCtx[E], fa: PureConc[E, A]): PureConc[E, A] = ctx.self.interruptible(ctx)(fa) private def unsafeRef[A](mVar: MVar[A]): Ref[PureConc[E, *], A] = new Ref[PureConc[E, *], A] { - override def get: PureConc[E, A] = mVar.read[PureConc[E, *]] + override def get: PureConc[E, A] = readMVar(mVar) override def set(a: A): PureConc[E, Unit] = modify(_ => (a, ())) override def access: PureConc[E, (A, A => PureConc[E, Boolean])] = uncancelable { _ => - mVar.read[PureConc[E, *]].flatMap { a => - MVar.empty[PureConc[E, *], Unit].map { called => + M.flatMap(readMVar(mVar)) { a => + M.map(emptyMVar[Unit]) { called => val setter = (au: A) => - called - .tryPut[PureConc[E, *]](()) - .ifM( - pure(false), - mVar.take[PureConc[E, *]].flatMap { ay => - if (a == ay) mVar.put[PureConc[E, *]](au).as(true) else pure(false) - }) + M.flatMap(tryPutMVar(called, ())) { alreadyCalled => + if (alreadyCalled) M.pure(false) + else + M.flatMap(takeMVar(mVar)) { ay => + if (a == ay) M.as(putMVar(mVar, au), true) + else M.pure(false) + } + } (a, setter) } } } override def tryUpdate(f: A => A): PureConc[E, Boolean] = - update(f).as(true) + M.as(update(f), true) override def tryModify[B](f: A => (A, B)): PureConc[E, Option[B]] = - modify(f).map(Some(_)) + M.map(modify(f))(Some(_)) override def update(f: A => A): PureConc[E, Unit] = - uncancelable { _ => - mVar.take[PureConc[E, *]].flatMap(a => mVar.put[PureConc[E, *]](f(a))) - } + uncancelable { _ => M.flatMap(takeMVar(mVar))(a => putMVar(mVar, f(a))) } override def modify[B](f: A => (A, B)): PureConc[E, B] = uncancelable { _ => - mVar.take[PureConc[E, *]].flatMap { a => + M.flatMap(takeMVar(mVar)) { a => val (a2, b) = f(a) - mVar.put[PureConc[E, *]](a2).as(b) + M.as(putMVar(mVar, a2), b) } } @@ -362,38 +412,37 @@ object pure { private def unsafeDeferred[A](mVar: MVar[A]): Deferred[PureConc[E, *], A] = new Deferred[PureConc[E, *], A] { override def get: PureConc[E, A] = - withCtx { ctx => interruptible(ctx, mVar.read[PureConc[E, *]]) } + withCtx { ctx => interruptible(ctx, readMVar(mVar)) } - override def complete(a: A): PureConc[E, Boolean] = mVar.tryPut[PureConc[E, *]](a) + override def complete(a: A): PureConc[E, Boolean] = tryPutMVar(mVar, a) - override def tryGet: PureConc[E, Option[A]] = mVar.tryRead[PureConc[E, *]] + override def tryGet: PureConc[E, Option[A]] = tryReadMVar(mVar) } def start[A](fa: PureConc[E, A]): PureConc[E, Fiber[PureConc[E, *], E, A]] = Thread.annotate("start", true) { - MVar.empty[PureConc[E, *], Outcome[PureConc[E, *], E, A]].flatMap { state => - MVar.empty[PureConc[E, *], Unit] flatMap { canceled => - MVar[PureConc[E, *], List[MaskFrame]](Nil) flatMap { masks => - MVar[PureConc[E, *], List[CancelationListener[E]]](Nil) flatMap { - cancelationListeners => - MVar[PureConc[E, *], Boolean](false) flatMap { finalizing => - MVar[PureConc[E, *], List[List[Finalizer[E]]]](Nil) flatMap { - activePolls => - val fiber = - new PureFiber[E, A]( - state, - canceled, - masks, - cancelationListeners, - finalizing, - activePolls) - - // the tryPut here is interesting: it encodes first-wins semantics on cancelation/completion - val body = guaranteeCase(fa)(state.tryPut[PureConc[E, *]](_).void) - val identified = localCtx(FiberCtx(fiber), body) - Thread.start(identified.attempt.void).as(fiber) - } - } + M.flatMap(emptyMVar[Outcome[PureConc[E, *], E, A]]) { state => + M.flatMap(emptyMVar[Unit]) { canceled => + M.flatMap(mvar(FiberState[E](Nil, 0, Nil))) { fiberState => + M.flatMap(mvar(List.empty[CancelationListener[E]])) { cancelationListeners => + M.flatMap(mvar(false)) { finalizing => + val fiber = + new PureFiber[E, A]( + state, + canceled, + fiberState, + cancelationListeners, + finalizing) + + // This is the RunTerminusK analogue: completion is not a user continuation. + val body = + M.handleErrorWith( + M.flatMap(fa)(a => fiber.complete(Outcome.Succeeded(M.pure(a)))))(e => + fiber.complete(Outcome.Errored(e))) + + val identified = localCtx(FiberCtx(fiber), body) + M.as(Thread.start(M.void(M.attempt(identified))), fiber) + } } } } @@ -451,82 +500,99 @@ object pure { Thread.annotate("uncancelable", true) { withCtx { ctx => val mask = new MaskId - val self = ctx.self - def updateMasks[B](f: List[MaskFrame] => (List[MaskFrame], B)): PureConc[E, B] = - self.masks.read[PureConc[E, *]].flatMap { ms => - val (updated, b) = f(ms) - self.masks.swap[PureConc[E, *]](updated).as(b) + + def updateState[B](stateCtx: FiberCtx[E])( + f: FiberState[E] => (FiberState[E], B)): PureConc[E, B] = + M.flatMap(readMVar(stateCtx.self.fiberState)) { state => + val (updated, b) = f(state) + M.as(swapMVar(stateCtx.self.fiberState, updated), b) } - val addF = updateMasks(ms => (MaskFrame(mask) :: ms, ())) - val removeF = updateMasks { - case MaskFrame(`mask`) :: ms => (ms, MaskUpdate.Removed) - case ms if ms.exists(_.id === mask) => (ms, MaskUpdate.Shadowed) - case ms => (ms, MaskUpdate.Absent) - } + val addF = + updateState(ctx)(state => + (state.copy(frames = MaskFrame(mask) :: state.frames), ())) + + val removeF = + updateState(ctx) { state => + state.frames match { + case MaskFrame(`mask`) :: frames => + (state.copy(frames = frames), MaskUpdate.Removed) + + case frames if frames.exists(_.id === mask) => + (state, MaskUpdate.Shadowed) - def restore(update: MaskUpdate) = + case _ => + (state, MaskUpdate.Absent) + } + } + + def enterPoll(callCtx: FiberCtx[E]) = + updateState(callCtx) { state => + state.frames match { + case MaskFrame(`mask`) :: frames => + ( + state.copy(frames = frames, activePolls = state.activePolls + 1), + MaskUpdate.Removed) + + case frames if frames.exists(_.id === mask) => + (state, MaskUpdate.Shadowed) + + case _ => + (state, MaskUpdate.Absent) + } + } + + def restore(callCtx: FiberCtx[E], update: MaskUpdate) = update match { - case MaskUpdate.Removed => self.exitPoll *> addF - case MaskUpdate.Shadowed | MaskUpdate.Absent => unit + case MaskUpdate.Removed => + updateState(callCtx) { state => + val activePolls = math.max(0, state.activePolls - 1) + + ( + state.copy( + frames = MaskFrame(mask) :: state.frames, + activePolls = activePolls), + ()) + } + case MaskUpdate.Shadowed | MaskUpdate.Absent => M.unit } val poll = new Poll[PureConc[E, *]] { def apply[a](fa: PureConc[E, a]) = withCtx { callCtx => if (callCtx.self eq self) - removeF.flatMap { update => - val restoreF = restore(update) - - val enterF = update match { - case MaskUpdate.Removed => self.enterPoll(callCtx.finalizers) - case MaskUpdate.Shadowed | MaskUpdate.Absent => unit + M.flatMap(enterPoll(callCtx)) { update => + val restoreF = restore(callCtx, update) + + update match { + case MaskUpdate.Removed => + onCancelWith( + cancelationBoundaryWith( + callCtx, + M.flatMap(M.attempt(fa)) { result => + M.flatMap(restoreF)(_ => M.rethrow(M.pure(result))) + }), + restoreF) + + case MaskUpdate.Shadowed | MaskUpdate.Absent => + fa } - - enterF *> - localCtx( - callCtx, - onCancel( - self - .realizeCancelationWith(callCtx) - .ifM( - Thread.done, - fa.attempt.flatMap { result => - self - .realizeCancelationWith(callCtx) - .ifM( - Thread.done, - restoreF *> result.pure[PureConc[E, *]].rethrow) - }), - restoreF - ) - ) } - else fa + else + fa } } + // UncancelableK and UnmaskK restore their frames before user continuations run. val runBody = - addF *> body(poll).attempt.flatMap { result => - removeF.flatMap { - case MaskUpdate.Removed => - val back = result.pure[PureConc[E, *]].rethrow - - self.currentPollFinalizers.flatMap { - case Some(finalizers) => - self - .realizeCancelationWith(ctx.withFinalizers(finalizers)) - .ifM(Thread.done, back) - case None => back - } - - case MaskUpdate.Shadowed | MaskUpdate.Absent => - result.pure[PureConc[E, *]].rethrow + M.flatMap(addF) { _ => + M.flatMap(M.attempt(body(poll))) { result => + M.flatMap(removeF)(_ => M.rethrow(M.pure(result))) } } - onCancel(runBody, removeF.void) + onCancelWith(runBody, M.void(removeF)) } } @@ -538,10 +604,10 @@ object pure { Thread.annotate("forceR")(productR(handleError(fa.void)(_ => ()))(fb)) def flatMap[A, B](fa: PureConc[E, A])(f: A => PureConc[E, B]): PureConc[E, B] = - M.flatMap(fa)(f) + M.flatMap(fa)(a => cancelationBoundary(f(a))) def tailRecM[A, B](a: A)(f: A => PureConc[E, Either[A, B]]): PureConc[E, B] = - M.tailRecM(a)(f) + M.tailRecM(a)(a => cancelationBoundary(f(a))) } implicit def eqPureConc[E: Eq, A: Eq]: Eq[PureConc[E, A]] = Eq.by(run(_)) @@ -558,6 +624,9 @@ object pure { private[this] def mvarLiftF[F[_], A](fa: F[A]): MVarR[F, A] = Kleisli.liftF[F, MVar.Universe, A](fa) + private[this] def rawMonad[E]: MonadError[PureConc[E, *], E] = + Kleisli.catsDataMonadErrorForKleisli + // this would actually be a very usful function for FreeT to have private[this] def flattenK[S[_]: Functor, M[_]: Monad, A]( ft: FreeT[S, FreeT[S, M, *], A]): FreeT[S, M, A] = @@ -569,7 +638,8 @@ object pure { // the type inferencer just... fails... completely here private[this] def withCtx[E, A](body: FiberCtx[E] => PureConc[E, A]): PureConc[E, A] = - mvarLiftF(ThreadT.liftF(Kleisli.ask[IdOC[E, *], FiberCtx[E]].map(body))).flatten + rawMonad[E].flatten( + mvarLiftF(ThreadT.liftF(Kleisli.ask[IdOC[E, *], FiberCtx[E]].map(body)))) // ApplicativeAsk[PureConc[E, *], FiberCtx[E]].ask.flatMap(body) private[this] def localCtx[E, A](ctx: FiberCtx[E], around: PureConc[E, A]): PureConc[E, A] = @@ -586,102 +656,135 @@ object pure { final class PureFiber[E, A]( val state0: MVar[Outcome[PureConc[E, *], E, A]], private[this] val canceled0: MVar[Unit], - private[pure] val masks: MVar[List[MaskFrame]], + private[pure] val fiberState: MVar[FiberState[E]], private[this] val cancelationListeners: MVar[List[CancelationListener[E]]], - private[this] val finalizing: MVar[Boolean], - private[this] val activePolls: MVar[List[List[Finalizer[E]]]]) + private[this] val finalizing: MVar[Boolean]) extends Fiber[PureConc[E, *], E, A] with Serializable { + private[this] val M: MonadError[PureConc[E, *], E] = rawMonad[E] + private[this] val Thread = ApplicativeThread[PureConc[E, *]] + private[this] val Ask = implicitly[MVar.Ask[PureConc[E, *]]] + + private[this] def emptyMVar[B]: PureConc[E, MVar[B]] = + MVar.empty[PureConc[E, *], B](M, Thread) + + private[this] def readMVar[B](mvar: MVar[B]): PureConc[E, B] = + mvar.read[PureConc[E, *]](M, Thread, Ask) + + private[this] def tryReadMVar[B](mvar: MVar[B]): PureConc[E, Option[B]] = + mvar.tryRead[PureConc[E, *]](M, Ask) + + private[this] def tryPutMVar[B](mvar: MVar[B], value: B): PureConc[E, Boolean] = + mvar.tryPut[PureConc[E, *]](value)(M, Thread, Ask) + + private[this] def swapMVar[B](mvar: MVar[B], value: B): PureConc[E, B] = + mvar.swap[PureConc[E, *]](value)(M, Thread, Ask) + def this(state0: MVar[Outcome[PureConc[E, *], E, A]]) = - this(state0, null, null, null, null, null) + this(state0, null, null, null, null) + + // Retained for binary compatibility with the former split mask/poll state constructor. + def this( + state0: MVar[Outcome[PureConc[E, *], E, A]], + canceled0: MVar[Unit], + _masks: MVar[List[MaskFrame]], + cancelationListeners: MVar[List[CancelationListener[E]]], + finalizing: MVar[Boolean], + _activePolls: MVar[List[List[Finalizer[E]]]]) = { + this(state0, canceled0, null, cancelationListeners, finalizing) + val _ = (_masks, _activePolls) + } - private[this] val state = state0[PureConc[E, *]] + private[this] val state = state0[PureConc[E, *]](M, Thread, Ask) private[pure] val currentMasks: PureConc[E, List[MaskFrame]] = - if (masks eq null) List.empty[MaskFrame].pure[PureConc[E, *]] - else masks.read[PureConc[E, *]] + if (fiberState eq null) M.pure(List.empty[MaskFrame]) + else M.map(readMVar(fiberState))(_.frames) private[pure] val hasActivePoll: PureConc[E, Boolean] = - if (activePolls eq null) false.pure[PureConc[E, *]] - else activePolls.read[PureConc[E, *]].map(_.nonEmpty) + if (fiberState eq null) M.pure(false) + else M.map(readMVar(fiberState))(_.activePolls > 0) - private[pure] def enterPoll(finalizers: List[Finalizer[E]]): PureConc[E, Unit] = - if (activePolls eq null) ().pure[PureConc[E, *]] + private[pure] def registerFinalizer(action: Finalizer[E]): PureConc[E, FinalizerId] = { + val id = new FinalizerId + + if (fiberState eq null) M.pure(id) else - activePolls.read[PureConc[E, *]].flatMap { polls => - activePolls.swap[PureConc[E, *]](finalizers :: polls).void + M.flatMap(readMVar(fiberState)) { state => + M.as( + swapMVar( + fiberState, + state.copy(finalizers = RegisteredFinalizer(id, action) :: state.finalizers)), + id) } + } - private[pure] val exitPoll: PureConc[E, Unit] = - if (activePolls eq null) ().pure[PureConc[E, *]] + private[pure] def removeFinalizer(id: FinalizerId): PureConc[E, Unit] = + if (fiberState eq null) M.unit else - activePolls.read[PureConc[E, *]].flatMap { - case _ :: polls => activePolls.swap[PureConc[E, *]](polls).void - case Nil => ().pure[PureConc[E, *]] + M.flatMap(readMVar(fiberState)) { state => + M.void( + swapMVar( + fiberState, + state.copy(finalizers = state.finalizers.filterNot(_.id eq id)))) } - private[pure] val currentPollFinalizers: PureConc[E, Option[List[Finalizer[E]]]] = - if (activePolls eq null) none[List[Finalizer[E]]].pure[PureConc[E, *]] - else activePolls.read[PureConc[E, *]].map(_.headOption) + private[pure] def complete(outcome: Outcome[PureConc[E, *], E, A]): PureConc[E, Unit] = + M.productR(setFinalizing(true))(M.void(tryPutMVar(state0, outcome))) private[pure] def registerCancelationListener( notify: PureConc[E, Unit]): PureConc[E, CancelationListenerId] = { val id = new CancelationListenerId - if (cancelationListeners eq null) id.pure[PureConc[E, *]] + if (cancelationListeners eq null) M.pure(id) else - cancelationListeners.read[PureConc[E, *]].flatMap { listeners => - cancelationListeners - .swap[PureConc[E, *]](CancelationListener(id, notify) :: listeners) - .as(id) + M.flatMap(readMVar(cancelationListeners)) { listeners => + M.as(swapMVar(cancelationListeners, CancelationListener(id, notify) :: listeners), id) } } private[pure] def removeCancelationListener(id: CancelationListenerId): PureConc[E, Unit] = - if (cancelationListeners eq null) ().pure[PureConc[E, *]] + if (cancelationListeners eq null) M.unit else - cancelationListeners.read[PureConc[E, *]].flatMap { listeners => - cancelationListeners.swap[PureConc[E, *]](listeners.filterNot(_.id === id)).void + M.flatMap(readMVar(cancelationListeners)) { listeners => + M.void(swapMVar(cancelationListeners, listeners.filterNot(_.id === id))) } private[this] def notifyCancelationListeners: PureConc[E, Unit] = - if (cancelationListeners eq null) ().pure[PureConc[E, *]] - else cancelationListeners.swap[PureConc[E, *]](Nil).flatMap(_.traverse_(_.action)) + if (cancelationListeners eq null) M.unit + else + M.flatMap(swapMVar(cancelationListeners, Nil))( + _.foldLeft(M.unit)((acc, listener) => M.productR(acc)(listener.action))) private[pure] def interruptible[B](ctx: FiberCtx[E])(fb: PureConc[E, B]): PureConc[E, B] = { - val Thread = ApplicativeThread[PureConc[E, *]] - - ctx.self.currentMasks.flatMap { + M.flatMap(ctx.self.currentMasks) { case Nil => - MVar.empty[PureConc[E, *], Option[B]].flatMap { signal => - val notifyCancelation = signal.tryPut[PureConc[E, *]](None).void + M.flatMap(emptyMVar[Option[B]]) { signal => + val notifyCancelation = M.void(tryPutMVar(signal, None)) - ctx.self.registerCancelationListener(notifyCancelation).flatMap { listener => + M.flatMap(ctx.self.registerCancelationListener(notifyCancelation)) { listener => val awaitCompletion = - Thread.start(fb.flatMap(b => signal.tryPut[PureConc[E, *]](Some(b)).void)) + Thread.start(M.flatMap(fb)(b => M.void(tryPutMVar(signal, Some(b))))) val checkCancelation = - signal.tryRead[PureConc[E, *]].flatMap { - case Some(_) => ().pure[PureConc[E, *]] + M.flatMap(tryReadMVar(signal)) { + case Some(_) => M.unit case None => - ctx - .self - .realizeCancelationWith(ctx) - .ifM(notifyCancelation, ().pure[PureConc[E, *]]) + M.flatMap(ctx.self.realizeCancelationWith(ctx)) { canceled => + if (canceled) notifyCancelation else M.unit + } } - awaitCompletion *> - checkCancelation *> - signal.read[PureConc[E, *]].flatMap { + M.productR(awaitCompletion)( + M.productR(checkCancelation)(M.flatMap(readMVar(signal)) { case Some(b) => - ctx.self.removeCancelationListener(listener).as(b) + M.as(ctx.self.removeCancelationListener(listener), b) case None => - ctx.self.removeCancelationListener(listener) *> - ctx.self.realizeCancelationWith(ctx) *> - Thread.done - } + M.productR(ctx.self.removeCancelationListener(listener))( + M.productR(ctx.self.realizeCancelationWith(ctx))(Thread.done)) + })) } } @@ -691,95 +794,116 @@ object pure { } private[pure] val isFinalizing: PureConc[E, Boolean] = - if (finalizing eq null) false.pure[PureConc[E, *]] - else finalizing.read[PureConc[E, *]] + if (finalizing eq null) M.pure(false) + else readMVar(finalizing) private[this] def setFinalizing(value: Boolean): PureConc[E, Unit] = - if (finalizing eq null) ().pure[PureConc[E, *]] - else finalizing.swap[PureConc[E, *]](value).void + if (finalizing eq null) M.unit + else M.void(swapMVar(finalizing, value)) private[this] def finalizeWith( ctx: FiberCtx[E], finalizers: List[PureConc[E, Unit]]): PureConc[E, Boolean] = localCtx( ctx.withFinalizers(Nil).withFinalizing(true), - allocateForPureConc[E].uncancelable(_ => finalizers.sequence_) *> - (state0.tryPut[PureConc[E, *]](Outcome.Canceled()).flatMap { - case true => true.pure[PureConc[E, *]] - case false => - state.read.map { - case Outcome.Canceled() => true - case _ => false - } - } <* setFinalizing(false)) + M.productR(allocateForPureConc[E].uncancelable(_ => + finalizers.foldLeft(M.unit)((acc, finalizer) => M.productR(acc)(finalizer))))( + M.productL( + M.flatMap(tryPutMVar(state0, Outcome.Canceled(): Outcome[PureConc[E, *], E, A])) { + case true => M.pure(true) + case false => + M.map(state.read) { + case Outcome.Canceled() => true + case _ => false + } + })(setFinalizing(false))) ) private[this] def whileFinalizing[B](ctx: FiberCtx[E])(fb: PureConc[E, B]): PureConc[E, B] = - localCtx(ctx.withFinalizers(Nil).withFinalizing(true), setFinalizing(true) *> fb) + localCtx( + ctx.withFinalizers(Nil).withFinalizing(true), + M.productR(setFinalizing(true))(fb)) private[this] def finalizationOutcome: PureConc[E, Boolean] = - state.read.map { + M.map(state.read) { case Outcome.Canceled() => true case _ => false } - private[pure] def realizeCancelationWith(ctx: FiberCtx[E]): PureConc[E, Boolean] = - if (ctx.finalizing) false.pure[PureConc[E, *]] + private[this] def realizeCancelationWith( + ctx: FiberCtx[E], + deferWhileFinalizersRegistered: Boolean): PureConc[E, Boolean] = + if (ctx.finalizing) M.pure(false) else - isFinalizing.ifM( - finalizationOutcome, - canceled0.tryRead[PureConc[E, *]].flatMap { - case Some(_) => - ctx - .self - .currentMasks - .map(_.isEmpty) - .ifM( - whileFinalizing(ctx)(finalizeWith(ctx, ctx.finalizers)), - false.pure[PureConc[E, *]] - ) - case None => false.pure[PureConc[E, *]] - } - ) + M.flatMap(isFinalizing) { finalizing => + if (finalizing) finalizationOutcome + else + M.flatMap(tryReadMVar(canceled0)) { + case Some(_) => + M.flatMap(readMVar(fiberState)) { state => + if (state.frames.isEmpty && + (!deferWhileFinalizersRegistered || state.finalizers.isEmpty)) + whileFinalizing(ctx)(finalizeWith(ctx, state.finalizers.map(_.action))) + else + M.pure(false) + } + case None => M.pure(false) + } + } + + private[pure] def realizeCancelationWith(ctx: FiberCtx[E]): PureConc[E, Boolean] = + realizeCancelationWith(ctx, deferWhileFinalizersRegistered = false) + + // The evaluator runs outside localCtx, so it must let successful finalizer frames unwind. + // Explicit user boundaries call realizeCancelationWith and do not defer. + private[pure] def realizeCancelationAtEvaluatorBoundaryWith( + ctx: FiberCtx[E]): PureConc[E, Boolean] = + realizeCancelationWith(ctx, deferWhileFinalizersRegistered = true) private[pure] def awaitCancelationWith(ctx: FiberCtx[E]): PureConc[E, Boolean] = { def blocked = - MVar.empty[PureConc[E, *], Unit].flatMap(_.read[PureConc[E, *]]).as(false) + M.as(M.flatMap(emptyMVar[Unit])(readMVar), false) if (ctx.finalizing) blocked else - ctx.self.currentMasks.flatMap { + M.flatMap(ctx.self.currentMasks) { case Nil => - isFinalizing.ifM( - canceled0.tryRead[PureConc[E, *]].map(_.isEmpty), - canceled0.read[PureConc[E, *]] *> realizeCancelationWith(ctx)) + M.flatMap(isFinalizing) { finalizing => + if (finalizing) + M.map(tryReadMVar(canceled0))(_.isEmpty) + else + M.productR(readMVar(canceled0))(realizeCancelationWith(ctx)) + } case _ => - isFinalizing.ifM(canceled0.tryRead[PureConc[E, *]].map(_.isEmpty), blocked) + M.flatMap(isFinalizing) { finalizing => + if (finalizing) + M.map(tryReadMVar(canceled0))(_.isEmpty) + else blocked + } } } private[pure] def cancelAndRealizeWith(ctx: FiberCtx[E]): PureConc[E, Boolean] = - if (ctx.finalizing) false.pure[PureConc[E, *]] + if (ctx.finalizing) M.pure(false) else - isFinalizing.ifM( - ctx.self.currentMasks.map(_.isEmpty), - ctx.self.currentMasks.flatMap { - case Nil => - whileFinalizing(ctx) { - requestCancelation *> finalizeWith(ctx, ctx.finalizers) - } - - case _ => - requestCancelation.as(false) - } - ) + M.flatMap(isFinalizing) { finalizing => + if (finalizing) M.map(ctx.self.currentMasks)(_.isEmpty) + else + M.flatMap(readMVar(fiberState)) { state => + if (state.frames.isEmpty) + whileFinalizing(ctx) { + M.productR(requestCancelation)( + finalizeWith(ctx, state.finalizers.map(_.action))) + } + else + M.as(requestCancelation, false) + } + } private[this] def requestCancelation: PureConc[E, Unit] = - canceled0 - .tryPut[PureConc[E, *]](()) - .flatMap(inserted => - if (inserted) notifyCancelationListeners else ().pure[PureConc[E, *]]) + M.flatMap(tryPutMVar(canceled0, ()))(inserted => + if (inserted) notifyCancelationListeners else M.unit) val join: PureConc[E, Outcome[PureConc[E, *], E, A]] = if (canceled0 eq null) state.read @@ -788,13 +912,14 @@ object pure { } val cancel: PureConc[E, Unit] = - if (canceled0 eq null) state.tryPut(Outcome.Canceled()).void + if (canceled0 eq null) + M.void(tryPutMVar(state0, Outcome.Canceled(): Outcome[PureConc[E, *], E, A])) else allocateForPureConc[E].uncancelable { _ => - state.tryRead.flatMap { - case Some(_) => ().pure[PureConc[E, *]] + M.flatMap(tryReadMVar(state0)) { + case Some(_) => M.unit case None => - requestCancelation *> state.read.void + M.productR(requestCancelation)(M.void(state.read)) } } } diff --git a/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala b/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala index 487f5612d7..7c9ca5edac 100644 --- a/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala +++ b/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala @@ -248,6 +248,57 @@ class PureConcSuite assertEquals(forked, Outcome.Succeeded[Option, Int, Unit](None)) } + test("observe pending cancelation before a pure polled action") { + val t = F.uncancelable(poll => F.canceled *> poll(F.unit)) + + assertEquals(pure.run(t), Outcome.Canceled[Option, Int, Unit]()) + } + + test("observe pending cancelation before an error handler") { + val t = for { + handlerRan <- F.ref(false) + finalizerCount <- F.ref(0) + fiber <- F.start { + F.onCancel( + F.handleErrorWith(F.uncancelable(_ => F.canceled *> F.raiseError[Unit](1)))(_ => + handlerRan.set(true)), + finalizerCount.update(_ + 1)) + } + outcome <- fiber.join + handled <- handlerRan.get + finalized <- finalizerCount.get + } yield (outcome === Outcome.canceled[F, Int, Unit], handled, finalized) + + assertEquals( + pure.run(t), + Outcome.Succeeded[Option, Int, (Boolean, Boolean, Int)](Some((true, false, 1)))) + } + + test("observe pending cancelation before the next tailRecM iteration") { + val t = for { + iterationRan <- F.ref(false) + finalizerCount <- F.ref(0) + fiber <- F.start { + F.onCancel( + F.tailRecM[Int, Unit](0) { + case 0 => + F.uncancelable(_ => F.canceled.as(Left(1): Either[Int, Unit])) + case _ => + iterationRan.set(true).as(Right(()): Either[Int, Unit]) + }, + finalizerCount.update(_ + 1) + ) + } + outcome <- fiber.join + iterated <- iterationRan.get + finalized <- finalizerCount.get + } yield (outcome === Outcome.canceled[F, Int, Unit], iterated, finalized) + + assertEquals( + pure.run(t), + Outcome.Succeeded[Option, Int, (Boolean, Boolean, Int)](Some((true, false, 1)))) + } + test("ignore poll from another fiber") { val t = for { started <- F.deferred[Unit] @@ -309,7 +360,7 @@ class PureConcSuite assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, Int](Some(1))) } - test("unregister finalizers before observing masked external cancelation") { + test("allow masked completion after unregistering external cancelation finalizers") { val t = for { masked <- F.deferred[Unit] gate <- F.deferred[Unit] @@ -325,7 +376,7 @@ class PureConcSuite _ <- releaser.join outcome <- fiber.join back <- finalized.get - } yield (outcome === Outcome.canceled[F, Int, Unit], back) + } yield (outcome.isSuccess, back) assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, (Boolean, Int)](Some((true, 0)))) } @@ -356,7 +407,7 @@ class PureConcSuite assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, Int](Some(1))) } - test("run outer finalizers when a masked self-cancel is observed inside poll") { + test("skip outer finalizers when a masked self-cancel reaches the fiber terminus") { val t = for { finalized <- F.ref(0) fiber <- F.start { @@ -364,11 +415,11 @@ class PureConcSuite F.uncancelable { poll => poll(F.uncancelable(_ => F.canceled)) }, finalized.update(_ + 1)) } - _ <- fiber.join + outcome <- fiber.join back <- finalized.get - } yield back + } yield (outcome.isSuccess, back) - assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, Int](Some(1))) + assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, (Boolean, Int)](Some((true, 0)))) } test("observe pending self-cancel before running a polled region") { @@ -412,7 +463,7 @@ class PureConcSuite Outcome.Succeeded[Option, Int, (Int, Int, Boolean)](Some((0, 1, false)))) } - test("select the innermost active poll finalizers") { + test("skip active poll finalizers when a masked self-cancel reaches the fiber terminus") { val t = for { finalized <- F.ref("") fiber <- F.start { @@ -428,14 +479,16 @@ class PureConcSuite finalized.update(_ + "A")) } } - _ <- fiber.join + outcome <- fiber.join back <- finalized.get - } yield back + } yield (outcome.isSuccess, back) - assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, String](Some("BA"))) + assertEquals( + pure.run(t), + Outcome.Succeeded[Option, Int, (Boolean, String)](Some((true, "")))) } - test("restore outer poll finalizers after an inner poll completes") { + test("skip restored poll finalizers when a masked self-cancel reaches the fiber terminus") { val t = for { outerFinalized <- F.ref(0) innerFinalized <- F.ref(0) @@ -451,12 +504,14 @@ class PureConcSuite ) } } - _ <- fiber.join + outcome <- fiber.join outer <- outerFinalized.get inner <- innerFinalized.get - } yield (outer, inner) + } yield (outcome.isSuccess, outer, inner) - assertEquals(pure.run(t), Outcome.Succeeded[Option, Int, (Int, Int)](Some((1, 0)))) + assertEquals( + pure.run(t), + Outcome.Succeeded[Option, Int, (Boolean, Int, Int)](Some((true, 0, 0)))) } test("observe nested self-cancel inside a polled region before continuing") { @@ -572,11 +627,13 @@ class PureConcSuite assertEquals( pure.run( TimeT.run(T.race(TimeT.liftF(F.uncancelable(_ => F.canceled.as(1))), T.never[Unit]))), - Outcome.Canceled[Option, Int, Either[Int, Unit]]()) + Outcome.Succeeded[Option, Int, Either[Int, Unit]](Some(Left(1))) + ) assertEquals( pure.run( TimeT.run(T.race(T.never[Unit], TimeT.liftF(F.uncancelable(_ => F.canceled.as(1)))))), - Outcome.Canceled[Option, Int, Either[Unit, Int]]()) + Outcome.Succeeded[Option, Int, Either[Unit, Int]](Some(Right(1))) + ) assertEquals( pure.run( TimeT.run(T.race(TimeT.liftF(F.start(F.unit).flatMap(_.join).as(1)), T.never[Unit]))), diff --git a/laws/shared/src/test/scala/cats/effect/laws/ResourcePureConcSuite.scala b/laws/shared/src/test/scala/cats/effect/laws/ResourcePureConcSuite.scala index 30acc9c2a0..4f6ee0d4b4 100644 --- a/laws/shared/src/test/scala/cats/effect/laws/ResourcePureConcSuite.scala +++ b/laws/shared/src/test/scala/cats/effect/laws/ResourcePureConcSuite.scala @@ -17,7 +17,7 @@ package cats.effect package laws -import cats.effect.kernel.{MonadCancel, Resource} +import cats.effect.kernel.{GenConcurrent, MonadCancel, Outcome, Resource} import cats.effect.kernel.testkit.{pure, OutcomeGenerators, PureConcGenerators, TestInstances} import cats.laws.discipline.arbitrary._ import cats.syntax.all._ @@ -27,10 +27,64 @@ import org.scalacheck.{Cogen, Prop} import munit.DisciplineSuite class ResourcePureConcSuite extends DisciplineSuite with BaseSuite with TestInstances { + import PureConcGenerators._ import OutcomeGenerators._ import pure._ + test("preserve masking when forceR discards a pure resource") { + type F[A] = PureConc[Throwable, A] + val F = GenConcurrent[F] + + def run(resource: Resource[F, Int]): Outcome[Option, Throwable, Int] = + pure.run(resource.use(F.pure)) + + val target = Resource(F.canceled.as(1 -> F.uncancelable(_ => F.canceled *> F.never[Unit]))) + val forced = Resource.pure[F, Unit](()).forceR(target) + val expected = Outcome.Succeeded[Option, Throwable, Int](None) + + assertEquals(run(target), expected) + assertEquals(run(forced), expected) + } + + test("derive race from racePair for a masked self-canceling resource") { + type F[A] = PureConc[Throwable, A] + val F = GenConcurrent[F] + type R[A] = Resource[F, A] + val R = GenConcurrent[R] + + val fa = Resource(F.canceled.as(1 -> F.uncancelable(_ => F.canceled *> F.never[Unit]))) + val expected = R.race(fa, R.never[Int]) + val received = R.uncancelable { poll => + R.racePair(fa, R.never[Int]).flatMap { + case Left((outcome, fiber)) => + outcome match { + case Outcome.Succeeded(value) => fiber.cancel *> value.map(_.asLeft[Int]) + case Outcome.Errored(error) => fiber.cancel *> R.raiseError(error) + case Outcome.Canceled() => + (fiber.cancel *> fiber.join).flatMap { + case Outcome.Succeeded(value) => value.map(_.asRight[Int]) + case Outcome.Errored(error) => R.raiseError(error) + case Outcome.Canceled() => poll(R.canceled) *> R.never + } + } + case Right((fiber, outcome)) => + outcome match { + case Outcome.Succeeded(value) => fiber.cancel *> value.map(_.asRight[Int]) + case Outcome.Errored(error) => fiber.cancel *> R.raiseError(error) + case Outcome.Canceled() => + (fiber.cancel *> fiber.join).flatMap { + case Outcome.Succeeded(value) => value.map(_.asLeft[Int]) + case Outcome.Errored(error) => R.raiseError(error) + case Outcome.Canceled() => poll(R.canceled) *> R.never + } + } + } + } + + assertEquals(pure.run(expected.use(F.pure)), pure.run(received.use(F.pure))) + } + implicit def exec(sbool: Resource[PureConc[Throwable, *], Boolean]): Prop = Prop( pure From e41274908d09c70f7b0ec98305e300dcc8812527 Mon Sep 17 00:00:00 2001 From: Daniel Spiewak Date: Mon, 3 Aug 2026 13:07:44 +0200 Subject: [PATCH 10/10] Fix PureConc finalizer cancelation masking --- .../cats/effect/kernel/testkit/pure.scala | 140 +++++++++++++----- .../cats/effect/laws/PureConcSuite.scala | 127 ++++++++++++++++ .../effect/laws/ResourcePureConcSuite.scala | 19 ++- 3 files changed, 245 insertions(+), 41 deletions(-) diff --git a/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala b/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala index a6b8d195f6..a2c10fa0c0 100644 --- a/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala +++ b/kernel-testkit/shared/src/main/scala/cats/effect/kernel/testkit/pure.scala @@ -41,17 +41,21 @@ object pure { implicit val eq: Eq[MaskId] = Eq.fromUniversalEquals[MaskId] } - private[pure] final case class MaskFrame(id: MaskId) + private[pure] final case class MaskFrame(id: MaskId, finalizerTail: Boolean = false) private[pure] final class FinalizerId - private[pure] final case class RegisteredFinalizer[E](id: FinalizerId, action: Finalizer[E]) + private[pure] final case class RegisteredFinalizer[E]( + id: FinalizerId, + action: Finalizer[E], + polledMasks: Set[MaskId] = Set.empty) // These are the purely functional analogue of IOFiber's mask and finalizer stacks. private[pure] final case class FiberState[E]( frames: List[MaskFrame], activePolls: Int, - finalizers: List[RegisteredFinalizer[E]]) + finalizers: List[RegisteredFinalizer[E]], + deferredAfterFinalizer: Boolean = false) private[pure] final class CancelationListenerId @@ -178,7 +182,7 @@ object pure { val completed = M.flatMap(canceled) { a => withCtx { ctx => - M.flatMap(ctx.self.realizeCancelationWith(ctx)) { canceled => + M.flatMap(ctx.self.realizeCancelationAtTerminusWith(ctx)) { canceled => if (canceled) ApplicativeThread[PureConc[E, *]].done[A] else M.pure(a) } @@ -514,25 +518,38 @@ object pure { (state.copy(frames = MaskFrame(mask) :: state.frames), ())) val removeF = - updateState(ctx) { state => - state.frames match { - case MaskFrame(`mask`) :: frames => - (state.copy(frames = frames), MaskUpdate.Removed) - - case frames if frames.exists(_.id === mask) => - (state, MaskUpdate.Shadowed) - - case _ => - (state, MaskUpdate.Absent) + M.flatMap(ctx.self.hasPendingCancelation) { canceled => + updateState(ctx) { state => + state.frames match { + case MaskFrame(`mask`, finalizerTail) :: frames => + ( + state.copy( + frames = frames, + deferredAfterFinalizer = + state.deferredAfterFinalizer || (canceled && finalizerTail)), + MaskUpdate.Removed) + + case frames if frames.exists(_.id === mask) => + (state, MaskUpdate.Shadowed) + + case _ => + (state, MaskUpdate.Absent) + } } } def enterPoll(callCtx: FiberCtx[E]) = updateState(callCtx) { state => state.frames match { - case MaskFrame(`mask`) :: frames => + case MaskFrame(`mask`, _) :: frames => ( - state.copy(frames = frames, activePolls = state.activePolls + 1), + state.copy( + frames = frames, + activePolls = state.activePolls + 1, + finalizers = state + .finalizers + .map(finalizer => + finalizer.copy(polledMasks = finalizer.polledMasks + mask))), MaskUpdate.Removed) case frames if frames.exists(_.id === mask) => @@ -706,6 +723,10 @@ object pure { if (fiberState eq null) M.pure(false) else M.map(readMVar(fiberState))(_.activePolls > 0) + private[pure] val hasPendingCancelation: PureConc[E, Boolean] = + if (canceled0 eq null) M.pure(false) + else M.map(tryReadMVar(canceled0))(_.nonEmpty) + private[pure] def registerFinalizer(action: Finalizer[E]): PureConc[E, FinalizerId] = { val id = new FinalizerId @@ -723,15 +744,25 @@ object pure { private[pure] def removeFinalizer(id: FinalizerId): PureConc[E, Unit] = if (fiberState eq null) M.unit else - M.flatMap(readMVar(fiberState)) { state => - M.void( - swapMVar( - fiberState, - state.copy(finalizers = state.finalizers.filterNot(_.id eq id)))) + M.flatMap(hasPendingCancelation) { canceled => + M.flatMap(readMVar(fiberState)) { state => + val removed = state.finalizers.find(_.id eq id) + val finalizers = state.finalizers.filterNot(_.id eq id) + val frames = + (removed, state.frames) match { + case (Some(finalizer), frame :: frames) + if !canceled && finalizer.polledMasks.contains(frame.id) => + frame.copy(finalizerTail = true) :: frames + + case _ => state.frames + } + + M.void(swapMVar(fiberState, state.copy(frames = frames, finalizers = finalizers))) + } } private[pure] def complete(outcome: Outcome[PureConc[E, *], E, A]): PureConc[E, Unit] = - M.productR(setFinalizing(true))(M.void(tryPutMVar(state0, outcome))) + M.productR(markFinalizing)(M.void(tryPutMVar(state0, outcome))) private[pure] def registerCancelationListener( notify: PureConc[E, Unit]): PureConc[E, CancelationListenerId] = { @@ -797,32 +828,31 @@ object pure { if (finalizing eq null) M.pure(false) else readMVar(finalizing) - private[this] def setFinalizing(value: Boolean): PureConc[E, Unit] = + private[this] def markFinalizing: PureConc[E, Unit] = if (finalizing eq null) M.unit - else M.void(swapMVar(finalizing, value)) + else M.void(swapMVar(finalizing, true)) + + private[this] def maskForFinalization: PureConc[E, Unit] = + if (fiberState eq null) M.unit + else + M.flatMap(readMVar(fiberState)) { state => + M.void( + swapMVar(fiberState, state.copy(frames = MaskFrame(new MaskId) :: state.frames))) + } private[this] def finalizeWith( ctx: FiberCtx[E], finalizers: List[PureConc[E, Unit]]): PureConc[E, Boolean] = localCtx( ctx.withFinalizers(Nil).withFinalizing(true), - M.productR(allocateForPureConc[E].uncancelable(_ => - finalizers.foldLeft(M.unit)((acc, finalizer) => M.productR(acc)(finalizer))))( - M.productL( - M.flatMap(tryPutMVar(state0, Outcome.Canceled(): Outcome[PureConc[E, *], E, A])) { - case true => M.pure(true) - case false => - M.map(state.read) { - case Outcome.Canceled() => true - case _ => false - } - })(setFinalizing(false))) + M.productR(finalizers.foldLeft(M.unit)((acc, finalizer) => M.productR(acc)(finalizer)))( + M.as(tryPutMVar(state0, Outcome.Canceled(): Outcome[PureConc[E, *], E, A]), true)) ) private[this] def whileFinalizing[B](ctx: FiberCtx[E])(fb: PureConc[E, B]): PureConc[E, B] = localCtx( ctx.withFinalizers(Nil).withFinalizing(true), - M.productR(setFinalizing(true))(fb)) + M.productR(maskForFinalization)(M.productR(markFinalizing)(fb))) private[this] def finalizationOutcome: PureConc[E, Boolean] = M.map(state.read) { @@ -832,7 +862,8 @@ object pure { private[this] def realizeCancelationWith( ctx: FiberCtx[E], - deferWhileFinalizersRegistered: Boolean): PureConc[E, Boolean] = + deferWhileFinalizersRegistered: Boolean, + deferAfterFinalizer: Boolean): PureConc[E, Boolean] = if (ctx.finalizing) M.pure(false) else M.flatMap(isFinalizing) { finalizing => @@ -842,7 +873,8 @@ object pure { case Some(_) => M.flatMap(readMVar(fiberState)) { state => if (state.frames.isEmpty && - (!deferWhileFinalizersRegistered || state.finalizers.isEmpty)) + (!deferWhileFinalizersRegistered || state.finalizers.isEmpty) && + (!deferAfterFinalizer || !state.deferredAfterFinalizer)) whileFinalizing(ctx)(finalizeWith(ctx, state.finalizers.map(_.action))) else M.pure(false) @@ -852,13 +884,41 @@ object pure { } private[pure] def realizeCancelationWith(ctx: FiberCtx[E]): PureConc[E, Boolean] = - realizeCancelationWith(ctx, deferWhileFinalizersRegistered = false) + if (fiberState eq null) + realizeCancelationWith( + ctx, + deferWhileFinalizersRegistered = false, + deferAfterFinalizer = false) + else + M.flatMap(readMVar(fiberState)) { state => + val consumeDeferred = state.frames.isEmpty && state.deferredAfterFinalizer + val clearDeferred = + if (consumeDeferred) + M.void(swapMVar(fiberState, state.copy(deferredAfterFinalizer = false))) + else + M.unit + + M.productR(clearDeferred)( + realizeCancelationWith( + ctx, + deferWhileFinalizersRegistered = false, + deferAfterFinalizer = false)) + } // The evaluator runs outside localCtx, so it must let successful finalizer frames unwind. // Explicit user boundaries call realizeCancelationWith and do not defer. private[pure] def realizeCancelationAtEvaluatorBoundaryWith( ctx: FiberCtx[E]): PureConc[E, Boolean] = - realizeCancelationWith(ctx, deferWhileFinalizersRegistered = true) + realizeCancelationWith( + ctx, + deferWhileFinalizersRegistered = true, + deferAfterFinalizer = true) + + private[pure] def realizeCancelationAtTerminusWith(ctx: FiberCtx[E]): PureConc[E, Boolean] = + realizeCancelationWith( + ctx, + deferWhileFinalizersRegistered = true, + deferAfterFinalizer = true) private[pure] def awaitCancelationWith(ctx: FiberCtx[E]): PureConc[E, Boolean] = { def blocked = diff --git a/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala b/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala index 7c9ca5edac..91c7f1cf82 100644 --- a/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala +++ b/laws/shared/src/test/scala/cats/effect/laws/PureConcSuite.scala @@ -545,12 +545,139 @@ class PureConcSuite Outcome.Canceled[Option, Int, Unit]()) } + test("allow the owning poll after unregistering a cancelation finalizer") { + val fa = F.uncancelable { poll => + F.onCancel(poll(F.unit), F.never[Unit]) *> poll(F.canceled) *> F.never[Unit] + } + + assertEquals(pure.run(fa), Outcome.Canceled[Option, Int, Unit]()) + } + test("run a guarantee finalizer around a masked self-cancel") { val fa = F.guarantee(F.uncancelable(_ => F.canceled), F.never[Unit]) assertEquals(pure.run(fa), Outcome.Succeeded[Option, Int, Unit](None)) } + test("preserve success when a finalizer self-cancels at the fiber terminus") { + val fa = F.guarantee(F.pure(1), F.canceled) + + assertEquals(pure.run(fa), Outcome.Succeeded[Option, Int, Int](Some(1))) + } + + test("preserve error when a finalizer self-cancels at the fiber terminus") { + val fa = F.guarantee(F.raiseError[Int](42), F.canceled) + + assertEquals(pure.run(fa), Outcome.Errored[Option, Int, Int](42)) + } + + test("run a self-canceling finalizer to completion") { + val fa = for { + finalized <- F.ref(false) + fiber <- F.start(F.guarantee(F.pure(1), F.canceled *> finalized.set(true))) + outcome <- fiber.join + result <- outcome.embedNever + didFinalize <- finalized.get + } yield (result, didFinalize) + + assertEquals( + pure.run(fa), + Outcome.Succeeded[Option, Int, (Int, Boolean)](Some((1, true)))) + } + + test("run a self-canceling error finalizer to completion") { + val fa = for { + finalized <- F.ref(false) + fiber <- F.start(F.guarantee(F.raiseError[Int](42), F.canceled *> finalized.set(true))) + outcome <- fiber.join + didFinalize <- finalized.get + } yield (outcome.fold(false, _ == 42, _ => false), didFinalize) + + assertEquals( + pure.run(fa), + Outcome.Succeeded[Option, Int, (Boolean, Boolean)](Some((true, true)))) + } + + test("observe finalizer self-cancel before the next unmasked continuation") { + val fa = for { + finalized <- F.ref(false) + continued <- F.ref(false) + fiber <- F.start( + F.guarantee(F.unit, F.canceled *> finalized.set(true)) *> + continued.set(true)) + outcome <- fiber.join + didFinalize <- finalized.get + didContinue <- continued.get + } yield (outcome.isCanceled, didFinalize, didContinue) + + assertEquals( + pure.run(fa), + Outcome.Succeeded[Option, Int, (Boolean, Boolean, Boolean)](Some((true, true, false)))) + } + + test("retain finalizer deferral through an enclosing mask") { + val fa = F.uncancelable { _ => + F.guarantee(F.pure(1), F.canceled).flatMap(i => F.pure(i + 1)) + } + + assertEquals(pure.run(fa), Outcome.Succeeded[Option, Int, Int](Some(2))) + } + + test("observe nested finalizer self-cancel before the next unmasked continuation") { + val fa = for { + finalized <- F.ref(false) + maskedContinuation <- F.ref(false) + unmaskedContinuation <- F.ref(false) + fiber <- F.start(F.uncancelable { _ => + F.guarantee(F.unit, F.canceled *> finalized.set(true)) *> + maskedContinuation.set(true) + } *> unmaskedContinuation.set(true)) + outcome <- fiber.join + didFinalize <- finalized.get + didRunMasked <- maskedContinuation.get + didRunUnmasked <- unmaskedContinuation.get + } yield (outcome.isCanceled, didFinalize, didRunMasked, didRunUnmasked) + + assertEquals( + pure.run(fa), + Outcome.Succeeded[Option, Int, (Boolean, Boolean, Boolean, Boolean)]( + Some((true, true, true, false)))) + } + + test("remain cancelable after a successful bracket release") { + val fa = + F.bracketFull(_ => F.unit)(_ => F.pure(1))((_, _) => F.unit) *> + F.canceled *> + F.never[Unit] + + assertEquals(pure.run(fa), Outcome.Canceled[Option, Int, Unit]()) + } + + test("remain cancelable after an errored bracket release") { + val fa = F.flatMap( + F.attempt(F.bracketFull(_ => F.unit)(_ => F.raiseError[Int](42))((_, _) => F.unit))) { + case Left(42) => F.canceled *> F.never[Unit] + case _ => F.raiseError[Unit](0) + } + + assertEquals(pure.run(fa), Outcome.Canceled[Option, Int, Unit]()) + } + + test("ignore the owning poll in a canceled bracket release") { + val fa = for { + finalized <- F.ref(false) + fiber <- F.start(F.bracketFull(poll => F.pure(poll))(_ => F.canceled) { (poll, _) => + poll(F.canceled) *> finalized.set(true) + }) + outcome <- fiber.join + didFinalize <- finalized.get + } yield (outcome.isCanceled, didFinalize) + + assertEquals( + pure.run(fa), + Outcome.Succeeded[Option, Int, (Boolean, Boolean)](Some((true, true)))) + } + test("associate finalizers across an uncancelable boundary") { val left = F.uncancelable(_ => F.onCancel(F.canceled, F.never[Unit])) val right = F.onCancel(F.uncancelable(_ => F.canceled), F.never[Unit]) diff --git a/laws/shared/src/test/scala/cats/effect/laws/ResourcePureConcSuite.scala b/laws/shared/src/test/scala/cats/effect/laws/ResourcePureConcSuite.scala index 4f6ee0d4b4..51e88eacbd 100644 --- a/laws/shared/src/test/scala/cats/effect/laws/ResourcePureConcSuite.scala +++ b/laws/shared/src/test/scala/cats/effect/laws/ResourcePureConcSuite.scala @@ -17,7 +17,8 @@ package cats.effect package laws -import cats.effect.kernel.{GenConcurrent, MonadCancel, Outcome, Resource} +import cats.CommutativeApplicative +import cats.effect.kernel.{GenConcurrent, MonadCancel, Outcome, ParallelF, Resource} import cats.effect.kernel.testkit.{pure, OutcomeGenerators, PureConcGenerators, TestInstances} import cats.laws.discipline.arbitrary._ import cats.syntax.all._ @@ -85,6 +86,22 @@ class ResourcePureConcSuite extends DisciplineSuite with BaseSuite with TestInst assertEquals(pure.run(expected.use(F.pure)), pure.run(received.use(F.pure))) } + test("ignore release self-cancelation through parallel applicative identity") { + type F[A] = PureConc[Throwable, A] + type R[A] = Resource[F, A] + type P[A] = ParallelF[R, A] + + val F = GenConcurrent[F] + val P = CommutativeApplicative[P] + + val resource = Resource(F.pure(1 -> F.canceled)) + val identityApplied = ParallelF.value(P.ap(P.pure((i: Int) => i))(ParallelF(resource))) + val expected = Outcome.Succeeded[Option, Throwable, Int](Some(1)) + + assertEquals(pure.run(resource.use(F.pure)), expected) + assertEquals(pure.run(identityApplied.use(F.pure)), expected) + } + implicit def exec(sbool: Resource[PureConc[Throwable, *], Boolean]): Prop = Prop( pure