diff --git a/scache/src/main/scala/com/evolution/scache/SerialMap.scala b/scache/src/main/scala/com/evolution/scache/SerialMap.scala index 9003b4c2..23ec59ef 100644 --- a/scache/src/main/scala/com/evolution/scache/SerialMap.scala +++ b/scache/src/main/scala/com/evolution/scache/SerialMap.scala @@ -1,8 +1,8 @@ package com.evolution.scache import cats.Applicative +import cats.effect.Concurrent import cats.effect.implicits.* -import cats.effect.{Concurrent, Ref} import cats.syntax.all.* import com.evolutiongaming.catshelper.{Runtime, SerialRef} @@ -122,44 +122,45 @@ object SerialMap { self => def modify[A](key: K)(f: Option[V] => F[(Option[V], A)]) = { - def remove = cache.remove(key) - - def adding(added: Ref[F, Boolean]) = { - for { - _ <- added.set(true) - serialRef <- SerialRef[F].of(State.empty[V]) - } yield serialRef + def remove(serialRef: SerialRef[F, State[V]]) = { + cache.modify(key) { + case Some(current) if current eq serialRef => ((), Cache.Directive.Remove) + case _ => ((), Cache.Directive.Ignore) + } } - def modify(serialRef: SerialRef[F, State[V]], added: Ref[F, Boolean]) = { + def modify(serialRef: SerialRef[F, State[V]]) = { def modify(state: State[V]) = { def onValue(value: Option[V]) = { - f(value).attempt.map { + f(value).attempt.map[(State[V], F[A])] { case Right((Some(value), a)) => val state = State.full(value) val fa = a.pure[F] (state, fa) case Right((None, a)) => - val state = State.removed - val fa = remove.as(a) + val state = State.removed[V] + val fa = remove(serialRef).as(a) + (state, fa) + + case Left(error) if state == State.Empty => + // let [possible] waiting next caller retry with new value + val state = State.removed[V] + val fa = remove(serialRef) *> error.raiseError[F, A] (state, fa) case Left(error) => - val fa = for { - added <- added.get - _ <- if (added) remove.void else ().pure[F] - a <- error.raiseError[F, A] - } yield a + val fa = error.raiseError[F, A] (state, fa) } } def onRemoving = { val state = State.removed[V] - val fa = self.modify(key)(f) + // remove own `serialRef` - it might have not been removed, if it was still in `Loading` state + val fa = remove(serialRef) *> Concurrent[F].cede *> self.modify(key)(f) (state, fa).pure[F] } @@ -176,11 +177,11 @@ object SerialMap { self => } yield a } - for { - added <- Ref[F].of(false) - serialRef <- cache.getOrUpdate(key) { adding(added) } - a <- modify(serialRef, added).uncancelable - } yield a + // uncancelable to not leak `State.Empty` entry + cache + .getOrUpdate(key) { SerialRef[F].of(State.empty[V]) } + .flatMap { serialRef => modify(serialRef) } + .uncancelable } def update[A](key: K)(f: Option[V] => F[Option[V]]) = { diff --git a/scache/src/test/scala/com/evolution/scache/SerialMapSpec.scala b/scache/src/test/scala/com/evolution/scache/SerialMapSpec.scala index 0040a68e..20d1ffe0 100644 --- a/scache/src/test/scala/com/evolution/scache/SerialMapSpec.scala +++ b/scache/src/test/scala/com/evolution/scache/SerialMapSpec.scala @@ -1,6 +1,8 @@ package com.evolution.scache -import cats.effect.{Async, Concurrent, Deferred, IO, Outcome} +import cats.MonadThrow +import cats.effect.{Async, Concurrent, Deferred, IO, Outcome, Ref} +import cats.kernel.CommutativeMonoid import cats.syntax.all.* import com.evolution.scache.IOSuite.* import com.evolutiongaming.catshelper.CatsHelper.* @@ -8,6 +10,7 @@ import com.evolutiongaming.catshelper.SerialRef import org.scalatest.funsuite.AsyncFunSuite import org.scalatest.matchers.should.Matchers +import scala.concurrent.duration.* import scala.util.control.NoStackTrace class SerialMapSpec extends AsyncFunSuite with Matchers { @@ -61,6 +64,34 @@ class SerialMapSpec extends AsyncFunSuite with Matchers { `not leak on failures`[IO].run() } + test("not lose concurrent update when entry creator fails") { + `not lose concurrent update when entry creator fails`[IO].run() + } + + test("not lose concurrent update when modify of existing value fails") { + `not lose concurrent update when modify of existing value fails`[IO].run() + } + + test("not lose concurrent put completed before entry creator acquired permit") { + `not lose concurrent put completed before entry creator acquired permit`.run() + } + + test("not lose put made after clear when entry creator fails") { + `not lose put made after clear when entry creator fails`[IO].run() + } + + test("not lose put made after clear when entry creator removes value") { + `not lose put made after clear when entry creator removes value`[IO].run() + } + + test("not leak entry when creator is canceled before acquiring permit") { + `not leak entry when creator is canceled before acquiring permit`.run() + } + + test("not get stuck on removed entry left in cache while it was loading") { + `not get stuck on removed entry left in cache while it was loading`.run() + } + test("modify serially for the same key") { `modify serially for the same key`[IO].run() } @@ -283,6 +314,191 @@ class SerialMapSpec extends AsyncFunSuite with Matchers { } } + private def `not lose concurrent update when entry creator fails`[F[_]: Async] = { + val key = "key" + for { + serialMap <- SerialMap.of[F, String, Int] + blocked <- Deferred[F, Unit] + acquired <- Deferred[F, Unit] + value0 = serialMap.modify(key) { _ => + for { + _ <- acquired.complete(()) + _ <- blocked.get + a <- TestError.raiseError[F, (Option[Int], Unit)] + } yield a + } + value0 <- value0.attempt.startEnsure + _ <- acquired.get + value1 <- serialMap.put(key, 1).startEnsure + // make `put` lock on soon to fail `modify` on `value0` + _ <- Async[F].sleep(100.millis) + _ <- blocked.complete(()) + value0 <- value0.join + value1 <- value1.join + value2 <- serialMap.get(key) + } yield { + value0 shouldEqual Outcome.succeeded(IO.pure(TestError.asLeft)) + value1 shouldEqual Outcome.succeeded(IO.pure(none[Int])) + value2 shouldEqual 1.some + } + } + + private def `not lose concurrent update when modify of existing value fails`[F[_]: Async] = { + val key = "key" + for { + serialMap <- SerialMap.of[F, String, Int] + _ <- serialMap.put(key, 0) + blocked <- Deferred[F, Unit] + acquired <- Deferred[F, Unit] + value0 = serialMap.modify(key) { _ => + for { + _ <- acquired.complete(()) + _ <- blocked.get + a <- TestError.raiseError[F, (Option[Int], Unit)] + } yield a + } + value0 <- value0.attempt.startEnsure + _ <- acquired.get + value1 <- serialMap.put(key, 1).startEnsure + // make `put` lock on soon to fail `modify` on `value0` + _ <- Async[F].sleep(100.millis) + _ <- blocked.complete(()) + value0 <- value0.join + value1 <- value1.join + value2 <- serialMap.get(key) + } yield { + value0 shouldEqual Outcome.succeeded(IO.pure(TestError.asLeft)) + value1 shouldEqual Outcome.succeeded(IO.pure(0.some)) + value2 shouldEqual 1.some + } + } + + private def `not lose concurrent put completed before entry creator acquired permit` = { + val key = "key" + Cache.loading[IO, String, SerialRef[IO, SerialMap.State[Int]]].use { cache => + for { + published <- Deferred[IO, Unit] + proceed <- Deferred[IO, Unit] + // pause the first caller between: setting value in cache AND it acquires the permit to modify it + pausing <- onFirstGetOrUpdate(cache) { published.complete(()) *> proceed.get.void } + serialMap = SerialMap(pausing) + value0 <- serialMap + .modify(key) { _ => TestError.raiseError[IO, (Option[Int], Unit)] } + .attempt + .start + _ <- published.get + value1 <- serialMap.put(key, 1) + _ <- proceed.complete(()) + value0 <- value0.joinWithNever + value2 <- serialMap.get(key) + } yield { + value0 shouldEqual TestError.asLeft + value1 shouldEqual none[Int] + value2 shouldEqual 1.some + } + } + } + + private def `not lose put made after clear when entry creator fails`[F[_]: Async] = { + val key = "key" + for { + serialMap <- SerialMap.of[F, String, Int] + blocked <- Deferred[F, Unit] + acquired <- Deferred[F, Unit] + value0 = serialMap.modify(key) { _ => + for { + _ <- acquired.complete(()) + _ <- blocked.get + a <- TestError.raiseError[F, (Option[Int], Unit)] + } yield a + } + value0 <- value0.attempt.startEnsure + _ <- acquired.get + // drop the `serialRef` of `value0` from the cache while `modify` is in not finished + _ <- serialMap.clear + // add new `serialRef` with value `Full(1)` + value1 <- serialMap.put(key, 1) + _ <- blocked.complete(()) + value0 <- value0.join + value2 <- serialMap.get(key) + } yield { + value0 shouldEqual Outcome.succeeded(IO.pure(TestError.asLeft)) + value1 shouldEqual none[Int] + value2 shouldEqual 1.some + } + } + + private def `not lose put made after clear when entry creator removes value`[F[_]: Async] = { + val key = "key" + for { + serialMap <- SerialMap.of[F, String, Int] + blocked <- Deferred[F, Unit] + acquired <- Deferred[F, Unit] + value0 = serialMap.modify(key) { _ => + for { + _ <- acquired.complete(()) + _ <- blocked.get + } yield (none[Int], ()) + } + value0 <- value0.startEnsure + _ <- acquired.get + // drop the `serialRef` of `value0` from the cache while `modify` is in not finished + _ <- serialMap.clear + // add new `serialRef` with value `Full(1)` + value1 <- serialMap.put(key, 1) + _ <- blocked.complete(()) + value0 <- value0.join + value2 <- serialMap.get(key) + } yield { + value0 shouldEqual Outcome.succeeded(IO.pure(())) + value1 shouldEqual none[Int] + value2 shouldEqual 1.some + } + } + + private def `not leak entry when creator is canceled before acquiring permit` = { + val key = "key" + Cache.loading[IO, String, SerialRef[IO, SerialMap.State[Int]]].use { cache => + for { + // cancel the first caller right after the entry is set in cache, before it acquires the permit to modify it + canceling <- onFirstGetOrUpdate(cache) { IO.canceled } + serialMap = SerialMap(canceling) + value0 <- serialMap + .modify(key) { _ => TestError.raiseError[IO, (Option[Int], Unit)] } + .start + value0 <- value0.join + keys <- serialMap.keys + size <- serialMap.size + } yield { + // creator is not interrupted, it proceeds to run `f` and to clean up on its failure + value0 shouldEqual Outcome.errored[IO, Throwable, Unit](TestError) + keys shouldEqual Set.empty + size shouldEqual 0 + } + } + } + + private def `not get stuck on removed entry left in cache while it was loading` = { + val key = "key" + Cache.loading[IO, String, SerialRef[IO, SerialMap.State[Int]]].use { cache => + for { + // `LoadingCache.modify` passes `None` to `f` for an entry still in `Loading` state, + // which happens when a waiter wakes up between deferred completion and entry being set to `Value` + loading <- onFirstModify(cache) { _ => None } + serialMap = SerialMap(loading) + // the `serialRef` becomes `Removed`, but removal from cache is ignored as entry seems to be loading + value0 <- serialMap.modify(key) { _ => (none[Int], ()).pure[IO] } + value1 <- serialMap.put(key, 1).start + value1 <- value1.join.timeout(1.second) + value2 <- serialMap.get(key) + } yield { + value0 shouldEqual (()) + value1 shouldEqual Outcome.succeeded(IO.pure(none[Int])) + value2 shouldEqual 1.some + } + } + } + private def `modify serially for the same key`[F[_]: Async] = { val key = "key" for { @@ -343,4 +559,63 @@ object SerialMapSpec { } case object TestError extends RuntimeException with NoStackTrace + + // run `hook` after the first `getOrUpdate` call returns, after the first entry has been stored in cache + def onFirstGetOrUpdate[K, V](cache: Cache[IO, K, V])(hook: IO[Unit]): IO[Cache[IO, K, V]] = { + Ref[IO].of(true).map { first => + new DelegatingCache(cache) { + override def getOrUpdate(key: K)(value: => IO[V]) = { + super.getOrUpdate(key)(value).flatTap { _ => + first.getAndSet(false).flatMap { first => hook.whenA(first) } + } + } + } + } + } + + // pass value transformed by `hook` to `f` in the first `modify` call + def onFirstModify[K, V](cache: Cache[IO, K, V])(hook: Option[V] => Option[V]): IO[Cache[IO, K, V]] = { + Ref[IO].of(true).map { first => + new DelegatingCache(cache) { + override def modify[A](key: K)(f: Option[V] => (A, Cache.Directive[IO, V])) = { + first.getAndSet(false).flatMap { first => + super.modify(key) { value => f(if (first) hook(value) else value) } + } + } + } + } + } + + class DelegatingCache[F[_]: MonadThrow, K, V](cache: Cache[F, K, V]) extends Cache.Abstract1[F, K, V] { + + def get(key: K) = cache.get(key) + + def get1(key: K) = cache.get1(key) + + def getOrUpdate(key: K)(value: => F[V]) = cache.getOrUpdate(key)(value) + + def getOrUpdate1[A](key: K)(value: => F[(A, V, Option[Release])]) = cache.getOrUpdate1(key)(value) + + def put(key: K, value: V, release: Option[Release]) = cache.put(key, value, release) + + def modify[A](key: K)(f: Option[V] => (A, Directive[F, V])) = cache.modify(key)(f) + + def contains(key: K) = cache.contains(key) + + def size = cache.size + + def keys = cache.keys + + def values = cache.values + + def values1 = cache.values1 + + def remove(key: K) = cache.remove(key) + + def clear: F[Released] = cache.clear + + def foldMap[A: CommutativeMonoid](f: (K, Either[F[V], V]) => F[A]) = cache.foldMap(f) + + def foldMapPar[A: CommutativeMonoid](f: (K, Either[F[V], V]) => F[A]) = cache.foldMapPar(f) + } }