diff --git a/benchmarks/ce/src/main/scala/sage/benchmarks/CeBenchmarks.scala b/benchmarks/ce/src/main/scala/sage/benchmarks/CeBenchmarks.scala index bbb137e4..44ac2457 100644 --- a/benchmarks/ce/src/main/scala/sage/benchmarks/CeBenchmarks.scala +++ b/benchmarks/ce/src/main/scala/sage/benchmarks/CeBenchmarks.scala @@ -1,46 +1,13 @@ package sage.benchmarks -import java.util.concurrent.TimeUnit +import org.openjdk.jmh.annotations.Param -import org.openjdk.jmh.annotations.* - -@State(Scope.Benchmark) -@BenchmarkMode(Array(Mode.Throughput)) -@OutputTimeUnit(TimeUnit.SECONDS) -@OperationsPerInvocation(1000) // = Payloads.KeyCount -@Fork(1) -@Warmup(iterations = 3, time = 3) -@Measurement(iterations = 5, time = 3) -class ThroughputBench extends RedisBenchState { +class ThroughputBench extends ThroughputBenchBase { @Param(Array("sage-ce", "redis4cats")) var client: String = "sage-ce" - @Param(Array("1", "8", "64", "256")) var concurrency: Int = 1 - @Param(Array("16", "1024")) var valueSize: Int = 16 - - protected def subjectName: String = client - override protected def seedValueBytes: Int = valueSize - protected def buildClient(host: String, port: Int, name: String): BenchClient = Clients.build(host, port, name) - - @Benchmark def get(): Long = subject.getAll(keys, concurrency) - @Benchmark def set(): Long = subject.setAll(keys, Payloads.value(valueSize), concurrency) } -// Run one large-reply command per invocation. The throughput benchmark covers concurrent requests. valueSize controls the seeded value size. -@State(Scope.Benchmark) -@BenchmarkMode(Array(Mode.Throughput)) -@OutputTimeUnit(TimeUnit.SECONDS) -@Fork(1) -@Warmup(iterations = 3, time = 3) -@Measurement(iterations = 5, time = 3) -class CollectionBench extends RedisBenchState { +class CollectionBench extends CollectionBenchBase { @Param(Array("sage-ce", "redis4cats")) var client: String = "sage-ce" - @Param(Array("16")) var valueSize: Int = 16 - - protected def subjectName: String = client - override protected def seedValueBytes: Int = valueSize - protected def buildClient(host: String, port: Int, name: String): BenchClient = Clients.build(host, port, name) - - @Benchmark def mget(): Long = subject.mget(keys) - @Benchmark def hgetall(): Long = subject.hgetall(Payloads.HashKey) } diff --git a/benchmarks/ce/src/main/scala/sage/benchmarks/Clients.scala b/benchmarks/ce/src/main/scala/sage/benchmarks/Clients.scala index 267497dd..1bdfda61 100644 --- a/benchmarks/ce/src/main/scala/sage/benchmarks/Clients.scala +++ b/benchmarks/ce/src/main/scala/sage/benchmarks/Clients.scala @@ -21,77 +21,29 @@ object Clients { } } -final class SageCeBench(host: String, port: Int) extends BenchClient { +final class SageCeBench(host: String, port: Int) extends SageBench[IO] { - private val client: SageClient = + protected val client: SageClient = SageClient.connect(SageConfig(topology = Topology.Standalone(Endpoint(host, port)))).unsafeRunSync() - def name: String = "sage-ce" + protected def run[A](effect: IO[A]): Unit = effect.unsafeRunSync(): Unit - def seed(prefix: String, count: Int, value: String, hashKey: String, fields: Int): Unit = { - val sets = (0 until count).toList.traverse_(i => client.set(s"$prefix:$i", value)) - val hash = (0 until fields).map(i => (s"f$i", value)).toList match { - case h :: t => client.hSet(hashKey, h, t*).void - case Nil => IO.unit - } - (sets *> hash).unsafeRunSync() - } - - def getAll(keys: Array[String], concurrency: Int): Long = - Payloads - .groups(keys, concurrency) - .toList - .parTraverse(_.toList.traverse(client.get[String])) - .map(_.flatten.flatten.map(_.length.toLong).sum) - .unsafeRunSync() - - def setAll(keys: Array[String], value: String, concurrency: Int): Long = - Payloads - .groups(keys, concurrency) - .toList - .parTraverse_(_.toList.traverse_(client.set(_, value))) - .as(keys.length.toLong) - .unsafeRunSync() - - def mget(keys: Array[String]): Long = - client.mGet[String](keys.head, keys.tail*).map(_.flatten.map(_.length.toLong).sum).unsafeRunSync() - - def hgetall(key: String): Long = client.hGetAll[String, String](key).map(_.size.toLong).unsafeRunSync() - - def close(): Unit = client.close.unsafeRunSync() + protected def inLanes[A](work: Payloads.Workload)(perKey: String => IO[A]): IO[Unit] = work.lanes.parTraverse_(_.traverse_(perKey)) } final class Redis4catsBench(host: String, port: Int) extends BenchClient { private val (redis, release) = Redis[IO].utf8(s"redis://$host:$port").allocated.unsafeRunSync() - def name: String = "redis4cats" - - def seed(prefix: String, count: Int, value: String, hashKey: String, fields: Int): Unit = { - val sets = (0 until count).toList.traverse_(i => redis.set(s"$prefix:$i", value)) - val hash = (0 until fields).toList.traverse_(i => redis.hSet(hashKey, s"f$i", value)) - (sets *> hash).unsafeRunSync() - } - - def getAll(keys: Array[String], concurrency: Int): Long = - Payloads - .groups(keys, concurrency) - .toList - .parTraverse(_.toList.traverse(redis.get)) - .map(_.flatten.flatten.map(_.length.toLong).sum) - .unsafeRunSync() + def getAll(work: Payloads.Workload): Unit = + work.lanes.parTraverse_(_.traverse_(redis.get)).unsafeRunSync() - def setAll(keys: Array[String], value: String, concurrency: Int): Long = - Payloads - .groups(keys, concurrency) - .toList - .parTraverse_(_.toList.traverse_(redis.set(_, value))) - .as(keys.length.toLong) - .unsafeRunSync() + def setAll(work: Payloads.Workload, value: String): Unit = + work.lanes.parTraverse_(_.traverse_(redis.set(_, value))).unsafeRunSync() - def mget(keys: Array[String]): Long = redis.mGet(keys.toSet).map(_.values.map(_.length.toLong).sum).unsafeRunSync() + def mget(): Unit = redis.mGet(Payloads.Keys.set).void.unsafeRunSync() - def hgetall(key: String): Long = redis.hGetAll(key).map(_.size.toLong).unsafeRunSync() + def hgetall(): Unit = redis.hGetAll(Payloads.HashKey).void.unsafeRunSync() def close(): Unit = release.unsafeRunSync() } diff --git a/benchmarks/future/src/main/scala/sage/benchmarks/Clients.scala b/benchmarks/future/src/main/scala/sage/benchmarks/Clients.scala new file mode 100644 index 00000000..b8142cd4 --- /dev/null +++ b/benchmarks/future/src/main/scala/sage/benchmarks/Clients.scala @@ -0,0 +1,5 @@ +package sage.benchmarks + +object Clients { + def build(host: String, port: Int, name: String): BenchClient = throw new IllegalArgumentException(s"unknown client: $name") +} diff --git a/benchmarks/kyo/src/main/scala/sage/benchmarks/Clients.scala b/benchmarks/kyo/src/main/scala/sage/benchmarks/Clients.scala index cd770963..00c0413f 100644 --- a/benchmarks/kyo/src/main/scala/sage/benchmarks/Clients.scala +++ b/benchmarks/kyo/src/main/scala/sage/benchmarks/Clients.scala @@ -22,37 +22,13 @@ private object Run { KyoApp.Unsafe.runAndBlock(Duration.Infinity)(program).getOrThrow } -final class SageKyoBench(host: String, port: Int) extends BenchClient { +final class SageKyoBench(host: String, port: Int) extends SageBench[[A] =>> A < (Abort[SageException] & Async)] { - private val client: SageClient = + protected val client: SageClient = Run(SageClient.connect(SageConfig(topology = Topology.Standalone(Endpoint(host, port))))) - def name: String = "sage-kyo" + protected def run[A](effect: A < (Abort[SageException] & Async)): Unit = Run(effect): Unit - def seed(prefix: String, count: Int, value: String, hashKey: String, fields: Int): Unit = - Run(for { - _ <- Kyo.foreachDiscard(0 until count)(i => client.set(s"$prefix:$i", value)) - _ <- Kyo.foreachDiscard(0 until fields)(i => client.hSet(hashKey, (s"f$i", value))) - } yield ()) - - def getAll(keys: Array[String], concurrency: Int): Long = - Run( - Async - .foreach(Payloads.groups(keys, concurrency).toList, concurrency)(g => Kyo.foreach(g.toList)(k => client.get[String](k)).map(_.toList)) - .map(_.toList.flatten.flatten.map(_.length.toLong).sum) - ) - - def setAll(keys: Array[String], value: String, concurrency: Int): Long = - Run( - Async - .foreachDiscard(Payloads.groups(keys, concurrency).toList, concurrency)(g => Kyo.foreachDiscard(g.toList)(k => client.set(k, value))) - .map(_ => keys.length.toLong) - ) - - def mget(keys: Array[String]): Long = - Run(client.mGet[String](keys.head, keys.tail*).map(_.flatten.map(_.length.toLong).sum)) - - def hgetall(key: String): Long = Run(client.hGetAll[String, String](key).map(_.size.toLong)) - - def close(): Unit = Run(client.close) + protected def inLanes[A](work: Payloads.Workload)(perKey: String => A < (Abort[SageException] & Async)): Unit < (Abort[SageException] & Async) = + Async.foreachDiscard(work.lanes, work.concurrency)(Kyo.foreachDiscard(_)(perKey)) } diff --git a/benchmarks/kyo/src/main/scala/sage/benchmarks/KyoBenchmarks.scala b/benchmarks/kyo/src/main/scala/sage/benchmarks/KyoBenchmarks.scala index ce699fbb..275cdc7f 100644 --- a/benchmarks/kyo/src/main/scala/sage/benchmarks/KyoBenchmarks.scala +++ b/benchmarks/kyo/src/main/scala/sage/benchmarks/KyoBenchmarks.scala @@ -1,46 +1,13 @@ package sage.benchmarks -import java.util.concurrent.TimeUnit +import org.openjdk.jmh.annotations.Param -import org.openjdk.jmh.annotations.* +class ThroughputBench extends ThroughputBenchBase { -@State(Scope.Benchmark) -@BenchmarkMode(Array(Mode.Throughput)) -@OutputTimeUnit(TimeUnit.SECONDS) -@OperationsPerInvocation(1000) // = Payloads.KeyCount -@Fork(1) -@Warmup(iterations = 3, time = 3) -@Measurement(iterations = 5, time = 3) -class ThroughputBench extends RedisBenchState { - - @Param(Array("sage-kyo")) var client: String = "sage-kyo" - @Param(Array("1", "8", "64", "256")) var concurrency: Int = 1 - @Param(Array("16", "1024")) var valueSize: Int = 16 - - protected def subjectName: String = client - override protected def seedValueBytes: Int = valueSize - protected def buildClient(host: String, port: Int, name: String): BenchClient = Clients.build(host, port, name) - - @Benchmark def get(): Long = subject.getAll(keys, concurrency) - @Benchmark def set(): Long = subject.setAll(keys, Payloads.value(valueSize), concurrency) + @Param(Array("sage-kyo")) var client: String = "sage-kyo" } -// Run one large-reply command per invocation. The throughput benchmark covers concurrent requests. valueSize controls the seeded value size. -@State(Scope.Benchmark) -@BenchmarkMode(Array(Mode.Throughput)) -@OutputTimeUnit(TimeUnit.SECONDS) -@Fork(1) -@Warmup(iterations = 3, time = 3) -@Measurement(iterations = 5, time = 3) -class CollectionBench extends RedisBenchState { +class CollectionBench extends CollectionBenchBase { @Param(Array("sage-kyo")) var client: String = "sage-kyo" - @Param(Array("16")) var valueSize: Int = 16 - - protected def subjectName: String = client - override protected def seedValueBytes: Int = valueSize - protected def buildClient(host: String, port: Int, name: String): BenchClient = Clients.build(host, port, name) - - @Benchmark def mget(): Long = subject.mget(keys) - @Benchmark def hgetall(): Long = subject.hgetall(Payloads.HashKey) } diff --git a/benchmarks/ox/src/main/scala/sage/benchmarks/Clients.scala b/benchmarks/ox/src/main/scala/sage/benchmarks/Clients.scala index 8ac15028..919be1ac 100644 --- a/benchmarks/ox/src/main/scala/sage/benchmarks/Clients.scala +++ b/benchmarks/ox/src/main/scala/sage/benchmarks/Clients.scala @@ -1,15 +1,14 @@ package sage.benchmarks -import java.util.concurrent.{CountDownLatch, Executors} -import java.util.concurrent.atomic.{AtomicInteger, AtomicLong, AtomicReference} +import java.util.concurrent.{CompletableFuture, CountDownLatch, Executors} +import java.util.concurrent.atomic.{AtomicInteger, AtomicReference} import scala.concurrent.{Await, Future} import scala.concurrent.duration.* -import scala.jdk.CollectionConverters.* -import scala.util.{Failure, Success} +import scala.util.Failure -import _root_.ox.{fork, supervised} -import io.lettuce.core.{RedisClient, RedisFuture} +import _root_.ox.{fork, supervised, Ox} +import io.lettuce.core.RedisClient import org.apache.pekko.actor.ActorSystem import redis.clients.jedis.{DefaultJedisClientConfig, HostAndPort, JedisPool, RedisProtocol} @@ -34,62 +33,27 @@ object Clients { * Sage's Ox API requires an `Ox` scope. A holder fiber keeps the client's scope open for the benchmark lifetime, while each benchmark * operation runs in a short-lived supervised scope and shares the same connection. */ -final class SageOxBench(host: String, port: Int) extends BenchClient { - - @volatile private var client: SageClient = null - private val ready = new CountDownLatch(1) - private val shutdown = new CountDownLatch(1) - - private val holder = Thread.ofVirtual().start { () => - supervised { - client = SageClient.connect(SageConfig(topology = Topology.Standalone(Endpoint(host, port)))) - ready.countDown() - shutdown.await() - try client.close - catch { case _: Throwable => () } - } - } - ready.await() - - def name: String = "sage-ox" - - def seed(prefix: String, count: Int, value: String, hashKey: String, fields: Int): Unit = supervised { - (0 until count).foreach { i => - client.set(s"$prefix:$i", value) - } - (0 until fields).foreach { i => - client.hSet(hashKey, (s"f$i", value)) - } - } +final class SageOxBench(host: String, port: Int) extends SageBench[[A] =>> Ox ?=> A] { - def getAll(keys: Array[String], concurrency: Int): Long = supervised { - Payloads - .groups(keys, concurrency) - .toList - .map(g => fork(g.foldLeft(0L)((t, k) => t + client.get[String](k).fold(0L)(_.length.toLong)))) - .map(_.join()) - .sum - } + private val opened = new CompletableFuture[SageClient] + private val shutdown = new CountDownLatch(1) - def setAll(keys: Array[String], value: String, concurrency: Int): Long = supervised { - Payloads - .groups(keys, concurrency) - .toList - .map(g => - fork(g.foldLeft(0L) { (n, k) => - client.set(k, value) - n + 1 - }) - ) - .map(_.join()) - .sum + private val holder = Thread.ofVirtual().start { () => + try + supervised { + opened.complete(SageClient.scoped(SageConfig(topology = Topology.Standalone(Endpoint(host, port))))): Unit + shutdown.await() + } + catch { case t: Throwable => opened.completeExceptionally(t): Unit } } + protected val client: SageClient = opened.join() - def mget(keys: Array[String]): Long = supervised(client.mGet[String](keys.head, keys.tail*).flatten.map(_.length.toLong).sum) + protected def run[A](effect: Ox ?=> A): Unit = supervised(effect): Unit - def hgetall(key: String): Long = supervised(client.hGetAll[String, String](key).size.toLong) + protected def inLanes[A](work: Payloads.Workload)(perKey: String => Ox ?=> A): Ox ?=> Unit = + work.lanes.map(g => fork(g.foreach(perKey(_): A))).foreach(_.join()) - def close(): Unit = { + override def close(): Unit = { shutdown.countDown() holder.join() } @@ -105,65 +69,15 @@ final class LettuceBench(host: String, port: Int) extends BenchClient { private val conn = client.connect() private val async = conn.async() - def name: String = "lettuce" + def getAll(work: Payloads.Workload): Unit = + SlidingWindow(work)((k, done) => async.get(k).whenComplete((_, t) => done(t)): Unit).run() - def seed(prefix: String, count: Int, value: String, hashKey: String, fields: Int): Unit = { - val writes = (0 until count).map(i => async.set(s"$prefix:$i", value)) ++ (0 until fields).map(i => async.hset(hashKey, s"f$i", value)) - writes.foreach { f => - f.get() - } - } + def setAll(work: Payloads.Workload, value: String): Unit = + SlidingWindow(work)((k, done) => async.set(k, value).whenComplete((_, t) => done(t)): Unit).run() - // Keep up to concurrency futures in flight, starting the next request whenever one completes. This avoids making all requests wait for - // the slowest member of a fixed batch. Completion callbacks run on Lettuce's event loop. - private def slidingWindow[T](keys: Array[String], concurrency: Int)(submit: String => RedisFuture[T])(score: T => Long): Long = { - val n = keys.length - val width = math.max(1, math.min(concurrency, n)) - val total = new AtomicLong(0L) - val nextIndex = new AtomicInteger(0) - val remaining = new CountDownLatch(n) - val failure = new AtomicReference[Throwable]() - def fireNext(): Unit = { - val i = nextIndex.getAndIncrement() - if (i < n) { - try - submit(keys(i)).whenComplete { (v, t) => - if (t != null) { failure.compareAndSet(null, t): Unit } - else if (v != null) { total.addAndGet(score(v)): Unit } - remaining.countDown() - fireNext() - }: Unit - catch { - case t: Throwable => - failure.compareAndSet(null, t) - remaining.countDown() - fireNext() - } - } - } - var k = 0 - while (k < width) { - fireNext() - k += 1 - } - remaining.await() - val t = failure.get() - if (t != null) throw t // never publish numbers for a run where commands failed - total.get() - } - - def getAll(keys: Array[String], concurrency: Int): Long = - slidingWindow(keys, concurrency)(async.get)(v => v.length.toLong) - - def setAll(keys: Array[String], value: String, concurrency: Int): Long = { - slidingWindow(keys, concurrency)(k => async.set(k, value))(_ => 0L) - keys.length.toLong - } - - def mget(keys: Array[String]): Long = - async.mget(keys*).get().asScala.iterator.filter(_.hasValue).map(_.getValue.length.toLong).sum + def mget(): Unit = async.mget(Payloads.Keys.all*).get(): Unit - def hgetall(key: String): Long = async.hgetall(key).get().size.toLong + def hgetall(): Unit = async.hgetall(Payloads.HashKey).get(): Unit def close(): Unit = { conn.close() @@ -181,63 +95,56 @@ final class RediscalaBench(host: String, port: Int) extends BenchClient { import system.dispatcher private val client = redis.RedisClient(host, port) - def name: String = "rediscala" - private def await[A](f: Future[A]): A = Await.result(f, 5.minutes) - def seed(prefix: String, count: Int, value: String, hashKey: String, fields: Int): Unit = { - val writes = (0 until count).map(i => client.set(s"$prefix:$i", value)) ++ (0 until fields).map(i => client.hset(hashKey, s"f$i", value)) - writes.foreach { f => - await(f) - } + def getAll(work: Payloads.Workload): Unit = + SlidingWindow(work)((k, done) => client.get[String](k).onComplete { case Failure(t) => done(t); case _ => done(null) }).run() + + def setAll(work: Payloads.Workload, value: String): Unit = + SlidingWindow(work)((k, done) => client.set(k, value).onComplete { case Failure(t) => done(t); case _ => done(null) }).run() + + def mget(): Unit = await(client.mget[String](Payloads.Keys.all*)): Unit + + def hgetall(): Unit = await(client.hgetall[String](Payloads.HashKey)): Unit + + def close(): Unit = { + client.stop() + Await.result(system.terminate(), 30.seconds): Unit } +} - // match LettuceBench.slidingWindow by keeping up to concurrency futures in flight and starting the next request after each completion. - private def slidingWindow[T](keys: Array[String], concurrency: Int)(submit: String => Future[T])(score: T => Long): Long = { - val n = keys.length - val width = math.max(1, math.min(concurrency, n)) - val total = new AtomicLong(0L) - val nextIndex = new AtomicInteger(0) - val remaining = new CountDownLatch(n) - val failure = new AtomicReference[Throwable]() - def fireNext(): Unit = { - val i = nextIndex.getAndIncrement() - if (i < n) - submit(keys(i)).onComplete { result => - result match { - case Success(v) => total.addAndGet(score(v)) - case Failure(t) => failure.compareAndSet(null, t) - } - remaining.countDown() - fireNext() - } +/** + * Keeps up to `concurrency` requests in flight and starts the next one whenever one completes, so no request waits for the slowest member + * of a fixed batch. `submit` starts the request for a key and calls `done(failure)` once, with a null `failure` on success. + */ +final private class SlidingWindow(work: Payloads.Workload)(submit: (String, Throwable => Unit) => Unit) { + private val nextIndex = new AtomicInteger(0) + private val remaining = new CountDownLatch(Payloads.Keys.all.length) + private val failure = new AtomicReference[Throwable]() + + private val done: Throwable => Unit = t => { + if (t != null) failure.compareAndSet(null, t): Unit + remaining.countDown() + fireNext() + } + + private def fireNext(): Unit = { + val i = nextIndex.getAndIncrement() + if (i < Payloads.Keys.all.length) { + try submit(Payloads.Keys.all(i), done) + catch { case t: Throwable => done(t) } } - var k = 0 - while (k < width) { + } + + def run(): Unit = { + var k = 0 + while (k < work.concurrency) { fireNext() k += 1 } remaining.await() - val t = failure.get() + val t = failure.get() if (t != null) throw t // never publish numbers for a run where commands failed - total.get() - } - - def getAll(keys: Array[String], concurrency: Int): Long = - slidingWindow(keys, concurrency)(k => client.get[String](k))(_.fold(0L)(_.length.toLong)) - - def setAll(keys: Array[String], value: String, concurrency: Int): Long = { - slidingWindow(keys, concurrency)(k => client.set(k, value))(_ => 0L) - keys.length.toLong - } - - def mget(keys: Array[String]): Long = await(client.mget[String](keys*)).flatten.map(_.length.toLong).sum - - def hgetall(key: String): Long = await(client.hgetall[String](key)).size.toLong - - def close(): Unit = { - client.stop() - Await.result(system.terminate(), 30.seconds): Unit } } @@ -257,58 +164,23 @@ final class JedisBench(host: String, port: Int) extends BenchClient { private val pool = new JedisPool(poolCfg, new HostAndPort(host, port), config) private val executor = Executors.newVirtualThreadPerTaskExecutor() - def name: String = "jedis" + // one lane per group on a virtual thread, each with its own borrowed connection running blocking commands sequentially + private def lanes(work: Payloads.Workload)(run: (redis.clients.jedis.Jedis, String) => Any): Unit = + work.lanes + .map(g => executor.submit[Unit](() => borrow(j => g.foreach(run(j, _))))) + .foreach(_.get()) - def seed(prefix: String, count: Int, value: String, hashKey: String, fields: Int): Unit = { - val j = pool.getResource() - try { - val p = j.pipelined() - (0 until count).foreach { i => - p.set(s"$prefix:$i", value) - } - (0 until fields).foreach { i => - p.hset(hashKey, s"f$i", value) - } - p.sync() - } finally j.close() - } + def getAll(work: Payloads.Workload): Unit = lanes(work)(_.get(_)) - // one lane per group on a virtual thread, each with its own borrowed connection running blocking commands sequentially - private def lanes(keys: Array[String], concurrency: Int)(run: (redis.clients.jedis.Jedis, Array[String]) => Long): Long = - Payloads - .groups(keys, concurrency) - .map(g => - executor.submit[Long] { () => - val j = pool.getResource() - try run(j, g) - finally j.close() - } - ) - .map(_.get()) - .sum - - def getAll(keys: Array[String], concurrency: Int): Long = - lanes(keys, concurrency)((j, g) => g.foldLeft(0L)((t, k) => t + Option(j.get(k)).fold(0L)(_.length.toLong))) - - def setAll(keys: Array[String], value: String, concurrency: Int): Long = { - lanes(keys, concurrency) { (j, g) => - g.foreach { k => - j.set(k, value) - } - g.length.toLong - } - keys.length.toLong - } + def setAll(work: Payloads.Workload, value: String): Unit = lanes(work)(_.set(_, value)) - def mget(keys: Array[String]): Long = { - val j = pool.getResource() - try j.mget(keys*).asScala.iterator.filter(_ != null).map(_.length.toLong).sum - finally j.close() - } + def mget(): Unit = borrow(_.mget(Payloads.Keys.all*)): Unit + + def hgetall(): Unit = borrow(_.hgetAll(Payloads.HashKey)): Unit - def hgetall(key: String): Long = { + private inline def borrow[A](inline f: redis.clients.jedis.Jedis => A): A = { val j = pool.getResource() - try j.hgetAll(key).size.toLong + try f(j) finally j.close() } diff --git a/benchmarks/ox/src/main/scala/sage/benchmarks/OxBenchmarks.scala b/benchmarks/ox/src/main/scala/sage/benchmarks/OxBenchmarks.scala index a78e7caa..ef2aa045 100644 --- a/benchmarks/ox/src/main/scala/sage/benchmarks/OxBenchmarks.scala +++ b/benchmarks/ox/src/main/scala/sage/benchmarks/OxBenchmarks.scala @@ -1,46 +1,13 @@ package sage.benchmarks -import java.util.concurrent.TimeUnit +import org.openjdk.jmh.annotations.Param -import org.openjdk.jmh.annotations.* - -@State(Scope.Benchmark) -@BenchmarkMode(Array(Mode.Throughput)) -@OutputTimeUnit(TimeUnit.SECONDS) -@OperationsPerInvocation(1000) // = Payloads.KeyCount -@Fork(1) -@Warmup(iterations = 3, time = 3) -@Measurement(iterations = 5, time = 3) -class ThroughputBench extends RedisBenchState { +class ThroughputBench extends ThroughputBenchBase { @Param(Array("sage-ox", "lettuce", "rediscala", "jedis")) var client: String = "sage-ox" - @Param(Array("1", "8", "64", "256")) var concurrency: Int = 1 - @Param(Array("16", "1024")) var valueSize: Int = 16 - - protected def subjectName: String = client - override protected def seedValueBytes: Int = valueSize - protected def buildClient(host: String, port: Int, name: String): BenchClient = Clients.build(host, port, name) - - @Benchmark def get(): Long = subject.getAll(keys, concurrency) - @Benchmark def set(): Long = subject.setAll(keys, Payloads.value(valueSize), concurrency) } -// Run one large-reply command per invocation. The throughput benchmark covers concurrent requests. valueSize controls the seeded value size. -@State(Scope.Benchmark) -@BenchmarkMode(Array(Mode.Throughput)) -@OutputTimeUnit(TimeUnit.SECONDS) -@Fork(1) -@Warmup(iterations = 3, time = 3) -@Measurement(iterations = 5, time = 3) -class CollectionBench extends RedisBenchState { +class CollectionBench extends CollectionBenchBase { @Param(Array("sage-ox", "lettuce", "rediscala", "jedis")) var client: String = "sage-ox" - @Param(Array("16")) var valueSize: Int = 16 - - protected def subjectName: String = client - override protected def seedValueBytes: Int = valueSize - protected def buildClient(host: String, port: Int, name: String): BenchClient = Clients.build(host, port, name) - - @Benchmark def mget(): Long = subject.mget(keys) - @Benchmark def hgetall(): Long = subject.hgetall(Payloads.HashKey) } diff --git a/benchmarks/pekko/src/main/scala/sage/benchmarks/Clients.scala b/benchmarks/pekko/src/main/scala/sage/benchmarks/Clients.scala index 38ae6cba..109d2557 100644 --- a/benchmarks/pekko/src/main/scala/sage/benchmarks/Clients.scala +++ b/benchmarks/pekko/src/main/scala/sage/benchmarks/Clients.scala @@ -20,56 +20,21 @@ object Clients { } } -final class SagePekkoBench(host: String, port: Int) extends BenchClient { +final class SagePekkoBench(host: String, port: Int) extends SageBench[Future] { private given system: ActorSystem[Nothing] = ActorSystem(Behaviors.empty, "sage-bench") private given ExecutionContext = system.executionContext - private val client: SageClient = + protected val client: SageClient = Await.result(SageClient.connect(SageConfig(topology = Topology.Standalone(Endpoint(host, port)))), 30.seconds) - // bounds in-flight commands by running each item of a group sequentially while groups run in parallel - private def seqTraverse[A, B](as: List[A])(f: A => Future[B]): Future[List[B]] = - as.foldLeft(Future.successful(List.empty[B]))((accF, a) => accF.flatMap(acc => f(a).map(_ :: acc))).map(_.reverse) + protected def run[A](effect: Future[A]): Unit = Await.result(effect, 120.seconds): Unit - private def seqRun[A](as: List[A])(f: A => Future[Any]): Future[Unit] = - as.foldLeft(Future.unit)((accF, a) => accF.flatMap(_ => f(a).map(_ => ()))) + protected def inLanes[A](work: Payloads.Workload)(perKey: String => Future[A]): Future[Any] = + Future.traverse(work.lanes)(_.foldLeft(Future.unit)((previous, key) => previous.flatMap(_ => perKey(key).map(_ => ())))) - def name: String = "sage-pekko" - - def seed(prefix: String, count: Int, value: String, hashKey: String, fields: Int): Unit = { - val sets = seqRun((0 until count).toList)(i => client.set(s"$prefix:$i", value)) - val hash = (0 until fields).map(i => (s"f$i", value)).toList match { - case h :: t => client.hSet(hashKey, h, t*).map(_ => ()) - case Nil => Future.unit - } - Await.result(sets.flatMap(_ => hash), 120.seconds) - } - - def getAll(keys: Array[String], concurrency: Int): Long = - Await.result( - Future - .traverse(Payloads.groups(keys, concurrency).toList)(g => seqTraverse(g.toList)(client.get[String])) - .map(_.flatten.flatten.map(_.length.toLong).sum), - 120.seconds - ) - - def setAll(keys: Array[String], value: String, concurrency: Int): Long = - Await.result( - Future - .traverse(Payloads.groups(keys, concurrency).toList)(g => seqRun(g.toList)(client.set(_, value))) - .map(_ => keys.length.toLong), - 120.seconds - ) - - def mget(keys: Array[String]): Long = - Await.result(client.mGet[String](keys.head, keys.tail*).map(_.flatten.map(_.length.toLong).sum), 120.seconds) - - def hgetall(key: String): Long = - Await.result(client.hGetAll[String, String](key).map(_.size.toLong), 120.seconds) - - def close(): Unit = { - Await.result(client.close, 30.seconds) + override def close(): Unit = { + super.close() system.terminate() Await.ready(system.whenTerminated, 10.seconds): Unit } diff --git a/benchmarks/pekko/src/main/scala/sage/benchmarks/PekkoBenchmarks.scala b/benchmarks/pekko/src/main/scala/sage/benchmarks/PekkoBenchmarks.scala index 835237fb..9d9df092 100644 --- a/benchmarks/pekko/src/main/scala/sage/benchmarks/PekkoBenchmarks.scala +++ b/benchmarks/pekko/src/main/scala/sage/benchmarks/PekkoBenchmarks.scala @@ -1,46 +1,13 @@ package sage.benchmarks -import java.util.concurrent.TimeUnit +import org.openjdk.jmh.annotations.Param -import org.openjdk.jmh.annotations.* +class ThroughputBench extends ThroughputBenchBase { -@State(Scope.Benchmark) -@BenchmarkMode(Array(Mode.Throughput)) -@OutputTimeUnit(TimeUnit.SECONDS) -@OperationsPerInvocation(1000) // = Payloads.KeyCount -@Fork(1) -@Warmup(iterations = 3, time = 3) -@Measurement(iterations = 5, time = 3) -class ThroughputBench extends RedisBenchState { - - @Param(Array("sage-pekko")) var client: String = "sage-pekko" - @Param(Array("1", "8", "64", "256")) var concurrency: Int = 1 - @Param(Array("16", "1024")) var valueSize: Int = 16 - - protected def subjectName: String = client - override protected def seedValueBytes: Int = valueSize - protected def buildClient(host: String, port: Int, name: String): BenchClient = Clients.build(host, port, name) - - @Benchmark def get(): Long = subject.getAll(keys, concurrency) - @Benchmark def set(): Long = subject.setAll(keys, Payloads.value(valueSize), concurrency) + @Param(Array("sage-pekko")) var client: String = "sage-pekko" } -// Run one large-reply command per invocation. The throughput benchmark covers concurrent requests. valueSize controls the seeded value size. -@State(Scope.Benchmark) -@BenchmarkMode(Array(Mode.Throughput)) -@OutputTimeUnit(TimeUnit.SECONDS) -@Fork(1) -@Warmup(iterations = 3, time = 3) -@Measurement(iterations = 5, time = 3) -class CollectionBench extends RedisBenchState { +class CollectionBench extends CollectionBenchBase { @Param(Array("sage-pekko")) var client: String = "sage-pekko" - @Param(Array("16")) var valueSize: Int = 16 - - protected def subjectName: String = client - override protected def seedValueBytes: Int = valueSize - protected def buildClient(host: String, port: Int, name: String): BenchClient = Clients.build(host, port, name) - - @Benchmark def mget(): Long = subject.mget(keys) - @Benchmark def hgetall(): Long = subject.hgetall(Payloads.HashKey) } diff --git a/benchmarks/shared/src/main/scala/sage/benchmarks/BenchClient.scala b/benchmarks/shared/src/main/scala/sage/benchmarks/BenchClient.scala index b7d38ea5..91247f73 100644 --- a/benchmarks/shared/src/main/scala/sage/benchmarks/BenchClient.scala +++ b/benchmarks/shared/src/main/scala/sage/benchmarks/BenchClient.scala @@ -1,35 +1,27 @@ package sage.benchmarks /** - * One client under benchmark. Each method waits for its effect to finish, allowing JMH to measure the full round trip. Each method also - * returns a checksum that the benchmark consumes, which gives JMH a result to report and prevents the JIT from eliminating the call. + * One client under benchmark. Each method waits for its effect to finish, allowing JMH to measure the full round trip. */ trait BenchClient extends AutoCloseable { - def name: String - - /** - * Seed `count` string keys `prefix:0..count-1` with `value`, plus one hash `hashKey` of `fields` field/value pairs. - */ - def seed(prefix: String, count: Int, value: String, hashKey: String, fields: Int): Unit - /** - * GET every key with `concurrency` commands in flight; returns the total length of the values read. + * GET every key with `work.concurrency` commands in flight. */ - def getAll(keys: Array[String], concurrency: Int): Long + def getAll(work: Payloads.Workload): Unit /** - * SET every key to `value` with `concurrency` commands in flight; returns the number of writes. + * SET every key to `value` with `work.concurrency` commands in flight. */ - def setAll(keys: Array[String], value: String, concurrency: Int): Long + def setAll(work: Payloads.Workload, value: String): Unit /** - * One MGET of all `keys`; returns the total length of the values read. + * One MGET of all `Payloads.Keys`. */ - def mget(keys: Array[String]): Long + def mget(): Unit /** - * One HGETALL of `key`; returns the number of fields read. + * One HGETALL of `Payloads.HashKey`. */ - def hgetall(key: String): Long + def hgetall(): Unit } diff --git a/benchmarks/shared/src/main/scala/sage/benchmarks/ClusterFormation.scala b/benchmarks/shared/src/main/scala/sage/benchmarks/ClusterFormation.scala deleted file mode 100644 index 7d865836..00000000 --- a/benchmarks/shared/src/main/scala/sage/benchmarks/ClusterFormation.scala +++ /dev/null @@ -1,60 +0,0 @@ -package sage.benchmarks - -import java.io.{BufferedReader, InputStreamReader, OutputStreamWriter} -import java.net.Socket -import java.nio.charset.StandardCharsets.UTF_8 - -/** - * Forms a single-node cluster on a freshly-started `--cluster-enabled` server: the node claims every slot and announces the - * testcontainers-mapped host/port so the address it reports in `CLUSTER SLOTS` is reachable from the benchmark (the same trick as the - * integration `ClusterSuite`). - */ -object ClusterFormation { - - def formSingleNodeCluster(host: String, port: Int): Unit = { - val socket = new Socket(host, port) - try { - val out = new OutputStreamWriter(socket.getOutputStream, UTF_8) - val in = new BufferedReader(new InputStreamReader(socket.getInputStream, UTF_8)) - def command(args: String*): String = { - out.write(args.mkString("", " ", "\r\n")) - out.flush() - reply(in) - } - def clusterOk: Boolean = command("CLUSTER", "INFO").contains("cluster_state:ok") - - command("CONFIG", "SET", "cluster-announce-ip", host) - command("CONFIG", "SET", "cluster-announce-port", port.toString) - if (!clusterOk) { command("CLUSTER", "ADDSLOTSRANGE", "0", "16383"): Unit } - var attempts = 100 - var ok = clusterOk - while (!ok && attempts > 0) { - Thread.sleep(100) - attempts -= 1 - ok = clusterOk - } - if (!ok) throw new IllegalStateException("single-node cluster did not converge") - } finally socket.close() - } - - private def reply(in: BufferedReader): String = { - val line = in.readLine() - if (line == null) throw new IllegalStateException("connection closed while forming the cluster") - else if (line.startsWith("-")) throw new IllegalStateException(s"server error: $line") - else if (line.startsWith("$")) { - val length = line.drop(1).toInt - if (length < 0) "" - else { - val payload = new Array[Char](length) - var read = 0 - while (read < length) { - val n = in.read(payload, read, length - read) - if (n < 0) throw new IllegalStateException("connection closed while forming the cluster") - read += n - } - in.readLine() - new String(payload) - } - } else line - } -} diff --git a/benchmarks/shared/src/main/scala/sage/benchmarks/Payloads.scala b/benchmarks/shared/src/main/scala/sage/benchmarks/Payloads.scala index 3698adfa..9bab6331 100644 --- a/benchmarks/shared/src/main/scala/sage/benchmarks/Payloads.scala +++ b/benchmarks/shared/src/main/scala/sage/benchmarks/Payloads.scala @@ -1,12 +1,14 @@ package sage.benchmarks +import scala.collection.immutable.ArraySeq + /** * Fixed shapes shared by every cell's benchmarks, so sage and the competitors are measured on identical data. */ object Payloads { /** - * The number of keys a throughput/MGET workload touches per invocation. Kept in sync with the literal in `@OperationsPerInvocation`. + * The number of keys a throughput/MGET workload touches per invocation. */ final val KeyCount = 1000 @@ -17,21 +19,20 @@ object Payloads { final val HashKey = "bench:hash" - def value(size: Int): String = "v" * size - - def keys(prefix: String): Array[String] = Array.tabulate(KeyCount)(i => s"$prefix:$i") + // The split matches MGET's `(first, rest*)` signature without copying the keys on each call. + object Keys { + val first: String = "bench:0" + val rest: ArraySeq[String] = ArraySeq.tabulate(KeyCount - 1)(i => s"bench:${i + 1}") + val all: Array[String] = (first +: rest).toArray + val set: Set[String] = all.toSet + } /** - * Split keys into `concurrency` near-equal groups; running each group sequentially in its own fiber bounds in-flight commands to `concurrency`. + * The split of `Keys` into `concurrency` lanes, key `i` going to lane `i % concurrency`. Running each lane sequentially in its own fiber + * bounds in-flight commands to `concurrency`. */ - def groups(keys: Array[String], concurrency: Int): Array[Array[String]] = { - val g = math.max(1, concurrency) - val lanes = Array.fill(g)(Array.newBuilder[String]) - var i = 0 - while (i < keys.length) { - lanes(i % g) += keys(i) - i += 1 - } - lanes.map(_.result()) + final class Workload(val concurrency: Int) { + require(concurrency > 0, s"concurrency must be positive, got $concurrency") + val lanes: List[List[String]] = List.tabulate(concurrency)(lane => List.range(lane, Keys.all.length, concurrency).map(Keys.all(_))) } } diff --git a/benchmarks/shared/src/main/scala/sage/benchmarks/RedisBenchState.scala b/benchmarks/shared/src/main/scala/sage/benchmarks/RedisBenchState.scala index ec78e2bb..99a1ac8e 100644 --- a/benchmarks/shared/src/main/scala/sage/benchmarks/RedisBenchState.scala +++ b/benchmarks/shared/src/main/scala/sage/benchmarks/RedisBenchState.scala @@ -1,43 +1,76 @@ package sage.benchmarks -import org.openjdk.jmh.annotations.{Level, Setup, TearDown} +import java.util.concurrent.TimeUnit + +import com.dimafeng.testcontainers.GenericContainer +import org.openjdk.jmh.annotations.* /** - * Shared JMH state for every cell's benchmarks. It starts Redis, builds the cell's clients (Sage and any competitors), seeds the data, and - * shuts everything down for each trial. Concrete subclasses define the `@State`, `@Param`, and `@Benchmark` annotations for their cell. + * Shared JMH state for every cell's benchmarks. It starts Redis, seeds the data, builds the client under test, and shuts everything down for + * each trial. Each cell defines a `client` `@Param` whose unique value, such as `sage-zio` or `redis4cats`, identifies the client in merged + * results. */ +@State(Scope.Benchmark) +@BenchmarkMode(Array(Mode.Throughput)) +@OutputTimeUnit(TimeUnit.SECONDS) +@Fork(1) +@Warmup(iterations = 3, time = 3) +@Measurement(iterations = 5, time = 3) abstract class RedisBenchState { - val fixture: RedisFixture = new RedisFixture - var subject: BenchClient = null - var keys: Array[String] = Array.empty + var subject: BenchClient = null + private var redis: GenericContainer = null /** - * The size of the seeded values. Subclasses with a `valueSize` parameter override this to match the values used by their GET benchmark. + * The size of the seeded values and of the values the SET benchmarks write. */ - protected def seedValueBytes: Int + def valueSize: Int - /** - * The client under test for this trial. Its unique name, such as `sage-zio` or `redis4cats`, identifies it in merged results. - */ - protected def subjectName: String + def client: String + + lazy val value: String = "v" * valueSize + + protected def clusterEnabled: Boolean = false /** - * Builds only the named client. Other clients and their runtimes remain stopped during the trial and cannot affect the result. + * Builds only the client under test. Other clients and their runtimes remain stopped during the trial and cannot affect the result. */ - protected def buildClient(host: String, port: Int, name: String): BenchClient + protected def buildClient(host: String, port: Int): BenchClient = Clients.build(host, port, client) @Setup(Level.Trial) def setupTrial(): Unit = { - fixture.start() - subject = buildClient(fixture.host, fixture.port, subjectName) - keys = Payloads.keys("bench") - subject.seed("bench", Payloads.KeyCount, Payloads.value(seedValueBytes), Payloads.HashKey, Payloads.HashFields) + redis = RedisFixture.start(clusterEnabled, value) + subject = buildClient(redis.host, redis.mappedPort(6379)) } @TearDown(Level.Trial) def tearDownTrial(): Unit = { if (subject != null) subject.close() - fixture.stop() + if (redis != null) redis.stop() } } + +@OperationsPerInvocation(Payloads.KeyCount) +abstract class ThroughputWorkload extends RedisBenchState { + + @Param(Array("1", "8", "64", "256")) var concurrency: Int = 1 + + lazy val work = new Payloads.Workload(concurrency) + + @Benchmark def get(): Unit = subject.getAll(work) + @Benchmark def set(): Unit = subject.setAll(work, value) +} + +abstract class ThroughputBenchBase extends ThroughputWorkload { + + @Param(Array("16", "1024")) var valueSize: Int = 16 +} + +// Run one large-reply command per invocation. The throughput benchmark covers concurrent requests. valueSize controls the seeded value size. +abstract class CollectionBenchBase extends RedisBenchState { + + @Param(Array("16")) var valueSize: Int = 16 + + @Benchmark def mget(): Unit = subject.mget() + @Benchmark def hgetall(): Unit = subject.hgetall() +} diff --git a/benchmarks/shared/src/main/scala/sage/benchmarks/RedisFixture.scala b/benchmarks/shared/src/main/scala/sage/benchmarks/RedisFixture.scala index 17974a45..b725490e 100644 --- a/benchmarks/shared/src/main/scala/sage/benchmarks/RedisFixture.scala +++ b/benchmarks/shared/src/main/scala/sage/benchmarks/RedisFixture.scala @@ -1,33 +1,32 @@ package sage.benchmarks import com.dimafeng.testcontainers.GenericContainer +import org.testcontainers.images.builder.Transferable /** * A self-provisioned Redis for one JMH trial: started in `@Setup(Level.Trial)` and stopped in `@TearDown`, so every trial measures against a * fresh, isolated server. Pinned to the same image as the integration tests. */ -final class RedisFixture { - - private var container: Option[GenericContainer] = None - - var host: String = "" - var port: Int = 0 +object RedisFixture { + val Image = "redis:8.8.0" - def start(clusterEnabled: Boolean = false): Unit = { - val command = if (clusterEnabled) Seq("redis-server", "--cluster-enabled", "yes") else Seq.empty - val c = GenericContainer(RedisFixture.Image, exposedPorts = Seq(6379), command = command) + def start(clusterEnabled: Boolean, value: String): GenericContainer = { + val command = if (clusterEnabled) Seq("redis-server", "--cluster-enabled", "yes") else Seq.empty + val c = GenericContainer(Image, exposedPorts = Seq(6379), command = command) c.start() - container = Some(c) - host = c.host - port = c.mappedPort(6379) - } - - def stop(): Unit = { - container.foreach(_.stop()) - container = None + // The node claims every slot and announces the mapped host and port, so the address it reports in CLUSTER SLOTS is reachable from the host. + // A cluster that never reaches cluster_state:ok fails the seed check below with CLUSTERDOWN replies. + val formCluster = + if (!clusterEnabled) "" + else + s"redis-cli config set cluster-announce-ip ${c.host} cluster-announce-port ${c.mappedPort(6379)}; redis-cli cluster addslotsrange 0 16383; " + + "for i in $(seq 100); do redis-cli cluster info | grep -q cluster_state:ok && break; sleep 0.1; done; " + // redis-cli runs one command per input line; one SET per key avoids CROSSSLOT errors on the cluster-enabled server + val commands = Payloads.Keys.all.map(k => s"SET $k $value") ++ (0 until Payloads.HashFields).map(i => s"HSET ${Payloads.HashKey} f$i $value") + c.container.copyFileToContainer(Transferable.of(commands.mkString("", "\n", "\n")), "/tmp/seed.txt") + val replies = c.execInContainer("sh", "-c", s"${formCluster}redis-cli < /tmp/seed.txt").getStdout + val failures = replies.linesIterator.filter(r => r.nonEmpty && r != "OK" && r.toLongOption.isEmpty).toVector + if (failures.nonEmpty) throw new IllegalStateException(s"setup failed: ${failures.take(3).mkString("; ")}") + c } } - -object RedisFixture { - val Image = "redis:8.8.0" -} diff --git a/benchmarks/shared/src/main/scala/sage/benchmarks/SageBench.scala b/benchmarks/shared/src/main/scala/sage/benchmarks/SageBench.scala new file mode 100644 index 00000000..eefc50bb --- /dev/null +++ b/benchmarks/shared/src/main/scala/sage/benchmarks/SageBench.scala @@ -0,0 +1,29 @@ +package sage.benchmarks + +import sage.client.internal.Client + +/** + * Sage through one backend's public `SageClient`. Every cell sends the same commands, and each cell defines only how its effect type runs and + * waits. + */ +abstract class SageBench[F[_]] extends BenchClient { + + protected val client: Client[F, String] + + protected def run[A](effect: F[A]): Unit + + /** + * Runs the lanes in parallel and the keys of each lane one after another. + */ + protected def inLanes[A](work: Payloads.Workload)(perKey: String => F[A]): F[Any] + + def getAll(work: Payloads.Workload): Unit = run(inLanes(work)(client.get[String](_))) + + def setAll(work: Payloads.Workload, value: String): Unit = run(inLanes(work)(client.set(_, value))) + + def mget(): Unit = run(client.mGet[String](Payloads.Keys.first, Payloads.Keys.rest*)) + + def hgetall(): Unit = run(client.hGetAll[String, String](Payloads.HashKey)) + + def close(): Unit = run(client.close) +} diff --git a/benchmarks/shared/src/main/scala/sage/benchmarks/TopologyBenchState.scala b/benchmarks/shared/src/main/scala/sage/benchmarks/TopologyBenchState.scala deleted file mode 100644 index 6cf63f64..00000000 --- a/benchmarks/shared/src/main/scala/sage/benchmarks/TopologyBenchState.scala +++ /dev/null @@ -1,37 +0,0 @@ -package sage.benchmarks - -import org.openjdk.jmh.annotations.{Level, Setup, TearDown} - -/** - * Shared JMH state for running the same workload against standalone, cluster, and master-replica clients. Each client uses one provisioned - * server. The cluster trial runs a `--cluster-enabled` server in a separate container. The result therefore includes server-mode and - * instance differences as well as client dispatch. - */ -abstract class TopologyBenchState { - - val fixture: RedisFixture = new RedisFixture - var subject: BenchClient = null - var keys: Array[String] = Array.empty - - protected def topologyName: String - - protected def seedValueBytes: Int - - protected def buildClient(host: String, port: Int, topology: String): BenchClient - - @Setup(Level.Trial) - def setupTrial(): Unit = { - val cluster = topologyName == "cluster" - fixture.start(clusterEnabled = cluster) - if (cluster) ClusterFormation.formSingleNodeCluster(fixture.host, fixture.port) - subject = buildClient(fixture.host, fixture.port, topologyName) - keys = Payloads.keys("bench") - subject.seed("bench", Payloads.KeyCount, Payloads.value(seedValueBytes), Payloads.HashKey, Payloads.HashFields) - } - - @TearDown(Level.Trial) - def tearDownTrial(): Unit = { - if (subject != null) subject.close() - fixture.stop() - } -} diff --git a/benchmarks/zio/src/main/scala/sage/benchmarks/Clients.scala b/benchmarks/zio/src/main/scala/sage/benchmarks/Clients.scala index 0a6e37da..24806382 100644 --- a/benchmarks/zio/src/main/scala/sage/benchmarks/Clients.scala +++ b/benchmarks/zio/src/main/scala/sage/benchmarks/Clients.scala @@ -22,13 +22,13 @@ object Clients { case other => throw new IllegalArgumentException(s"unknown client: $other") } - def buildTopology(host: String, port: Int, topology: String): BenchClient = { + def buildTopology(host: String, port: Int, name: String, topology: String): BenchClient = { val endpoint = Endpoint(host, port) - topology match { - case "standalone" => new SageZioBench(Topology.Standalone(endpoint)) - case "cluster" => new SageZioBench(Topology.Cluster(Vector(endpoint))) - case "master-replica" => new SageZioBench(Topology.MasterReplica(Vector(endpoint))) - case other => throw new IllegalArgumentException(s"unknown topology: $other") + (name, topology) match { + case ("sage-zio", "standalone") => new SageZioBench(Topology.Standalone(endpoint)) + case ("sage-zio", "cluster") => new SageZioBench(Topology.Cluster(Vector(endpoint))) + case ("sage-zio", "master-replica") => new SageZioBench(Topology.MasterReplica(Vector(endpoint))) + case other => throw new IllegalArgumentException(s"unknown client and topology: $other") } } } @@ -39,39 +39,14 @@ private object Run { Unsafe.unsafe(implicit u => runtime.unsafe.run(z).getOrThrowFiberFailure()) } -final class SageZioBench(topology: Topology) extends BenchClient { +final class SageZioBench(topology: Topology) extends SageBench[IO[SageException, *]] { - private val client: SageClient = - Run(SageClient.connect(SageConfig(topology = topology))) + protected val client: SageClient = Run(SageClient.connect(SageConfig(topology = topology))) - def name: String = "sage-zio" + protected def run[A](effect: IO[SageException, A]): Unit = Run(effect): Unit - def seed(prefix: String, count: Int, value: String, hashKey: String, fields: Int): Unit = - Run( - ZIO.foreachDiscard(0 until count)(i => client.set(s"$prefix:$i", value)) *> - ZIO.foreachDiscard(0 until fields)(i => client.hSet(hashKey, (s"f$i", value))) - ) - - def getAll(keys: Array[String], concurrency: Int): Long = - Run( - ZIO - .foreachPar(Payloads.groups(keys, concurrency).toList)(g => ZIO.foreach(g.toList)(k => client.get[String](k))) - .map(_.flatten.flatten.map(_.length.toLong).sum) - ) - - def setAll(keys: Array[String], value: String, concurrency: Int): Long = - Run( - ZIO - .foreachParDiscard(Payloads.groups(keys, concurrency).toList)(g => ZIO.foreachDiscard(g.toList)(k => client.set(k, value))) - .as(keys.length.toLong) - ) - - def mget(keys: Array[String]): Long = - Run(client.mGet[String](keys.head, keys.tail*).map(_.flatten.map(_.length.toLong).sum)) - - def hgetall(key: String): Long = Run(client.hGetAll[String, String](key).map(_.size.toLong)) - - def close(): Unit = Run(client.close) + protected def inLanes[A](work: Payloads.Workload)(perKey: String => IO[SageException, A]): IO[SageException, Unit] = + ZIO.foreachParDiscard(work.lanes)(ZIO.foreachDiscard(_)(perKey)) } final class ZioRedisBench(host: String, port: Int) extends BenchClient { @@ -99,32 +74,15 @@ final class ZioRedisBench(host: String, port: Int) extends BenchClient { ) ) - def name: String = "zio-redis" + def getAll(work: Payloads.Workload): Unit = + Run(ZIO.foreachParDiscard(work.lanes)(g => ZIO.foreachDiscard(g)(k => redis.get(k).returning[String]))) - def seed(prefix: String, count: Int, value: String, hashKey: String, fields: Int): Unit = - Run( - ZIO.foreachDiscard(0 until count)(i => redis.set(s"$prefix:$i", value)) *> - ZIO.foreachDiscard(0 until fields)(i => redis.hSet(hashKey, (s"f$i", value))) - ) - - def getAll(keys: Array[String], concurrency: Int): Long = - Run( - ZIO - .foreachPar(Payloads.groups(keys, concurrency).toList)(g => ZIO.foreach(g.toList)(k => redis.get(k).returning[String])) - .map(_.flatten.flatten.map(_.length.toLong).sum) - ) - - def setAll(keys: Array[String], value: String, concurrency: Int): Long = - Run( - ZIO - .foreachParDiscard(Payloads.groups(keys, concurrency).toList)(g => ZIO.foreachDiscard(g.toList)(k => redis.set(k, value))) - .as(keys.length.toLong) - ) + def setAll(work: Payloads.Workload, value: String): Unit = + Run(ZIO.foreachParDiscard(work.lanes)(g => ZIO.foreachDiscard(g)(k => redis.set(k, value)))) - def mget(keys: Array[String]): Long = - Run(redis.mGet(keys.head, keys.tail*).returning[String].map(_.flatten.map(_.length.toLong).sum)) + def mget(): Unit = Run(redis.mGet(Payloads.Keys.first, Payloads.Keys.rest*).returning[String].unit) - def hgetall(key: String): Long = Run(redis.hGetAll(key).returning[String, String].map(_.size.toLong)) + def hgetall(): Unit = Run(redis.hGetAll(Payloads.HashKey).returning[String, String].unit) def close(): Unit = Run(scope.close(Exit.unit)) } diff --git a/benchmarks/zio/src/main/scala/sage/benchmarks/ZioBenchmarks.scala b/benchmarks/zio/src/main/scala/sage/benchmarks/ZioBenchmarks.scala index 5ab69004..eaeab786 100644 --- a/benchmarks/zio/src/main/scala/sage/benchmarks/ZioBenchmarks.scala +++ b/benchmarks/zio/src/main/scala/sage/benchmarks/ZioBenchmarks.scala @@ -1,69 +1,25 @@ package sage.benchmarks -import java.util.concurrent.TimeUnit - import org.openjdk.jmh.annotations.* -@State(Scope.Benchmark) -@BenchmarkMode(Array(Mode.Throughput)) -@OutputTimeUnit(TimeUnit.SECONDS) -@OperationsPerInvocation(1000) // = Payloads.KeyCount -@Fork(1) -@Warmup(iterations = 3, time = 3) -@Measurement(iterations = 5, time = 3) -class ThroughputBench extends RedisBenchState { +class ThroughputBench extends ThroughputBenchBase { @Param(Array("sage-zio", "zio-redis")) var client: String = "sage-zio" - @Param(Array("1", "8", "64", "256")) var concurrency: Int = 1 - @Param(Array("16", "1024")) var valueSize: Int = 16 +} - protected def subjectName: String = client - override protected def seedValueBytes: Int = valueSize - protected def buildClient(host: String, port: Int, name: String): BenchClient = Clients.build(host, port, name) +class CollectionBench extends CollectionBenchBase { - @Benchmark def get(): Long = subject.getAll(keys, concurrency) - @Benchmark def set(): Long = subject.setAll(keys, Payloads.value(valueSize), concurrency) + @Param(Array("sage-zio", "zio-redis")) var client: String = "sage-zio" } -// Measures the throughput workload across Sage client topologies, including server mode. See TopologyBenchState. -@State(Scope.Benchmark) -@BenchmarkMode(Array(Mode.Throughput)) -@OutputTimeUnit(TimeUnit.SECONDS) -@OperationsPerInvocation(1000) // = Payloads.KeyCount -@Fork(1) -@Warmup(iterations = 3, time = 3) -@Measurement(iterations = 5, time = 3) -class TopologyBench extends TopologyBenchState { +// Measures the throughput workload across Sage client topologies. The cluster trial runs a `--cluster-enabled` server in a separate +// container, so the result includes server-mode and instance differences as well as client dispatch. +class TopologyBench extends ThroughputWorkload { @Param(Array("sage-zio")) var client: String = "sage-zio" @Param(Array("standalone", "cluster", "master-replica")) var topology: String = "standalone" - @Param(Array("1", "8", "64", "256")) var concurrency: Int = 1 @Param(Array("16")) var valueSize: Int = 16 - protected def topologyName: String = topology - override protected def seedValueBytes: Int = valueSize - protected def buildClient(host: String, port: Int, name: String): BenchClient = Clients.buildTopology(host, port, name) - - @Benchmark def get(): Long = subject.getAll(keys, concurrency) - @Benchmark def set(): Long = subject.setAll(keys, Payloads.value(valueSize), concurrency) -} - -// Run one large-reply command per invocation. The throughput benchmark covers concurrent requests. valueSize controls the seeded value size. -@State(Scope.Benchmark) -@BenchmarkMode(Array(Mode.Throughput)) -@OutputTimeUnit(TimeUnit.SECONDS) -@Fork(1) -@Warmup(iterations = 3, time = 3) -@Measurement(iterations = 5, time = 3) -class CollectionBench extends RedisBenchState { - - @Param(Array("sage-zio", "zio-redis")) var client: String = "sage-zio" - @Param(Array("16")) var valueSize: Int = 16 - - protected def subjectName: String = client - override protected def seedValueBytes: Int = valueSize - protected def buildClient(host: String, port: Int, name: String): BenchClient = Clients.build(host, port, name) - - @Benchmark def mget(): Long = subject.mget(keys) - @Benchmark def hgetall(): Long = subject.hgetall(Payloads.HashKey) + override protected def clusterEnabled: Boolean = topology == "cluster" + override protected def buildClient(host: String, port: Int): BenchClient = Clients.buildTopology(host, port, client, topology) } diff --git a/build.sbt b/build.sbt index 109687f4..ab7a997e 100644 --- a/build.sbt +++ b/build.sbt @@ -200,6 +200,8 @@ lazy val integrationTests = (projectMatrix in file("integration-tests")) "io.circe" %% "circe-parser" % circeVersion % Test, "io.circe" %% "circe-generic" % circeVersion % Test ), + // a test step that lost its `>>` would otherwise compile as a discarded statement + Test / scalacOptions += "-Wnonunit-statement", // the Future anchor rows compile but don't boot containers Test / testOptions += { val isAnchor = moduleName.value.endsWith("-future") diff --git a/docs/client-side-caching.md b/docs/client-side-caching.md index 9d5bc7a3..ac43e329 100644 --- a/docs/client-side-caching.md +++ b/docs/client-side-caching.md @@ -32,7 +32,7 @@ Sage enables server-assisted client tracking. When a cached key changes, the ser A read is cacheable when its result depends only on the current value of its keys. The server can then invalidate the result whenever one of those keys changes. Reads that vary with time (`TTL`, `OBJECT IDLETIME`) or are non-deterministic (`SRANDMEMBER`) are read-only but not cacheable because a key change cannot reliably invalidate them. ::: warning -`cached` rejects writes and keyless reads with `NotCacheable`. The server cannot invalidate a keyless read when data changes, which could leave a stale value in the cache. +`cached` rejects writes, keyless reads, and blocking commands with `NotCacheable`. The server cannot invalidate a keyless read when data changes, which could leave a stale value in the cache. A blocking command would hold the shared connection that cached reads use. ::: Tune cache sizing and behavior through `clientCache` on [`SageConfig`](/configuration). diff --git a/docs/configuration.md b/docs/configuration.md index affd5e05..c4778963 100644 --- a/docs/configuration.md +++ b/docs/configuration.md @@ -88,8 +88,8 @@ val config = SageConfig( ) ``` -Sage refreshes on a timer only when you configure this setting. Each refresh costs one `CLUSTER SLOTS`, and ticks arriving within `minRefreshInterval` of the last refresh are -skipped, so a short interval cannot flood the cluster. `MasterReplicaConfig` has the same setting. +Sage refreshes on a timer only when you configure this setting. Each refresh costs one `CLUSTER SLOTS`. Ticks arriving within `minRefreshInterval` of the last refresh +are deferred and run as one refresh at the end of that interval, so a short interval cannot flood the cluster. `MasterReplicaConfig` has the same setting. ## Master-replica @@ -140,7 +140,7 @@ val config = SageConfig( ) ``` -`TrustSource.System` uses the system trust store. Use `TrustSource.Pem` or `TrustSource.TrustStore` for a private CA. Use `TrustSource.Custom(sslContext)` to supply your own `SSLContext`, including for mutual TLS. `AuthConfig` redacts its password in logs and in any printed `SageConfig`. +`TrustSource.System` uses the system trust store. Use `TrustSource.Pem` or `TrustSource.TrustStore` for a private CA. Sage reads that file before it first connects to a node, and the node's reconnects and dedicated connections reuse what it read. A standalone client therefore reads the file once, when it connects. A cluster or master-replica client reads it again whenever it adds a node or opens a pub/sub or discovery connection, so a replaced CA bundle applies to those connections only. Use `TrustSource.Custom(sslContext)` to supply your own `SSLContext`, including for mutual TLS. `AuthConfig` redacts its password in logs and in any printed `SageConfig`. ::: warning `TrustSource.Insecure` is for local development only. It trusts every certificate and skips hostname verification, leaving the connection open to machine-in-the-middle attacks. Never use it in production. diff --git a/docs/error-handling.md b/docs/error-handling.md index 474d0239..690b388b 100644 --- a/docs/error-handling.md +++ b/docs/error-handling.md @@ -9,16 +9,16 @@ Every Sage failure is a `SageException` in a sealed hierarchy that you can match | `ProtocolError(message)` | Malformed RESP3 on the wire; the connection is discarded. | | `DecodeError(expected, actual)` | A reply was well-formed but not the shape a decoder or codec required (the built-in codecs decode strictly). | | `ServerError(code, detail)` | An error reply from the server. `code` is the leading token (`WRONGTYPE`, `NOSCRIPT`, `BUSYGROUP`, the generic `ERR`, …). | -| `ConnectionFailed(message)` | The initial connection could not be established (host unreachable, connection refused, or connect timeout). Distinct from `ConnectionLost`, which is a live connection dropping. | +| `ConnectionFailed(message)` | The initial connection could not be established: the host is unreachable, the connection is refused or times out, the server does not answer `HELLO` within `connectTimeout`, or it closes the socket during setup. Distinct from `ConnectionLost`, which is a live connection dropping. | | `ConnectionLost(mayHaveExecuted)` | The connection dropped around this command. | | `NotConnected()` | The client was never started, or has been closed. | | `UnsupportedServer(message)` | The server rejected `HELLO 3` (it predates RESP3, or is a RESP2-only proxy). | | `TlsError(message)` | TLS could not be established (rejected certificate or unusable trust material). | | `CrossSlot(message)` | An unsupported multi-key command or a transaction touched keys in more than one cluster slot. `MGET`, `MSET`, `EXISTS`, `DEL`, `UNLINK`, and `TOUCH` are transparently split outside transactions. | -| `TimedOut(message)` | A pooled connection wait exceeded `dedicatedPool.acquireTimeout`, a topology probe exceeded `connectTimeout`, or distributed lock acquisition exceeded its wait, lease, or replica confirmation budget. | +| `TimedOut(message)` | A pooled connection wait exceeded `dedicatedPool.acquireTimeout`, a topology probe exceeded `connectTimeout` (a master-replica `connect` fails this way when its `ROLE` probes time out), or distributed lock acquisition exceeded its wait, lease, or replica confirmation budget. | | `LockLost(message)` | A distributed lock expired, changed owner, or could not confirm renewal or release within its deadline. Acquisition can raise this before the body starts if the granting master changes role or the lease expires. Sage attempts to cancel a running body. Cancellation depends on the backend and cannot undo completed work. | | `TransactionDiscarded(message)` | A transaction was discarded server-side (`EXECABORT`); nothing ran. | -| `NotCacheable(message)` | `cached` was given a command that cannot be safely cached. | +| `NotCacheable(message)` | `cached` was given a command that cannot be safely cached: a write, a keyless or time-varying read, or a blocking command. | | `InvalidArgument(message)` | A programming error, rejected before any server call: an invalid configuration or rate-limit policy, a blocking command inside a pipeline or transaction, or a command a cluster client cannot route as written. | Regular commands have no per-command timeout. Use your backend's timeout combinator to bound their duration. See [Distributed locks](/distributed-locks) for lock deadlines and cancellation behavior. diff --git a/docs/observability.md b/docs/observability.md index b976590a..d256d5fb 100644 --- a/docs/observability.md +++ b/docs/observability.md @@ -19,7 +19,7 @@ Register one or more `SageListener` instances on `SageConfig`. Each listener rec | `Cache.Hit(command)` / `Cache.Miss(command)` | A `cached` read was served locally, or had to fetch from the server. | | `TopologyChanged(masters)` | The cluster's slot-owning master set changed (a failover, or scaling a shard in or out). | -Events omit command arguments and payloads. This keeps secrets such as `AUTH` credentials and user values out of listeners. Where an event carries `node`, it is `Some` for cluster and master-replica clients and `None` for a standalone client. A node that keeps failing to connect produces one `ConnectFailed` rather than one event per attempt. +Events omit command arguments and payloads. This keeps secrets such as `AUTH` credentials and user values out of listeners. Where an event carries `node`, it is `Some` for cluster and master-replica clients and `None` for a standalone client. A cluster command that runs on several nodes, such as a cross-slot `MGET` or a command sent to every master, reports `None`. A node that keeps failing to connect produces one `ConnectFailed` rather than one event per attempt. ### Registering a listener diff --git a/docs/pubsub.md b/docs/pubsub.md index ac9d497f..6aeb3adf 100644 --- a/docs/pubsub.md +++ b/docs/pubsub.md @@ -1,6 +1,6 @@ # Pub/Sub -Subscribing returns the backend's native stream type. Ox returns `Flow`, ZIO returns `ZStream`, Cats Effect returns an fs2 `Stream`, Kyo returns `Stream`, and Pekko returns `Source`. Each message contains its channel and a payload decoded by a `ValueCodec`. Ending the stream or closing its scope unsubscribes. +Subscribing returns the backend's native stream type. Ox returns `Flow`, ZIO returns `ZStream`, Cats Effect returns an fs2 `Stream`, Kyo returns `Stream`, and Pekko returns `Source`. Each message contains its channel and a payload decoded by a `ValueCodec`. Ending the stream or closing its scope unsubscribes. If the server refuses a subscribed name, for example with `NOPERM` from an ACL, the stream fails with that `ServerError`, including when a reconnect subscribes the name again. Temporary errors such as `BUSY` or `LOADING` do not end the stream; Sage subscribes again after a backoff. ## Classic channels diff --git a/examples/pekko/src/main/scala/sage/examples/pekko/PubSubExample.scala b/examples/pekko/src/main/scala/sage/examples/pekko/PubSubExample.scala index 189a541d..cf1f60ab 100644 --- a/examples/pekko/src/main/scala/sage/examples/pekko/PubSubExample.scala +++ b/examples/pekko/src/main/scala/sage/examples/pekko/PubSubExample.scala @@ -4,7 +4,6 @@ import scala.concurrent.ExecutionContext import scala.concurrent.Future import org.apache.pekko.actor.typed.ActorSystem -import org.apache.pekko.stream.{Materializer, SystemMaterializer} import org.apache.pekko.stream.scaladsl.{Keep, Sink} import sage.* @@ -17,8 +16,7 @@ import sage.backend.* */ object PubSubExample { - def run(client: SageClient)(using system: ActorSystem[?], ec: ExecutionContext): Future[Unit] = { - given Materializer = SystemMaterializer(system).materializer + def run(client: SageClient)(using ActorSystem[?], ExecutionContext): Future[Unit] = { val (confirmed, received) = client.subscribe[String]("news").take(3).toMat(Sink.seq)(Keep.both).run() for { diff --git a/integration-tests/ce/src/test/scala/sage/integration/CeSmokeSuite.scala b/integration-tests/ce/src/test/scala/sage/integration/CeSmokeSuite.scala index 90d7a751..8592934f 100644 --- a/integration-tests/ce/src/test/scala/sage/integration/CeSmokeSuite.scala +++ b/integration-tests/ce/src/test/scala/sage/integration/CeSmokeSuite.scala @@ -1,144 +1,33 @@ package sage.integration -import scala.concurrent.duration.* - import cats.effect.IO import cats.effect.unsafe.implicits.global import cats.syntax.all.* +import kyo.compat.* +import munit.{Location, TestOptions} import sage.* import sage.backend.* -class CeSmokeSuite extends ServerSuite(Images.redis) { - - private def withNativeClient(body: SageClient => IO[Unit]): Unit = - withContainers(server => SageClient.resource(configOf(server)).use(body).unsafeRunSync()) - - test("a distributed lock scopes native effects and skips contended bodies") { - withNativeClient { client => - val locks = client.lock[String]() - var evaluated = false - for { - busy <- locks.withLock("native-lock", 2.seconds) { - locks.tryWithLock("native-lock") { - evaluated = true - client.ping() - } - } - acquired <- locks.tryWithLock("native-lock")(client.ping()) - } yield { - assertEquals(busy, None) - assertEquals(acquired, Some("PONG")) - assert(!evaluated) - } - } - } - - test("an end user connects and round-trips with native Cats Effect") { - withNativeClient { client => - for { - pong <- client.ping() - _ <- (1 to 50).toList.parTraverse_(i => client.set(s"key-$i", s"value-$i")) - values <- (1 to 50).toList.parTraverse(i => client.get[String](s"key-$i")) - } yield { - assertEquals(pong, "PONG") - assertEquals(values, (1 to 50).toList.map(i => Some(s"value-$i"))) - } - } - } - - test("a pipeline returns a typed tuple natively, surfacing failures per position") { - withNativeClient { client => - for { - _ <- client.set("pipe:a", "x") - _ <- client.set("pipe:n", 10) - out <- client.pipeline((Commands.get[String, String]("pipe:a"), Commands.incrBy[String]("pipe:n", 5))) - _ <- client.set("pipe:str", "hello") - attempt <- client.pipelineAttempt((Commands.get[String, String]("pipe:str"), Commands.incr[String]("pipe:str"))) - } yield { - assertEquals(out, (Some("x"), 15L)) - assert(attempt._1 == Right(Some("hello")), attempt._1) - assert(attempt._2.isLeft, attempt._2) - } - } - } - - test("a transaction commits atomically with native Cats Effect, guarded by WATCH") { - withNativeClient { client => - for { - _ <- client.set("tx:n", 1) - out <- client.transaction { tx => - for { - _ <- tx.watch("tx:n") - _ <- tx.get[Int]("tx:n") - res <- tx.exec((Commands.incr[String]("tx:n"), Commands.incrBy[String]("tx:n", 4))) - } yield res - } - } yield assertEquals(out, Some((2L, 6L))) - } - } - - test("scanAll streams every key as a native fs2 Stream") { - withNativeClient { client => - for { - _ <- (1 to 50).toList.parTraverse_(i => client.set(s"scan-$i", "v")) - keys <- client.scanAll(pattern = Some("scan-*"), count = Some(10L)).compile.toVector - } yield assertEquals(keys.toSet, (1 to 50).map(i => s"scan-$i").toSet) - } - } +class CeSmokeSuite extends SmokeSuite { - test("subscribe delivers published messages as a native fs2 Stream") { - withNativeClient { client => - client.subscribeResource[String]("smoke").use { stream => - for { - _ <- (1 to 3).toList.traverse_(i => client.publish("smoke", s"m$i")) - messages <- stream.take(3).compile.toVector - } yield { - assertEquals(messages.map(_.channel).toSet, Set("smoke")) - assertEquals(messages.map(_.payload).toList, List("m1", "m2", "m3")) - } - } - } - } + private def nativeTest(options: TestOptions)(body: SageClient => IO[Unit])(using Location): Unit = + test(options)(withContainers(server => SageClient.resource(configOf(server)).use(body).unsafeRunSync())) - test("hScanAll streams every field/value pair as a native fs2 Stream") { - withNativeClient { client => - for { - _ <- (1 to 50).toList.parTraverse_(i => client.hSet("hscan", (s"f$i", s"v$i"))) - pairs <- client.hScanAll[String, String]("hscan", count = Some(10L)).compile.toVector - } yield assertEquals(pairs.toMap, (1 to 50).map(i => s"f$i" -> s"v$i").toMap) - } - } + private val lift: [A] => IO[A] => CIO[A] = [A] => (io: IO[A]) => CIO.lift(io) - test("sScanAll streams every member as a native fs2 Stream") { - withNativeClient { client => - for { - _ <- (1 to 50).toList.parTraverse_(i => client.sAdd("sscan", s"m$i")) - members <- client.sScanAll[String]("sscan", count = Some(10L)).compile.toVector - } yield assertEquals(members.toSet, (1 to 50).map(i => s"m$i").toSet) - } - } + nativeTest("a distributed lock scopes native effects and skips contended bodies")(lockScopesNativeEffects(_)(lift).lower) - test("zScanAll streams every member/score pair as a native fs2 Stream") { - withNativeClient { client => - for { - _ <- (1 to 50).toList.parTraverse_(i => client.zAdd("zscan")((s"m$i", i.toDouble))) - pairs <- client.zScanAll[String]("zscan", count = Some(10L)).compile.toVector - } yield assertEquals(pairs.toMap, (1 to 50).map(i => s"m$i" -> i.toDouble).toMap) - } + nativeTest("scanAll streams every key as a native fs2 Stream") { client => + scanAllFindsEveryKey(client)(lift)(client.scanAll(pattern = Some("scan-*"), count = Some(10L)).compile.toVector).lower } - test("client.rateLimiter admits up to capacity then denies") { - withNativeClient { client => - val rl = client.rateLimiter[String](RateLimit(capacity = 2, refillTokens = 1, refillPeriod = 1.hour)) + nativeTest("subscribe delivers published messages as a native fs2 Stream") { client => + client.subscribeResource[String]("smoke").use { stream => for { - first <- rl.tryAcquire("smoke") - second <- rl.tryAcquire("smoke") - denied <- rl.tryAcquire("smoke") - } yield { - assert(first.isAllowed && second.isAllowed, "the first two are admitted") - assert(!denied.isAllowed, "the third is denied once the bucket empties") - } + _ <- (1 to 3).toList.traverse_(i => client.publish("smoke", s"m$i")) + messages <- stream.take(3).compile.toVector + } yield assertEquals(messages.toList, List("m1", "m2", "m3").map(Message("smoke", _))) } } } diff --git a/integration-tests/kyo/src/test/scala/sage/integration/KyoSmokeSuite.scala b/integration-tests/kyo/src/test/scala/sage/integration/KyoSmokeSuite.scala index 808ed896..bc3068b9 100644 --- a/integration-tests/kyo/src/test/scala/sage/integration/KyoSmokeSuite.scala +++ b/integration-tests/kyo/src/test/scala/sage/integration/KyoSmokeSuite.scala @@ -1,180 +1,58 @@ package sage.integration -import java.util.concurrent.TimeUnit - -import scala.concurrent.duration.FiniteDuration - import kyo.* +import kyo.compat.* +import munit.{Location, TestOptions} import sage.* import sage.backend.* -class KyoSmokeSuite extends ServerSuite(Images.redis) { +class KyoSmokeSuite extends SmokeSuite { - private def withNativeClient(body: SageClient => Unit < (Scope & Abort[Throwable] & Async)): Unit = - withBoundedClient(Duration.Infinity)(body) - - private def withBoundedClient(timeout: Duration)(body: SageClient => Unit < (Scope & Abort[Throwable] & Async)): Unit = - withContainers { server => + private def nativeTest(options: TestOptions, timeout: Duration = Duration.Infinity)(body: SageClient => Unit < (Scope & Abort[Throwable] & Async))( + using Location + ): Unit = + test(options)(withContainers { server => val program: Unit < (Scope & Abort[Throwable] & Async) = SageClient.scoped(configOf(server)).map(body) import AllowUnsafe.embrace.danger KyoApp.Unsafe.runAndBlock(timeout)(program).getOrThrow - } + }) - test("a distributed lock scopes native effects and skips contended bodies") { - withNativeClient { client => - val locks = client.lock[String]() - var evaluated = false - for { - busy <- locks.withLock("native-lock", FiniteDuration(2L, TimeUnit.SECONDS)) { - locks.tryWithLock("native-lock") { - evaluated = true - client.ping() - } - } - acquired <- locks.tryWithLock("native-lock")(client.ping()) - } yield { - assertEquals(busy, None) - assertEquals(acquired, Some("PONG")) - assert(!evaluated) - } - } - } + private val lift: [A] => (A < (Abort[SageException] & Async)) => CIO[A] = [A] => (v: A < (Abort[SageException] & Async)) => CIO.lift(v) - test("an end user connects and round-trips with native Kyo") { - withNativeClient { client => - for { - pong <- client.ping() - _ <- Async.foreachDiscard(1 to 50)(i => client.set(s"key-$i", s"value-$i")) - values <- Async.foreach((1 to 50).toList)(i => client.get[String](s"key-$i")) - } yield { - assertEquals(pong, "PONG") - assertEquals(values.toList, (1 to 50).toList.map(i => Some(s"value-$i"))) - } - } - } + nativeTest("a distributed lock scopes native effects and skips contended bodies")(lockScopesNativeEffects(_)(lift).lower) - test("a distributed lock preserves a native panic and releases ownership") { - withBoundedClient(3L.seconds) { client => - val failure = new IllegalStateException("body panic") - val locks = client.lock[String]() - for { - result <- Abort.run[SageException](locks.tryWithLock[Int]("native-panic")(Abort.panic(failure))) - _ = result match { - case Result.Panic(error) => assert(error eq failure) - case other => fail(s"expected the original panic, got $other") - } - exists <- client.exists("4:lock:native-panic") - acquired <- locks.tryWithLock("native-panic")(client.ping()) - } yield { - assertEquals(exists, 0L) - assertEquals(acquired, Some("PONG")) - } + nativeTest("a distributed lock preserves a native panic and releases ownership", 3L.seconds) { client => + val failure = new IllegalStateException("body panic") + val locks = client.lock[String]() + for { + result <- Abort.run[SageException](locks.tryWithLock[Int]("native-panic")(Abort.panic(failure))) + _ = assertEquals(result, Result.Panic(failure)) + exists <- client.exists("4:lock:native-panic") + acquired <- locks.tryWithLock("native-panic")(client.ping()) + } yield { + assertEquals(exists, 0L) + assertEquals(acquired, Some("PONG")) } } - test("a pipeline returns a typed tuple natively, surfacing failures per position") { - withNativeClient { client => - for { - _ <- client.set("pipe:a", "x") - _ <- client.set("pipe:n", 10) - out <- client.pipeline((Commands.get[String, String]("pipe:a"), Commands.incrBy[String]("pipe:n", 5))) - _ <- client.set("pipe:str", "hello") - attempt <- client.pipelineAttempt((Commands.get[String, String]("pipe:str"), Commands.incr[String]("pipe:str"))) - } yield { - assertEquals(out, (Some("x"), 15L)) - assert(attempt._1 == Right(Some("hello")), attempt._1) - assert(attempt._2.isLeft, attempt._2) - } - } + nativeTest("scanAll streams every key as a native Kyo Stream") { client => + scanAllFindsEveryKey(client)(lift)(client.scanAll(pattern = Some("scan-*"), count = Some(10L)).run).lower } - test("a transaction commits atomically with native Kyo, guarded by WATCH") { - withNativeClient { client => - for { - _ <- client.set("tx:n", 1) - out <- client.transaction { tx => - for { - _ <- tx.watch("tx:n") - _ <- tx.get[Int]("tx:n") - res <- tx.exec((Commands.incr[String]("tx:n"), Commands.incrBy[String]("tx:n", 4))) - } yield res - } - } yield assertEquals(out, Some((2L, 6L))) - } - } - - test("scanAll streams every key as a native Kyo Stream") { - withNativeClient { client => - for { - _ <- Async.foreachDiscard(1 to 50)(i => client.set(s"scan-$i", "v")) - keys <- client.scanAll(pattern = Some("scan-*"), count = Some(10L)).run - } yield assertEquals(keys.toSet, (1 to 50).map(i => s"scan-$i").toSet) - } - } - - test("subscribe delivers published messages as a native Kyo Stream") { - withNativeClient { client => - for { - stream <- client.subscribeScoped[String]("smoke") - _ <- Kyo.foreachDiscard(1 to 3)(i => client.publish("smoke", s"m$i")) - chunk <- stream.take(3).run - } yield { - val messages = chunk.toList - assertEquals(messages.map(_.channel).toSet, Set("smoke")) - assertEquals(messages.map(_.payload), List("m1", "m2", "m3")) - } - } - } - - test("hScanAll streams every field/value pair as a native Kyo Stream") { - withNativeClient { client => - for { - _ <- Async.foreachDiscard(1 to 50)(i => client.hSet("hscan", (s"f$i", s"v$i"))) - pairs <- client.hScanAll[String, String]("hscan", count = Some(10L)).run - } yield assertEquals(pairs.toMap, (1 to 50).map(i => s"f$i" -> s"v$i").toMap) - } - } - - test("sScanAll streams every member as a native Kyo Stream") { - withNativeClient { client => - for { - _ <- Async.foreachDiscard(1 to 50)(i => client.sAdd("sscan", s"m$i")) - members <- client.sScanAll[String]("sscan", count = Some(10L)).run - } yield assertEquals(members.toSet, (1 to 50).map(i => s"m$i").toSet) - } - } - - test("zScanAll streams every member/score pair as a native Kyo Stream") { - withNativeClient { client => - for { - _ <- Async.foreachDiscard(1 to 50)(i => client.zAdd("zscan")((s"m$i", i.toDouble))) - pairs <- client.zScanAll[String]("zscan", count = Some(10L)).run - } yield assertEquals(pairs.toMap, (1 to 50).map(i => s"m$i" -> i.toDouble).toMap) - } + nativeTest("subscribe delivers published messages as a native Kyo Stream") { client => + for { + stream <- client.subscribeScoped[String]("smoke") + _ <- Kyo.foreachDiscard(1 to 3)(i => client.publish("smoke", s"m$i")) + chunk <- stream.take(3).run + } yield assertEquals(chunk.toList, List("m1", "m2", "m3").map(Message("smoke", _))) } // regression for the 4096-page rechunk that buffered unbounded streams; the timeout makes a recurrence fail rather than hang - test("xTail emits replayed entries immediately instead of buffering them") { - withBoundedClient(15L.seconds) { client => - for { - _ <- Kyo.foreachDiscard(1 to 3)(i => client.xAdd("xtail", XAddId.Explicit(StreamId(i.toLong, 0L)))(("f", s"v$i"))) - entries <- client.xTail[String, String]("xtail").take(3).run - } yield assertEquals(entries.toList.map(_.fields.head._2), List("v1", "v2", "v3")) - } - } - - test("client.rateLimiter admits up to capacity then denies") { - withBoundedClient(15L.seconds) { client => - val rl = client.rateLimiter[String](RateLimit(capacity = 2, refillTokens = 1, refillPeriod = FiniteDuration(1L, TimeUnit.HOURS))) - for { - first <- rl.tryAcquire("smoke") - second <- rl.tryAcquire("smoke") - denied <- rl.tryAcquire("smoke") - } yield { - assert(first.isAllowed && second.isAllowed, "the first two are admitted") - assert(!denied.isAllowed, "the third is denied once the bucket empties") - } - } + nativeTest("xTail emits replayed entries immediately instead of buffering them", 15L.seconds) { client => + for { + _ <- Kyo.foreachDiscard(1 to 3)(i => client.xAdd("xtail", XAddId.Explicit(StreamId(i.toLong, 0L)))(("f", s"v$i"))) + entries <- client.xTail[String, String]("xtail").take(3).run + } yield assertEquals(entries.toList.flatMap(_.fields.map(_._2)), List("v1", "v2", "v3")) } } diff --git a/integration-tests/ox/src/test/scala/sage/integration/OxSmokeSuite.scala b/integration-tests/ox/src/test/scala/sage/integration/OxSmokeSuite.scala index 45c22bed..281dd916 100644 --- a/integration-tests/ox/src/test/scala/sage/integration/OxSmokeSuite.scala +++ b/integration-tests/ox/src/test/scala/sage/integration/OxSmokeSuite.scala @@ -1,159 +1,79 @@ package sage.integration -import scala.concurrent.duration.* - +import kyo.compat.* +import munit.{Location, TestOptions} import ox.{fork, supervised, Ox} import sage.* import sage.backend.* -class OxSmokeSuite extends ServerSuite(Images.redis) { - - private def withNativeClient(body: Ox ?=> SageClient => Unit): Unit = - withContainers(server => supervised(body(SageClient.scoped(configOf(server))))) - - test("a distributed lock scopes native effects and skips contended bodies") { - withNativeClient { client => - val locks = client.lock[String]() - var evaluated = false - val busy = locks.withLock("native-lock", 2.seconds) { - locks.tryWithLock("native-lock") { - evaluated = true - client.ping() - } - } - val acquired = locks.tryWithLock("native-lock")(client.ping()) - assertEquals(busy, None) - assertEquals(acquired, Some("PONG")) - assert(!evaluated) - } - } - - test("an end user connects and round-trips with direct-style Ox") { - withNativeClient { client => - assertEquals(client.ping(), "PONG") - val values = (1 to 50).toList - .map(i => - fork { - client.set(s"key-$i", s"value-$i") - client.get[String](s"key-$i") - } - ) - .map(_.join()) - assertEquals(values, (1 to 50).toList.map(i => Some(s"value-$i"))) - } - } - - test("a pipeline returns a typed tuple natively, surfacing failures per position") { - withNativeClient { client => - client.set("pipe:a", "x") - client.set("pipe:n", 10) - val out = client.pipeline((Commands.get[String, String]("pipe:a"), Commands.incrBy[String]("pipe:n", 5))) - assertEquals(out, (Some("x"), 15L)) - client.set("pipe:str", "hello") - val attempt = client.pipelineAttempt((Commands.get[String, String]("pipe:str"), Commands.incr[String]("pipe:str"))) - assert(attempt._1 == Right(Some("hello")), attempt._1) - assert(attempt._2.isLeft, attempt._2) - } - } +class OxSmokeSuite extends SmokeSuite { - test("a transaction commits atomically with direct-style Ox, guarded by WATCH") { - withNativeClient { client => - client.set("tx:n", 1) - val out = client.transaction { tx => - tx.watch("tx:n") - tx.get[Int]("tx:n") - tx.exec((Commands.incr[String]("tx:n"), Commands.incrBy[String]("tx:n", 4))) - } - assertEquals(out, Some((2L, 6L))) - } - } + private def nativeTest(options: TestOptions)(body: Ox ?=> SageClient => Unit)(using Location): Unit = + test(options)(withContainers(server => supervised(body(SageClient.scoped(configOf(server)))))) - test("scanAll streams every key as a native Ox Flow") { - withNativeClient { client => - (1 to 50).foreach(i => client.set(s"scan-$i", "v")) - val keys = client.scanAll(pattern = Some("scan-*"), count = Some(10L)).runToList() - assertEquals(keys.toSet, (1 to 50).map(i => s"scan-$i").toSet) - } - } + private val lift: [A] => (Ox ?=> A) => CIO[A] = [A] => (run: Ox ?=> A) => CIO.lift(run) - test("subscribe delivers published messages as a native Ox Flow") { - withNativeClient { client => - val stream = client.subscribeScoped[String]("smoke") - (1 to 3).foreach(i => client.publish("smoke", s"m$i")) - val messages = stream.take(3).runToList() - assertEquals(messages.map(_.channel).toSet, Set("smoke")) - assertEquals(messages.map(_.payload), List("m1", "m2", "m3")) - } - } + nativeTest("a distributed lock scopes native effects and skips contended bodies")(lockScopesNativeEffects(_)(lift).lower) - test("a scoped subscription survives a flow run ending — take(1) must not unsubscribe, only scope close does") { - withNativeClient { client => - val stream = client.subscribeScoped[String]("scoped-live") - client.publish("scoped-live", "a") - val first = stream.take(1).runToList() - client.publish("scoped-live", "b") - val second = stream.take(1).runToList() - assertEquals(first.map(_.payload), List("a")) - assertEquals(second.map(_.payload), List("b")) - } + nativeTest("an end user connects and round-trips with direct-style Ox") { client => + assertEquals(client.ping(), "PONG") + val values = (1 to 50).toList + .map(i => + fork { + val _ = client.set(s"key-$i", s"value-$i") + client.get[String](s"key-$i") + } + ) + .map(_.join()) + assertEquals(values, (1 to 50).toList.map(i => Some(s"value-$i"))) } - test("a plain subscribe Flow resubscribes on every run instead of yielding an empty stream on re-run") { - withNativeClient { client => - val stream = client.subscribe[String]("rerun") - // publish continuously so each run's subscription receives one; the second run must resubscribe, not complete empty - val running = new java.util.concurrent.atomic.AtomicBoolean(true) - val publisher = fork { - while (running.get()) { - client.publish("rerun", "tick") - Thread.sleep(50) - } - } - try { - val first = stream.take(1).runToList() - val second = stream.take(1).runToList() - assertEquals(first.map(_.payload), List("tick")) - assertEquals(second.map(_.payload), List("tick")) - } finally { - running.set(false) - publisher.join() - } + nativeTest("a transaction commits atomically with direct-style Ox, guarded by WATCH") { client => + val _ = client.set("tx:n", 1) + val out = client.transaction { tx => + tx.watch("tx:n") + val _ = tx.get[Int]("tx:n") + tx.exec((Commands.incr[String]("tx:n"), Commands.incrBy[String]("tx:n", 4))) } + assertEquals(out, Some((2L, 6L))) } - test("hScanAll streams every field/value pair as a native Ox Flow") { - withNativeClient { client => - (1 to 50).foreach(i => client.hSet("hscan", (s"f$i", s"v$i"))) - val pairs = client.hScanAll[String, String]("hscan", count = Some(10L)).runToList() - assertEquals(pairs.toMap, (1 to 50).map(i => s"f$i" -> s"v$i").toMap) - } + nativeTest("scanAll streams every key as a native Ox Flow") { client => + scanAllFindsEveryKey(client)(lift)(client.scanAll(pattern = Some("scan-*"), count = Some(10L)).runToList()).lower } - test("sScanAll streams every member as a native Ox Flow") { - withNativeClient { client => - (1 to 50).foreach(i => client.sAdd("sscan", s"m$i")) - val members = client.sScanAll[String]("sscan", count = Some(10L)).runToList() - assertEquals(members.toSet, (1 to 50).map(i => s"m$i").toSet) - } + nativeTest("subscribe delivers published messages as a native Ox Flow") { client => + val stream = client.subscribeScoped[String]("smoke") + (1 to 3).foreach(i => client.publish("smoke", s"m$i")) + val messages = stream.take(3).runToList() + assertEquals(messages, List("m1", "m2", "m3").map(Message("smoke", _))) } - test("zScanAll streams every member/score pair as a native Ox Flow") { - withNativeClient { client => - (1 to 50).foreach(i => client.zAdd("zscan")((s"m$i", i.toDouble))) - val pairs = client.zScanAll[String]("zscan", count = Some(10L)).runToList() - assertEquals(pairs.toMap, (1 to 50).map(i => s"m$i" -> i.toDouble).toMap) - } + nativeTest("a scoped subscription survives a flow run ending — take(1) must not unsubscribe, only scope close does") { client => + val stream = client.subscribeScoped[String]("scoped-live") + val _ = client.publish("scoped-live", "a") + val first = stream.take(1).runToList() + val _ = client.publish("scoped-live", "b") + val second = stream.take(1).runToList() + assertEquals(first.map(_.payload), List("a")) + assertEquals(second.map(_.payload), List("b")) } - test("client.rateLimiter admits up to capacity then denies") { - withNativeClient { client => - val rl = client.rateLimiter[String](RateLimit(capacity = 2, refillTokens = 1, refillPeriod = 1.hour)) - val first = rl.tryAcquire("smoke") - val second = rl.tryAcquire("smoke") - val denied = rl.tryAcquire("smoke") - assert(first.isAllowed && second.isAllowed, "the first two are admitted") - assert(!denied.isAllowed, "the third is denied once the bucket empties") + nativeTest("a plain subscribe Flow resubscribes on every run instead of yielding an empty stream on re-run") { client => + val stream = client.subscribe[String]("rerun") + def awaitSubscribers(count: Long): Unit = + Eventually(100)(CIO.deferLift(client.pubsubNumSub("rerun").getOrElse("rerun", 0L)).is(count)).lower + // each run subscribes on its own, so publish once that run's subscription is on the server + def runOnce(): List[String] = { + val run = fork(stream.take(1).runToList()) + awaitSubscribers(1) + val _ = client.publish("rerun", "tick") + run.join().map(_.payload) } + val first = runOnce() + awaitSubscribers(0) + assertEquals(first, List("tick")) + assertEquals(runOnce(), List("tick")) } } diff --git a/integration-tests/pekko/src/test/scala/sage/integration/PekkoSmokeSuite.scala b/integration-tests/pekko/src/test/scala/sage/integration/PekkoSmokeSuite.scala index 1e9f1b4b..8ec0290b 100644 --- a/integration-tests/pekko/src/test/scala/sage/integration/PekkoSmokeSuite.scala +++ b/integration-tests/pekko/src/test/scala/sage/integration/PekkoSmokeSuite.scala @@ -3,187 +3,57 @@ package sage.integration import scala.concurrent.{Await, Future} import scala.concurrent.duration.* +import kyo.compat.* +import munit.{Location, TestOptions} import org.apache.pekko.actor.typed.ActorSystem import org.apache.pekko.actor.typed.scaladsl.Behaviors -import org.apache.pekko.stream.Materializer import org.apache.pekko.stream.scaladsl.{Keep, Sink} import sage.* import sage.backend.* -class PekkoSmokeSuite extends ServerSuite(Images.redis) { +class PekkoSmokeSuite extends SmokeSuite { - private def withNativeClient[A](body: (SageClient, Materializer, ActorSystem[Nothing]) => Future[A]): A = - withContainers { server => - given system: ActorSystem[Nothing] = ActorSystem(Behaviors.empty, "pekko-smoke") - given scala.concurrent.ExecutionContext = system.executionContext - val mat = Materializer(system) - try - Await.result( - SageClient.connect(configOf(server)).flatMap { client => - body(client, mat, system).transformWith(result => client.close.recover { case _ => () }.transform(_ => result)) - }, - 30.seconds - ) + private def nativeTest(options: TestOptions)(body: ActorSystem[Nothing] ?=> SageClient => Future[Unit])(using Location): Unit = + test(options)(withContainers { server => + given system: ActorSystem[Nothing] = ActorSystem(Behaviors.empty, "pekko-smoke") + try Await.result(SageClient.use(configOf(server))(body)(using system.executionContext), 30.seconds) finally { system.terminate() Await.ready(system.whenTerminated, 10.seconds): Unit } - } - - test("a distributed lock scopes native effects and skips contended bodies") { - withNativeClient { (client, _, _) => - val locks = client.lock[String]() - var evaluated = false - for { - busy <- locks.withLock("native-lock", 2.seconds) { - locks.tryWithLock("native-lock") { - evaluated = true - client.ping() - } - } - acquired <- locks.tryWithLock("native-lock")(client.ping()) - } yield { - assertEquals(busy, None) - assertEquals(acquired, Some("PONG")) - assert(!evaluated) - } - } - } - - test("an end user connects and round-trips with scala.concurrent.Future") { - val values = withNativeClient { (client, _, _) => - for { - pong <- client.ping() - _ <- Future.traverse(1 to 50)(i => client.set(s"key-$i", s"value-$i")) - got <- Future.traverse(1 to 50)(i => client.get[String](s"key-$i")) - } yield (pong, got) - } - assertEquals(values._1, "PONG") - assertEquals(values._2.toList, (1 to 50).toList.map(i => Option(s"value-$i"))) - } + }) - test("a pipeline returns a typed tuple natively, surfacing failures per position") { - val (out, attempt) = withNativeClient { (client, _, _) => - for { - _ <- client.set("pipe:a", "x") - _ <- client.set("pipe:n", 10) - out <- client.pipeline((Commands.get[String, String]("pipe:a"), Commands.incrBy[String]("pipe:n", 5))) - _ <- client.set("pipe:str", "hello") - attempt <- client.pipelineAttempt((Commands.get[String, String]("pipe:str"), Commands.incr[String]("pipe:str"))) - } yield (out, attempt) - } - assertEquals(out, (Some("x"), 15L)) - assert(attempt._1 == Right(Some("hello")), attempt._1) - assert(attempt._2.isLeft, attempt._2) - } + private val lift: [A] => Future[A] => CIO[A] = [A] => (future: Future[A]) => CIO.lift(future) - test("a transaction commits atomically with Future, guarded by WATCH") { - val out = withNativeClient { (client, _, _) => - for { - _ <- client.set("tx:n", 1) - res <- client.transaction { tx => - for { - _ <- tx.watch("tx:n") - _ <- tx.get[Int]("tx:n") - res <- tx.exec((Commands.incr[String]("tx:n"), Commands.incrBy[String]("tx:n", 4))) - } yield res - } - } yield res - } - assertEquals(out, Some((2L, 6L))) - } + nativeTest("a distributed lock scopes native effects and skips contended bodies")(lockScopesNativeEffects(_)(lift).unsafeRun) - test("scanAll streams every key as a native Pekko Source") { - val keys = withNativeClient { (client, mat, _) => - given Materializer = mat - for { - _ <- Future.traverse(1 to 50)(i => client.set(s"scan-$i", "v")) - keys <- client.scanAll(pattern = Some("scan-*"), count = Some(10L)).runWith(Sink.seq) - } yield keys - } - assertEquals(keys.toSet, (1 to 50).map(i => s"scan-$i").toSet) + nativeTest("scanAll streams every key as a native Pekko Source") { client => + scanAllFindsEveryKey(client)(lift)(client.scanAll(pattern = Some("scan-*"), count = Some(10L)).runWith(Sink.seq)).unsafeRun } - test("subscribe delivers published messages as a native Pekko Source") { - val messages = withNativeClient { (client, mat, _) => - given Materializer = mat - val (confirmed, received) = - client.subscribe[String]("smoke").take(3).toMat(Sink.seq)(Keep.both).run() - for { - _ <- confirmed - // publish sequentially so the asserted m1/m2/m3 order is deterministic - _ <- (1 to 3).foldLeft(Future.successful(0L))((acc, i) => acc.flatMap(_ => client.publish("smoke", s"m$i"))) - messages <- received - } yield messages - } - assertEquals(messages.map(_.channel).toSet, Set("smoke")) - assertEquals(messages.map(_.payload).toList, List("m1", "m2", "m3")) - } - - test("hScanAll streams every field/value pair as a native Pekko Source") { - val pairs = withNativeClient { (client, mat, _) => - given Materializer = mat - for { - _ <- Future.traverse(1 to 50)(i => client.hSet("hscan", (s"f$i", s"v$i"))) - pairs <- client.hScanAll[String, String]("hscan", count = Some(10L)).runWith(Sink.seq) - } yield pairs - } - assertEquals(pairs.toMap, (1 to 50).map(i => s"f$i" -> s"v$i").toMap) - } - - test("sScanAll streams every member as a native Pekko Source") { - val members = withNativeClient { (client, mat, _) => - given Materializer = mat - for { - _ <- Future.traverse(1 to 50)(i => client.sAdd("sscan", s"m$i")) - members <- client.sScanAll[String]("sscan", count = Some(10L)).runWith(Sink.seq) - } yield members - } - assertEquals(members.toSet, (1 to 50).map(i => s"m$i").toSet) - } - - test("zScanAll streams every member/score pair as a native Pekko Source") { - val pairs = withNativeClient { (client, mat, _) => - given Materializer = mat - for { - _ <- Future.traverse(1 to 50)(i => client.zAdd("zscan")((s"m$i", i.toDouble))) - pairs <- client.zScanAll[String]("zscan", count = Some(10L)).runWith(Sink.seq) - } yield pairs - } - assertEquals(pairs.toMap, (1 to 50).map(i => s"m$i" -> i.toDouble).toMap) - } - - test("tailing helpers surface an infinite block timeout through the effect because Future cannot interrupt it") { - withNativeClient { (client, mat, system) => - given ActorSystem[Nothing] = system - given Materializer = mat - given scala.concurrent.ExecutionContext = system.executionContext - val tailed = - client.xTail[String, String]("stream:forever", block = BlockTimeout.Forever).runWith(Sink.ignore).failed - val consumed = - client.xConsume[String, String]("workers", "w1", "stream:forever", block = BlockTimeout.Forever)(_ => Future.unit).completion.failed - for { - e1 <- tailed - e2 <- consumed - } yield { - assert(e1.isInstanceOf[SageException.InvalidArgument], e1) - assert(e2.isInstanceOf[SageException.InvalidArgument], e2) - } - } + nativeTest("subscribe delivers published messages as a native Pekko Source") { client => + val (confirmed, received) = + client.subscribe[String]("smoke").take(3).toMat(Sink.seq)(Keep.both).run() + for { + _ <- confirmed + // publish sequentially so the asserted m1/m2/m3 order is deterministic + _ <- (1 to 3).foldLeft(Future.successful(0L))((acc, i) => acc.flatMap(_ => client.publish("smoke", s"m$i"))) + messages <- received + } yield assertEquals(messages.toList, List("m1", "m2", "m3").map(Message("smoke", _))) } - test("client.rateLimiter admits up to capacity then denies") { - val (first, second, denied) = withNativeClient { (client, _, system) => - given scala.concurrent.ExecutionContext = system.executionContext - val rl = client.rateLimiter[String](RateLimit(capacity = 2, refillTokens = 1, refillPeriod = 1.hour)) - for { - a <- rl.tryAcquire("smoke") - b <- rl.tryAcquire("smoke") - c <- rl.tryAcquire("smoke") - } yield (a, b, c) + nativeTest("tailing helpers surface an infinite block timeout through the effect because Future cannot interrupt it") { client => + val tailed = + client.xTail[String, String]("stream:forever", block = BlockTimeout.Forever).runWith(Sink.ignore).failed + val consumed = + client.xConsume[String, String]("workers", "w1", "stream:forever", block = BlockTimeout.Forever)(_ => Future.unit).completion.failed + for { + e1 <- tailed + e2 <- consumed + } yield { + assert(e1.isInstanceOf[SageException.InvalidArgument], e1) + assert(e2.isInstanceOf[SageException.InvalidArgument], e2) } - assert(first.isAllowed && second.isAllowed, "the first two are admitted") - assert(!denied.isAllowed, "the third is denied once the bucket empties") } } diff --git a/integration-tests/shared/src/test/scala/sage/integration/ContainerClient.scala b/integration-tests/shared/src/test/scala/sage/integration/ContainerClient.scala index c16e9190..a14253d4 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/ContainerClient.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/ContainerClient.scala @@ -1,26 +1,150 @@ package sage.integration +import java.net.{InetSocketAddress, Socket} +import java.nio.charset.StandardCharsets.US_ASCII +import java.util.concurrent.TimeUnit + +import scala.concurrent.{Await, ExecutionContext} +import scala.concurrent.duration.* +import scala.reflect.{classTag, ClassTag} +import scala.util.{Failure, Using} + import com.dimafeng.testcontainers.GenericContainer +import com.dimafeng.testcontainers.munit.TestContainersSuite import kyo.compat.* +import munit.{Location, TestOptions} +import org.rnorth.ducttape.ratelimits.RateLimiterBuilder +import org.rnorth.ducttape.unreliables.Unreliables +import org.testcontainers.containers.wait.strategy.AbstractWaitStrategy -import sage.client.{Endpoint, SageConfig, Topology} -import sage.client.internal.Client +import sage.Bytes +import sage.client.{Endpoint, LockClient, SageConfig, Topology} +import sage.client.internal.{Client, Paged, Subscription} +import sage.commands.{Command, Commands} /** * Shared connect-and-teardown helpers for the testcontainers suites: build a config for a started container, and run a body against a * client that is always closed afterwards. */ -trait ContainerClient { +trait ContainerClient extends munit.FunSuite with TestContainersSuite { + + // The Ox cell's unsafeRun uses this value. Keeping it non-private avoids unused-private warnings in the other cells. + given ExecutionContext = munitExecutionContext + + protected def serverDef(image: String, exposedPorts: Seq[Int] = Seq(6379), command: Seq[String] = Seq()): GenericContainer.Def[GenericContainer] = + GenericContainer.Def(image, exposedPorts = exposedPorts, command = command, waitStrategy = new AnswersPing) + + // container hooks run outside munitTimeout + protected def prepare(setup: CIO[Unit]): Unit = Await.result(setup.unsafeRun, 3.minutes) + + protected def containerTest(options: TestOptions)(body: Containers => CIO[Any])(using Location): Unit = + test(options)(withContainers(body(_).unsafeRun)) protected def configOf(server: GenericContainer): SageConfig = SageConfig(topology = Topology.Standalone(Endpoint(server.host, server.mappedPort(6379)))) + protected def connectAndUse[A](config: SageConfig)(body: Client[CIO, String] => CIO[A]): CIO[A] = opened(Client.connect(config))(_.close)(body) + + // SHUTDOWN reports a connection failure because the server closes the socket, so its result is ignored. Later checks verify its effect. + protected def shutdown(config: SageConfig): CIO[Unit] = connectAndUse(config)(_.run(admin("SHUTDOWN", "NOSAVE"))).liftToTry.unit + + protected def withSubscription[M, A](subscribe: CIO[Subscription[CIO, M]])(body: Subscription[CIO, M] => CIO[A]): CIO[A] = + opened(subscribe)(_.close)(body) + // CIO.acquireReleaseWith fails to compile on the Ox/Future cells when its type argument nests CIO (Client[CIO, String]); fold instead - protected def connectAndUse[A](config: SageConfig)(body: Client[CIO, String] => CIO[A]): CIO[A] = - Client.connect(config).flatMap { client => - body(client).fold( - result => client.close.map(_ => result), - error => client.close.flatMap(_ => CIO.fail(error)) - ) + private def opened[R, A](open: CIO[R])(close: R => CIO[Unit])(body: R => CIO[A]): CIO[A] = + open.flatMap(resource => body(resource).fold(result => close(resource).map(_ => result), error => close(resource).flatMap(_ => CIO.fail(error)))) + + extension [A](call: CIO[A]) { + protected def is(expected: A)(using Location): CIO[Unit] = call.flatMap(value => CIO.defer(assertEquals(value, expected))) + protected def satisfies(holds: A => Boolean)(using Location): CIO[Unit] = call.flatMap(value => CIO.defer(assert(holds(value), value))) + protected def >>[B](next: CIO[B]): CIO[B] = call.flatMap(_ => next) + } + + protected def failsWith[E <: Throwable: ClassTag](attempt: CIO[Any]): CIO[E] = + attempt.liftToTry.flatMap { + case Failure(error: E) => CIO.value(error) + case other => CIO.fail(new AssertionError(s"expected ${classTag[E].runtimeClass.getSimpleName}, got $other")) + } + + protected def drain[S, A](pages: Paged.Pages[S, A]): CIO[Set[A]] = { + def loop(state: S, found: Set[A]): CIO[Set[A]] = + pages.step(state).flatMap { + case Some((items, next)) => loop(next, found ++ items) + case None => CIO.value(found) + } + loop(pages.init, Set.empty) + } + + // CIO.foreachDiscard with an explicit concurrency does not compile on the Future cell, so one-at-a-time traversals fold instead + protected def inSequence[A](items: Iterable[A])(f: A => CIO[Any]): CIO[Unit] = + items.foldLeft(CIO.unit)((previous, item) => previous.flatMap(_ => f(item).unit)) + + protected def required[A](what: String, value: Option[A]): CIO[A] = + CIO.get(value.toRight(new AssertionError(s"$what returned nothing")).toTry) + + // For server commands the library does not model, such as REPLICAOF, CLIENT PAUSE, and SHUTDOWN. The reply is ignored. + protected def admin(name: String, args: String*): Command[Unit] = + Command(name, Command.NoKeys, args.toVector.map(Bytes.utf8), _ => Right(())) + + // The writer uses a separate connection, so each write is a server-side change for the reader's cache. + protected def cachedReadIsInvalidated(reader: Client[CIO, String], writer: Client[CIO, String], key: String): CIO[Unit] = { + val read = reader.cached(Commands.get[String, String](key), 1.minute) + writer.set(key, "v1") >> + read.is(Some("v1")) >> + read.is(Some("v1")) >> + writer.set(key, "v2") >> + // a cached read does not contact the server again. Poll until the invalidation message from an external write has been processed. + Eventually(50)(read.is(Some("v2"))) + } + + protected def contend( + holder: LockClient[CIO, String], + contender: Client[CIO, String], + key: String, + whileHeld: CIO[Unit] = CIO.unit + ): CIO[Unit] = { + val waiting = contender.lock[String]() + holder + .withLock(key, 2.seconds)(whileHeld.flatMap(_ => waiting.tryWithLock(key)(CIO.fail(new AssertionError("contended body ran"))))) + .is(None) >> + waiting.tryWithLock(key)(CIO.value(42)).is(Some(42)) >> + contender.exists(s"4:lock:$key").is(0L) + } + + protected def awaitRenewal(client: Client[CIO, String], key: String): CIO[Unit] = { + val ttl = client.pTtl(key) + Eventually(30, Duration.Zero)(ttl.flatMap(before => CIO.sleep(50.millis).flatMap(_ => ttl).satisfies(Ttls.renewed(before, _)))) + } + + // Reads `command`'s call count, then gives `next` a step that waits until the count has grown by `by`. + protected def awaitCalls[A](commandStats: CIO[String], command: String, by: Long)(next: CIO[Unit] => CIO[A]): CIO[A] = + commandStats.flatMap(before => + next(Eventually(30, 50.millis)(commandStats.map(commandCalls(_, command)).satisfies(_ >= commandCalls(before, command) + by))) + ) + + protected def commandCalls(commandStats: String, command: String): Long = + commandStats.linesIterator + .find(_.startsWith(s"cmdstat_${command.toLowerCase}:calls=")) + .fold(0L)(_.dropWhile(_ != '=').drop(1).takeWhile(_ != ',').toLong) +} + +// Docker Desktop accepts connections on a published port before the server behind it is reachable, so wait for a PING reply through it. +final private class AnswersPing extends AbstractWaitStrategy { + + withRateLimiter(RateLimiterBuilder.newBuilder().withRate(10, TimeUnit.SECONDS).withConstantThroughput().build()) + + override protected def waitUntilReady(): Unit = + waitStrategyTarget.getExposedPorts.forEach { port => + val address = new InetSocketAddress(waitStrategyTarget.getHost, waitStrategyTarget.getMappedPort(port)) + Unreliables.retryUntilTrue(startupTimeout.getSeconds.toInt, TimeUnit.SECONDS, () => getRateLimiter.getWhenReady(() => pong(address))) + } + + private def pong(address: InetSocketAddress): Boolean = + Using.resource(new Socket()) { socket => + socket.connect(address, 1000) + socket.setSoTimeout(1000) + socket.getOutputStream.write("PING\r\n".getBytes(US_ASCII)) + new String(socket.getInputStream.readNBytes(7), US_ASCII) == "+PONG\r\n" } } diff --git a/integration-tests/shared/src/test/scala/sage/integration/Eventually.scala b/integration-tests/shared/src/test/scala/sage/integration/Eventually.scala index c77c4dc6..5ffc287c 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/Eventually.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/Eventually.scala @@ -1,47 +1,34 @@ package sage.integration import scala.concurrent.duration.* +import scala.util.Failure import kyo.compat.* /** - * Polling for state a server reaches on its own schedule. The action is passed as a thunk. A by-name `CIO` parameter would erase to the same - * JVM signature as the effect value used by the Future and Ox cells, leading to runtime casts. + * Polling for state a server reaches on its own schedule. Each attempt runs the same `CIO` value again. */ object Eventually { /** - * Runs `action` up to `attempts` times, `interval` apart, yielding the first result `holds` accepts or the last one seen. + * Runs `check` up to `attempts` times, `interval` apart, until no assertion in it fails, and returns its result. Any other failure, such as a + * command error, fails at once. */ - def value[A](attempts: Int, interval: FiniteDuration = 100.millis)(action: () => CIO[A])(holds: A => Boolean): CIO[A] = - action().flatMap { seen => - if (holds(seen) || attempts <= 1) CIO.value(seen) - else CIO.sleep(interval).flatMap(_ => value(attempts - 1, interval)(action)(holds)) + def apply[A](attempts: Int, interval: FiniteDuration = 100.millis)(check: CIO[A]): CIO[A] = + retry(attempts, interval)(check) { + case _: AssertionError => true + case _ => false } /** - * Like [[value]], but fails with `orFail`'s message when the state never arrives. + * Runs `action` up to `attempts` times, `interval` apart, until it succeeds, and returns its result. Otherwise fails with the last failure. */ - def converges[A](attempts: Int, interval: FiniteDuration = 100.millis)(action: () => CIO[A])(holds: A => Boolean)( - orFail: A => String - ): CIO[Unit] = - value(attempts, interval)(action)(holds).flatMap { seen => - if (holds(seen)) CIO.value(()) else CIO.fail(new RuntimeException(orFail(seen))) - } + def succeeds[A](attempts: Int, interval: FiniteDuration = 100.millis)(action: CIO[A]): CIO[A] = + retry(attempts, interval)(action)(_ => true) - /** - * Polls until two consecutive results satisfy `changed`. - */ - def changes[A](attempts: Int, interval: FiniteDuration = 100.millis)(action: () => CIO[A])(changed: (A, A) => Boolean)( - orFail: A => String - ): CIO[Unit] = { - def loop(remaining: Int, previous: A): CIO[Unit] = - if (remaining <= 0) CIO.fail(new RuntimeException(orFail(previous))) - else - CIO.sleep(interval).flatMap(_ => action()).flatMap { current => - if (changed(previous, current)) CIO.unit else loop(remaining - 1, current) - } - - action().flatMap(loop(attempts - 1, _)) - } + private def retry[A](attempts: Int, interval: FiniteDuration)(action: CIO[A])(retries: Throwable => Boolean): CIO[A] = + action.liftToTry.flatMap { + case Failure(error) if attempts > 1 && retries(error) => CIO.sleep(interval).flatMap(_ => retry(attempts - 1, interval)(action)(retries)) + case result => CIO.get(result) + } } diff --git a/integration-tests/shared/src/test/scala/sage/integration/LegacyServerConnectSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/LegacyServerConnectSuite.scala index 95f24987..b119c5e5 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/LegacyServerConnectSuite.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/LegacyServerConnectSuite.scala @@ -7,12 +7,7 @@ import kyo.compat.* */ class LegacyServerConnectSuite extends ServerSuite(Images.legacyRedis) { - test("connects and round-trips against a pre-7.2 server that lacks CLIENT SETINFO") { - withClient { client => - for { - _ <- client.set("legacy", "ok") - value <- client.get[String]("legacy") - } yield assertEquals(value, Some("ok")) - } + clientTest("connects and round-trips against a pre-7.2 server that lacks CLIENT SETINFO") { client => + client.set("legacy", "ok").flatMap(_ => client.get[String]("legacy").is(Some("ok"))) } } diff --git a/integration-tests/shared/src/test/scala/sage/integration/RoundTripSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/RoundTripSuite.scala index d554842d..76381111 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/RoundTripSuite.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/RoundTripSuite.scala @@ -2,205 +2,107 @@ package sage.integration import kyo.compat.* -import sage.Bytes -import sage.SageException.DecodeError +import sage.SageException.ServerError import sage.client.internal.Client -import sage.commands.{Command, Commands} -import sage.protocol.Frame +import sage.commands.Commands -abstract class RoundTripSuite(image: String) extends ServerSuite(image) { +class RoundTripSuite extends BothServersSuite { - test("ping round-trips") { - withClient(client => client.ping().map(reply => assertEquals(reply, "PONG"))) - } - - test("values round-trip per call type, and a missing key is None") { - withClient { client => - for { - _ <- client.set("greeting", "hello") - _ <- client.set("count", 42) - _ <- client.set("flag", true) - greeting <- client.get[String]("greeting") - count <- client.get[Int]("count") - flag <- client.get[Boolean]("flag") - missing <- client.get[String]("missing-key") - } yield { - assertEquals(greeting, Some("hello")) - assertEquals(count, Some(42)) - assertEquals(flag, Some(true)) - assertEquals(missing, None) - } - } - } + clientTest("ping round-trips")(client => client.ping().is("PONG")) - test("concurrent fibers pipeline onto the Multiplexed Connection and match FIFO") { - withClient { client => - CIO - .foreach(1 to 200) { i => - for { - _ <- client.set(s"key-$i", s"value-$i") - value <- client.get[String](s"key-$i") - } yield assertEquals(value, Some(s"value-$i")) - } - .unit - } + clientTest("values round-trip per call type, and a missing key is None") { client => + client.set("greeting", "hello") >> + client.set("count", 42) >> + client.set("flag", true) >> + client.get[String]("greeting").is(Some("hello")) >> + client.get[Int]("count").is(Some(42)) >> + client.get[Boolean]("flag").is(Some(true)) >> + client.get[String]("missing-key").is(None) } - test("no reply misattribution under high fiber concurrency") { - def pingLoop(client: Client[CIO, String], fiber: Int, i: Int): CIO[Unit] = + clientTest("no reply misattribution under high fiber concurrency") { client => + def pingLoop(fiber: Int, i: Int): CIO[Unit] = if (i > 100) CIO.value(()) else { val token = s"$fiber-$i" - client.ping(Some(token)).flatMap { reply => - assertEquals(reply, token) - pingLoop(client, fiber, i + 1) - } + client.ping(Some(token)).is(token).flatMap(_ => pingLoop(fiber, i + 1)) } - withClient(client => CIO.foreachDiscard(1 to 500)(fiber => pingLoop(client, fiber, 1))) + CIO.foreachDiscard(1 to 500)(fiber => pingLoop(fiber, 1)) } - test("a pipeline yields one typed result per command in a single round-trip") { - withClient { client => - for { - _ <- client.set("p:a", "x") - _ <- client.set("p:n", 10) - out <- client.pipeline((Commands.get[String, String]("p:a"), Commands.incrBy[String]("p:n", 5))) - } yield assertEquals(out, (Some("x"), 15L)) - } + clientTest("a pipeline yields one typed result per command in a single round-trip") { client => + client.set("p:a", "x") >> + client.set("p:n", 10) >> + client.pipeline((Commands.get[String, String]("p:a"), Commands.incrBy[String]("p:n", 5))).is((Some("x"), 15L)) } - test("a command failure in a pipeline surfaces per-position without poisoning the rest") { - withClient { client => - for { - _ <- client.set("p:str", "hello") - results <- client.pipelineAttempt( - ( - Commands.get[String, String]("p:str"), - Commands.incr[String]("p:str"), - Commands.get[String, String]("p:str") - ) - ) - } yield { - val (a, b, c) = results - assertEquals(a, Right(Some("hello"))) - assert(b.isLeft, s"expected the INCR on a string to fail, got $b") - assertEquals(c, Right(Some("hello"))) - } - } + clientTest("a command failure in a pipeline surfaces per-position without poisoning the rest") { client => + client.set("p:str", "hello") >> + client + .pipelineAttempt( + ( + Commands.get[String, String]("p:str"), + Commands.incr[String]("p:str"), + Commands.get[String, String]("p:str") + ) + ) + .is((Right(Some("hello")), Left(ServerError("ERR", "value is not an integer or out of range")), Right(Some("hello")))) } - test("a large pipeline runs every command and returns one result per position") { - withClient { client => - val n = 200 - client.pipeline(Vector.fill(n)(Commands.incr[String]("p:rtt"))).flatMap { results => - client.get[Int]("p:rtt").map { stored => - assertEquals(results.length, n) - assertEquals(results, (1 to n).map(_.toLong).toVector) - assertEquals(stored, Some(n)) - } - } - } + clientTest("a large pipeline runs every command and returns one result per position") { client => + val n = 200 + client.pipeline(Vector.fill(n)(Commands.incr[String]("p:rtt"))).is((1 to n).map(_.toLong).toVector) >> + client.get[Int]("p:rtt").is(Some(n)) } - test("a transaction commits atomically and returns typed results") { - withClient { client => - for { - _ <- client.set("t:n", 10) - out <- client.transaction(tx => tx.exec((Commands.incr[String]("t:n"), Commands.incrBy[String]("t:n", 5)))) - } yield assertEquals(out, Some((11L, 16L))) - } + clientTest("a transaction commits atomically and returns typed results") { client => + client.set("t:n", 10) >> + client.transaction(tx => tx.exec((Commands.incr[String]("t:n"), Commands.incrBy[String]("t:n", 5)))).is(Some((11L, 16L))) } - test("a read-modify-write transaction commits when the watched key is unchanged") { - withClient { client => - for { - _ <- client.set("t:rmw", 5) - out <- client.transaction { tx => - for { - _ <- tx.watch("t:rmw") - cur <- tx.get[Int]("t:rmw") - res <- tx.exec(Vector(Commands.set[String, Int]("t:rmw", cur.getOrElse(0) + 1))) - } yield res - } - stored <- client.get[Int]("t:rmw") - } yield { - assert(out.isDefined, s"expected a committed transaction, got $out") - assertEquals(stored, Some(6)) - } - } + clientTest("a read-modify-write transaction commits when the watched key is unchanged") { client => + client.set("t:rmw", 5) >> + client + .transaction { tx => + for { + _ <- tx.watch("t:rmw") + cur <- tx.get[Int]("t:rmw") + res <- tx.exec(Vector(Commands.set[String, Int]("t:rmw", cur.getOrElse(0) + 1))) + } yield res + } + .is(Some(Vector(true))) >> + client.get[Int]("t:rmw").is(Some(6)) } - test("WATCH aborts the transaction when a watched key is modified concurrently") { - withContainers { server => - connectAndUse(configOf(server)) { client => - for { - other <- Client.connect(configOf(server)) - _ <- client.set("t:w", 1) - out <- client.transaction { tx => - for { - _ <- tx.watch("t:w") - _ <- tx.get[Int]("t:w") - _ <- other.set("t:w", 99) // a different connection changes the watched key before EXEC - res <- tx.exec(Vector(Commands.incr[String]("t:w"))) - } yield res - } - _ <- other.close - stored <- client.get[Int]("t:w") - } yield { - assertEquals(out, None) // aborted - assertEquals(stored, Some(99)) // the INCR never ran + clientsTest("WATCH aborts the transaction when a watched key is modified concurrently") { (client, other) => + client.set("t:w", 1) >> + client + .transaction { tx => + for { + _ <- tx.watch("t:w") + _ <- tx.get[Int]("t:w") + _ <- other.set("t:w", 99) // a different connection changes the watched key before EXEC + res <- tx.exec(Vector(Commands.incr[String]("t:w"))) + } yield res } - }.unsafeRun - } + .is(None) >> // aborted + client.get[Int]("t:w").is(Some(99)) // the INCR never ran } - test("an execution-phase error surfaces per-position while the other commands commit") { - withClient { client => - for { - _ <- client.set("t:str", "x") - res <- client.transaction(tx => tx.execAttempt((Commands.incr[String]("t:fresh"), Commands.incr[String]("t:str")))) - ok <- client.get[Int]("t:fresh") - } yield { - val (a, b) = res.getOrElse(fail("expected a committed transaction")) - assertEquals(a, Right(1L)) - assert(b.isLeft, s"expected the INCR on a string to fail, got $b") - assertEquals(ok, Some(1)) // Redis does not roll back, so the first INCR remains committed after the second one fails. - } - } + clientTest("an execution-phase error surfaces per-position while the other commands commit") { client => + client.set("t:str", "x") >> + client + .transaction(tx => tx.execAttempt((Commands.incr[String]("t:fresh"), Commands.incr[String]("t:str")))) + .is(Some((Right(1L), Left(ServerError("ERR", "value is not an integer or out of range"))))) >> + client.get[Int]("t:fresh").is(Some(1)) // Redis does not roll back, so the first INCR remains committed after the second one fails. } - test("closing the client releases its server connection") { - withContainers { server => - connectAndUse(configOf(server)) { observer => - for { - subject <- Client.connect(configOf(server)) - before <- connectionCount(observer) - _ <- subject.close - _ <- awaitConnectionCount(observer, before - 1, attempts = 50) - } yield () - }.unsafeRun + serverTest("closing the client releases its server connection") { server => + connectAndUse(configOf(server)) { observer => + connectAndUse(configOf(server))(_ => connectionCount(observer)).flatMap(before => Eventually(50)(connectionCount(observer).is(before - 1))) } } - private val clientList: Command[String] = - Command( - "CLIENT", - keyIndices = Command.NoKeys, - args = Vector(Bytes.utf8("LIST")), - decode = { - case Frame.BulkString(value) => Right(value.asUtf8String) - case Frame.VerbatimString(_, value) => Right(value.asUtf8String) - case other => Left(DecodeError("bulk or verbatim string", Frame.describe(other))) - } - ) - private def connectionCount(client: Client[CIO, String]): CIO[Int] = - client.run(clientList).map(_.linesIterator.count(_.nonEmpty)) - - private def awaitConnectionCount(client: Client[CIO, String], expected: Int, attempts: Int): CIO[Unit] = - Eventually.converges(attempts)(() => connectionCount(client))(_ == expected)(count => s"expected $expected connections, still $count") + client.clientList.map(_.linesIterator.count(_.nonEmpty)) } - -class RedisRoundTripSuite extends RoundTripSuite(Images.redis) - -class ValkeyRoundTripSuite extends RoundTripSuite(Images.valkey) diff --git a/integration-tests/shared/src/test/scala/sage/integration/ServerSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/ServerSuite.scala index 3adecd2f..95e6d24f 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/ServerSuite.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/ServerSuite.scala @@ -1,20 +1,54 @@ package sage.integration -import scala.concurrent.{ExecutionContext, Future} - import com.dimafeng.testcontainers.GenericContainer -import com.dimafeng.testcontainers.munit.TestContainerForAll +import com.dimafeng.testcontainers.lifecycle.and +import com.dimafeng.testcontainers.munit.{TestContainerForAll, TestContainersForAll} import kyo.compat.* +import munit.{Location, TestOptions} import sage.client.internal.Client -abstract class ServerSuite(image: String) extends munit.FunSuite with TestContainerForAll with ContainerClient { +trait ServerTests extends ContainerClient { + + protected def serverTest(options: TestOptions)(body: GenericContainer => CIO[Any])(using Location): Unit + + protected def clientTest(options: TestOptions)(body: Client[CIO, String] => CIO[Any])(using Location): Unit = + serverTest(options)(server => connectAndUse(configOf(server))(body)) + + protected def clientsTest(options: TestOptions)(body: (Client[CIO, String], Client[CIO, String]) => CIO[Any])(using Location): Unit = + serverTest(options)(server => connectAndUse(configOf(server))(first => connectAndUse(configOf(server))(body(first, _)))) +} + +abstract class ServerSuite(image: String) extends ServerTests with TestContainerForAll { + + override val containerDef: GenericContainer.Def[GenericContainer] = serverDef(image) + + protected def serverTest(options: TestOptions)(body: GenericContainer => CIO[Any])(using Location): Unit = containerTest(options)(body) +} + +abstract class BothServersSuite extends ServerTests with TestContainersForAll { + + override type Containers = GenericContainer and GenericContainer + + protected def redisDef: GenericContainer.Def[GenericContainer] = serverDef(Images.redis) + protected def valkeyDef: GenericContainer.Def[GenericContainer] = serverDef(Images.valkey) + + override def startContainers(): Containers = redisDef.start() and valkeyDef.start() + + protected def serverTest(options: TestOptions)(body: GenericContainer => CIO[Any])(using Location): Unit = { + onRedis(options)(body) + onValkey(options)(body) + } + + protected def onRedis(options: TestOptions)(body: GenericContainer => CIO[Any])(using Location): Unit = + containerTest(options.withName(s"${options.name} (redis)")) { case redis and _ => body(redis) } - override val containerDef: GenericContainer.Def[GenericContainer] = GenericContainer.Def(image, exposedPorts = Seq(6379)) + protected def onValkey(options: TestOptions)(body: GenericContainer => CIO[Any])(using Location): Unit = + containerTest(options.withName(s"${options.name} (valkey)")) { case _ and valkey => body(valkey) } - // The Ox cell's unsafeRun uses this value. Keeping it non-private avoids unused-private warnings in the other cells. - given ExecutionContext = munitExecutionContext + protected def redisTest(options: TestOptions)(body: Client[CIO, String] => CIO[Any])(using Location): Unit = + onRedis(options)(server => connectAndUse(configOf(server))(body)) - protected def withClient[A](body: Client[CIO, String] => CIO[A]): Future[A] = - withContainers(server => connectAndUse(configOf(server))(body).unsafeRun) + protected def valkeyTest(options: TestOptions)(body: Client[CIO, String] => CIO[Any])(using Location): Unit = + onValkey(options)(server => connectAndUse(configOf(server))(body)) } diff --git a/integration-tests/shared/src/test/scala/sage/integration/SmokeSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/SmokeSuite.scala new file mode 100644 index 00000000..3e58c08f --- /dev/null +++ b/integration-tests/shared/src/test/scala/sage/integration/SmokeSuite.scala @@ -0,0 +1,23 @@ +package sage.integration + +import scala.concurrent.duration.* + +import kyo.compat.* + +import sage.client.internal.Client + +// A Pekko Future starts when it is built, so each later step is built inside flatMap instead of chained with >>. +abstract class SmokeSuite extends ServerSuite(Images.redis) { + + protected def lockScopesNativeEffects[F[_]](client: Client[F, String])(lift: [A] => F[A] => CIO[A]): CIO[Unit] = { + val locks = client.lock[String]() + lift(locks.withLock("native-lock", 2.seconds)(locks.tryWithLock("native-lock")(fail("contended body ran")))) + .is(None) + .flatMap(_ => lift(locks.tryWithLock("native-lock")(client.ping())).is(Some("PONG"))) + } + + protected def scanAllFindsEveryKey[F[_]](client: Client[F, String])(lift: [A] => F[A] => CIO[A])(scan: => F[Iterable[String]]): CIO[Unit] = { + val keys = (1 to 50).map(i => s"scan-$i") + inSequence(keys)(key => lift(client.set(key, "v"))).flatMap(_ => lift(scan).map(_.toSet).is(keys.toSet)) + } +} diff --git a/integration-tests/shared/src/test/scala/sage/integration/Ttls.scala b/integration-tests/shared/src/test/scala/sage/integration/Ttls.scala index 73d3f163..a08a6b77 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/Ttls.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/Ttls.scala @@ -9,23 +9,16 @@ import sage.commands.{FieldTtl, Ttl} */ object Ttls { - def remaining(ttl: Ttl): Option[FiniteDuration] = - ttl match { - case Ttl.Expires(value) => Some(value) - case _ => None - } - - private def remaining(ttl: FieldTtl): Option[FiniteDuration] = + private def remaining(ttl: Ttl | FieldTtl): Option[FiniteDuration] = ttl match { + case Ttl.Expires(value) => Some(value) case FieldTtl.Expires(value) => Some(value) case _ => None } - def expires(ttl: Ttl): Boolean = remaining(ttl).exists(_ > Duration.Zero) - - def expires(ttl: FieldTtl): Boolean = remaining(ttl).exists(_ > Duration.Zero) - - def expiresWithin(ttl: Ttl, bound: FiniteDuration): Boolean = remaining(ttl).exists(value => value > Duration.Zero && value <= bound) + def expiresWithin(ttl: Ttl | FieldTtl, bound: FiniteDuration, above: FiniteDuration = Duration.Zero): Boolean = + remaining(ttl).exists(value => value > above && value <= bound) - def expiresWithin(ttl: FieldTtl, bound: FiniteDuration): Boolean = remaining(ttl).exists(value => value > Duration.Zero && value <= bound) + def renewed(before: Ttl, after: Ttl): Boolean = + remaining(before).zip(remaining(after)).exists((previous, current) => current > previous + 100.millis) } diff --git a/integration-tests/shared/src/test/scala/sage/integration/cluster/ClusterFailoverSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/cluster/ClusterFailoverSuite.scala index 90757f13..f1974c3c 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/cluster/ClusterFailoverSuite.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/cluster/ClusterFailoverSuite.scala @@ -1,19 +1,16 @@ package sage.integration.cluster -import scala.concurrent.ExecutionContext import scala.concurrent.duration.* -import scala.util.Try import com.dimafeng.testcontainers.FixedHostPortGenericContainer -import com.dimafeng.testcontainers.munit.TestContainerForEach +import com.dimafeng.testcontainers.munit.TestContainersForEach import kyo.compat.* import sage.Bytes -import sage.client.{ClusterConfig, Endpoint, SageConfig, Topology} import sage.client.internal.Client -import sage.cluster.Slot +import sage.cluster.{Node, Slot} import sage.commands.{Commands, Connection} -import sage.integration.{ContainerClient, Eventually, Images} +import sage.integration.{Eventually, Images} /** * Drives cluster failover recovery against a real multi-node cluster: three masters and their replicas in one container. When a master @@ -23,94 +20,44 @@ import sage.integration.{ContainerClient, Eventually, Images} * * The fixed host ports are 7100-7105; see [[MultiNodeCluster]] for why a multi-node cluster cannot use mapped random ports. */ -abstract class ClusterFailoverSuite(val image: String, val serverBinary: String) - extends munit.FunSuite - with TestContainerForEach - with ContainerClient - with MultiNodeCluster { - - protected val ports: Seq[Int] = 7100 to 7105 - private val victim = 7100 // redis-cli --cluster-create makes the first nodes masters, so 7100 is a master with a replica to promote - - override protected val replicasPerMaster: Option[Int] = Some(1) - - // start the replica's initial sync immediately, not after the default 5s window, so it is a live copy before the failover - override protected val extraServerFlags: Vector[String] = Vector("--repl-diskless-sync-delay 0") - - override val containerDef: FixedHostPortGenericContainer.Def = clusterContainerDef - - given ExecutionContext = munitExecutionContext - - // forming a six-node cluster and waiting out an election runs past munit's 30s default on a loaded CI box - override def munitTimeout: Duration = 120.seconds - - private def parseKeys(out: String): Vector[String] = - out.split("\n").iterator.map(_.trim).filter(_.nonEmpty).toVector - - private def victimReplicaPort(container: FixedHostPortGenericContainer): Int = { - val nodes = clusterNodes(container, victim) - val myId = - nodes.collectFirst { case node if node.isMyself => node.id }.getOrElse(throw new RuntimeException(s"victim $victim has no myself line")) - nodes - .collectFirst { case node if node.isReplica && node.masterId == myId => node.port } - .getOrElse(throw new RuntimeException(s"no replica found for victim $victim among ${nodes.map(_.port)}")) - } - - // The replication barrier: poll the victim's own replica until it holds every victim-owned key. WAIT keys off the calling connection's last - // write, so it would not cover writes the cluster client sent on its own routed connections; reading the replica directly proves recovery. - private def awaitReplicated(container: FixedHostPortGenericContainer, replicaPort: Int, expected: Set[String], attempts: Int): CIO[Unit] = - Eventually.converges(attempts)(() => CIO.blocking(parseKeys(cli(container, replicaPort, "keys", "*")).toSet))(expected.subsetOf)(have => - s"victim's replica did not catch up; missing ${(expected -- have).take(5)}" - ) +abstract class ClusterFailoverSuite(image: String, serverBinary: String) + extends MultiNodeCluster(image, serverBinary, nodeCount = 6, replicasPerMaster = 1) + with TestContainersForEach { + + private val victim = basePort // redis-cli --cluster-create makes the first nodes masters, so 7100 is a master with a replica to promote + + private def crashVictim(onVictim: Vector[String]): CIO[Int] = + for { + topology <- clusterTopology(victim) + replica <- required(s"a replica of victim $victim", topology.replicasForMaster(Node("127.0.0.1", victim)).headOption) + // The replication barrier: poll the victim's own replica until it holds every victim-owned key. WAIT keys off the calling connection's last + // write, so it would not cover writes the cluster client sent on its own routed connections; reading the replica directly proves recovery. + _ <- Eventually(100)(onNode(replica.port)(_.keys("*")).satisfies(keys => onVictim.toSet.subsetOf(keys.toSet))) + _ <- shutdown(node(victim)) + } yield replica.port // `cluster_state:ok` flips before every node will actually serve writes, so a freshly formed cluster can briefly answer CLUSTERDOWN; retry // each write across that warm-up window, as a real application would, so the failover the test means to exercise is not masked by a startup race - private def writeKey(client: Client[CIO, String], key: String, attempts: Int): CIO[Unit] = - client - .set(key, key) - .fold( - _ => CIO.value(()), - error => if (attempts <= 0) CIO.fail(error) else CIO.sleep(200.millis).flatMap(_ => writeKey(client, key, attempts - 1)) - ) - - private def writeAll(client: Client[CIO, String], keys: Vector[String]): CIO[Unit] = - keys.foldLeft(CIO.value(()))((acc, key) => acc.flatMap(_ => writeKey(client, key, 150))) - - // retries a transport error, but fails at once on a `reject` value: retrying a stale success would let it slip through behind a later refresh - private def awaitReadEquals(read: () => CIO[Option[String]], expected: String, reject: Option[String], attempts: Int): CIO[Boolean] = - read().fold( - { - case Some(v) if v == expected => CIO.value(true) - case Some(v) if reject.contains(v) => CIO.value(false) - case _ if attempts <= 0 => CIO.value(false) - case _ => CIO.sleep(200.millis).flatMap(_ => awaitReadEquals(read, expected, reject, attempts - 1)) - }, - _ => if (attempts <= 0) CIO.value(false) else CIO.sleep(200.millis).flatMap(_ => awaitReadEquals(read, expected, reject, attempts - 1)) - ) - - private def recoverKey(client: Client[CIO, String], key: String, attempts: Int): CIO[Boolean] = - awaitReadEquals(() => client.get[String](key), key, None, attempts) - - private def recoverCached(client: Client[CIO, String], key: String, expected: String, stale: String, attempts: Int): CIO[Boolean] = - awaitReadEquals(() => client.cached(Commands.get[String, String](key), 1.minute), expected, Some(stale), attempts) - - private def recoverAll(client: Client[CIO, String], keys: Vector[String]): CIO[Boolean] = - keys.foldLeft(CIO.value(true))((acc, key) => acc.flatMap(ok => if (!ok) CIO.value(false) else recoverKey(client, key, 150))) + private def seedVictim(client: Client[CIO, String], prefix: String): CIO[(String, Vector[String])] = + inSequence((1 to 30).map(i => s"$prefix:$i"))(key => Eventually.succeeds(150, 200.millis)(client.set(key, key))) + .flatMap(_ => onNode(victim)(_.keys("*"))) + .flatMap { + case keys @ (probe +: _) => CIO.value((probe, keys)) + case _ => CIO.fail(new AssertionError("no keys landed on the victim master; cannot prove failover recovery")) + } + + // retries a transport error, but stops at once on the stale value: retrying a stale success would let it slip through behind a later refresh + private def awaitRead(read: CIO[Option[String]], expected: String, stale: Option[String]): CIO[Unit] = + Eventually + .succeeds(150, 200.millis)(read.map(_.filter(v => v == expected || stale.contains(v))).flatMap(required("a fresh or stale read", _))) + .is(expected) + + private def recoverAll(client: Client[CIO, String], keys: Vector[String]): CIO[Unit] = + inSequence(keys)(key => awaitRead(client.get[String](key), key, None)) // retries until the node accepts the write (e.g. a replica once it is promoted to master) - private def writeDirect(container: FixedHostPortGenericContainer, port: Int, key: String, value: String, attempts: Int): CIO[Unit] = - Eventually.converges(attempts, 200.millis)(() => CIO.blocking(cli(container, port, "set", key, value)))(_.contains("OK"))(out => - s"could not write $key on $port: $out" - ) - - private def masterId(container: FixedHostPortGenericContainer, port: Int): String = cli(container, port, "cluster", "myid").trim - - private def masterPortsExcludingVictim(container: FixedHostPortGenericContainer): Vector[Int] = - clusterNodes(container, victim).filter(_.isMaster).map(_.port).filter(_ != victim) - - // Query the node directly because its own CLUSTER NODES entry contains the `myself` flag. - private def ownSlotCount(container: FixedHostPortGenericContainer, port: Int): Int = - clusterNodes(container, port).collectFirst { case node if node.isMyself => node.ownedSlotCount }.getOrElse(0) + private def writeDirect(port: Int, key: String, value: String): CIO[Unit] = + Eventually.succeeds(150, 200.millis)(onNode(port)(_.set(key, value)).unit) private def reshard(container: FixedHostPortGenericContainer, fromId: String, toId: String, slots: Int): String = exec( @@ -128,145 +75,63 @@ abstract class ClusterFailoverSuite(val image: String, val serverBinary: String) "--cluster-yes" ) - private def replicationWaits(container: FixedHostPortGenericContainer, port: Int): Long = - cli(container, port, "info", "commandstats").linesIterator - .find(_.startsWith("cmdstat_wait:calls=")) - .fold(0L)(_.stripPrefix("cmdstat_wait:calls=").takeWhile(_ != ',').toLong) - - test("distributed locks confirm acquisition and renewal on each master's replica") { - withContainers { container => - val config = SageConfig(topology = Topology.Cluster(Vector(Endpoint("127.0.0.1", ports.head)))) - formCluster(container).flatMap { _ => - connectAndUse(config) { client => - val keys = Vector("orders", "delta", "epsilon").map(tag => s"replicated-lock:{$tag}") - CIO.blocking(clusterNodes(container, ports.head)).flatMap { nodes => - val owners = keys.map(key => nodes.find(_.owns(Slot.of(Bytes.utf8(s"4:lock:$key")).value)).get) - assertEquals(owners.map(_.port).distinct.size, 3) - keys.zip(owners).foldLeft(CIO.unit) { case (previous, (key, owner)) => - previous.flatMap { _ => - val replica = nodes.find(node => node.isReplica && node.masterId == owner.id).get - val replicaConfig = SageConfig(topology = Topology.Standalone(Endpoint("127.0.0.1", replica.port))) - connectAndUse(replicaConfig) { reader => - for { - _ <- reader.run(Connection.readonly) - // Initial replica synchronization can finish after the cluster starts accepting writes. - _ <- Eventually.converges(100, 100.millis)(() => reader.run(Commands.role))(_.isConnectedReplica)(_ => - "cluster replica did not connect" - ) - waitsBefore <- CIO.blocking(replicationWaits(container, owner.port)) - _ <- client.lock[String](900.millis).withLock(key, 5.seconds) { - for { - before <- reader.exists(s"4:lock:$key") - _ <- - Eventually.converges(30, 50.millis)(() => CIO.blocking(replicationWaits(container, owner.port)))(_ >= waitsBefore + 2)( - waits => s"acquisition and renewal did not both wait for replication: $waits" - ) - after <- reader.exists(s"4:lock:$key") - waitsAfter <- CIO.blocking(replicationWaits(container, owner.port)) - } yield { - assertEquals(before, 1L) - assertEquals(after, 1L) - assert(waitsAfter >= waitsBefore + 2, "acquisition and renewal did not wait for replication") - } - } - } yield () - } - } - } - } + clusterTest("distributed locks confirm acquisition and renewal on each master's replica") { (_, client) => + val keys = Vector("orders", "delta", "epsilon").map(tag => s"replicated-lock:{$tag}") + clusterTopology(basePort).flatMap { topology => + val placed = keys.flatMap(key => + topology.nodeForSlot(Slot.of(Bytes.utf8(s"4:lock:$key"))).flatMap(owner => topology.replicasForMaster(owner).headOption.map((key, owner, _))) + ) + assertEquals(placed.map(_._2).distinct.size, 3) + inSequence(placed) { (key, owner, replica) => + onNode(replica.port) { reader => + for { + _ <- reader.run(Connection.readonly) + // Initial replica synchronization can finish after the cluster starts accepting writes. + _ <- Eventually(100)(reader.run(Commands.role).satisfies(_.isConnectedReplica)) + _ <- awaitCalls(onNode(owner.port)(_.info("commandstats")), "wait", 2)(replicated => + client.lock[String](900.millis).withLock(key, 5.seconds) { + reader.exists(s"4:lock:$key").is(1L) >> replicated >> reader.exists(s"4:lock:$key").is(1L) + } + ) + } yield () } - }.unsafeRun + } } } - test("the client recovers reads after a master crashes and its replica is promoted") { - withContainers { container => - val seeds = ports.map(p => Endpoint("127.0.0.1", p)).toVector - // short refresh interval so the topology refresh keeps pace with the caller's retries during the election - val config = SageConfig(topology = Topology.Cluster(seeds, ClusterConfig(minRefreshInterval = 500.millis))) - val keys = (1 to 30).map(i => s"failover:$i").toVector - - val program = - formCluster(container).flatMap { _ => - connectAndUse(config) { client => - for { - _ <- writeAll(client, keys) - // reading the failed node's keys proves that the client connected to the promoted replica. - onVictim <- CIO.blocking(parseKeys(cli(container, victim, "keys", "*"))) - replicaPort <- CIO.blocking(victimReplicaPort(container)) - _ <- awaitReplicated(container, replicaPort, onVictim.toSet, 100) - _ <- CIO.blocking(Try(cli(container, victim, "shutdown", "nosave"))) - recovered <- recoverAll(client, onVictim) - } yield { - assert(onVictim.nonEmpty, "no keys landed on the victim master; cannot prove failover recovery") - assert(recovered, "client did not recover the victim master's keys after its replica was promoted") - } - } - } - program.unsafeRun - } + clusterTest("the client recovers reads after a master crashes and its replica is promoted") { (_, client) => + for { + // reading the failed node's keys proves that the client connected to the promoted replica. + (_, onVictim) <- seedVictim(client, "failover") + _ <- crashVictim(onVictim) + _ <- recoverAll(client, onVictim) + } yield () } - test("a cached read follows MOVED to the new owner after its slot is resharded off its master") { - withContainers { container => - val seeds = ports.map(p => Endpoint("127.0.0.1", p)).toVector - val config = SageConfig(topology = Topology.Cluster(seeds, ClusterConfig(minRefreshInterval = 500.millis))) - val keys = (1 to 30).map(i => s"reshard:$i").toVector - - val program = - formCluster(container).flatMap { _ => - connectAndUse(config) { client => - for { - _ <- writeAll(client, keys) - onVictim <- CIO.blocking(parseKeys(cli(container, victim, "keys", "*"))) - probe = onVictim.head - first <- client.cached(Commands.get[String, String](probe), 1.minute) - dest <- CIO.blocking(masterPortsExcludingVictim(container).head) - fromId <- CIO.blocking(masterId(container, victim)) - toId <- CIO.blocking(masterId(container, dest)) - moved <- CIO.blocking(ownSlotCount(container, victim)) - _ <- CIO.blocking(reshard(container, fromId, toId, moved)) - _ <- awaitClusterOk(container, 60) - _ <- writeDirect(container, dest, probe, "reshard-fresh", 50) - recovered <- recoverCached(client, probe, "reshard-fresh", probe, 150) - } yield { - assert(onVictim.nonEmpty, "no keys landed on the victim master; cannot prove reshard recovery") - assertEquals(first, Some(probe)) - assert(recovered, "cached read did not observe the resharded slot's new owner's current value") - } - } - } - program.unsafeRun - } + clusterTest("a cached read follows MOVED to the new owner after its slot is resharded off its master") { (container, client) => + for { + (probe, _) <- seedVictim(client, "reshard") + _ <- client.cached(Commands.get[String, String](probe), 1.minute).is(Some(probe)) + topology <- clusterTopology(victim) + dest <- required("another master", topology.masters.find(_.port != victim)) + owned = (0 until Slot.Count).flatMap(Slot.at).count(topology.nodeForSlot(_).exists(_.port == victim)) + fromId <- onNode(victim)(_.clusterMyId) + toId <- onNode(dest.port)(_.clusterMyId) + _ <- CIO.blocking(reshard(container, fromId, toId, owned)) + _ <- awaitClusterOk + _ <- writeDirect(dest.port, probe, "reshard-fresh") + _ <- awaitRead(client.cached(Commands.get[String, String](probe), 1.minute), "reshard-fresh", Some(probe)) + } yield () } - test("a cached read recovers from the promoted master after a failover, never serving the dead master's entry") { - withContainers { container => - val seeds = ports.map(p => Endpoint("127.0.0.1", p)).toVector - val config = SageConfig(topology = Topology.Cluster(seeds, ClusterConfig(minRefreshInterval = 500.millis))) - val keys = (1 to 30).map(i => s"cachedfailover:$i").toVector - - val program = - formCluster(container).flatMap { _ => - connectAndUse(config) { client => - for { - _ <- writeAll(client, keys) - onVictim <- CIO.blocking(parseKeys(cli(container, victim, "keys", "*"))) - probe = onVictim.head - _ <- client.cached(Commands.get[String, String](probe), 1.minute) - replicaPort <- CIO.blocking(victimReplicaPort(container)) - _ <- awaitReplicated(container, replicaPort, onVictim.toSet, 100) - _ <- CIO.blocking(Try(cli(container, victim, "shutdown", "nosave"))) - _ <- writeDirect(container, replicaPort, probe, "failover-fresh", 150) - recovered <- recoverCached(client, probe, "failover-fresh", probe, 150) - } yield { - assert(onVictim.nonEmpty, "no keys landed on the victim master; cannot prove cached failover recovery") - assert(recovered, "cached read did not observe the promoted master's current value after the victim crashed") - } - } - } - program.unsafeRun - } + clusterTest("a cached read recovers from the promoted master after a failover, never serving the dead master's entry") { (_, client) => + for { + (probe, onVictim) <- seedVictim(client, "cachedfailover") + _ <- client.cached(Commands.get[String, String](probe), 1.minute) + replicaPort <- crashVictim(onVictim) + _ <- writeDirect(replicaPort, probe, "failover-fresh") + _ <- awaitRead(client.cached(Commands.get[String, String](probe), 1.minute), "failover-fresh", Some(probe)) + } yield () } } diff --git a/integration-tests/shared/src/test/scala/sage/integration/cluster/ClusterMultiMasterSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/cluster/ClusterMultiMasterSuite.scala index d6ab33c3..4a2583c6 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/cluster/ClusterMultiMasterSuite.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/cluster/ClusterMultiMasterSuite.scala @@ -1,179 +1,90 @@ package sage.integration.cluster -import scala.concurrent.{ExecutionContext, Future} import scala.concurrent.duration.* import com.dimafeng.testcontainers.FixedHostPortGenericContainer -import com.dimafeng.testcontainers.munit.TestContainerForAll +import com.dimafeng.testcontainers.munit.TestContainersForAll import kyo.compat.* +import sage.Bytes import sage.SageException.InvalidArgument -import sage.client.{Endpoint, SageConfig, Topology} -import sage.client.internal.Client +import sage.cluster.Slot import sage.commands.Commands -import sage.integration.{ContainerClient, Eventually, Images} +import sage.integration.{Eventually, Images} /** * Drives the keyless broadcast routing against a real multi-master cluster: three masters, no replicas, in one container. A node answers * `PUBSUB` introspection only for the subscribers attached to it, so the merge across masters needs more than one master to be visible; * [[ClusterSuite]] forms a single-node cluster owning every slot and cannot express it. Subscribers are attached with `redis-cli` on chosen * nodes, so no assertion depends on which master Sage pins its own subscription connection to. - * - * Each backend takes its own port range so the Redis and Valkey rows cannot collide on the host. */ -abstract class ClusterMultiMasterSuite(val image: String, val serverBinary: String, basePort: Int) - extends munit.FunSuite - with TestContainerForAll - with ContainerClient - with MultiNodeCluster { +abstract class ClusterMultiMasterSuite(image: String, serverBinary: String) + extends MultiNodeCluster(image, serverBinary, nodeCount = 3) + with TestContainersForAll { - protected val ports: Seq[Int] = basePort to (basePort + 2) + private def subscribeOn(container: FixedHostPortGenericContainer, port: Int, verb: String, name: String): Unit = + exec(container, "sh", "-c", s"nohup redis-cli -p $port $verb $name > /dev/null 2>&1 &"): Unit - override val containerDef: FixedHostPortGenericContainer.Def = clusterContainerDef + private def subscribed(container: FixedHostPortGenericContainer, port: Int, channel: String): CIO[Unit] = + CIO.blocking(subscribeOn(container, port, "subscribe", channel)) >> + Eventually(50)(onNode(port)(_.pubsubChannels()).satisfies(_.contains(channel))) - given ExecutionContext = munitExecutionContext - - override def munitTimeout: Duration = 120.seconds - - private val config = SageConfig(topology = Topology.Cluster(Vector(Endpoint("127.0.0.1", basePort)))) - - private def subscribeOn(container: FixedHostPortGenericContainer, port: Int, verb: String, name: String): Unit = { - exec(container, "sh", "-c", s"nohup redis-cli -p $port $verb $name > /dev/null 2>&1 &") - () + clusterTest("distributed locks acquire and release keys owned by different masters") { (_, first) => + val keys = Vector("orders", "delta", "epsilon").map(tag => s"distributed:{$tag}") + slotOwner.map(owner => keys.map(key => owner(s"4:lock:$key")).toSet).is(ports.toSet) >> + connectAndUse(clusterConfig)(second => inSequence(keys)(contend(first.lock(), second, _))) } - private def awaitOn(container: FixedHostPortGenericContainer, port: Int, form: String, name: String): CIO[Unit] = - Eventually.converges(50, 100.millis)(() => CIO.blocking(cli(container, port, "pubsub", form)))(_.contains(name))(out => - s"$name did not appear in $form on $port: $out" + clusterTest("PUBSUB CHANNELS returns a channel whose only subscriber sits on a master the client never picked") { (container, client) => + val expected = ports.map(p => s"only-$p").toSet + inSequence(ports)(p => subscribed(container, p, s"only-$p")).flatMap(_ => + client.pubsubChannels().satisfies(channels => expected.subsetOf(channels.toSet)) ) - - private def onCluster[A](body: (FixedHostPortGenericContainer, Client[CIO, String]) => CIO[A]): Future[A] = - withContainers { container => - formCluster(container).flatMap(_ => connectAndUse(config)(client => body(container, client))).unsafeRun - } - - test("distributed locks acquire and release keys owned by different masters") { - onCluster { (container, first) => - val keys = Vector("orders", "delta", "epsilon").map(tag => s"distributed:{$tag}") - assertEquals(keys.map(key => slotOwner(container, s"4:lock:$key")).toSet, ports.toSet) - connectAndUse(config) { second => - CIO - .foreach(keys) { key => - val holder = first.lock[String](leaseDuration = 600.millis) - val contender = second.lock[String]() - for { - denied <- holder.withLock(key, 2.seconds)(contender.tryWithLock(key)(CIO.value(1))) - acquired <- contender.tryWithLock(key)(CIO.value(42)) - exists <- first.exists(s"4:lock:$key") - } yield { - assertEquals(denied, None) - assertEquals(acquired, Some(42)) - assertEquals(exists, 0L) - } - } - .unit - } - } } - test("PUBSUB CHANNELS returns a channel whose only subscriber sits on a master the client never picked") { - onCluster { (container, client) => - ports.foreach(p => subscribeOn(container, p, "subscribe", s"only-$p")) - val expected = ports.map(p => s"only-$p").toSet - for { - _ <- ports.foldLeft(CIO.value(()))((acc, p) => acc.flatMap(_ => awaitOn(container, p, "channels", s"only-$p"))) - channels <- client.pubsubChannels() - } yield assert(expected.subsetOf(channels.toSet), s"swept channels $channels missed ${expected -- channels.toSet}") - } + clusterTest("PUBSUB CHANNELS reports a channel held on two masters once, rather than once per master") { (container, client) => + inSequence(ports.take(2))(subscribed(container, _, "twice")).flatMap(_ => client.pubsubChannels().map(_.count(_ == "twice")).is(1)) } - test("PUBSUB CHANNELS reports a channel held on two masters once, rather than once per master") { - onCluster { (container, client) => - subscribeOn(container, ports(0), "subscribe", "twice") - subscribeOn(container, ports(1), "subscribe", "twice") - for { - _ <- awaitOn(container, ports(0), "channels", "twice") - _ <- awaitOn(container, ports(1), "channels", "twice") - channels <- client.pubsubChannels() - } yield assertEquals(channels.count(_ == "twice"), 1, s"expected one entry for a channel held on two masters, got $channels") - } + clusterTest("PUBSUB NUMSUB sums a channel's subscribers across masters instead of reporting one master's count") { (container, client) => + inSequence(ports.take(2))(subscribed(container, _, "summed")).flatMap(_ => client.pubsubNumSub("summed").map(_.get("summed")).is(Some(2L))) } - test("PUBSUB NUMSUB sums a channel's subscribers across masters instead of reporting one master's count") { - onCluster { (container, client) => - subscribeOn(container, ports(0), "subscribe", "summed") - subscribeOn(container, ports(1), "subscribe", "summed") - for { - _ <- awaitOn(container, ports(0), "channels", "summed") - _ <- awaitOn(container, ports(1), "channels", "summed") - counts <- client.pubsubNumSub("summed") - } yield assertEquals(counts.get("summed"), Some(2L)) - } + clusterTest("PUBSUB SHARDCHANNELS concatenates the shard channels of every master, one per shard") { (container, client) => + val oneChannelPerShard = Vector("orders", "delta", "epsilon") + slotOwner.map(owner => oneChannelPerShard.foreach(channel => subscribeOn(container, owner(channel), "ssubscribe", channel))) >> + Eventually(50, 200.millis)(client.pubsubShardChannels().satisfies(found => oneChannelPerShard.forall(found.contains))) } - test("PUBSUB SHARDCHANNELS concatenates the shard channels of every master, one per shard") { - onCluster { (container, client) => - val oneChannelPerShard = Vector("orders", "delta", "epsilon") - oneChannelPerShard.foreach(channel => subscribeOn(container, slotOwner(container, channel), "ssubscribe", channel)) - val sweptAll: Vector[String] => Boolean = found => oneChannelPerShard.forall(found.contains) - Eventually.converges(50, 200.millis)(() => client.pubsubShardChannels())(sweptAll)(found => - s"swept shard channels $found missed ${oneChannelPerShard.filterNot(found.contains)}" - ) - } + clusterTest("PUBSUB SHARDNUMSUB attributes each shard channel's subscribers to it, across all three owners") { (container, client) => + // Use one channel per slot range and give each a distinct subscriber count. The expected result therefore requires replies from all masters. + val expected = Map("sn-d" -> 1L, "sn-a" -> 2L, "sn-c" -> 3L) + slotOwner.map(owner => + expected.foreach((channel, subscribers) => (1L to subscribers).foreach(_ => subscribeOn(container, owner(channel), "ssubscribe", channel))) + ) >> + Eventually(50, 200.millis)(client.pubsubShardNumSub(expected.keys.toSeq*).is(expected)) } - test("PUBSUB SHARDNUMSUB attributes each shard channel's subscribers to it, across all three owners") { - onCluster { (container, client) => - // Use one channel per slot range and give each a distinct subscriber count. The expected result therefore requires replies from all masters. - val expected = Map("sn-d" -> 1L, "sn-a" -> 2L, "sn-c" -> 3L) - expected.foreach { case (channel, subscribers) => - val owner = slotOwner(container, channel) - (1L to subscribers).foreach(_ => subscribeOn(container, owner, "ssubscribe", channel)) - } - Eventually.converges(50, 200.millis)(() => client.pubsubShardNumSub(expected.keys.toSeq*))(_ == expected)(counts => - s"expected $expected, got $counts" - ) - } + clusterTest("PUBSUB NUMPAT sums the distinct patterns of every master") { (container, client) => + subscribeOn(container, basePort, "psubscribe", "pat-a.*") + subscribeOn(container, basePort + 1, "psubscribe", "pat-b.*") + Eventually(50, 200.millis)(client.pubsubNumPat.satisfies(_ >= 2L)) } - test("PUBSUB NUMPAT sums the distinct patterns of every master") { - onCluster { (container, client) => - subscribeOn(container, ports(0), "psubscribe", "pat-a.*") - subscribeOn(container, ports(1), "psubscribe", "pat-b.*") - Eventually.converges(50, 200.millis)(() => client.pubsubNumPat)(_ >= 2L)(total => s"expected both masters' patterns to be counted, got $total") - } + clusterTest("MEMORY PURGE runs on every master, not just the one the client picked") { (_, client) => + inSequence(ports)(onNode(_)(_.run(admin("CONFIG", "RESETSTAT")))) >> + client.memoryPurge >> + inSequence(ports)(onNode(_)(_.info("commandstats")).map(commandCalls(_, "memory|purge")).is(1L)) } - test("MEMORY PURGE runs on every master, not just the one the client picked") { - onCluster { (container, client) => - ports.foreach(p => cli(container, p, "config", "resetstat")) - client.memoryPurge.map { _ => - val unpurged = ports.filterNot(p => cli(container, p, "info", "commandstats").contains("cmdstat_memory|purge:calls=")) - assert(unpurged.isEmpty, s"masters that never saw MEMORY PURGE: $unpurged") - } - } + clusterTest("a Pipeline rejects PUBSUB introspection, since a broadcast cannot be batched onto one node") { (_, client) => + failsWith[InvalidArgument](client.pipeline((Commands.pubsubNumPat, Commands.get[String, String]("unused")))) } - test("a Pipeline rejects PUBSUB introspection, since a broadcast cannot be batched onto one node") { - onCluster { (_, client) => - client - .pipeline((Commands.pubsubNumPat, Commands.get[String, String]("unused"))) - .fold( - results => CIO.value(fail(s"expected the pipeline to be rejected, got $results")), - { - case _: InvalidArgument => CIO.value(()) - case other => CIO.value(fail(s"expected InvalidArgument, got $other")) - } - ) - } - } - - private def slotOwner(container: FixedHostPortGenericContainer, key: String): Int = { - val slot = cli(container, basePort, "cluster", "keyslot", key).trim.toInt - clusterNodes(container, basePort).find(node => node.isMaster && node.owns(slot)).map(_.port).getOrElse(fail(s"no master owns slot $slot")) - } + private def slotOwner: CIO[String => Int] = + clusterTopology(basePort).map(topology => key => topology.nodeForSlot(Slot.of(Bytes.utf8(key))).fold(fail(s"no master owns $key"))(_.port)) } -class RedisClusterMultiMasterSuite extends ClusterMultiMasterSuite(Images.redis, "redis-server", 7200) +class RedisClusterMultiMasterSuite extends ClusterMultiMasterSuite(Images.redis, "redis-server") -class ValkeyClusterMultiMasterSuite extends ClusterMultiMasterSuite(Images.valkey, "valkey-server", 7210) +class ValkeyClusterMultiMasterSuite extends ClusterMultiMasterSuite(Images.valkey, "valkey-server") diff --git a/integration-tests/shared/src/test/scala/sage/integration/cluster/ClusterSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/cluster/ClusterSuite.scala index 4aa29634..5a33cacb 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/cluster/ClusterSuite.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/cluster/ClusterSuite.scala @@ -1,329 +1,136 @@ package sage.integration.cluster -import scala.concurrent.ExecutionContext -import scala.concurrent.duration.* - import com.dimafeng.testcontainers.GenericContainer -import com.dimafeng.testcontainers.munit.TestContainerForAll +import com.dimafeng.testcontainers.lifecycle.and import kyo.compat.* +import munit.{Location, TestOptions} -import sage.{Bytes, Message} -import sage.SageException.{DecodeError, ServerError} +import sage.Message +import sage.SageException.ServerError import sage.client.{Endpoint, SageConfig, Topology} -import sage.client.internal.{Client, ScanTarget} -import sage.commands.{Command, Commands, FlushMode, ScanCursor} -import sage.integration.{ContainerClient, Eventually, Images} -import sage.protocol.Frame +import sage.client.internal.{Client, Paged} +import sage.commands.{Commands, FlushMode} +import sage.integration.{BothServersSuite, Eventually, Images} +import sage.protocol.Frames /** * Drives the cluster runtime against a real cluster-enabled server. One node owns all 16384 slots, with `cluster-announce` pointed at the * testcontainers-mapped host port so the address the node reports in `CLUSTER SLOTS` is reachable from the test. This exercises topology * discovery, the `CLUSTER SLOTS` decoder against real wire output, and single-key routing; redirects and failover need multiple nodes. */ -abstract class ClusterSuite(image: String, serverBinary: String, supportsNumberedDatabases: Boolean = false) - extends munit.FunSuite - with TestContainerForAll - with ContainerClient { - - override val containerDef: GenericContainer.Def[GenericContainer] = { - val numberedDatabases = if (supportsNumberedDatabases) Seq("--cluster-databases", "16") else Seq.empty - GenericContainer.Def(image, exposedPorts = Seq(6379), command = Seq(serverBinary, "--cluster-enabled", "yes") ++ numberedDatabases) - } +class ClusterSuite extends BothServersSuite { - given ExecutionContext = munitExecutionContext + override protected def redisDef: GenericContainer.Def[GenericContainer] = serverDef(Images.redis, command = Seq("--cluster-enabled", "yes")) + override protected def valkeyDef: GenericContainer.Def[GenericContainer] = + serverDef(Images.valkey, command = Seq("--cluster-enabled", "yes", "--cluster-databases", "16")) - private def admin(name: String, args: String*): Command[Unit] = - Command(name, Command.NoKeys, args.toVector.map(Bytes.utf8), _ => Right(())) - - private val clusterInfo: Command[String] = - Command( - "CLUSTER", - Command.NoKeys, - Vector(Bytes.utf8("INFO")), - { - case Frame.BulkString(bytes) => Right(bytes.asUtf8String) - case Frame.VerbatimString(_, bytes) => Right(bytes.asUtf8String) - case other => Left(DecodeError("bulk or verbatim string", Frame.describe(other))) - } - ) + override def afterContainersStart(containers: Containers): Unit = containers match { + case redis and valkey => prepare(formCluster(redis) >> formCluster(valkey)) + } // a single node owning every slot, announcing the host-mapped endpoint so the address it reports is reachable from the test - private def formSingleNodeCluster(admin0: Client[CIO, String], host: String, port: Int): CIO[Unit] = - for { - _ <- admin0.run(admin("CONFIG", "SET", "cluster-announce-ip", host)) - _ <- admin0.run(admin("CONFIG", "SET", "cluster-announce-port", port.toString)) - // idempotent across tests sharing one container: only claim the slots if they are not already assigned - info <- admin0.run(clusterInfo) - _ <- if (info.contains("cluster_state:ok")) CIO.value(()) else admin0.run(admin("CLUSTER", "ADDSLOTSRANGE", "0", "16383")) - _ <- awaitClusterOk(admin0, 50) - } yield () - - private def awaitClusterOk(admin0: Client[CIO, String], attempts: Int): CIO[Unit] = - Eventually.converges(attempts)(() => admin0.run(clusterInfo))(_.contains("cluster_state:ok"))(info => s"cluster did not converge: $info") - - private def awaitCached(client: Client[CIO, String], key: String, expected: String, attempts: Int): CIO[Option[String]] = - Eventually.value(attempts)(() => client.cached(Commands.get[String, String](key), 1.minute))(_.contains(expected)) - - // one shared container per suite, so the cluster is formed once and all routing exercised in a single test - test("single-key commands, pipelines, and transactions route against a real cluster") { - withContainers { server => - val host = server.host - val port = server.mappedPort(6379) - val standalone = SageConfig(topology = Topology.Standalone(Endpoint(host, port))) - val clustered = SageConfig(topology = Topology.Cluster(Vector(Endpoint(host, port)))) - - val program = - connectAndUse(standalone)(formSingleNodeCluster(_, host, port)).flatMap { _ => - connectAndUse(clustered) { client => - for { - _ <- client.set("greeting", "hello") - value <- client.get[String]("greeting") - count <- client.incr("counter") - _ <- client.set("{t}a", "1") - _ <- client.set("{t}b", "2") - piped <- client.pipeline((Commands.get[String, String]("{t}a"), Commands.get[String, String]("{t}b"))) - commit <- client.transaction(tx => tx.exec(Vector(Commands.incr[String]("{t}c"), Commands.incr[String]("{t}c")))) - } yield { - assertEquals(value, Some("hello")) - assertEquals(count, 1L) - assertEquals(piped, (Some("1"), Some("2"))) - assertEquals(commit, Some(Vector(1L, 2L))) - } - } - } - program.unsafeRun + private def formCluster(server: GenericContainer): CIO[Unit] = + connectAndUse(configOf(server)) { admin0 => + admin0.configSet("cluster-announce-ip" -> server.host, "cluster-announce-port" -> server.mappedPort(6379).toString) >> + admin0.run(admin("CLUSTER", "ADDSLOTSRANGE", "0", "16383")) >> + Eventually(50)(admin0.clusterInfo.satisfies(_.contains("cluster_state:ok"))) } - } - test("supported cross-slot commands are transparently split and merged against a real cluster") { - withContainers { server => - val host = server.host - val port = server.mappedPort(6379) - val standalone = SageConfig(topology = Topology.Standalone(Endpoint(host, port))) - val clustered = SageConfig(topology = Topology.Cluster(Vector(Endpoint(host, port)))) - val keyA = "{mget-a}value" - val keyB = "{mget-b}value" - val missingA = "{mget-a}missing" - val missingB = "{mget-b}missing" - val msetA = "{mset-a}value" - val msetB = "{mset-b}value" + private def clusterConfig(server: GenericContainer, database: Int = 0): SageConfig = + SageConfig(topology = Topology.Cluster(Vector(Endpoint(server.host, server.mappedPort(6379)))), database = database) + + private def clusterTest(options: TestOptions)(body: Client[CIO, String] => CIO[Any])(using Location): Unit = + serverTest(options)(server => connectAndUse(clusterConfig(server))(body)) + + // one shared container per server, so each cluster is formed once and all routing exercised in a single test + clusterTest("single-key commands, pipelines, and transactions route against a real cluster") { client => + client.set("greeting", "hello") >> + client.get[String]("greeting").is(Some("hello")) >> + client.incr("counter").is(1L) >> + client.set("{t}a", "1") >> + client.set("{t}b", "2") >> + client.pipeline((Commands.get[String, String]("{t}a"), Commands.get[String, String]("{t}b"))).is((Some("1"), Some("2"))) >> + client.transaction(tx => tx.exec(Vector(Commands.incr[String]("{t}c"), Commands.incr[String]("{t}c")))).is(Some(Vector(1L, 2L))) + } - val program = - connectAndUse(standalone)(formSingleNodeCluster(_, host, port)).flatMap { _ => - connectAndUse(clustered) { client => - for { - _ <- client.set(keyA, "a") - _ <- client.set(keyB, "b") - values <- client.mGet[String](keyA, keyB, missingA, keyB) - piped <- client.pipeline((Commands.mGet[String, String](keyA, keyB), Commands.get[String, String](keyA))) - exists <- client.exists(keyA, keyB, missingA, keyB) - touched <- client.touch(keyA, keyB, missingA) - _ <- client.mSet(msetA -> "set-a", msetB -> "set-b") - setValues <- client.mGet[String](msetA, msetB) - deleted <- client.del(keyA, missingB) - unlinked <- client.unlink(keyB, missingA) - } yield { - assertEquals(values, Vector(Some("a"), Some("b"), None, Some("b"))) - assertEquals(piped, (Vector(Some("a"), Some("b")), Some("a"))) - assertEquals(exists, 3L) - assertEquals(touched, 2L) - assertEquals(setValues, Vector(Some("set-a"), Some("set-b"))) - assertEquals(deleted, 1L) - assertEquals(unlinked, 1L) - } - } - } - program.unsafeRun - } + clusterTest("supported cross-slot commands are transparently split and merged against a real cluster") { client => + val keyA = "{mget-a}value" + val keyB = "{mget-b}value" + val missingA = "{mget-a}missing" + val missingB = "{mget-b}missing" + val msetA = "{mset-a}value" + val msetB = "{mset-b}value" + + client.set(keyA, "a") >> + client.set(keyB, "b") >> + client.mGet[String](keyA, keyB, missingA, keyB).is(Vector(Some("a"), Some("b"), None, Some("b"))) >> + client + .pipeline((Commands.mGet[String, String](keyA, keyB), Commands.get[String, String](keyA))) + .is((Vector(Some("a"), Some("b")), Some("a"))) >> + client.exists(keyA, keyB, missingA, keyB).is(3L) >> + client.touch(keyA, keyB, missingA).is(2L) >> + client.mSet(msetA -> "set-a", msetB -> "set-b") >> + client.mGet[String](msetA, msetB).is(Vector(Some("set-a"), Some("set-b"))) >> + client.del(keyA, missingB).is(1L) >> + client.unlink(keyB, missingA).is(1L) } // Sharded and classic pub/sub against a real (single-node) cluster: SSUBSCRIBE/SPUBLISH route by slot and coexist with classic SUBSCRIBE. // Resubscription on slot migration needs multiple nodes and is covered deterministically by ClusterClientSpec. - test("sharded and classic pub/sub coexist on a cluster client") { - withContainers { server => - val host = server.host - val port = server.mappedPort(6379) - val standalone = SageConfig(topology = Topology.Standalone(Endpoint(host, port))) - val clustered = SageConfig(topology = Topology.Cluster(Vector(Endpoint(host, port)))) - - val program = - connectAndUse(standalone)(formSingleNodeCluster(_, host, port)).flatMap { _ => - connectAndUse(clustered) { client => - for { - shard <- client.subscribeShardChannels[String]("orders") - classic <- client.subscribeChannels[String]("news") - sCount <- client.sPublish("orders", "placed") - cCount <- client.publish("news", "hello") - sMsg <- shard.next - cMsg <- classic.next - channels <- client.pubsubShardChannels() - _ <- shard.close - _ <- classic.close - } yield { - assertEquals(sCount, 1L) - assertEquals(cCount, 1L) - assertEquals(sMsg, Some(Message("orders", "placed"))) - assertEquals(cMsg, Some(Message("news", "hello"))) - assert(channels.contains("orders"), channels) - } - } - } - program.unsafeRun + clusterTest("sharded and classic pub/sub coexist on a cluster client") { client => + withSubscription(client.subscribeShardChannels[String]("orders")) { shard => + withSubscription(client.subscribeChannels[String]("news")) { classic => + client.sPublish("orders", "placed").is(1L) >> + client.publish("news", "hello").is(1L) >> + shard.next.is(Some(Message("orders", "placed"))) >> + classic.next.is(Some(Message("news", "hello"))) >> + client.pubsubShardChannels().satisfies(_.contains("orders")) + } } } - // The scanTargets method lists the slot-owning masters, and runOn keeps each page on the node that created its cursor. This fixture has one - // master for all slots, so one target is expected. - test("scanAll sweeps every slot-owning master via node-pinned runOn") { - withContainers { server => - val host = server.host - val port = server.mappedPort(6379) - val standalone = SageConfig(topology = Topology.Standalone(Endpoint(host, port))) - val clustered = SageConfig(topology = Topology.Cluster(Vector(Endpoint(host, port)))) - val expected = (1 to 50).map(i => s"cscan:$i").toSet - - def writeKeys(client: Client[CIO, String], i: Int): CIO[Unit] = - if (i > 50) CIO.value(()) else client.set(s"cscan:$i", i.toString).flatMap(_ => writeKeys(client, i + 1)) - - def scanNode(client: Client[CIO, String], target: ScanTarget, cursor: ScanCursor, found: Set[String]): CIO[Set[String]] = - client.runOn(target, Commands.scan[String](cursor, pattern = Some("cscan:*"), count = Some(10L))).flatMap { page => - page.next match { - case Some(next) => scanNode(client, target, next, found ++ page.items) - case None => CIO.value(found ++ page.items) - } - } - - def sweep(client: Client[CIO, String], targets: Vector[ScanTarget], found: Set[String]): CIO[Set[String]] = - targets match { - case head +: tail => scanNode(client, head, ScanCursor.start, found).flatMap(sweep(client, tail, _)) - case _ => CIO.value(found) - } - - val program = - connectAndUse(standalone)(formSingleNodeCluster(_, host, port)).flatMap { _ => - connectAndUse(clustered) { client => - for { - _ <- writeKeys(client, 1) - targets <- client.scanTargets - found <- sweep(client, targets, Set.empty[String]) - } yield { - assert(targets.forall(_.node.isDefined), s"cluster scan targets must be node-pinned: $targets") - assertEquals(found, expected) - } - } - } - program.unsafeRun - } + // scanTargets returns one target per slot-owning master, and each target runs every page on the node that created its cursor. This fixture + // has one master for all slots, so one target is expected. + clusterTest("scanAll sweeps every slot-owning master through node-pinned scan targets") { client => + val expected = (1 to 50).map(i => s"cscan:$i").toSet + + CIO.foreachDiscard(1 to 50)(i => client.set(s"cscan:$i", i.toString)) >> + client.runner.scanTargets.satisfies(targets => targets.nonEmpty && !targets.contains(client.runner)) >> + drain(Paged.scanAll[String](client.runner, Some("cscan:*"), Some(10L), None)).is(expected) } // SCRIPT LOAD and FUNCTION LOAD run on every master, allowing key-routed EVALSHA and FCALL to find them. This fixture has one master for // all slots, so the broadcast has one target. Multi-master clusters use the same dispatch logic. - test("SCRIPT LOAD and FUNCTION LOAD broadcast so a key-routed EVALSHA and FCALL resolve") { - withContainers { server => - val host = server.host - val port = server.mappedPort(6379) - val standalone = SageConfig(topology = Topology.Standalone(Endpoint(host, port))) - val clustered = SageConfig(topology = Topology.Cluster(Vector(Endpoint(host, port)))) - val library = - """#!lua name=clib - |redis.register_function('clib_get', function(keys, args) return redis.call('get', keys[1]) end) - |""".stripMargin + clusterTest("SCRIPT LOAD and FUNCTION LOAD broadcast so a key-routed EVALSHA and FCALL resolve") { client => + val library = + """#!lua name=clib + |redis.register_function('clib_get', function(keys, args) return redis.call('get', keys[1]) end) + |""".stripMargin - val program = - connectAndUse(standalone)(formSingleNodeCluster(_, host, port)).flatMap { _ => - connectAndUse(clustered) { client => - for { - sha <- client.scriptLoad("return redis.call('get', KEYS[1])") - _ <- client.set("bcast-key", "v") - eval <- client.evalSha(sha, Seq("bcast-key")) - _ <- client.functionFlush(Some(FlushMode.Sync)) - name <- client.functionLoad(library) - fcall <- client.fCall("clib_get", Seq("bcast-key")) - } yield { - assertEquals(sha.length, 40) - assertEquals(name, "clib") - eval match { - case Frame.BulkString(b) => assertEquals(b.asUtf8String, "v") - case other => fail(s"expected bulk string, got $other") - } - fcall match { - case Frame.BulkString(b) => assertEquals(b.asUtf8String, "v") - case other => fail(s"expected bulk string, got $other") - } - } - } - } - program.unsafeRun - } + for { + sha <- client.scriptLoad("return redis.call('get', KEYS[1])") + _ <- client.set("bcast-key", "v") + _ <- client.evalSha(sha, Seq("bcast-key")).is(Frames.bulk("v")) + _ <- client.functionFlush(Some(FlushMode.Sync)) + _ <- client.functionLoad(library).is("clib") + _ <- client.fCall("clib_get", Seq("bcast-key")).is(Frames.bulk("v")) + } yield assertEquals(sha.length, 40) } - test("a cluster cached read is served locally and a server-side write evicts it via invalidation") { - withContainers { server => - val host = server.host - val port = server.mappedPort(6379) - val standalone = SageConfig(topology = Topology.Standalone(Endpoint(host, port))) - val clustered = SageConfig(topology = Topology.Cluster(Vector(Endpoint(host, port)))) - - val program = - connectAndUse(standalone)(formSingleNodeCluster(_, host, port)).flatMap { _ => - connectAndUse(clustered) { reader => - connectAndUse(clustered) { writer => - for { - _ <- writer.set("csc:cluster", "v1") - first <- reader.cached(Commands.get[String, String]("csc:cluster"), 1.minute) - hit <- reader.cached(Commands.get[String, String]("csc:cluster"), 1.minute) - _ <- writer.set("csc:cluster", "v2") - evicted <- awaitCached(reader, "csc:cluster", "v2", attempts = 50) - } yield { - assertEquals(first, Some("v1")) - assertEquals(hit, Some("v1")) - assertEquals(evicted, Some("v2")) - } - } - } - } - program.unsafeRun - } + serverTest("a cluster cached read is served locally and a server-side write evicts it via invalidation") { server => + connectAndUse(clusterConfig(server))(reader => connectAndUse(clusterConfig(server))(cachedReadIsInvalidated(reader, _, "csc:cluster"))) } - if (supportsNumberedDatabases) - test("a numbered database is selected on a Valkey cluster connection") { - withContainers { server => - val host = server.host - val port = server.mappedPort(6379) - val standalone = SageConfig(topology = Topology.Standalone(Endpoint(host, port))) - val clustered = SageConfig(topology = Topology.Cluster(Vector(Endpoint(host, port))), database = 2) - val key = "cluster-numbered-database" - - val program = - connectAndUse(standalone)(formSingleNodeCluster(_, host, port)).flatMap { _ => - connectAndUse(standalone)(_.set(key, "database-0")).flatMap { _ => - connectAndUse(clustered) { client => - client.set(key, "database-2").flatMap(_ => client.get[String](key)).map(value => assertEquals(value, Some("database-2"))) - }.flatMap { _ => - connectAndUse(standalone)(_.get[String](key)).map(value => assertEquals(value, Some("database-0"))) - } - } - } - program.unsafeRun - } + onValkey("a numbered database is selected on a Valkey cluster connection") { server => + val key = "cluster-numbered-database" + connectAndUse(configOf(server))(_.set(key, "database-0")).flatMap { _ => + connectAndUse(clusterConfig(server, database = 2))(client => + client.set(key, "database-2").flatMap(_ => client.get[String](key).is(Some("database-2"))) + ) + .flatMap(_ => connectAndUse(configOf(server))(_.get[String](key).is(Some("database-0")))) } + } - if (!supportsNumberedDatabases) - test("an unsupported cluster server rejects a numbered database during bootstrap") { - withContainers { server => - val endpoint = Endpoint(server.host, server.mappedPort(6379)) - val clustered = SageConfig(topology = Topology.Cluster(Vector(endpoint)), database = 2) - - val attempted: CIO[Unit] = connectAndUse(clustered)(_ => CIO.value(())) - val program: CIO[Unit] = attempted.fold( - _ => CIO.fail(new AssertionError("expected the cluster connection to reject database 2")), - error => CIO.value(assert(error.isInstanceOf[ServerError], s"expected the server's SELECT error, got $error")) - ) - program.unsafeRun - } - } + onRedis("an unsupported cluster server rejects a numbered database during bootstrap") { server => + failsWith[ServerError](connectAndUse(clusterConfig(server, database = 2))(_ => CIO.unit)) + } } - -class RedisClusterSuite extends ClusterSuite(Images.redis, "redis-server") - -class ValkeyClusterSuite extends ClusterSuite(Images.valkey, "valkey-server", supportsNumberedDatabases = true) diff --git a/integration-tests/shared/src/test/scala/sage/integration/cluster/MultiNodeCluster.scala b/integration-tests/shared/src/test/scala/sage/integration/cluster/MultiNodeCluster.scala index eaa281a7..556ac969 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/cluster/MultiNodeCluster.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/cluster/MultiNodeCluster.scala @@ -4,20 +4,13 @@ import scala.concurrent.duration.* import com.dimafeng.testcontainers.FixedHostPortGenericContainer import kyo.compat.* +import munit.{Location, TestOptions} -import sage.integration.Eventually - -/** - * One `CLUSTER NODES` row. The wire format is positional: ` ...`. Owned slot ranges start at field - * 8. Parsing the row here avoids positional indexing at each use site. - */ -final case class ClusterNode(id: String, port: Int, flags: Set[String], masterId: String, slots: Vector[Range]) { - def isMaster: Boolean = flags.contains("master") - def isReplica: Boolean = flags.contains("slave") - def isMyself: Boolean = flags.contains("myself") - def owns(slot: Int): Boolean = slots.exists(_.contains(slot)) - def ownedSlotCount: Int = slots.map(_.size).sum -} +import sage.client.{ClusterConfig, Endpoint, SageConfig, Topology} +import sage.client.internal.Client +import sage.cluster.{ClusterTopology, Node} +import sage.commands.Cluster +import sage.integration.{ContainerClient, Eventually} /** * Boots several cluster nodes in one container and forms them into a cluster. @@ -25,98 +18,78 @@ final case class ClusterNode(id: String, port: Int, flags: Set[String], masterId * A cluster node announces a single address used for both gossip and clients, so the testcontainers-mapped random ports the single-node * suites use cannot work here: gossip needs an address the nodes reach each other on. The escape is a fixed 1:1 host-port mapping plus * `cluster-announce-ip 127.0.0.1`, so `127.0.0.1:` resolves to the same node inside the container (gossip) and from the host (the - * test). The cost is fixed host ports, which must be free on the host and must not overlap between suites. + * test). The cost is fixed host ports, which must be free on the host. */ -trait MultiNodeCluster { +trait MultiNodeCluster(image: String, serverBinary: String, nodeCount: Int, replicasPerMaster: Int = 0) extends ContainerClient { - protected def image: String - protected def serverBinary: String - protected def ports: Seq[Int] + override type Containers = FixedHostPortGenericContainer - /** - * Replicas per master for `--cluster create`; `None` makes every node a master. - */ - protected def replicasPerMaster: Option[Int] = None + final protected val basePort = 7100 - /** - * Server flags this suite needs on top of the shared ones. - */ - protected def extraServerFlags: Vector[String] = Vector.empty + final protected val ports: Range = basePort until basePort + nodeCount + + // waiting out an election runs past munit's 30s default on a loaded CI box + override def munitTimeout: Duration = 120.seconds + + // short refresh interval so the topology refresh keeps pace with the caller's retries during an election + protected val clusterConfig: SageConfig = + SageConfig(topology = Topology.Cluster(ports.map(p => Endpoint("127.0.0.1", p)).toVector, ClusterConfig(minRefreshInterval = 500.millis))) + + protected def clusterTest(options: TestOptions)(body: (FixedHostPortGenericContainer, Client[CIO, String]) => CIO[Any])(using Location): Unit = + containerTest(options)(container => connectAndUse(clusterConfig)(body(container, _))) // each node needs its own cluster-config-file (they otherwise collide on nodes.conf); a low node-timeout keeps a failover election short - final protected def clusterContainerDef: FixedHostPortGenericContainer.Def = { + override def startContainers(): FixedHostPortGenericContainer = { val starts = ports .map(p => - (Vector( + Vector( serverBinary, s"--port $p", "--cluster-enabled yes", s"--cluster-config-file nodes-$p.conf", "--cluster-node-timeout 2000", - "--cluster-announce-ip 127.0.0.1" - ) ++ extraServerFlags ++ Vector("--save ''", "--appendonly no", "--protected-mode no", "--daemonize yes")).mkString(" ") + "--cluster-announce-ip 127.0.0.1", + // start a replica's initial sync immediately, not after the default 5s window, so it is a live copy before a failover + "--repl-diskless-sync-delay 0", + // on an idle master, a replica's replication offset first moves with this PING (every 10s by default) + "--repl-ping-replica-period 1", + "--save ''", + "--appendonly no", + "--protected-mode no", + "--daemonize yes" + ).mkString(" ") ) .mkString("; ") - FixedHostPortGenericContainer.Def(image, command = Seq("sh", "-c", s"$starts; tail -f /dev/null"), portBindings = ports.map(p => (p, p)).toSeq) + FixedHostPortGenericContainer + .Def(image, command = Seq("sh", "-c", s"$starts; tail -f /dev/null"), portBindings = ports.map(p => (p, p)).toSeq) + .start() } + override def afterContainersStart(container: FixedHostPortGenericContainer): Unit = prepare(formCluster(container)) + final protected def exec(container: FixedHostPortGenericContainer, args: String*): String = { val result = container.execInContainer(args*) result.getStdout + result.getStderr } - final protected def cli(container: FixedHostPortGenericContainer, port: Int, args: String*): String = - exec(container, ("redis-cli" +: "-p" +: port.toString +: args)*) - - /** - * The cluster as `port` currently sees it, one entry per node. - */ - final protected def clusterNodes(container: FixedHostPortGenericContainer, port: Int): Vector[ClusterNode] = - cli(container, port, "cluster", "nodes").linesIterator.map(_.trim).filter(_.nonEmpty).flatMap(parseNode).toVector - - // a replica owns no slots, so its row stops at the link-state field and the slot tokens are absent - private def parseNode(line: String): Option[ClusterNode] = { - val fields = line.split("\\s+") - Option.when(fields.length >= 8)( - ClusterNode( - id = fields(0), - port = fields(1).split("@")(0).split(":").last.toInt, - flags = fields(2).split(",").toSet, - masterId = fields(3), - // a `[slot-<-nodeid]` marker is an in-flight migration, not an owned slot - slots = fields.drop(8).filterNot(_.startsWith("[")).toVector.flatMap(slotRange) - ) - ) - } + final protected def node(port: Int): SageConfig = SageConfig(topology = Topology.Standalone(Endpoint("127.0.0.1", port))) + + final protected def onNode[A](port: Int)(body: Client[CIO, String] => CIO[A]): CIO[A] = connectAndUse(node(port))(body) + + // the slot owners and replicas as `port` sees them, decoded the way the client decodes them + final protected def clusterTopology(port: Int): CIO[ClusterTopology] = + onNode(port)(_.run(Cluster.slots(Node("127.0.0.1", port)))).map(ClusterTopology.from) - private def slotRange(token: String): Option[Range] = - token.split("-") match { - case Array(only) => only.toIntOption.map(slot => slot to slot) - case Array(from, to) => from.toIntOption.zip(to.toIntOption).map((low, high) => low to high) - case _ => None - } - - final protected def awaitPortsUp(container: FixedHostPortGenericContainer, attempts: Int): CIO[Unit] = - Eventually.converges(attempts, 300.millis)(() => CIO.blocking(ports.forall(p => cli(container, p, "ping").contains("PONG"))))(identity)(_ => - "cluster nodes did not start" - ) - - final protected def awaitClusterOk(container: FixedHostPortGenericContainer, attempts: Int): CIO[Unit] = - Eventually.converges(attempts, 500.millis)(() => CIO.blocking(clusterIsOk(container)))(identity)(_ => "cluster did not converge") - - private def clusterIsOk(container: FixedHostPortGenericContainer): Boolean = - cli(container, ports.head, "cluster", "info").contains("cluster_state:ok") - - // idempotent, so a suite sharing one container across tests forms the cluster on the first test only - final protected def formCluster(container: FixedHostPortGenericContainer, attempts: Int = 60): CIO[Unit] = - awaitPortsUp(container, attempts).flatMap { _ => - CIO.blocking(clusterIsOk(container)).flatMap { ok => - if (ok) CIO.value(()) - else { - val replicas = replicasPerMaster.toVector.flatMap(n => Vector("--cluster-replicas", n.toString)) - val create = Vector("redis-cli", "--cluster", "create") ++ ports.map(p => s"127.0.0.1:$p") ++ replicas ++ Vector("--cluster-yes") - CIO.blocking(exec(container, create*)).flatMap(_ => awaitClusterOk(container, attempts)) - } - } - } + private def awaitPortsUp: CIO[Unit] = Eventually.succeeds(60, 300.millis)(inSequence(ports)(onNode(_)(_.ping()))) + + final protected def awaitClusterOk: CIO[Unit] = + Eventually(60, 500.millis)(onNode(basePort)(_.clusterInfo).satisfies(_.contains("cluster_state:ok"))) + + private def formCluster(container: FixedHostPortGenericContainer): CIO[Unit] = { + val create = Vector("redis-cli", "--cluster", "create") ++ ports.map(p => s"127.0.0.1:$p") ++ + Vector("--cluster-replicas", replicasPerMaster.toString, "--cluster-yes") + awaitPortsUp.flatMap(_ => CIO.blocking(exec(container, create*))).flatMap(_ => awaitClusterOk) >> + // CLUSTER SLOTS lists a replica only once its replication offset is nonzero + Eventually(100)(clusterTopology(basePort).satisfies(_.shards.forall(_.replicas.size == replicasPerMaster))) + } } diff --git a/integration-tests/shared/src/test/scala/sage/integration/commands/BitmapsSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/commands/BitmapsSuite.scala index 5e21bba5..f4d09ab7 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/commands/BitmapsSuite.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/commands/BitmapsSuite.scala @@ -1,83 +1,44 @@ package sage.integration.commands -import kyo.compat.* - import sage.commands.{BitFieldOffset, BitFieldOp, BitFieldOverflow, BitFieldType, BitRange, BitUnit} -import sage.integration.{Images, ServerSuite} +import sage.integration.BothServersSuite -abstract class BitmapsSuite(image: String) extends ServerSuite(image) { +class BitmapsSuite extends BothServersSuite { - test("SETBIT GETBIT and BITCOUNT track individual bits") { - withClient { client => - for { - prev <- client.setBit("bits", 7L, true) - bit7 <- client.getBit("bits", 7L) - bit6 <- client.getBit("bits", 6L) - count <- client.bitCount("bits") - } yield { - assertEquals(prev, false) - assertEquals(bit7, true) - assertEquals(bit6, false) - assertEquals(count, 1L) - } - } + clientTest("SETBIT GETBIT and BITCOUNT track individual bits") { client => + client.setBit("bits", 7L, true).is(false) >> + client.getBit("bits", 7L).is(true) >> + client.getBit("bits", 6L).is(false) >> + client.bitCount("bits").is(1L) } - test("BITCOUNT and BITPOS honor byte and bit ranges") { - withClient { client => - for { - _ <- client.set("bitstr", "foobar") - whole <- client.bitCount("bitstr") - byteRange <- client.bitCount("bitstr", Some(BitRange(1L, 1L))) - bitRange <- client.bitCount("bitstr", Some(BitRange(5L, 30L, BitUnit.Bit))) - firstSet <- client.bitPos("bitstr", true) - firstZero <- client.bitPos("bitstr", false) - } yield { - assertEquals(whole, 26L) - assertEquals(byteRange, 6L) - assertEquals(bitRange, 17L) - assertEquals(firstSet, 1L) - assertEquals(firstZero, 0L) - } - } + clientTest("BITCOUNT and BITPOS honor byte and bit ranges") { client => + client.set("bitstr", "foobar") >> + client.bitCount("bitstr").is(26L) >> + client.bitCount("bitstr", Some(BitRange(1L, 1L))).is(6L) >> + client.bitCount("bitstr", Some(BitRange(5L, 30L, BitUnit.Bit))).is(17L) >> + client.bitPos("bitstr", true).is(1L) >> + client.bitPos("bitstr", false).is(0L) } - test("BITOP combines bitmaps into a destination") { - withClient { client => - for { - _ <- client.set("opa", "abc") - _ <- client.set("opb", "abd") - andN <- client.bitOpAnd("opand", "opa", "opb") - orN <- client.bitOpOr("opor", "opa", "opb") - xorN <- client.bitOpXor("opxor", "opa", "opb") - notN <- client.bitOpNot("opnot", "opa") - } yield { - assertEquals(andN, 3L) - assertEquals(orN, 3L) - assertEquals(xorN, 3L) - assertEquals(notN, 3L) - } - } + clientTest("BITOP combines bitmaps into a destination") { client => + client.set("opa", "abc") >> + client.set("opb", "abd") >> + client.bitOpAnd("opand", "opa", "opb").is(3L) >> + client.bitOpOr("opor", "opa", "opb").is(3L) >> + client.bitOpXor("opxor", "opa", "opb").is(3L) >> + client.bitOpNot("opnot", "opa").is(3L) } - test("BITFIELD runs typed sub-operations with overflow control") { - withClient { client => - for { - results <- client.bitField( - "bf", - BitFieldOp.Set(BitFieldType.Unsigned(8), BitFieldOffset.Absolute(0L), 255L), - BitFieldOp.Overflow(BitFieldOverflow.Fail), - BitFieldOp.IncrBy(BitFieldType.Unsigned(8), BitFieldOffset.Absolute(0L), 10L) - ) - readBack <- client.bitFieldRo("bf", BitFieldOp.Get(BitFieldType.Unsigned(8), BitFieldOffset.Absolute(0L))) - } yield { - assertEquals(results, Vector(Some(0L), None)) - assertEquals(readBack, Vector(255L)) - } - } + clientTest("BITFIELD runs typed sub-operations with overflow control") { client => + client + .bitField( + "bf", + BitFieldOp.Set(BitFieldType.Unsigned(8), BitFieldOffset.Absolute(0L), 255L), + BitFieldOp.Overflow(BitFieldOverflow.Fail), + BitFieldOp.IncrBy(BitFieldType.Unsigned(8), BitFieldOffset.Absolute(0L), 10L) + ) + .is(Vector(Some(0L), None)) >> + client.bitFieldRo("bf", BitFieldOp.Get(BitFieldType.Unsigned(8), BitFieldOffset.Absolute(0L))).is(Vector(255L)) } } - -class RedisBitmapsSuite extends BitmapsSuite(Images.redis) - -class ValkeyBitmapsSuite extends BitmapsSuite(Images.valkey) diff --git a/integration-tests/shared/src/test/scala/sage/integration/commands/CachingSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/commands/CachingSuite.scala index b6140f8e..a01fa63b 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/commands/CachingSuite.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/commands/CachingSuite.scala @@ -1,60 +1,16 @@ package sage.integration.commands -import scala.concurrent.Future import scala.concurrent.duration.* -import kyo.compat.* - import sage.SageException.NotCacheable -import sage.client.internal.Client import sage.commands.Commands -import sage.integration.{Eventually, Images, ServerSuite} - -abstract class CachingSuite(image: String) extends ServerSuite(image) { - - // a reader (whose cache we observe) plus a writer on a separate connection, so a write is a genuine server-side change - private def withReaderAndWriter[A](body: (Client[CIO, String], Client[CIO, String]) => CIO[A]): Future[A] = - withContainers { server => - val config = configOf(server) - connectAndUse(config)(reader => connectAndUse(config)(writer => body(reader, writer))).unsafeRun - } +import sage.integration.BothServersSuite - // a cached read does not contact the server again. Poll until the invalidation message from an external write has been processed. - private def awaitCached(client: Client[CIO, String], key: String, expected: String, attempts: Int): CIO[Option[String]] = - Eventually.value(attempts)(() => client.cached(Commands.get[String, String](key), 1.minute))(_.contains(expected)) +class CachingSuite extends BothServersSuite { - test("a repeated cached read is served locally and a server-side write evicts it via invalidation") { - withReaderAndWriter { (reader, writer) => - for { - _ <- writer.set("csc:key", "v1") - first <- reader.cached(Commands.get[String, String]("csc:key"), 1.minute) // fetch + cache - cached <- reader.cached(Commands.get[String, String]("csc:key"), 1.minute) // local hit - _ <- writer.set("csc:key", "v2") // server-side write -> invalidation push - evicted <- awaitCached(reader, "csc:key", "v2", attempts = 50) - } yield { - assertEquals(first, Some("v1")) - assertEquals(cached, Some("v1")) - assertEquals(evicted, Some("v2")) - } - } - } + clientsTest("a repeated cached read is served locally and a server-side write evicts it via invalidation")(cachedReadIsInvalidated(_, _, "csc:key")) - test("cached rejects a non-read-only command with NotCacheable") { - withClient { client => - client - .cached(Commands.set[String, String]("csc:write", "v"), 1.minute) - .fold( - _ => CIO.value(false), - { - case _: NotCacheable => CIO.value(true) - case _ => CIO.value(false) - } - ) - .map(rejected => assert(rejected, "expected cached on SET to fail with NotCacheable")) - } + clientTest("cached rejects a non-read-only command with NotCacheable") { client => + failsWith[NotCacheable](client.cached(Commands.set[String, String]("csc:write", "v"), 1.minute)) } } - -class RedisCachingSuite extends CachingSuite(Images.redis) - -class ValkeyCachingSuite extends CachingSuite(Images.valkey) diff --git a/integration-tests/shared/src/test/scala/sage/integration/commands/Coverage.scala b/integration-tests/shared/src/test/scala/sage/integration/commands/Coverage.scala index 2670c040..834c6268 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/commands/Coverage.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/commands/Coverage.scala @@ -4,69 +4,87 @@ package sage.integration.commands * The expected command coverage. Every command reported by the server is either implemented with a [[CommandSamples]] sample or listed in * `skipped` with a reason. The coverage test fails when the server reports an unexpected command. Subcommands represented as arguments to a * base command are omitted; the test ignores server command names containing spaces unless Sage implements that full name directly. + * + * JSON commands come from a loadable module and keep their subcommands, so each `JSON.DEBUG` subcommand must be implemented or skipped. The + * two module-bearing servers report different JSON sets: RedisJSON adds `MERGE` and `NUMPOWBY` ([[redisOnly]]), while valkey-json enumerates + * every `JSON.DEBUG` subcommand where RedisJSON reports only the bare container. */ object Coverage { + val redisOnly: Set[String] = Set("JSON.MERGE", "JSON.NUMPOWBY") + val skipped: Map[String, String] = Map( - "GETSET" -> "deprecated: SET with setGet subsumes it", - "SETNX" -> "deprecated: SET with SetCondition.IfNotExists", - "SETEX" -> "deprecated: SET with SetExpiry.In", - "PSETEX" -> "deprecated: SET with SetExpiry.In", - "SUBSTR" -> "deprecated: GETRANGE", - "HMSET" -> "deprecated: HSET takes multiple field/value pairs", - "RPOPLPUSH" -> "deprecated: LMOVE", - "ZRANGEBYSCORE" -> "deprecated: ZRANGE with ByScore", - "ZRANGEBYLEX" -> "deprecated: ZRANGE with ByLex", - "ZREVRANGE" -> "deprecated: ZRANGE with ByRank and rev", - "ZREVRANGEBYSCORE" -> "deprecated: ZRANGE with ByScore and rev", - "ZREVRANGEBYLEX" -> "deprecated: ZRANGE with ByLex and rev", - "BRPOPLPUSH" -> "deprecated: BLMOVE", - "GEORADIUS" -> "deprecated: GEOSEARCH/GEOSEARCHSTORE", - "GEORADIUS_RO" -> "deprecated: GEOSEARCH", - "GEORADIUSBYMEMBER" -> "deprecated: GEOSEARCH/GEOSEARCHSTORE", - "GEORADIUSBYMEMBER_RO" -> "deprecated: GEOSEARCH", - "SUBSCRIBE" -> "delivered via the subscribe stream API, not a runnable command", - "PSUBSCRIBE" -> "delivered via the pSubscribe stream API, not a runnable command", - "SSUBSCRIBE" -> "delivered via the sSubscribe stream API, not a runnable command", - "UNSUBSCRIBE" -> "sent by the subscribe stream's scope closure, not a runnable command", - "PUNSUBSCRIBE" -> "sent by the pSubscribe stream's scope closure, not a runnable command", - "SUNSUBSCRIBE" -> "sent by the sSubscribe stream's scope closure, not a runnable command", - "MULTI" -> "driven by the transaction scope API, not a runnable command", - "EXEC" -> "driven by the transaction scope API, not a runnable command", - "WATCH" -> "driven by the transaction scope API, not a runnable command", - "UNWATCH" -> "driven by the transaction scope API, not a runnable command", - "DISCARD" -> "the transaction scope abandons by dropping its connection, not via DISCARD", - "ASKING" -> "issued by cluster ASK-redirect handling, not a runnable command", - "RESTORE-ASKING" -> "issued by cluster ASK-redirect handling during slot migration, not a runnable command", - "XGROUP" -> "subcommand container; each XGROUP subcommand is its own Command name", - "XINFO" -> "subcommand container; each XINFO subcommand is its own Command name", - "XIDMPRECORD" -> "internal: replayed during AOF loading, not for client use", - "PFDEBUG" -> "internal HyperLogLog debugging command, out of scope", - "PFSELFTEST" -> "internal HyperLogLog self-test command, out of scope", - "PSYNC" -> "replication protocol between servers, never issued by a client", - "SYNC" -> "replication protocol between servers, never issued by a client", - "REPLCONF" -> "replication protocol between servers, never issued by a client", - "REPLICAOF" -> "server replication administration, out of scope", - "SLAVEOF" -> "deprecated replication administration (REPLICAOF), out of scope", - "FAILOVER" -> "server failover administration, out of scope", - "SHUTDOWN" -> "server lifecycle administration, out of scope", - "MONITOR" -> "debugging firehose that hijacks the connection, out of scope", - "DEBUG" -> "server debugging command, out of scope", - "LOLWUT" -> "easter-egg command, no client API", - "SELECT" -> "issued at connection setup from the configured database, not a runnable command (a per-call SELECT on the shared connection would affect all fibers)", - "SWAPDB" -> "whole-database administration, out of scope", - "QUIT" -> "deprecated: close the client instead of sending QUIT", - "RESET" -> "would reset state on the shared Multiplexed Connection across all fibers; not exposed", - "AUTH" -> "established at handshake via HELLO AUTH; re-auth on the shared Multiplexed Connection would affect all fibers", - "READONLY" -> "issued internally by the cluster read-routing runtime when reading from replicas, never as a user command", - "READWRITE" -> "issued internally by the cluster read-routing runtime, never as a user command", - "MODULE" -> "loadable modules are out of scope", - "CLUSTERSCAN" -> "cluster operator administration; sage consumes topology internally and is not a cluster-admin tool", - "TRIMSLOTS" -> "cluster operator administration; sage consumes topology internally and is not a cluster-admin tool", - "SAVE" -> "persistence administration; synchronous save blocks the server, operator/cron tooling not an app-client concern", - "BGSAVE" -> "persistence administration; backup is operator/cron tooling, not an app-client concern", - "BGREWRITEAOF" -> "persistence administration; AOF rewrite is operator/cron tooling, not an app-client concern", - "LASTSAVE" -> "persistence introspection paired with SAVE/BGSAVE, which are operator tooling", - "HOTKEYS" -> "hot-key profiling session (START/STOP/GET/RESET); operator diagnostic tooling like MONITOR, not exposed" + "GETSET" -> "deprecated: SET with setGet subsumes it", + "SETNX" -> "deprecated: SET with SetCondition.IfNotExists", + "SETEX" -> "deprecated: SET with SetExpiry.In", + "PSETEX" -> "deprecated: SET with SetExpiry.In", + "SUBSTR" -> "deprecated: GETRANGE", + "HMSET" -> "deprecated: HSET takes multiple field/value pairs", + "RPOPLPUSH" -> "deprecated: LMOVE", + "ZRANGEBYSCORE" -> "deprecated: ZRANGE with ByScore", + "ZRANGEBYLEX" -> "deprecated: ZRANGE with ByLex", + "ZREVRANGE" -> "deprecated: ZRANGE with ByRank and rev", + "ZREVRANGEBYSCORE" -> "deprecated: ZRANGE with ByScore and rev", + "ZREVRANGEBYLEX" -> "deprecated: ZRANGE with ByLex and rev", + "BRPOPLPUSH" -> "deprecated: BLMOVE", + "GEORADIUS" -> "deprecated: GEOSEARCH/GEOSEARCHSTORE", + "GEORADIUS_RO" -> "deprecated: GEOSEARCH", + "GEORADIUSBYMEMBER" -> "deprecated: GEOSEARCH/GEOSEARCHSTORE", + "GEORADIUSBYMEMBER_RO" -> "deprecated: GEOSEARCH", + "SUBSCRIBE" -> "delivered via the subscribe stream API, not a runnable command", + "PSUBSCRIBE" -> "delivered via the pSubscribe stream API, not a runnable command", + "SSUBSCRIBE" -> "delivered via the sSubscribe stream API, not a runnable command", + "UNSUBSCRIBE" -> "sent by the subscribe stream's scope closure, not a runnable command", + "PUNSUBSCRIBE" -> "sent by the pSubscribe stream's scope closure, not a runnable command", + "SUNSUBSCRIBE" -> "sent by the sSubscribe stream's scope closure, not a runnable command", + "MULTI" -> "driven by the transaction scope API, not a runnable command", + "EXEC" -> "driven by the transaction scope API, not a runnable command", + "WATCH" -> "driven by the transaction scope API, not a runnable command", + "UNWATCH" -> "driven by the transaction scope API, not a runnable command", + "DISCARD" -> "the transaction scope abandons by dropping its connection, not via DISCARD", + "ASKING" -> "issued by cluster ASK-redirect handling, not a runnable command", + "RESTORE-ASKING" -> "issued by cluster ASK-redirect handling during slot migration, not a runnable command", + "XGROUP" -> "subcommand container; each XGROUP subcommand is its own Command name", + "XINFO" -> "subcommand container; each XINFO subcommand is its own Command name", + "XIDMPRECORD" -> "internal: replayed during AOF loading, not for client use", + "PFDEBUG" -> "internal HyperLogLog debugging command, out of scope", + "PFSELFTEST" -> "internal HyperLogLog self-test command, out of scope", + "PSYNC" -> "replication protocol between servers, never issued by a client", + "SYNC" -> "replication protocol between servers, never issued by a client", + "REPLCONF" -> "replication protocol between servers, never issued by a client", + "REPLICAOF" -> "server replication administration, out of scope", + "SLAVEOF" -> "deprecated replication administration (REPLICAOF), out of scope", + "FAILOVER" -> "server failover administration, out of scope", + "SHUTDOWN" -> "server lifecycle administration, out of scope", + "MONITOR" -> "debugging firehose that hijacks the connection, out of scope", + "DEBUG" -> "server debugging command, out of scope", + "LOLWUT" -> "easter-egg command, no client API", + "SELECT" -> "issued at connection setup from the configured database, not a runnable command (a per-call SELECT on the shared connection would affect all fibers)", + "SWAPDB" -> "whole-database administration, out of scope", + "QUIT" -> "deprecated: close the client instead of sending QUIT", + "RESET" -> "would reset state on the shared Multiplexed Connection across all fibers; not exposed", + "AUTH" -> "established at handshake via HELLO AUTH; re-auth on the shared Multiplexed Connection would affect all fibers", + "READONLY" -> "issued internally by the cluster read-routing runtime when reading from replicas, never as a user command", + "READWRITE" -> "issued internally by the cluster read-routing runtime, never as a user command", + "MODULE" -> "loadable modules are out of scope", + "CLUSTERSCAN" -> "cluster operator administration; sage consumes topology internally and is not a cluster-admin tool", + "TRIMSLOTS" -> "cluster operator administration; sage consumes topology internally and is not a cluster-admin tool", + "SAVE" -> "persistence administration; synchronous save blocks the server, operator/cron tooling not an app-client concern", + "BGSAVE" -> "persistence administration; backup is operator/cron tooling, not an app-client concern", + "BGREWRITEAOF" -> "persistence administration; AOF rewrite is operator/cron tooling, not an app-client concern", + "LASTSAVE" -> "persistence introspection paired with SAVE/BGSAVE, which are operator tooling", + "HOTKEYS" -> "hot-key profiling session (START/STOP/GET/RESET); operator diagnostic tooling like MONITOR, not exposed", + "JSON.FORGET" -> "alias of JSON.DEL", + "JSON.NUMPOWBY" -> "niche exponentiation, Redis-only with no Valkey equivalent", + "JSON.DEBUG" -> "subcommand container; JSON.DEBUG MEMORY is modeled as its own Command name", + "JSON.DEBUG DEPTH" -> "diagnostic subcommand, out of scope", + "JSON.DEBUG FIELDS" -> "diagnostic subcommand, out of scope", + "JSON.DEBUG HELP" -> "help text, not a runnable operation", + "JSON.DEBUG KEYTABLE-CHECK" -> "internal keytable diagnostic, out of scope", + "JSON.DEBUG KEYTABLE-CORRUPT" -> "internal keytable diagnostic, out of scope", + "JSON.DEBUG KEYTABLE-DISTRIBUTION" -> "internal keytable diagnostic, out of scope", + "JSON.DEBUG MAX-DEPTH-KEY" -> "internal diagnostic, out of scope", + "JSON.DEBUG MAX-SIZE-KEY" -> "internal diagnostic, out of scope", + "JSON.DEBUG TEST-SHARED-API" -> "internal test hook, out of scope" ) } diff --git a/integration-tests/shared/src/test/scala/sage/integration/commands/CoverageSpec.scala b/integration-tests/shared/src/test/scala/sage/integration/commands/CoverageSpec.scala index 22494295..1927bfa2 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/commands/CoverageSpec.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/commands/CoverageSpec.scala @@ -1,67 +1,106 @@ package sage.integration.commands -import scala.concurrent.ExecutionContext - import com.dimafeng.testcontainers.GenericContainer import com.dimafeng.testcontainers.lifecycle.and import com.dimafeng.testcontainers.munit.TestContainersForAll import kyo.compat.* +import sage.Bytes +import sage.SageException.DecodeError import sage.client.SageConfig -import sage.commands.CommandSamples -import sage.integration.Images +import sage.client.internal.Client +import sage.commands.{Command, CommandFilterBy, CommandSamples} +import sage.integration.{ContainerClient, Images} +import sage.protocol.Frame /** - * Diffs the implemented commands (the core sample fixtures) against the command list each live server reports, after subtracting module - * commands. The partition must be exact: every drift fails with the offending names. The JSON extension family is a loadable module the - * subtraction removes here; it is validated separately by [[JsonCoverageSpec]]. + * Diffs the implemented commands (the sample fixtures) against the command list each live server reports. The partition must be exact: every + * drift fails with the offending names. + * + * Module commands are subtracted except JSON, which the module-bearing images report (Redis bundles RedisJSON, and Valkey Bundle ships + * valkey-json). JSON names keep their subcommands, so a new `JSON.DEBUG` subcommand fails until acknowledged. */ -class CoverageSpec extends munit.FunSuite with TestContainersForAll with CoverageSupport { +class CoverageSpec extends ContainerClient with TestContainersForAll { - override type Containers = GenericContainer and GenericContainer + override type Containers = GenericContainer and GenericContainer and GenericContainer override def startContainers(): Containers = - GenericContainer.Def(Images.redis, exposedPorts = Seq(6379)).start() and - GenericContainer.Def(Images.valkey, exposedPorts = Seq(6379)).start() + serverDef(Images.redis).start() and serverDef(Images.valkey).start() and serverDef(Images.valkeyBundle).start() - given ExecutionContext = munitExecutionContext + private val implemented = CommandSamples.all.map(_.command.name).toSet - private val implemented: Set[String] = CommandSamples.all.map(_.command.name).toSet.filterNot(_.startsWith("JSON.")) + private def json(names: Set[String]): Set[String] = names.filter(_.startsWith("JSON.")) test("implemented commands never overlap the acknowledged gaps") { assertEquals(implemented.intersect(Coverage.skipped.keySet), Set.empty[String]) } - test("the partition is exact against both live servers") { - withContainers { case redis and valkey => - (for { - redisCore <- coreCommands(configOf(redis)) - valkeyCore <- coreCommands(configOf(valkey)) + containerTest("the partition is exact against every live server and the JSON backend differences are acknowledged") { + case redis and valkey and bundle => + for { + redisListing <- listing(configOf(redis)) + valkeyListing <- listing(configOf(valkey)) + bundleListing <- listing(configOf(bundle)) } yield { // Count a subcommand modeled as an argument under its base command. Track a space-separated command name only when Sage implements // that complete name, as it does for XINFO and XGROUP. - val serverUnion = (redisCore ++ valkeyCore).filterNot(name => name.contains(' ') && !implemented.contains(name)) - assertExactPartition("core", serverUnion, implemented, Coverage.skipped.keySet) - report("redis", redisCore) - report("valkey", valkeyCore) - }).unsafeRun - } + val serverUnion = + (redisListing ++ valkeyListing ++ bundleListing).filterNot(name => name.contains(' ') && !name.startsWith("JSON.") && !implemented(name)) + assertExactPartition(serverUnion, implemented, Coverage.skipped.keySet) + report("redis", redisListing) + report("valkey", valkeyListing) + + val (redisJson, valkeyJson) = (json(redisListing), json(bundleListing)) + assertEquals(redisJson -- valkeyJson, Coverage.redisOnly, "Redis-only JSON commands drifted from the acknowledged set") + assertEquals(valkeyJson -- redisJson, json(serverUnion).filter(_.startsWith("JSON.DEBUG ")), "only valkey-json lists JSON.DEBUG subcommands") + } } - private def report(server: String, core: Set[String]): Unit = + private def report(server: String, names: Set[String]): Unit = println( - s"[coverage] $server: ${core.size} commands, ${core.intersect(implemented).size} implemented, " + - s"${core.intersect(Coverage.skipped.keySet).size} skipped" + s"[coverage] $server: ${names.size} commands, ${names.intersect(implemented).size} implemented, " + + s"${names.intersect(Coverage.skipped.keySet).size} skipped" ) - private def coreCommands(config: SageConfig): CIO[Set[String]] = + private def listing(config: SageConfig): CIO[Set[String]] = connectAndUse(config) { client => for { - all <- client.run(commandList) + all <- commandList(client) modules <- client.run(moduleNames) - module <- modules.foldLeft(CIO.value(Set.empty[String])) { (acc, name) => - acc.flatMap(commands => client.run(commandListForModule(name)).map(commands ++ _)) - } - } yield all.toSet -- module + module <- CIO.foreach(modules)(name => commandList(client, Some(CommandFilterBy.Module(name)))).map(_.iterator.flatten.toSet) + } yield all -- module ++ json(all) } + + private def commandList(client: Client[CIO, String], filterBy: Option[CommandFilterBy] = None): CIO[Set[String]] = + client.commandList(filterBy).map(_.iterator.map(normalize).toSet) + + private val moduleNames: Command[Vector[String]] = + Command( + "MODULE", + keyIndices = Command.NoKeys, + args = Vector(Bytes.utf8("LIST")), + decode = { + case Frame.Array(modules) => + Right(modules.collect { case Frame.Map(entries) => + entries.collectFirst { + case (Frame.BulkString(key), Frame.BulkString(value)) if key.asUtf8String == "name" => value.asUtf8String + } + }.flatten) + case other => Left(DecodeError("array of module maps", Frame.describe(other))) + } + ) + + // the server reports lowercase names with pipe-separated subcommands; Command names are uppercase words + private def normalize(name: String): String = name.toUpperCase.replace('|', ' ') + + private def assertExactPartition(serverUnion: Set[String], implemented: Set[String], skipped: Set[String]): Unit = { + val unacknowledged = serverUnion -- implemented -- skipped + assert(unacknowledged.isEmpty, s"unacknowledged server commands: ${unacknowledged.toVector.sorted.mkString(", ")}") + + val unknownImplemented = implemented -- serverUnion + assert(unknownImplemented.isEmpty, s"implemented commands unknown to both servers: ${unknownImplemented.toVector.sorted.mkString(", ")}") + + val stale = skipped -- serverUnion + assert(stale.isEmpty, s"skipped entries unknown to both servers: ${stale.toVector.sorted.mkString(", ")}") + } } diff --git a/integration-tests/shared/src/test/scala/sage/integration/commands/CoverageSupport.scala b/integration-tests/shared/src/test/scala/sage/integration/commands/CoverageSupport.scala deleted file mode 100644 index cae9a2d6..00000000 --- a/integration-tests/shared/src/test/scala/sage/integration/commands/CoverageSupport.scala +++ /dev/null @@ -1,68 +0,0 @@ -package sage.integration.commands - -import sage.Bytes -import sage.SageException.DecodeError -import sage.commands.Command -import sage.integration.ContainerClient -import sage.protocol.Frame - -/** - * Shared machinery for the command-coverage specs: the `COMMAND LIST` decoders, name normalization, and the exact-partition assertion. The - * connect/teardown and config helpers come from [[ContainerClient]]. The core and JSON specs supply only their images, filters, and skip maps. - */ -trait CoverageSupport extends ContainerClient { this: munit.FunSuite => - - protected val commandList: Command[Vector[String]] = rawCommandList(Vector(Bytes.utf8("LIST"))) - - protected def commandListForModule(module: String): Command[Vector[String]] = - rawCommandList(Vector("LIST", "FILTERBY", "MODULE", module).map(Bytes.utf8)) - - protected val moduleNames: Command[Vector[String]] = - Command( - "MODULE", - keyIndices = Command.NoKeys, - args = Vector(Bytes.utf8("LIST")), - decode = { - case Frame.Array(modules) => - Right(modules.collect { case Frame.Map(entries) => - entries.collectFirst { - case (Frame.BulkString(key), Frame.BulkString(value)) if key.asUtf8String == "name" => value.asUtf8String - } - }.flatten) - case other => Left(DecodeError("array of module maps", Frame.describe(other))) - } - ) - - // the server reports lowercase names with pipe-separated subcommands; Command names are uppercase words - protected def normalize(name: String): String = name.toUpperCase.replace('|', ' ') - - protected def assertExactPartition(label: String, serverUnion: Set[String], implemented: Set[String], skipped: Set[String]): Unit = { - val unacknowledged = serverUnion -- implemented -- skipped - assert(unacknowledged.isEmpty, s"$label unacknowledged server commands: ${unacknowledged.toVector.sorted.mkString(", ")}") - - val unknownImplemented = implemented -- serverUnion - assert(unknownImplemented.isEmpty, s"$label implemented commands unknown to both servers: ${unknownImplemented.toVector.sorted.mkString(", ")}") - - val stale = skipped -- serverUnion - assert(stale.isEmpty, s"$label skipped entries unknown to both servers: ${stale.toVector.sorted.mkString(", ")}") - } - - private def rawCommandList(args: Vector[Bytes]): Command[Vector[String]] = - Command( - "COMMAND", - keyIndices = Command.NoKeys, - args = args, - decode = { - case Frame.Array(elements) => - elements.foldLeft[Either[DecodeError, Vector[String]]](Right(Vector.empty)) { (acc, frame) => - acc.flatMap { names => - frame match { - case Frame.BulkString(name) => Right(names :+ normalize(name.asUtf8String)) - case other => Left(DecodeError("bulk string command name", Frame.describe(other))) - } - } - } - case other => Left(DecodeError("array of command names", Frame.describe(other))) - } - ) -} diff --git a/integration-tests/shared/src/test/scala/sage/integration/commands/FunctionsSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/commands/FunctionsSuite.scala index a43d50ac..1a97728e 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/commands/FunctionsSuite.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/commands/FunctionsSuite.scala @@ -3,10 +3,10 @@ package sage.integration.commands import kyo.compat.* import sage.commands.FlushMode -import sage.integration.{Images, ServerSuite} -import sage.protocol.Frame +import sage.integration.BothServersSuite +import sage.protocol.{Frame, Frames} -abstract class FunctionsSuite(image: String) extends ServerSuite(image) { +class FunctionsSuite extends BothServersSuite { private val library = """#!lua name=saregg @@ -14,80 +14,45 @@ abstract class FunctionsSuite(image: String) extends ServerSuite(image) { |redis.register_function('saregg_one', function(keys, args) return 1 end) |""".stripMargin - test("FUNCTION LOAD registers a library that FCALL then invokes") { - withClient { client => - for { - _ <- client.functionFlush(Some(FlushMode.Sync)) - name <- client.functionLoad(library) - one <- client.fCall("saregg_one") - echo <- client.fCall("saregg_echo", Seq.empty[String], Seq("hello")) - } yield { - assertEquals(name, "saregg") - assertEquals(one, Frame.Integer(1L)) - echo match { - case Frame.BulkString(b) => assertEquals(b.asUtf8String, "hello") - case other => fail(s"expected bulk string, got $other") - } - } - } + clientTest("FUNCTION LOAD registers a library that FCALL then invokes") { client => + client.functionFlush(Some(FlushMode.Sync)) >> + client.functionLoad(library).is("saregg") >> + client.fCall("saregg_one").is(Frame.Integer(1L)) >> + client.fCall("saregg_echo", Seq.empty[String], Seq("hello")).is(Frames.bulk("hello")) } - test("FCALL_RO invokes a function flagged no-writes") { - withClient { client => - for { - _ <- client.functionFlush(Some(FlushMode.Sync)) - _ <- client.functionLoad( - """#!lua name=saregg - |redis.register_function{function_name='saregg_ro', callback=function(keys, args) return args[1] end, flags={'no-writes'}} - |""".stripMargin - ) - ro <- client.fCallRo("saregg_ro", Seq.empty[String], Seq("hi")) - } yield ro match { - case Frame.BulkString(b) => assertEquals(b.asUtf8String, "hi") - case other => fail(s"expected bulk string, got $other") - } - } + clientTest("FCALL_RO invokes a function flagged no-writes") { client => + client.functionFlush(Some(FlushMode.Sync)) >> + client.functionLoad( + """#!lua name=saregg + |redis.register_function{function_name='saregg_ro', callback=function(keys, args) return args[1] end, flags={'no-writes'}} + |""".stripMargin + ) >> + client.fCallRo("saregg_ro", Seq.empty[String], Seq("hi")).is(Frames.bulk("hi")) } - test("FUNCTION LIST and STATS describe loaded libraries; DELETE removes them") { - withClient { client => - for { - _ <- client.functionFlush(Some(FlushMode.Sync)) - _ <- client.functionLoad(library) - libraries <- client.functionList() - withCode <- client.functionList(Some("saregg"), withCode = true) - stats <- client.functionStats - _ <- client.functionDelete("saregg") - afterDel <- client.functionList() - } yield { - val lib = libraries.find(_.libraryName == "saregg") - assert(lib.isDefined, libraries.toString) - assertEquals(lib.map(_.engine), Some("LUA")) - assertEquals(lib.map(_.functions.map(_.name).toSet), Some(Set("saregg_echo", "saregg_one"))) - assertEquals(withCode.flatMap(_.code).headOption.isDefined, true) - assert(stats.engines.contains("LUA"), stats.toString) - assert(afterDel.forall(_.libraryName != "saregg")) - } - } + clientTest("FUNCTION LIST and STATS describe loaded libraries; DELETE removes them") { client => + client.functionFlush(Some(FlushMode.Sync)) >> + client.functionLoad(library) >> + client + .functionList() + .map(_.map(l => (l.libraryName, l.engine, l.functions.map(_.name).toSet))) + .is(Vector(("saregg", "LUA", Set("saregg_echo", "saregg_one")))) >> + client.functionList(Some("saregg"), withCode = true).map(_.map(_.code)).is(Vector(Some(library))) >> + client.functionStats.satisfies(_.engines.contains("LUA")) >> + client.functionDelete("saregg") >> + client.functionList().is(Vector.empty) } - test("FUNCTION DUMP and RESTORE round-trip the library payload") { - withClient { client => - for { - _ <- client.functionFlush(Some(FlushMode.Sync)) - _ <- client.functionLoad(library) - payload <- client.functionDump - _ <- client.functionFlush(Some(FlushMode.Sync)) - empty <- client.functionList() - _ <- client.functionRestore(payload) - restored <- client.functionList() - } yield { - assert(empty.isEmpty) - assert(restored.exists(_.libraryName == "saregg")) - } - } + clientTest("FUNCTION DUMP and RESTORE round-trip the library payload") { client => + for { + _ <- client.functionFlush(Some(FlushMode.Sync)) + _ <- client.functionLoad(library) + payload <- client.functionDump + _ <- client.functionFlush(Some(FlushMode.Sync)) + _ <- client.functionList().is(Vector.empty) + _ <- client.functionRestore(payload) + _ <- client.functionList().map(_.map(_.libraryName)).is(Vector("saregg")) + } yield () } } - -class RedisFunctionsSuite extends FunctionsSuite(Images.redis) -class ValkeyFunctionsSuite extends FunctionsSuite(Images.valkey) diff --git a/integration-tests/shared/src/test/scala/sage/integration/commands/GeoSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/commands/GeoSuite.scala index 796b51ca..5fbbf1dc 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/commands/GeoSuite.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/commands/GeoSuite.scala @@ -3,101 +3,59 @@ package sage.integration.commands import kyo.compat.* import sage.commands.{GeoCoordinates, GeoCount, GeoOrigin, GeoShape, GeoSort, GeoUnit} -import sage.integration.{Images, ServerSuite} +import sage.integration.BothServersSuite -abstract class GeoSuite(image: String) extends ServerSuite(image) { +class GeoSuite extends BothServersSuite { private val palermo = GeoCoordinates(13.361389, 38.115556) private val catania = GeoCoordinates(15.087269, 37.502669) - test("GEOADD GEOPOS GEODIST and GEOHASH store and read positions") { - withClient { client => - for { - added <- client.geoAdd("Sicily")(("Palermo", palermo), ("Catania", catania)) - dup <- client.geoAdd("Sicily", changed = true)(("Palermo", palermo)) - pos <- client.geoPos("Sicily", "Palermo", "NonExisting") - dist <- client.geoDist("Sicily", "Palermo", "Catania", GeoUnit.Kilometers) - hash <- client.geoHash("Sicily", "Palermo", "NonExisting") - } yield { - assertEquals(added, 2L) - assertEquals(dup, 0L) - assertEquals(pos.size, 2) - assert(pos(0).exists(c => Math.abs(c.longitude - palermo.longitude) < 0.001 && Math.abs(c.latitude - palermo.latitude) < 0.001)) - assertEquals(pos(1), None) - assert(dist.exists(d => d > 166.0 && d < 167.0)) - assert(hash(0).exists(_.startsWith("sqc8b49rny"))) - assertEquals(hash(1), None) - } - } + clientTest("GEOADD GEOPOS GEODIST and GEOHASH store and read positions") { client => + client.geoAdd("Sicily")(("Palermo", palermo), ("Catania", catania)).is(2L) >> + client.geoAdd("Sicily", changed = true)(("Palermo", palermo)).is(0L) >> + client + .geoPos("Sicily", "Palermo", "NonExisting") + .map(_.map(_.map(c => (c.longitude - palermo.longitude).abs < 0.001 && (c.latitude - palermo.latitude).abs < 0.001))) + .is(Vector(Some(true), None)) >> + client.geoDist("Sicily", "Palermo", "Catania", GeoUnit.Kilometers).satisfies(_.exists(d => d > 166.0 && d < 167.0)) >> + client.geoHash("Sicily", "Palermo", "NonExisting").map(_.map(_.map(_.startsWith("sqc8b49rny")))).is(Vector(Some(true), None)) } - test("GEOSEARCH returns members within an area, ordered and limited") { - withClient { client => - for { - _ <- client.geoAdd("Sicily2")(("Palermo", palermo), ("Catania", catania)) - members <- client.geoSearch[String]( - "Sicily2", - GeoOrigin.FromLonLat(GeoCoordinates(15.0, 37.0)), - GeoShape.ByRadius(200.0, GeoUnit.Kilometers), - sort = Some(GeoSort.Asc) - ) - nearest <- client.geoSearch( - "Sicily2", - GeoOrigin.FromMember("Palermo"), - GeoShape.ByBox(400.0, 400.0, GeoUnit.Kilometers), - count = Some(GeoCount(1)) - ) - } yield { - assertEquals(members, Vector("Catania", "Palermo")) - assertEquals(nearest, Vector("Palermo")) - } - } + clientTest("GEOSEARCH returns members within an area, ordered and limited") { client => + client.geoAdd("Sicily2")(("Palermo", palermo), ("Catania", catania)) >> + client + .geoSearch[String]( + "Sicily2", + GeoOrigin.FromLonLat(GeoCoordinates(15.0, 37.0)), + GeoShape.ByRadius(200.0, GeoUnit.Kilometers), + sort = Some(GeoSort.Asc) + ) + .is(Vector("Catania", "Palermo")) >> + client + .geoSearch("Sicily2", GeoOrigin.FromMember("Palermo"), GeoShape.ByBox(400.0, 400.0, GeoUnit.Kilometers), count = Some(GeoCount(1))) + .is(Vector("Palermo")) } - test("GEOSEARCH with projections returns coordinates, distance and hash") { - withClient { client => - for { - _ <- client.geoAdd("Sicily3")(("Palermo", palermo), ("Catania", catania)) - hits <- client.geoSearchWith( - "Sicily3", - GeoOrigin.FromMember("Palermo"), - GeoShape.ByRadius(200.0, GeoUnit.Kilometers), - withCoord = true, - withDist = true, - withHash = true, - sort = Some(GeoSort.Asc) - ) - } yield { - assertEquals(hits.map(_.member), Vector("Palermo", "Catania")) - assert(hits.forall(h => h.distance.isDefined && h.hash.isDefined && h.coordinates.isDefined)) - assert(hits.head.distance.exists(_ < 1.0)) - } - } + clientTest("GEOSEARCH with projections returns coordinates, distance and hash") { client => + client.geoAdd("Sicily3")(("Palermo", palermo), ("Catania", catania)) >> + client + .geoSearchWith( + "Sicily3", + GeoOrigin.FromMember("Palermo"), + GeoShape.ByRadius(200.0, GeoUnit.Kilometers), + withCoord = true, + withDist = true, + withHash = true, + sort = Some(GeoSort.Asc) + ) + .map(_.map(h => (h.member, h.distance.map(_.round), h.hash.isDefined, h.coordinates.isDefined))) + .is(Vector(("Palermo", Some(0L), true, true), ("Catania", Some(166L), true, true))) } - test("GEOSEARCHSTORE writes matches into a destination key") { - withClient { client => - for { - _ <- client.geoAdd("Sicily4")(("Palermo", palermo), ("Catania", catania)) - stored <- client.geoSearchStore[String]( - "Sicily4Store", - "Sicily4", - GeoOrigin.FromLonLat(GeoCoordinates(15.0, 37.0)), - GeoShape.ByRadius(200.0, GeoUnit.Kilometers) - ) - members <- client.geoSearch[String]( - "Sicily4Store", - GeoOrigin.FromLonLat(GeoCoordinates(15.0, 37.0)), - GeoShape.ByRadius(200.0, GeoUnit.Kilometers) - ) - } yield { - assertEquals(stored, 2L) - assertEquals(members.toSet, Set("Palermo", "Catania")) - } - } + clientTest("GEOSEARCHSTORE writes matches into a destination key") { client => + val area = GeoShape.ByRadius(200.0, GeoUnit.Kilometers) + client.geoAdd("Sicily4")(("Palermo", palermo), ("Catania", catania)) >> + client.geoSearchStore[String]("Sicily4Store", "Sicily4", GeoOrigin.FromLonLat(GeoCoordinates(15.0, 37.0)), area).is(2L) >> + client.geoSearch[String]("Sicily4Store", GeoOrigin.FromLonLat(GeoCoordinates(15.0, 37.0)), area).map(_.toSet).is(Set("Palermo", "Catania")) } } - -class RedisGeoSuite extends GeoSuite(Images.redis) - -class ValkeyGeoSuite extends GeoSuite(Images.valkey) diff --git a/integration-tests/shared/src/test/scala/sage/integration/commands/HashesSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/commands/HashesSuite.scala index 09e9b6e9..d5813971 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/commands/HashesSuite.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/commands/HashesSuite.scala @@ -1,111 +1,126 @@ package sage.integration.commands +import java.time.Instant + +import scala.concurrent.duration.* + import kyo.compat.* -import sage.commands.ScanCursor -import sage.integration.{Images, ServerSuite} - -abstract class HashesSuite(image: String) extends ServerSuite(image) { - - test("HSET writes fields, HGET and HMGET read them back, HEXISTS and HDEL remove them") { - withClient { client => - for { - added <- client.hSet("hash-basic", ("f1", "v1"), ("f2", "v2")) - one <- client.hGet[String, String]("hash-basic", "f1") - many <- client.hmGet[String, String]("hash-basic", "f1", "missing", "f2") - present <- client.hExists("hash-basic", "f1") - removed <- client.hDel("hash-basic", "f1", "missing") - gone <- client.hExists("hash-basic", "f1") - } yield { - assertEquals(added, 2L) - assertEquals(one, Some("v1")) - assertEquals(many, Vector(Some("v1"), None, Some("v2"))) - assertEquals(present, true) - assertEquals(removed, 1L) - assertEquals(gone, false) - } - } +import sage.commands.* +import sage.integration.BothServersSuite +import sage.integration.Ttls.expiresWithin + +class HashesSuite extends BothServersSuite { + + clientTest("HSET writes fields, HGET and HMGET read them back, HEXISTS and HDEL remove them") { client => + client.hSet("hash-basic", ("f1", "v1"), ("f2", "v2")).is(2L) >> + client.hGet[String, String]("hash-basic", "f1").is(Some("v1")) >> + client.hmGet[String, String]("hash-basic", "f1", "missing", "f2").is(Vector(Some("v1"), None, Some("v2"))) >> + client.hExists("hash-basic", "f1").is(true) >> + client.hDel("hash-basic", "f1", "missing").is(1L) >> + client.hExists("hash-basic", "f1").is(false) } - test("HSETNX only writes an absent field") { - withClient { client => - for { - first <- client.hSetNx("hash-setnx", "f", "one") - second <- client.hSetNx("hash-setnx", "f", "two") - value <- client.hGet[String, String]("hash-setnx", "f") - } yield { - assertEquals(first, true) - assertEquals(second, false) - assertEquals(value, Some("one")) - } - } + clientTest("HSETNX only writes an absent field") { client => + client.hSetNx("hash-setnx", "f", "one").is(true) >> + client.hSetNx("hash-setnx", "f", "two").is(false) >> + client.hGet[String, String]("hash-setnx", "f").is(Some("one")) } - test("HGETALL HKEYS HVALS HLEN HSTRLEN view the whole hash") { - withClient { client => - for { - _ <- client.hSet("hash-view", ("a", "1"), ("b", "22")) - all <- client.hGetAll[String, String]("hash-view") - keys <- client.hKeys[String]("hash-view") - vals <- client.hVals[String]("hash-view") - len <- client.hLen("hash-view") - strLen <- client.hStrLen("hash-view", "b") - } yield { - assertEquals(all, Map("a" -> "1", "b" -> "22")) - assertEquals(keys.toSet, Set("a", "b")) - assertEquals(vals.toSet, Set("1", "22")) - assertEquals(len, 2L) - assertEquals(strLen, 2L) - } - } + clientTest("HGETALL HKEYS HVALS HLEN HSTRLEN view the whole hash") { client => + client.hSet("hash-view", ("a", "1"), ("b", "22")) >> + client.hGetAll[String, String]("hash-view").is(Map("a" -> "1", "b" -> "22")) >> + client.hKeys[String]("hash-view").map(_.toSet).is(Set("a", "b")) >> + client.hVals[String]("hash-view").map(_.toSet).is(Set("1", "22")) >> + client.hLen("hash-view").is(2L) >> + client.hStrLen("hash-view", "b").is(2L) + } + + clientTest("HINCRBY and HINCRBYFLOAT count atomically on a field") { client => + client.hSet("hash-incr", ("n", "10")) >> + client.hIncrBy("hash-incr", "n", 5L).is(15L) >> + client.hIncrByFloat("hash-incr", "n", 0.5).is(15.5) } - test("HINCRBY and HINCRBYFLOAT count atomically on a field") { - withClient { client => - for { - _ <- client.hSet("hash-incr", ("n", "10")) - byInt <- client.hIncrBy("hash-incr", "n", 5L) - byFlt <- client.hIncrByFloat("hash-incr", "n", 0.5) - } yield { - assertEquals(byInt, 15L) - assertEquals(byFlt, 15.5) - } + clientTest("HRANDFIELD returns a member, a count of members, and field/value pairs") { client => + for { + _ <- client.hSet("hash-rand", ("a", "1"), ("b", "2"), ("c", "3")) + _ <- client.hRandField[String]("hash-rand").satisfies(_.exists(Set("a", "b", "c"))) + few <- client.hRandField[String]("hash-rand", 2L) + pairs <- client.hRandFieldWithValues[String, String]("hash-rand", -5L) + _ <- client.hRandField[String]("hash-rand-missing").is(None) + } yield { + assertEquals(few.size, 2) + assert(few.toSet.subsetOf(Set("a", "b", "c"))) + assertEquals(pairs.size, 5) + assert(pairs.forall { case (f, v) => Map("a" -> "1", "b" -> "2", "c" -> "3").get(f).contains(v) }) } } - test("HRANDFIELD returns a member, a count of members, and field/value pairs") { - withClient { client => - for { - _ <- client.hSet("hash-rand", ("a", "1"), ("b", "2"), ("c", "3")) - single <- client.hRandField[String]("hash-rand") - few <- client.hRandField[String]("hash-rand", 2L) - pairs <- client.hRandFieldWithValues[String, String]("hash-rand", -5L) - empty <- client.hRandField[String]("hash-rand-missing") - } yield { - assert(single.exists(Set("a", "b", "c"))) - assertEquals(few.size, 2) - assert(few.toSet.subsetOf(Set("a", "b", "c"))) - assertEquals(pairs.size, 5) - assert(pairs.forall { case (f, v) => Map("a" -> "1", "b" -> "2", "c" -> "3").get(f).contains(v) }) - assertEquals(empty, None) - } + clientTest("HSCAN streams field/value pairs and NOVALUES streams bare fields") { client => + client.hSet("hash-scan", ("a", "1"), ("b", "2"), ("c", "3")) >> + client.hScan[String, String]("hash-scan", ScanCursor.start).map(_.items.toMap).is(Map("a" -> "1", "b" -> "2", "c" -> "3")) >> + client.hScanNoValues[String]("hash-scan", ScanCursor.start).map(_.items.toSet).is(Set("a", "b", "c")) + } + + // Hash field expiration exists only on Redis. + redisTest("HEXPIRE/HTTL/HPERSIST set, read, and clear per-field TTLs") { client => + for { + _ <- client.hSet("hfe-ttl", ("a", "1"), ("b", "2")) + _ <- client.hExpire("hfe-ttl", 100.seconds)("a", "missing").is(Vector(FieldExpiry.Updated, FieldExpiry.NoField)) + ttl <- client.hTtl("hfe-ttl")("a", "b", "missing") + _ <- client.hPersist("hfe-ttl")("a", "b").is(Vector(FieldPersist.Persisted, FieldPersist.NoExpiry)) + _ <- client.hTtl("hfe-ttl")("a").is(Vector(FieldTtl.NoExpiry)) + } yield { + assert(expiresWithin(ttl(0), 100.seconds)) + assertEquals(ttl(1), FieldTtl.NoExpiry) + assertEquals(ttl(2), FieldTtl.NoField) } } - test("HSCAN streams field/value pairs and NOVALUES streams bare fields") { - withClient { client => - for { - _ <- client.hSet("hash-scan", ("a", "1"), ("b", "2"), ("c", "3")) - page <- client.hScan[String, String]("hash-scan", ScanCursor.start) - fieldPage <- client.hScanNoValues[String]("hash-scan", ScanCursor.start) - } yield { - assertEquals(page.items.toMap, Map("a" -> "1", "b" -> "2", "c" -> "3")) - assertEquals(fieldPage.items.toSet, Set("a", "b", "c")) - } + redisTest("a field-TTL command on a missing key reports NoField per field, not a null") { client => + client.hExpire("hfe-missing", 100.seconds)("a", "b").is(Vector(FieldExpiry.NoField, FieldExpiry.NoField)) + } + + redisTest("HEXPIREAT pins an absolute deadline that HEXPIRETIME reads back") { client => + val at = Instant.ofEpochSecond(Instant.now().getEpochSecond + 3600) + client.hSet("hfe-at", ("a", "1")) >> + client.hExpireAt("hfe-at", at)("a").is(Vector(FieldExpiry.Updated)) >> + client.hExpireTime("hfe-at")("a").is(Vector(FieldExpiryTime.At(at))) + } + + redisTest("HPTTL and HPEXPIRETIME read the millisecond-precision TTL and absolute deadline") { client => + for { + now <- client.time + at = Instant.ofEpochSecond(now.getEpochSecond + 3600) + _ <- client.hSet("hfe-px", ("a", "1")) + _ <- client.hExpireAt("hfe-px", at)("a") + ttl <- client.hpTtl("hfe-px")("a", "missing") + _ <- client.hpExpireTime("hfe-px")("a").is(Vector(FieldExpiryTime.At(at))) + } yield { + assert(expiresWithin(ttl(0), 3600.seconds)) + assertEquals(ttl(1), FieldTtl.NoField) } } -} -class RedisHashesSuite extends HashesSuite(Images.redis) + redisTest("HGETDEL returns field values and removes them") { client => + client.hSet("hfe-getdel", ("a", "1"), ("b", "2")) >> + client.hGetDel[String, String]("hfe-getdel")("a", "missing").is(Vector(Some("1"), None)) >> + client.hGetAll[String, String]("hfe-getdel").is(Map("b" -> "2")) + } + + redisTest("HGETEX returns field values and sets their TTL") { client => + client.hSet("hfe-getex", ("a", "1")) >> + client.hGetEx[String, String]("hfe-getex", GetExpiry.In(100.seconds))("a").is(Vector(Some("1"))) >> + client.hTtl("hfe-getex")("a").satisfies(ttl => expiresWithin(ttl(0), 100.seconds)) + } -class ValkeyHashesSuite extends HashesSuite(Images.valkey) + redisTest("HSETEX sets fields with a shared TTL and honors FNX/FXX") { client => + client.hSetEx("hfe-setex", SetExpiry.In(100.seconds), HSetExCondition.IfNoneExist)(("a", "1"), ("b", "2")).is(true) >> + client.hTtl("hfe-setex")("a").satisfies(ttl => expiresWithin(ttl(0), 100.seconds)) >> + client.hSetEx("hfe-setex", condition = HSetExCondition.IfNoneExist)(("a", "9")).is(false) >> + client.hGet[String, String]("hfe-setex", "a").is(Some("1")) >> + client.hSetEx("hfe-setex", SetExpiry.KeepTtl, HSetExCondition.IfAllExist)(("a", "10")).is(true) >> + client.hGet[String, String]("hfe-setex", "a").is(Some("10")) + } +} diff --git a/integration-tests/shared/src/test/scala/sage/integration/commands/HyperLogLogSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/commands/HyperLogLogSuite.scala index 62501fff..8d72b4dd 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/commands/HyperLogLogSuite.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/commands/HyperLogLogSuite.scala @@ -1,43 +1,21 @@ package sage.integration.commands -import kyo.compat.* +import sage.integration.BothServersSuite -import sage.integration.{Images, ServerSuite} +class HyperLogLogSuite extends BothServersSuite { -abstract class HyperLogLogSuite(image: String) extends ServerSuite(image) { - - test("PFADD reports change and PFCOUNT estimates cardinality") { - withClient { client => - for { - first <- client.pfAdd("hll", "a", "b", "c") - again <- client.pfAdd("hll", "a") - count <- client.pfCount("hll") - empty <- client.pfAdd[String]("hll-empty") - } yield { - assertEquals(first, true) - assertEquals(again, false) - assertEquals(count, 3L) - assertEquals(empty, true) - } - } + clientTest("PFADD reports change and PFCOUNT estimates cardinality") { client => + client.pfAdd("hll", "a", "b", "c").is(true) >> + client.pfAdd("hll", "a").is(false) >> + client.pfCount("hll").is(3L) >> + client.pfAdd[String]("hll-empty").is(true) } - test("PFMERGE unions HyperLogLogs and PFCOUNT spans multiple keys") { - withClient { client => - for { - _ <- client.pfAdd("hll-a", "a", "b", "c") - _ <- client.pfAdd("hll-b", "c", "d", "e") - _ <- client.pfMerge("hll-merged", "hll-a", "hll-b") - merged <- client.pfCount("hll-merged") - both <- client.pfCount("hll-a", "hll-b") - } yield { - assertEquals(merged, 5L) - assertEquals(both, 5L) - } - } + clientTest("PFMERGE unions HyperLogLogs and PFCOUNT spans multiple keys") { client => + client.pfAdd("hll-a", "a", "b", "c") >> + client.pfAdd("hll-b", "c", "d", "e") >> + client.pfMerge("hll-merged", "hll-a", "hll-b") >> + client.pfCount("hll-merged").is(5L) >> + client.pfCount("hll-a", "hll-b").is(5L) } } - -class RedisHyperLogLogSuite extends HyperLogLogSuite(Images.redis) - -class ValkeyHyperLogLogSuite extends HyperLogLogSuite(Images.valkey) diff --git a/integration-tests/shared/src/test/scala/sage/integration/commands/JsonCoverage.scala b/integration-tests/shared/src/test/scala/sage/integration/commands/JsonCoverage.scala deleted file mode 100644 index 11bf681c..00000000 --- a/integration-tests/shared/src/test/scala/sage/integration/commands/JsonCoverage.scala +++ /dev/null @@ -1,41 +0,0 @@ -package sage.integration.commands - -/** - * The acknowledged coverage partition for the JSON extension module, kept separate from [[Coverage]] because JSON commands come from a - * loadable module the core spec subtracts. Every JSON command a module-bearing server reports, including each `JSON.DEBUG` subcommand, is - * either implemented (a [[sage.commands.CommandSamples]] sample) or listed here with a reason. The two servers report different command - * sets: RedisJSON adds `MERGE` and `NUMPOWBY` ([[redisOnly]]), while valkey-json enumerates every `JSON.DEBUG` subcommand where RedisJSON - * reports only the bare container ([[valkeyOnly]]). - */ -object JsonCoverage { - - val redisOnly: Set[String] = Set("JSON.MERGE", "JSON.NUMPOWBY") - - val valkeyOnly: Set[String] = Set( - "JSON.DEBUG DEPTH", - "JSON.DEBUG FIELDS", - "JSON.DEBUG HELP", - "JSON.DEBUG KEYTABLE-CHECK", - "JSON.DEBUG KEYTABLE-CORRUPT", - "JSON.DEBUG KEYTABLE-DISTRIBUTION", - "JSON.DEBUG MAX-DEPTH-KEY", - "JSON.DEBUG MAX-SIZE-KEY", - "JSON.DEBUG MEMORY", - "JSON.DEBUG TEST-SHARED-API" - ) - - val skipped: Map[String, String] = Map( - "JSON.FORGET" -> "alias of JSON.DEL", - "JSON.NUMPOWBY" -> "niche exponentiation, Redis-only with no Valkey equivalent", - "JSON.DEBUG" -> "subcommand container; JSON.DEBUG MEMORY is modeled as its own Command name", - "JSON.DEBUG DEPTH" -> "diagnostic subcommand, out of scope", - "JSON.DEBUG FIELDS" -> "diagnostic subcommand, out of scope", - "JSON.DEBUG HELP" -> "help text, not a runnable operation", - "JSON.DEBUG KEYTABLE-CHECK" -> "internal keytable diagnostic, out of scope", - "JSON.DEBUG KEYTABLE-CORRUPT" -> "internal keytable diagnostic, out of scope", - "JSON.DEBUG KEYTABLE-DISTRIBUTION" -> "internal keytable diagnostic, out of scope", - "JSON.DEBUG MAX-DEPTH-KEY" -> "internal diagnostic, out of scope", - "JSON.DEBUG MAX-SIZE-KEY" -> "internal diagnostic, out of scope", - "JSON.DEBUG TEST-SHARED-API" -> "internal test hook, out of scope" - ) -} diff --git a/integration-tests/shared/src/test/scala/sage/integration/commands/JsonCoverageSpec.scala b/integration-tests/shared/src/test/scala/sage/integration/commands/JsonCoverageSpec.scala deleted file mode 100644 index 19831d45..00000000 --- a/integration-tests/shared/src/test/scala/sage/integration/commands/JsonCoverageSpec.scala +++ /dev/null @@ -1,52 +0,0 @@ -package sage.integration.commands - -import scala.concurrent.ExecutionContext - -import com.dimafeng.testcontainers.GenericContainer -import com.dimafeng.testcontainers.lifecycle.and -import com.dimafeng.testcontainers.munit.TestContainersForAll -import kyo.compat.* - -import sage.client.SageConfig -import sage.commands.CommandSamples -import sage.integration.Images - -/** - * The extension-aware coverage spec for the JSON module. It runs the module-bearing images (Redis, which bundles RedisJSON, and Valkey Bundle, - * which ships valkey-json), takes the union of every `JSON.*` command each server reports, and requires an exact partition against the - * implemented samples plus [[JsonCoverage.skipped]]. Unlike the core spec it does not drop subcommand names, so each `JSON.DEBUG` subcommand - * must be implemented or explicitly skipped and a new one fails until acknowledged. Redis-only commands are covered through Redis and - * Valkey-only ones through the bundle. - */ -class JsonCoverageSpec extends munit.FunSuite with TestContainersForAll with CoverageSupport { - - override type Containers = GenericContainer and GenericContainer - - override def startContainers(): Containers = - GenericContainer.Def(Images.redis, exposedPorts = Seq(6379)).start() and - GenericContainer.Def(Images.valkeyBundle, exposedPorts = Seq(6379)).start() - - given ExecutionContext = munitExecutionContext - - private val implemented: Set[String] = CommandSamples.all.map(_.command.name).toSet.filter(_.startsWith("JSON.")) - - test("implemented JSON commands never overlap the acknowledged gaps") { - assertEquals(implemented.intersect(JsonCoverage.skipped.keySet), Set.empty[String]) - } - - test("the JSON partition is exact per backend and the backend differences are acknowledged") { - withContainers { case redis and valkey => - (for { - redisJson <- jsonCommands(configOf(redis)) - valkeyJson <- jsonCommands(configOf(valkey)) - } yield { - assertExactPartition("JSON", redisJson ++ valkeyJson, implemented, JsonCoverage.skipped.keySet) - assertEquals(redisJson -- valkeyJson, JsonCoverage.redisOnly, "Redis-only JSON commands drifted from the acknowledged set") - assertEquals(valkeyJson -- redisJson, JsonCoverage.valkeyOnly, "Valkey-only JSON commands drifted from the acknowledged set") - }).unsafeRun - } - } - - private def jsonCommands(config: SageConfig): CIO[Set[String]] = - connectAndUse(config)(_.run(commandList).map(_.toSet.filter(_.startsWith("JSON.")))) -} diff --git a/integration-tests/shared/src/test/scala/sage/integration/commands/JsonSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/commands/JsonSuite.scala index 5bb3796e..eabf980d 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/commands/JsonSuite.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/commands/JsonSuite.scala @@ -1,129 +1,75 @@ package sage.integration.commands +import com.dimafeng.testcontainers.GenericContainer import kyo.compat.* import sage.SageException.DecodeError import sage.codec.ValueCodec import sage.commands.{JsonPath, JsonSetCondition, JsonType} -import sage.integration.{Images, ServerSuite} +import sage.integration.{BothServersSuite, Images} import sage.protocol.Frame final case class JsonAddress(city: String, zip: String) final case class JsonPerson(name: String, age: Int, address: JsonAddress) -abstract class JsonSuite(image: String) extends ServerSuite(image) { +class JsonSuite extends BothServersSuite { - test("JSON.SET and JSON.GET store and read a document, honoring NX/XX") { - withClient { client => - for { - set <- client.jsonSet("doc", JsonPath.root, """{"a":1,"s":"hi"}""") - nxSkip <- client.jsonSet("doc", JsonPath.root, """{"a":2}""", JsonSetCondition.IfNotExists) - whole <- client.jsonGet[String]("doc") - field <- client.jsonGet[String]("doc", JsonPath("$.a")) - absent <- client.jsonGet[String]("missing") - } yield { - assert(set) - assert(!nxSkip) - assert(whole.exists(_.contains("\"a\""))) - assert(field.exists(_.contains("1"))) - assertEquals(absent, None) - } - } + override protected def valkeyDef: GenericContainer.Def[GenericContainer] = serverDef(Images.valkeyBundle) + + clientTest("JSON.SET and JSON.GET store and read a document, honoring NX/XX") { client => + client.jsonSet("doc", JsonPath.root, """{"a":1,"s":"hi"}""").is(true) >> + client.jsonSet("doc", JsonPath.root, """{"a":2}""", JsonSetCondition.IfNotExists).is(false) >> + client.jsonGet[String]("doc").is(Some("""{"a":1,"s":"hi"}""")) >> + client.jsonGet[String]("doc", JsonPath("$.a")).is(Some("[1]")) >> + client.jsonGet[String]("missing").is(None) } - test("JSON.TYPE, JSON.OBJKEYS, JSON.OBJLEN inspect structure") { - withClient { client => - for { - _ <- client.jsonSet("shape", JsonPath.root, """{"a":1,"b":true,"c":"x"}""") - tpe <- client.jsonType("shape", JsonPath("$.a")) - keys <- client.jsonObjKeys("shape") - len <- client.jsonObjLen("shape") - none <- client.jsonType("shape", JsonPath("$.missing")) - } yield { - assertEquals(tpe, Vector(Some(JsonType.Integer))) - assert(keys.headOption.flatten.exists(_.toSet == Set("a", "b", "c"))) - assertEquals(len, Vector(Some(3L))) - assertEquals(none, Vector.empty) - } - } + clientTest("JSON.TYPE, JSON.OBJKEYS, JSON.OBJLEN inspect structure") { client => + client.jsonSet("shape", JsonPath.root, """{"a":1,"b":true,"c":"x"}""") >> + client.jsonType("shape", JsonPath("$.a")).is(Vector(Some(JsonType.Integer))) >> + client.jsonObjKeys("shape").is(Vector(Some(Vector("a", "b", "c")))) >> + client.jsonObjLen("shape").is(Vector(Some(3L))) >> + client.jsonType("shape", JsonPath("$.missing")).is(Vector.empty) } - test("numeric, string, and boolean mutations return per-match results") { - withClient { client => - for { - _ <- client.jsonSet("scalars", JsonPath.root, """{"n":1,"s":"ab","b":false}""") - incr <- client.jsonNumIncrBy("scalars", JsonPath("$.n"), 4.0) - mult <- client.jsonNumMultBy("scalars", JsonPath("$.n"), 2.0) - strLen <- client.jsonStrAppend("scalars", JsonPath("$.s"), "\"cd\"") - len <- client.jsonStrLen("scalars", JsonPath("$.s")) - toggle <- client.jsonToggle("scalars", JsonPath("$.b")) - } yield { - assertEquals(incr, Vector(Some(5.0))) - assertEquals(mult, Vector(Some(10.0))) - assertEquals(strLen, Vector(Some(4L))) - assertEquals(len, Vector(Some(4L))) - assertEquals(toggle, Vector(Some(true))) - } - } + clientTest("numeric, string, and boolean mutations return per-match results") { client => + client.jsonSet("scalars", JsonPath.root, """{"n":1,"s":"ab","b":false}""") >> + client.jsonNumIncrBy("scalars", JsonPath("$.n"), 4.0).is(Vector(Some(5.0))) >> + client.jsonNumMultBy("scalars", JsonPath("$.n"), 2.0).is(Vector(Some(10.0))) >> + client.jsonStrAppend("scalars", JsonPath("$.s"), "\"cd\"").is(Vector(Some(4L))) >> + client.jsonStrLen("scalars", JsonPath("$.s")).is(Vector(Some(4L))) >> + client.jsonToggle("scalars", JsonPath("$.b")).is(Vector(Some(true))) } - test("a multi-match path returns one entry per match for JSON.TYPE and JSON.NUMINCRBY") { - withClient { client => - for { - _ <- client.jsonSet("multi", JsonPath.root, """{"a":{"x":1},"b":{"x":"s"}}""") - tpe <- client.jsonType("multi", JsonPath("$..x")) - _ <- client.jsonSet("nums", JsonPath.root, """{"a":{"x":1},"b":{"x":2}}""") - incr <- client.jsonNumIncrBy("nums", JsonPath("$..x"), 5.0) - } yield { - assertEquals(tpe.toSet, Set(Option(JsonType.Integer), Option(JsonType.String))) - assertEquals(incr.flatten.toSet, Set(6.0, 7.0)) - } - } + clientTest("a multi-match path returns one entry per match for JSON.TYPE and JSON.NUMINCRBY") { client => + client.jsonSet("multi", JsonPath.root, """{"a":{"x":1},"b":{"x":"s"}}""") >> + client.jsonType("multi", JsonPath("$..x")).map(_.toSet).is(Set(Option(JsonType.Integer), Option(JsonType.String))) >> + client.jsonSet("nums", JsonPath.root, """{"a":{"x":1},"b":{"x":2}}""") >> + client.jsonNumIncrBy("nums", JsonPath("$..x"), 5.0).map(_.flatten.toSet).is(Set(6.0, 7.0)) } - test("array commands append, index, insert, pop, trim, and length") { - withClient { client => - for { - _ <- client.jsonSet("arr", JsonPath.root, """{"xs":[1,2,3]}""") - appLen <- client.jsonArrAppend("arr", JsonPath("$.xs"), "4", "5") - idx <- client.jsonArrIndex("arr", JsonPath("$.xs"), "3") - insLen <- client.jsonArrInsert("arr", JsonPath("$.xs"), 0L, "0") - len <- client.jsonArrLen("arr", JsonPath("$.xs")) - popped <- client.jsonArrPop[String]("arr", JsonPath("$.xs")) - trim <- client.jsonArrTrim("arr", JsonPath("$.xs"), 0L, 1L) - } yield { - assertEquals(appLen, Vector(Some(5L))) - assertEquals(idx, Vector(Some(2L))) - assertEquals(insLen, Vector(Some(6L))) - assertEquals(len, Vector(Some(6L))) - assert(popped.headOption.flatten.exists(_.contains("5"))) - assertEquals(trim, Vector(Some(2L))) - } - } + clientTest("array commands append, index, insert, pop, trim, and length") { client => + client.jsonSet("arr", JsonPath.root, """{"xs":[1,2,3]}""") >> + client.jsonArrAppend("arr", JsonPath("$.xs"), "4", "5").is(Vector(Some(5L))) >> + client.jsonArrIndex("arr", JsonPath("$.xs"), "3").is(Vector(Some(2L))) >> + client.jsonArrInsert("arr", JsonPath("$.xs"), 0L, "0").is(Vector(Some(6L))) >> + client.jsonArrLen("arr", JsonPath("$.xs")).is(Vector(Some(6L))) >> + client.jsonArrPop[String]("arr", JsonPath("$.xs")).is(Vector(Some("5"))) >> + client.jsonArrTrim("arr", JsonPath("$.xs"), 0L, 1L).is(Vector(Some(2L))) } - test("JSON.MGET, JSON.MSET, JSON.DEL, JSON.CLEAR, JSON.DEBUG MEMORY, JSON.RESP") { - withClient { client => - for { - _ <- client.jsonMSet(("m1", JsonPath.root, """{"v":1}"""), ("m2", JsonPath.root, """{"v":2}""")) - mget <- client.jsonMGet[String](JsonPath("$.v"))("m1", "m2", "m3") - del <- client.jsonDel("m1", JsonPath("$.v")) - clear <- client.jsonClear("m2", JsonPath.root) - mem <- client.jsonDebugMemory("m2") - resp <- client.jsonResp("m2") - } yield { - assert(mget(0).exists(_.contains("1"))) - assert(mget(1).exists(_.contains("2"))) - assertEquals(mget(2), None) - assertEquals(del, 1L) - assertEquals(clear, 1L) - assert(mem.headOption.flatten.exists(_ > 0L)) - assertNotEquals(resp, Frame.Null: Frame) - } - } + clientTest("JSON.MGET, JSON.MSET, JSON.DEL, JSON.CLEAR, JSON.DEBUG MEMORY, JSON.RESP") { client => + for { + _ <- client.jsonMSet(("m1", JsonPath.root, """{"v":1}"""), ("m2", JsonPath.root, """{"v":2}""")) + _ <- client.jsonMGet[String](JsonPath("$.v"))("m1", "m2", "m3").is(Vector(Some("[1]"), Some("[2]"), None)) + _ <- client.jsonDel("m1", JsonPath("$.v")).is(1L) + _ <- client.jsonClear("m2", JsonPath.root).is(1L) + _ <- client.jsonDebugMemory("m2").satisfies(_.headOption.flatten.exists(_ > 0L)) + resp <- client.jsonResp("m2") + } yield assertNotEquals(resp, Frame.Null: Frame) } - test("a user-supplied JSON codec (circe) round-trips typed documents") { + clientTest("a user-supplied JSON codec (circe) round-trips typed documents") { client => import io.circe.generic.auto.* import io.circe.parser.decode import io.circe.syntax.* @@ -131,27 +77,20 @@ abstract class JsonSuite(image: String) extends ServerSuite(image) { ValueCodec.string.emap(s => decode[A](s).left.map(DecodeError.fromThrowable))(_.asJson.noSpaces) val alice = JsonPerson("Alice", 30, JsonAddress("NYC", "10001")) - withClient { client => - for { - _ <- client.jsonSet("person:1", JsonPath.root, alice) - whole <- client.jsonGet[JsonPerson]("person:1") - ages <- client.jsonGet[Vector[Int]]("person:1", JsonPath("$.age")) - } yield { - assertEquals(whole, Some(alice)) - assertEquals(ages, Some(Vector(30))) - } - } + client.jsonSet("person:1", JsonPath.root, alice) >> + client.jsonGet[JsonPerson]("person:1").is(Some(alice)) >> + client.jsonGet[Vector[Int]]("person:1", JsonPath("$.age")).is(Some(Vector(30))) } - test("a legacy (non-$) path fails with a clear typed error, not silent data") { - withClient { client => + clientTest("a legacy (non-$) path fails with a clear typed error, not silent data") { client => + failsWith[DecodeError]( client.jsonSet("legacy", JsonPath.root, """{"xs":[1,2,3]}""").flatMap(_ => client.jsonArrLen("legacy", JsonPath(".xs"))) - }.failed.map { error => - assert(error.isInstanceOf[DecodeError], s"expected a DecodeError, got $error") - assert(error.getMessage.contains("legacy"), s"error should name the legacy-path cause: ${error.getMessage}") - } + ).satisfies(_.getMessage.contains("legacy")) } -} - -class RedisJsonSuite extends JsonSuite(Images.redis) -class ValkeyJsonSuite extends JsonSuite(Images.valkeyBundle) + // JSON.MERGE exists on Redis but not on Valkey Bundle. + redisTest("JSON.MERGE updates existing and creates new members") { client => + client.jsonSet("merge", JsonPath.root, """{"a":1,"b":2}""") >> + client.jsonMerge("merge", JsonPath.root, """{"b":20,"c":3}""") >> + client.jsonGet[String]("merge").is(Some("""{"a":1,"b":20,"c":3}""")) + } +} diff --git a/integration-tests/shared/src/test/scala/sage/integration/commands/KeysSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/commands/KeysSuite.scala index 6effb26d..673f7f1a 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/commands/KeysSuite.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/commands/KeysSuite.scala @@ -6,353 +6,189 @@ import scala.concurrent.duration.* import kyo.compat.* -import sage.client.internal.Client +import sage.client.internal.Paged import sage.commands.* -import sage.integration.{Images, ServerSuite} -import sage.integration.Ttls.{expiresWithin, remaining} +import sage.integration.BothServersSuite +import sage.integration.Ttls.expiresWithin -abstract class KeysSuite(image: String) extends ServerSuite(image) { +class KeysSuite extends BothServersSuite { - test("COPY copies and only overwrites with replace") { - withClient { client => - for { - _ <- client.set("keys-copy-src", "v1") - _ <- client.set("keys-copy-taken", "v2") - fresh <- client.copy("keys-copy-src", "keys-copy-dst") - ontoTaken <- client.copy("keys-copy-src", "keys-copy-taken") - replaced <- client.copy("keys-copy-src", "keys-copy-taken", replace = true) - copied <- client.get[String]("keys-copy-dst") - } yield { - assertEquals(fresh, true) - assertEquals(ontoTaken, false) - assertEquals(replaced, true) - assertEquals(copied, Some("v1")) - } - } + clientTest("COPY copies and only overwrites with replace") { client => + client.set("keys-copy-src", "v1") >> + client.set("keys-copy-taken", "v2") >> + client.copy("keys-copy-src", "keys-copy-dst").is(true) >> + client.copy("keys-copy-src", "keys-copy-taken").is(false) >> + client.copy("keys-copy-src", "keys-copy-taken", replace = true).is(true) >> + client.get[String]("keys-copy-dst").is(Some("v1")) } - test("EXISTS TOUCH DEL UNLINK count the keys they hit") { - withClient { client => - for { - _ <- client.mSet(("keys-cnt-a", "1"), ("keys-cnt-b", "2"), ("keys-cnt-c", "3")) - present <- client.exists("keys-cnt-a", "keys-cnt-b", "keys-cnt-c", "keys-cnt-missing") - touched <- client.touch("keys-cnt-a", "keys-cnt-b") - deleted <- client.del("keys-cnt-a", "keys-cnt-b") - removed <- client.unlink("keys-cnt-c", "keys-cnt-missing") - } yield { - assertEquals(present, 3L) - assertEquals(touched, 2L) - assertEquals(deleted, 2L) - assertEquals(removed, 1L) - } - } + clientTest("EXISTS TOUCH DEL UNLINK count the keys they hit") { client => + client.mSet(("keys-cnt-a", "1"), ("keys-cnt-b", "2"), ("keys-cnt-c", "3")) >> + client.exists("keys-cnt-a", "keys-cnt-b", "keys-cnt-c", "keys-cnt-missing").is(3L) >> + client.touch("keys-cnt-a", "keys-cnt-b").is(2L) >> + client.del("keys-cnt-a", "keys-cnt-b").is(2L) >> + client.unlink("keys-cnt-c", "keys-cnt-missing").is(1L) } - test("EXPIRE sets a ttl and PERSIST clears it") { - withClient { client => - for { - _ <- client.set("keys-expire", "v") - applied <- client.expire("keys-expire", 60.seconds) - ttl <- client.ttl("keys-expire") - persisted <- client.persist("keys-expire") - cleared <- client.ttl("keys-expire") - missing <- client.expire("keys-expire-missing", 60.seconds) - } yield { - assertEquals(applied, true) - assert(expiresWithin(ttl, 60.seconds)) - assertEquals(persisted, true) - assertEquals(cleared, Ttl.NoExpiry) - assertEquals(missing, false) - } - } + clientTest("EXPIRE sets a ttl and PERSIST clears it") { client => + client.set("keys-expire", "v") >> + client.expire("keys-expire", 60.seconds).is(true) >> + client.ttl("keys-expire").satisfies(expiresWithin(_, 60.seconds)) >> + client.persist("keys-expire").is(true) >> + client.ttl("keys-expire").is(Ttl.NoExpiry) >> + client.expire("keys-expire-missing", 60.seconds).is(false) } - test("a sub-second duration takes the millisecond path end to end") { - withClient { client => - for { - _ <- client.set("keys-pexpire", "v") - _ <- client.expire("keys-pexpire", 90500.millis) - ttl <- client.pTtl("keys-pexpire") - } yield assert(remaining(ttl).exists(r => r > 89.seconds && r <= 90500.millis)) - } + clientTest("a sub-second duration takes the millisecond path end to end") { client => + client.set("keys-pexpire", "v") >> + client.expire("keys-pexpire", 90500.millis) >> + client.pTtl("keys-pexpire").satisfies(expiresWithin(_, 90500.millis, above = 89.seconds)) } - test("EXPIRE conditions guard against the current ttl") { - withClient { client => - for { - _ <- client.set("keys-cond", "v") - noExpiry <- client.expire("keys-cond", 60.seconds, ExpireCondition.IfNoExpiry) - notLonger <- client.expire("keys-cond", 30.seconds, ExpireCondition.IfGreater) - longer <- client.expire("keys-cond", 120.seconds, ExpireCondition.IfGreater) - shorter <- client.expire("keys-cond", 60.seconds, ExpireCondition.IfLess) - hasExpiry <- client.expire("keys-cond", 90.seconds, ExpireCondition.IfHasExpiry) - notWithout <- client.expire("keys-cond", 30.seconds, ExpireCondition.IfNoExpiry) - } yield { - assertEquals(noExpiry, true) - assertEquals(notLonger, false) - assertEquals(longer, true) - assertEquals(shorter, true) - assertEquals(hasExpiry, true) - assertEquals(notWithout, false) - } - } + clientTest("EXPIRE conditions guard against the current ttl") { client => + client.set("keys-cond", "v") >> + client.expire("keys-cond", 60.seconds, ExpireCondition.IfNoExpiry).is(true) >> + client.expire("keys-cond", 30.seconds, ExpireCondition.IfGreater).is(false) >> + client.expire("keys-cond", 120.seconds, ExpireCondition.IfGreater).is(true) >> + client.expire("keys-cond", 60.seconds, ExpireCondition.IfLess).is(true) >> + client.expire("keys-cond", 90.seconds, ExpireCondition.IfHasExpiry).is(true) >> + client.expire("keys-cond", 30.seconds, ExpireCondition.IfNoExpiry).is(false) } - test("EXPIREAT and EXPIRETIME round-trip an absolute deadline") { - withClient { client => - val deadline = Instant.ofEpochSecond(Instant.now().getEpochSecond + 3600) - for { - _ <- client.set("keys-at", "v") - applied <- client.expireAt("keys-at", deadline) - seconds <- client.expireTime("keys-at") - millis <- client.pExpireTime("keys-at") - _ <- client.set("keys-at-plain", "v") - noExpiry <- client.expireTime("keys-at-plain") - noKey <- client.expireTime("keys-at-missing") - } yield { - assertEquals(applied, true) - assertEquals(seconds, ExpiryTime.At(deadline)) - assertEquals(millis, ExpiryTime.At(deadline)) - assertEquals(noExpiry, ExpiryTime.NoExpiry) - assertEquals(noKey, ExpiryTime.NoKey) - } - } + clientTest("EXPIREAT and EXPIRETIME round-trip an absolute deadline") { client => + val deadline = Instant.ofEpochSecond(Instant.now().getEpochSecond + 3600) + client.set("keys-at", "v") >> + client.expireAt("keys-at", deadline).is(true) >> + client.expireTime("keys-at").is(ExpiryTime.At(deadline)) >> + client.pExpireTime("keys-at").is(ExpiryTime.At(deadline)) >> + client.set("keys-at-plain", "v") >> + client.expireTime("keys-at-plain").is(ExpiryTime.NoExpiry) >> + client.expireTime("keys-at-missing").is(ExpiryTime.NoKey) } - test("TTL distinguishes a missing key from a key without expiry") { - withClient { client => - for { - noKey <- client.ttl("keys-ttl-missing") - _ <- client.set("keys-ttl-plain", "v") - noExpiry <- client.pTtl("keys-ttl-plain") - } yield { - assertEquals(noKey, Ttl.NoKey) - assertEquals(noExpiry, Ttl.NoExpiry) - } - } + clientTest("TTL distinguishes a missing key from a key without expiry") { client => + client.ttl("keys-ttl-missing").is(Ttl.NoKey) >> + client.set("keys-ttl-plain", "v") >> + client.pTtl("keys-ttl-plain").is(Ttl.NoExpiry) } - test("KEYS returns the keys matching a pattern") { - withClient { client => - for { - _ <- client.mSet(("keys-glob:1", "a"), ("keys-glob:2", "b"), ("keys-other", "c")) - matched <- client.keys("keys-glob:*") - } yield assertEquals(matched.toSet, Set("keys-glob:1", "keys-glob:2")) - } + clientTest("KEYS returns the keys matching a pattern") { client => + client.mSet(("keys-glob:1", "a"), ("keys-glob:2", "b"), ("keys-other", "c")) >> + client.keys("keys-glob:*").map(_.toSet).is(Set("keys-glob:1", "keys-glob:2")) } - test("RANDOMKEY returns a key once data exists") { - withClient { client => - for { - _ <- client.set("keys-random", "v") - random <- client.randomKey - } yield assert(random.isDefined) - } + clientTest("RANDOMKEY returns a key once data exists") { client => + client.set("keys-random", "v") >> + client.randomKey.satisfies(_.isDefined) } - test("RENAME moves a key and RENAMENX refuses an occupied destination") { - withClient { client => - for { - _ <- client.set("keys-ren-a", "v") - _ <- client.set("keys-ren-taken", "w") - _ <- client.rename("keys-ren-a", "keys-ren-b") - moved <- client.get[String]("keys-ren-b") - refused <- client.renameNx("keys-ren-b", "keys-ren-taken") - accepted <- client.renameNx("keys-ren-b", "keys-ren-c") - } yield { - assertEquals(moved, Some("v")) - assertEquals(refused, false) - assertEquals(accepted, true) - } - } + clientTest("RENAME moves a key and RENAMENX refuses an occupied destination") { client => + client.set("keys-ren-a", "v") >> + client.set("keys-ren-taken", "w") >> + client.rename("keys-ren-a", "keys-ren-b") >> + client.get[String]("keys-ren-b").is(Some("v")) >> + client.renameNx("keys-ren-b", "keys-ren-taken").is(false) >> + client.renameNx("keys-ren-b", "keys-ren-c").is(true) } - test("TYPE reports the key's type and None for a missing key") { - withClient { client => - for { - _ <- client.set("keys-type-str", "v") - _ <- client.lPush("keys-type-list", "v") - str <- client.typeOf("keys-type-str") - list <- client.typeOf("keys-type-list") - missing <- client.typeOf("keys-type-missing") - } yield { - assertEquals(str, Some(RedisType.String)) - assertEquals(list, Some(RedisType.List)) - assertEquals(missing, None) - } - } + clientTest("TYPE reports the key's type and None for a missing key") { client => + client.set("keys-type-str", "v") >> + client.lPush("keys-type-list", "v") >> + client.typeOf("keys-type-str").is(Some(RedisType.String)) >> + client.typeOf("keys-type-list").is(Some(RedisType.List)) >> + client.typeOf("keys-type-missing").is(None) } - test("SCAN visits every key, terminating on the zero cursor rather than an empty page") { - withClient { client => - val pairs = (1 to 100).map(i => (s"keys-scan:$i", "v")).toVector - for { - _ <- client.mSet(pairs.head, pairs.tail*) - found <- scanAll(client, pattern = Some("keys-scan:*"), count = Some(10L), ofType = None) - } yield assertEquals(found, pairs.map(_._1).toSet) - } + clientTest("SCAN visits every key, terminating on the zero cursor rather than an empty page") { client => + val first = ("keys-scan:0", "v") + val rest = (1 to 99).map(i => (s"keys-scan:$i", "v")) + client.mSet(first, rest*) >> + drain(Paged.scanAll[String](client.runner, Some("keys-scan:*"), Some(10L), None)).is((first +: rest).map(_._1).toSet) } - test("SCAN filters by type") { - withClient { client => - for { - _ <- client.set("keys-scant-str", "v") - _ <- client.lPush("keys-scant-list", "v") - found <- scanAll(client, pattern = Some("keys-scant-*"), count = None, ofType = Some(RedisType.List)) - } yield assertEquals(found, Set("keys-scant-list")) - } + clientTest("SCAN filters by type") { client => + client.set("keys-scant-str", "v") >> + client.lPush("keys-scant-list", "v") >> + drain(Paged.scanAll[String](client.runner, Some("keys-scant-*"), None, Some(RedisType.List))).is(Set("keys-scant-list")) } - test("SORT orders a list numerically and alphabetically, with LIMIT and DESC") { - withClient { client => - for { - _ <- client.rPush("keys-sort", "3", "1", "2") - numeric <- client.sort[String]("keys-sort") - desc <- client.sort[String]("keys-sort", order = SortOrder.Desc, limit = Some(Limit(0L, 2L))) - alpha <- client.sort[String]("keys-sort", alpha = true, order = SortOrder.Desc) - } yield { - assertEquals(numeric, Vector(Some("1"), Some("2"), Some("3"))) - assertEquals(desc, Vector(Some("3"), Some("2"))) - assertEquals(alpha, Vector(Some("3"), Some("2"), Some("1"))) - } - } + clientTest("SORT orders a list numerically and alphabetically, with LIMIT and DESC") { client => + client.rPush("keys-sort", "3", "1", "2") >> + client.sort[String]("keys-sort").is(Vector(Some("1"), Some("2"), Some("3"))) >> + client.sort[String]("keys-sort", order = SortOrder.Desc, limit = Some(Limit(0L, 2L))).is(Vector(Some("3"), Some("2"))) >> + client.sort[String]("keys-sort", alpha = true, order = SortOrder.Desc).is(Vector(Some("3"), Some("2"), Some("1"))) } - test("SORT BY/GET sorts by external weights and projects external values, nil for missing") { - withClient { client => - for { - _ <- client.rPush("keys-sortby", "1", "2", "3") - _ <- client.mSet(("keys-w-1", "30"), ("keys-w-2", "10"), ("keys-w-3", "20")) - _ <- client.mSet(("keys-d-1", "A"), ("keys-d-3", "C")) - got <- client.sort[String]("keys-sortby", by = Some("keys-w-*"), get = Vector("keys-d-*", "#")) - } yield assertEquals(got, Vector(None, Some("2"), Some("C"), Some("3"), Some("A"), Some("1"))) - } + clientTest("SORT BY/GET sorts by external weights and projects external values, nil for missing") { client => + client.rPush("keys-sortby", "1", "2", "3") >> + client.mSet(("keys-w-1", "30"), ("keys-w-2", "10"), ("keys-w-3", "20")) >> + client.mSet(("keys-d-1", "A"), ("keys-d-3", "C")) >> + client + .sort[String]("keys-sortby", by = Some("keys-w-*"), get = Vector("keys-d-*", "#")) + .is(Vector(None, Some("2"), Some("C"), Some("3"), Some("A"), Some("1"))) } - test("SORT_RO reads without storing; SORT STORE writes the result and returns the count") { - withClient { client => - for { - _ <- client.rPush("keys-sortstore", "b", "a", "c") - ro <- client.sortRo[String]("keys-sortstore", alpha = true) - stored <- client.sortStore("keys-sortstore-dst", "keys-sortstore", alpha = true) - dst <- client.lRange[String]("keys-sortstore-dst", 0L, -1L) - } yield { - assertEquals(ro, Vector(Some("a"), Some("b"), Some("c"))) - assertEquals(stored, 3L) - assertEquals(dst, Vector("a", "b", "c")) - } - } + clientTest("SORT_RO reads without storing; SORT STORE writes the result and returns the count") { client => + client.rPush("keys-sortstore", "b", "a", "c") >> + client.sortRo[String]("keys-sortstore", alpha = true).is(Vector(Some("a"), Some("b"), Some("c"))) >> + client.sortStore("keys-sortstore-dst", "keys-sortstore", alpha = true).is(3L) >> + client.lRange[String]("keys-sortstore-dst", 0L, -1L).is(Vector("a", "b", "c")) } - test("MOVE relocates a key out of the current database") { - withClient { client => - for { - _ <- client.set("keys-move", "v") - moved <- client.move("keys-move", 1) - gone <- client.exists("keys-move") - again <- client.move("keys-move", 1) - } yield { - assertEquals(moved, true) - assertEquals(gone, 0L) - assertEquals(again, false) - } - } + clientTest("MOVE relocates a key out of the current database") { client => + client.set("keys-move", "v") >> + client.move("keys-move", 1).is(true) >> + client.exists("keys-move").is(0L) >> + client.move("keys-move", 1).is(false) } - test("DUMP and RESTORE round-trip a value through its serialized form") { - withClient { client => - for { - _ <- client.set("keys-dump", "payload") - dumped <- client.dump("keys-dump") - restored <- dumped.fold(CIO.value(()))(bytes => client.restore("keys-restore", bytes)) - value <- client.get[String]("keys-restore") - replaced <- dumped.fold(CIO.value(false))(bytes => client.restore("keys-restore", bytes, replace = true).map(_ => true)) - missing <- client.dump("keys-dump-missing") - } yield { - assert(dumped.isDefined) - assertEquals(value, Some("payload")) - assertEquals(replaced, true) - assertEquals(missing, None) - } - } + clientTest("DUMP and RESTORE round-trip a value through its serialized form") { client => + for { + _ <- client.set("keys-dump", "payload") + bytes <- client.dump("keys-dump").flatMap(required("DUMP", _)) + _ <- client.restore("keys-restore", bytes) + _ <- client.get[String]("keys-restore").is(Some("payload")) + _ <- client.restore("keys-restore", bytes, replace = true) + _ <- client.dump("keys-dump-missing").is(None) + } yield () } - test("MIGRATE reports NOKEY when the source key is absent") { - withClient { client => - for { - result <- client.migrate("localhost", 6379, 0, 1.second)("keys-migrate-ghost") - } yield assertEquals(result, MigrateResult.NoKey) - } + clientTest("MIGRATE reports NOKEY when the source key is absent") { client => + client.migrate("localhost", 6379, 0, 1.second)("keys-migrate-ghost").is(MigrateResult.NoKey) } - test("OBJECT exposes encoding, refcount, and idle time; None for a missing key") { - withClient { client => - for { - _ <- client.set("keys-object", "12345") - encoding <- client.objectEncoding("keys-object") - refCount <- client.objectRefCount("keys-object") - idle <- client.objectIdleTime("keys-object") - missEnc <- client.objectEncoding("keys-object-missing") - missRef <- client.objectRefCount("keys-object-missing") - missIdle <- client.objectIdleTime("keys-object-missing") - } yield { - assertEquals(encoding, Some("int")) - assert(refCount.exists(_ >= 1L)) - assert(idle.exists(_ >= Duration.Zero)) - assertEquals(missEnc, None) - assertEquals(missRef, None) - assertEquals(missIdle, None) - } - } + clientTest("OBJECT exposes encoding, refcount, and idle time; None for a missing key") { client => + client.set("keys-object", "12345") >> + client.objectEncoding("keys-object").is(Some("int")) >> + client.objectRefCount("keys-object").satisfies(_.exists(_ >= 1L)) >> + client.objectIdleTime("keys-object").satisfies(_.exists(_ >= Duration.Zero)) >> + client.objectEncoding("keys-object-missing").is(None) >> + client.objectRefCount("keys-object-missing").is(None) >> + client.objectIdleTime("keys-object-missing").is(None) } - test("OBJECT FREQ reports access frequency under an LFU policy, None for a missing key") { - withClient { client => - for { - _ <- client.configSet("maxmemory-policy" -> "allkeys-lfu") - _ <- client.set("keys-freq", "v") - freq <- client.objectFreq("keys-freq") - missing <- client.objectFreq("keys-freq-missing") - _ <- client.configSet("maxmemory-policy" -> "noeviction") - } yield { - assert(freq.exists(_ >= 0L)) - assertEquals(missing, None) - } - } + clientTest("OBJECT FREQ reports access frequency under an LFU policy, None for a missing key") { client => + client.configSet("maxmemory-policy" -> "allkeys-lfu") >> + client.set("keys-freq", "v") >> + client.objectFreq("keys-freq").satisfies(_.exists(_ >= 0L)) >> + client.objectFreq("keys-freq-missing").is(None) >> + client.configSet("maxmemory-policy" -> "noeviction") } - test("FLUSHALL empties the keyspace") { - withClient { client => - for { - _ <- client.set("keys-flush", "v") - before <- client.exists("keys-flush") - _ <- client.flushAll() - after <- client.exists("keys-flush") - } yield { - assertEquals(before, 1L) - assertEquals(after, 0L) - } - } + clientTest("FLUSHALL empties the keyspace") { client => + client.set("keys-flush", "v") >> + client.exists("keys-flush").is(1L) >> + client.flushAll() >> + client.exists("keys-flush").is(0L) } - private def scanAll( - client: Client[CIO, String], - pattern: Option[String], - count: Option[Long], - ofType: Option[RedisType] - ): CIO[Set[String]] = { - def loop(cursor: ScanCursor, found: Set[String]): CIO[Set[String]] = - client.scan(cursor, pattern, count, ofType).flatMap { page => - val collected = found ++ page.items - page.next match { - case Some(next) => loop(next, collected) - case None => CIO.value(collected) - } - } - loop(ScanCursor.start, Set.empty) + // DELIFEQ exists only on Valkey. + valkeyTest("DELIFEQ deletes only when the current value matches") { client => + client.set("vk-lock", "token-1") >> + client.delIfEq("vk-lock", "other").is(false) >> + client.exists("vk-lock").is(1L) >> + client.delIfEq("vk-lock", "token-1").is(true) >> + client.exists("vk-lock").is(0L) >> + client.delIfEq("vk-lock-absent", "x").is(false) } } - -class RedisKeysSuite extends KeysSuite(Images.redis) - -class ValkeyKeysSuite extends KeysSuite(Images.valkey) diff --git a/integration-tests/shared/src/test/scala/sage/integration/commands/ListsSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/commands/ListsSuite.scala index 4faef76c..fd47c127 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/commands/ListsSuite.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/commands/ListsSuite.scala @@ -5,178 +5,93 @@ import scala.concurrent.duration.* import kyo.compat.* import sage.commands.{BlockTimeout, InsertPosition, ListSide} -import sage.integration.{Images, ServerSuite} +import sage.integration.BothServersSuite -abstract class ListsSuite(image: String) extends ServerSuite(image) { +class ListsSuite extends BothServersSuite { - test("RPUSH and LPUSH build a list that LRANGE LLEN and LINDEX read") { - withClient { client => - for { - _ <- client.rPush("list-build", "b", "c") - _ <- client.lPush("list-build", "a") - all <- client.lRange[String]("list-build", 0L, -1L) - len <- client.lLen("list-build") - second <- client.lIndex[String]("list-build", 1L) - } yield { - assertEquals(all, Vector("a", "b", "c")) - assertEquals(len, 3L) - assertEquals(second, Some("b")) - } - } + clientTest("RPUSH and LPUSH build a list that LRANGE LLEN and LINDEX read") { client => + client.rPush("list-build", "b", "c") >> + client.lPush("list-build", "a") >> + client.lRange[String]("list-build", 0L, -1L).is(Vector("a", "b", "c")) >> + client.lLen("list-build").is(3L) >> + client.lIndex[String]("list-build", 1L).is(Some("b")) } - test("LPUSHX and RPUSHX only extend an existing list") { - withClient { client => - for { - absent <- client.rPushX("list-x-missing", "v") - _ <- client.rPush("list-x", "a") - present <- client.rPushX("list-x", "b") - head <- client.lPushX("list-x", "z") - all <- client.lRange[String]("list-x", 0L, -1L) - } yield { - assertEquals(absent, 0L) - assertEquals(present, 2L) - assertEquals(head, 3L) - assertEquals(all, Vector("z", "a", "b")) - } - } + clientTest("LPUSHX and RPUSHX only extend an existing list") { client => + client.rPushX("list-x-missing", "v").is(0L) >> + client.rPush("list-x", "a") >> + client.rPushX("list-x", "b").is(2L) >> + client.lPushX("list-x", "z").is(3L) >> + client.lRange[String]("list-x", 0L, -1L).is(Vector("z", "a", "b")) } - test("LPOP and RPOP pop one or several elements, empty when the list is gone") { - withClient { client => - for { - _ <- client.rPush("list-pop", "a", "b", "c", "d") - head <- client.lPop[String]("list-pop") - tail <- client.rPop[String]("list-pop") - twoHead <- client.lPopCount[String]("list-pop", 2L) - drained <- client.lPopCount[String]("list-pop", 2L) - none <- client.lPop[String]("list-pop") - } yield { - assertEquals(head, Some("a")) - assertEquals(tail, Some("d")) - assertEquals(twoHead, Vector("b", "c")) - assertEquals(drained, Vector.empty[String]) - assertEquals(none, None) - } - } + clientTest("LPOP and RPOP pop one or several elements, empty when the list is gone") { client => + client.rPush("list-pop", "a", "b", "c", "d") >> + client.lPop[String]("list-pop").is(Some("a")) >> + client.rPop[String]("list-pop").is(Some("d")) >> + client.lPopCount[String]("list-pop", 2L).is(Vector("b", "c")) >> + client.lPopCount[String]("list-pop", 2L).is(Vector.empty[String]) >> + client.lPop[String]("list-pop").is(None) } - test("RPOP with a count pops several elements from the tail in pop order") { - withClient { client => - for { - _ <- client.rPush("list-rpopc", "a", "b", "c", "d") - tail <- client.rPopCount[String]("list-rpopc", 2L) - rest <- client.lRange[String]("list-rpopc", 0L, -1L) - } yield { - assertEquals(tail, Vector("d", "c")) - assertEquals(rest, Vector("a", "b")) - } - } + clientTest("RPOP with a count pops several elements from the tail in pop order") { client => + client.rPush("list-rpopc", "a", "b", "c", "d") >> + client.rPopCount[String]("list-rpopc", 2L).is(Vector("d", "c")) >> + client.lRange[String]("list-rpopc", 0L, -1L).is(Vector("a", "b")) } - test("LSET LINSERT LREM and LTRIM edit the list in place") { - withClient { client => - for { - _ <- client.rPush("list-edit", "a", "b", "b", "c") - _ <- client.lSet("list-edit", 0L, "A") - inserted <- client.lInsert("list-edit", InsertPosition.Before, "c", "x") - noPivot <- client.lInsert("list-edit", InsertPosition.After, "zzz", "y") - removed <- client.lRem("list-edit", 0L, "b") - _ <- client.lTrim("list-edit", 0L, 1L) - all <- client.lRange[String]("list-edit", 0L, -1L) - } yield { - assertEquals(inserted, 5L) - assertEquals(noPivot, -1L) - assertEquals(removed, 2L) - assertEquals(all, Vector("A", "x")) - } - } + clientTest("LSET LINSERT LREM and LTRIM edit the list in place") { client => + client.rPush("list-edit", "a", "b", "b", "c") >> + client.lSet("list-edit", 0L, "A") >> + client.lInsert("list-edit", InsertPosition.Before, "c", "x").is(5L) >> + client.lInsert("list-edit", InsertPosition.After, "zzz", "y").is(-1L) >> + client.lRem("list-edit", 0L, "b").is(2L) >> + client.lTrim("list-edit", 0L, 1L) >> + client.lRange[String]("list-edit", 0L, -1L).is(Vector("A", "x")) } - test("LPOS finds the first match, all matches, and reports None when absent") { - withClient { client => - for { - _ <- client.rPush("list-pos", "a", "b", "a", "c", "a") - first <- client.lPos("list-pos", "a") - last <- client.lPos("list-pos", "a", rank = Some(-1L)) - every <- client.lPosCount("list-pos", "a", 0L) - none <- client.lPos("list-pos", "zzz") - } yield { - assertEquals(first, Some(0L)) - assertEquals(last, Some(4L)) - assertEquals(every, Vector(0L, 2L, 4L)) - assertEquals(none, None) - } - } + clientTest("LPOS finds the first match, all matches, and reports None when absent") { client => + client.rPush("list-pos", "a", "b", "a", "c", "a") >> + client.lPos("list-pos", "a").is(Some(0L)) >> + client.lPos("list-pos", "a", rank = Some(-1L)).is(Some(4L)) >> + client.lPosCount("list-pos", "a", 0L).is(Vector(0L, 2L, 4L)) >> + client.lPos("list-pos", "zzz").is(None) } - test("LMOVE shifts an element between ends and LMPOP pops from the first non-empty key") { - withClient { client => - for { - _ <- client.rPush("list-src", "a", "b", "c") - moved <- client.lMove[String]("list-src", "list-dst", ListSide.Left, ListSide.Right) - dst <- client.lRange[String]("list-dst", 0L, -1L) - popped <- client.lMpop[String]("list-empty", "list-src")(ListSide.Left, count = Some(2L)) - emptyOk <- client.lMpop[String]("list-empty")(ListSide.Left) - } yield { - assertEquals(moved, Some("a")) - assertEquals(dst, Vector("a")) - assertEquals(popped, Some(("list-src", Vector("b", "c")))) - assertEquals(emptyOk, None) - } - } + clientTest("LMOVE shifts an element between ends and LMPOP pops from the first non-empty key") { client => + client.rPush("list-src", "a", "b", "c") >> + client.lMove[String]("list-src", "list-dst", ListSide.Left, ListSide.Right).is(Some("a")) >> + client.lRange[String]("list-dst", 0L, -1L).is(Vector("a")) >> + client.lMpop[String]("list-empty", "list-src")(ListSide.Left, count = Some(2L)).is(Some(("list-src", Vector("b", "c")))) >> + client.lMpop[String]("list-empty")(ListSide.Left).is(None) } - test("BLPOP and BRPOP return a present element and time out to None on an empty key") { - withClient { client => - for { - _ <- client.rPush("blpop-data", "a", "b") - head <- client.blPop[String]("blpop-data")(BlockTimeout.After(1.second)) - tail <- client.brPop[String]("blpop-data")(BlockTimeout.After(1.second)) - empty <- client.blPop[String]("blpop-empty")(BlockTimeout.After(100.millis)) - } yield { - assertEquals(head, Some(("blpop-data", "a"))) - assertEquals(tail, Some(("blpop-data", "b"))) - assertEquals(empty, None) - } - } + clientTest("BLPOP and BRPOP return a present element and time out to None on an empty key") { client => + client.rPush("blpop-data", "a", "b") >> + client.blPop[String]("blpop-data")(BlockTimeout.After(1.second)).is(Some(("blpop-data", "a"))) >> + client.brPop[String]("blpop-data")(BlockTimeout.After(1.second)).is(Some(("blpop-data", "b"))) >> + client.blPop[String]("blpop-empty")(BlockTimeout.After(100.millis)).is(None) } - test("BLMOVE and BLMPOP move and pop, timing out to None when nothing is available") { - withClient { client => - for { - _ <- client.rPush("blmove-src", "a", "b") - moved <- client.blMove[String]("blmove-src", "blmove-dst", ListSide.Left, ListSide.Right, BlockTimeout.After(1.second)) - dst <- client.lRange[String]("blmove-dst", 0L, -1L) - popped <- client.blMpop[String]("blmpop-empty", "blmove-src")(ListSide.Left, BlockTimeout.After(1.second), count = Some(2L)) - none <- client.blMpop[String]("blmpop-empty")(ListSide.Left, BlockTimeout.After(100.millis)) - } yield { - assertEquals(moved, Some("a")) - assertEquals(dst, Vector("a")) - assertEquals(popped, Some(("blmove-src", Vector("b")))) - assertEquals(none, None) - } - } + clientTest("BLMOVE and BLMPOP move and pop, timing out to None when nothing is available") { client => + client.rPush("blmove-src", "a", "b") >> + client.blMove[String]("blmove-src", "blmove-dst", ListSide.Left, ListSide.Right, BlockTimeout.After(1.second)).is(Some("a")) >> + client.lRange[String]("blmove-dst", 0L, -1L).is(Vector("a")) >> + client + .blMpop[String]("blmpop-empty", "blmove-src")(ListSide.Left, BlockTimeout.After(1.second), count = Some(2L)) + .is(Some(("blmove-src", Vector("b")))) >> + client.blMpop[String]("blmpop-empty")(ListSide.Left, BlockTimeout.After(100.millis)).is(None) } - test("a blocking command does not stall ordinary commands on the multiplexed connection") { - withClient { client => - CIO - .zip( - client.blPop[String]("nonstall-queue")(BlockTimeout.After(5.seconds)), - for { - pong <- client.ping() - _ <- client.rPush("nonstall-queue", "payload") - } yield pong - ) - .map { case (popped, pong) => - assertEquals(popped, Some(("nonstall-queue", "payload"))) - assertEquals(pong, "PONG") - } - } + clientTest("a blocking command does not stall ordinary commands on the multiplexed connection") { client => + CIO + .zip( + client.blPop[String]("nonstall-queue")(BlockTimeout.After(5.seconds)), + for { + pong <- client.ping() + _ <- client.rPush("nonstall-queue", "payload") + } yield pong + ) + .is((Some(("nonstall-queue", "payload")), "PONG")) } } - -class RedisListsSuite extends ListsSuite(Images.redis) - -class ValkeyListsSuite extends ListsSuite(Images.valkey) diff --git a/integration-tests/shared/src/test/scala/sage/integration/commands/PubsubSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/commands/PubsubSuite.scala index 162cf0a7..0c84f523 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/commands/PubsubSuite.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/commands/PubsubSuite.scala @@ -1,87 +1,40 @@ package sage.integration.commands -import kyo.compat.* - import sage.{Message, PatternMessage} -import sage.integration.{Images, ServerSuite} - -abstract class PubsubSuite(image: String) extends ServerSuite(image) { - - // a subscribe blocks until the server confirms it, but UNSUBSCRIBE is fire-and-forget, so let it propagate before re-checking PUBSUB - private def settle: CIO[Unit] = CIO.blocking(Thread.sleep(300)) - - test("PUBLISH delivers to a channel subscriber; PUBSUB introspection reflects the subscription") { - withClient { client => - for { - sub <- client.subscribeChannels[String]("news") - channels <- client.pubsubChannels() - numSub <- client.pubsubNumSub("news") - numPat <- client.pubsubNumPat - received <- client.publish("news", "hello") - first <- sub.next - _ <- client.publish("news", "world") - second <- sub.next - _ <- sub.close - } yield { - assert(channels.contains("news"), channels) - assertEquals(numSub, Map("news" -> 1L)) - assertEquals(numPat, 0L) - assertEquals(received, 1L) - assertEquals(first, Some(Message("news", "hello"))) - assertEquals(second, Some(Message("news", "world"))) - } +import sage.integration.{BothServersSuite, Eventually} + +class PubsubSuite extends BothServersSuite { + + clientTest("PUBLISH delivers to a channel subscriber; PUBSUB introspection reflects the subscription") { client => + withSubscription(client.subscribeChannels[String]("news")) { sub => + client.pubsubChannels().satisfies(_.contains("news")) >> + client.pubsubNumSub("news").is(Map("news" -> 1L)) >> + client.pubsubNumPat.is(0L) >> + client.publish("news", "hello").is(1L) >> + sub.next.is(Some(Message("news", "hello"))) >> + client.publish("news", "world") >> + sub.next.is(Some(Message("news", "world"))) } } - test("PSUBSCRIBE matches by pattern, naming the pattern and the concrete channel; PUBSUB NUMPAT counts it") { - withClient { client => - for { - sub <- client.subscribePatterns[String]("news.*") - numPat <- client.pubsubNumPat - _ <- client.publish("news.sports", "goal") - msg <- sub.next - _ <- sub.close - } yield { - assertEquals(numPat, 1L) - assertEquals(msg, Some(PatternMessage("news.*", "news.sports", "goal"))) - } + clientTest("PSUBSCRIBE matches by pattern, naming the pattern and the concrete channel; PUBSUB NUMPAT counts it") { client => + withSubscription(client.subscribePatterns[String]("news.*")) { sub => + client.pubsubNumPat.is(1L) >> client.publish("news.sports", "goal") >> sub.next.is(Some(PatternMessage("news.*", "news.sports", "goal"))) } } - test("SSUBSCRIBE delivers a sharded message; SPUBLISH returns the receiver count; PUBSUB SHARDCHANNELS reflects it") { - withClient { client => - for { - sub <- client.subscribeShardChannels[String]("orders") - channels <- client.pubsubShardChannels() - numSub <- client.pubsubShardNumSub("orders") - received <- client.sPublish("orders", "placed") - first <- sub.next - _ <- sub.close - } yield { - assert(channels.contains("orders"), channels) - assertEquals(numSub, Map("orders" -> 1L)) - assertEquals(received, 1L) - assertEquals(first, Some(Message("orders", "placed"))) - } + clientTest("SSUBSCRIBE delivers a sharded message; SPUBLISH returns the receiver count; PUBSUB SHARDCHANNELS reflects it") { client => + withSubscription(client.subscribeShardChannels[String]("orders")) { sub => + client.pubsubShardChannels().satisfies(_.contains("orders")) >> + client.pubsubShardNumSub("orders").is(Map("orders" -> 1L)) >> + client.sPublish("orders", "placed").is(1L) >> + sub.next.is(Some(Message("orders", "placed"))) } } - test("closing the last subscriber unsubscribes on the server") { - withClient { client => - for { - sub <- client.subscribeChannels[String]("bye") - active <- client.pubsubChannels() - _ <- sub.close - _ <- settle - inactive <- client.pubsubChannels() - } yield { - assert(active.contains("bye"), active) - assert(!inactive.contains("bye"), inactive) - } - } + clientTest("closing the last subscriber unsubscribes on the server") { client => + withSubscription(client.subscribeChannels[String]("bye"))(_ => client.pubsubChannels().satisfies(_.contains("bye"))) >> + // a subscribe blocks until the server confirms it, but UNSUBSCRIBE is fire-and-forget, so poll until it propagates + Eventually(50)(client.pubsubChannels().satisfies(!_.contains("bye"))) } } - -class RedisPubsubSuite extends PubsubSuite(Images.redis) - -class ValkeyPubsubSuite extends PubsubSuite(Images.valkey) diff --git a/integration-tests/shared/src/test/scala/sage/integration/commands/RedisArraysSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/commands/RedisArraysSuite.scala index 72f32ee2..43c6fb18 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/commands/RedisArraysSuite.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/commands/RedisArraysSuite.scala @@ -10,148 +10,78 @@ import sage.integration.{Images, ServerSuite} */ class RedisArraysSuite extends ServerSuite(Images.redis) { - test("ARSET/ARGET/ARLEN/ARCOUNT cover writes, reads, and length") { - withClient { client => - for { - filled <- client.arSet("ar-basic", 0L, "a", "b", "c") - b <- client.arGet[String]("ar-basic", 1L) - missing <- client.arGet[String]("ar-basic", 99L) - len <- client.arLen("ar-basic") - count <- client.arCount("ar-basic") - } yield { - assertEquals(filled, 3L) - assertEquals(b, Some("b")) - assertEquals(missing, None) - assertEquals(len, 3L) - assertEquals(count, 3L) - } - } + clientTest("ARSET/ARGET/ARLEN/ARCOUNT cover writes, reads, and length") { client => + client.arSet("ar-basic", 0L, "a", "b", "c").is(3L) >> + client.arGet[String]("ar-basic", 1L).is(Some("b")) >> + client.arGet[String]("ar-basic", 99L).is(None) >> + client.arLen("ar-basic").is(3L) >> + client.arCount("ar-basic").is(3L) } - test("ARMSET/ARMGET/ARGETRANGE keep sparse gaps as None") { - withClient { client => - for { - _ <- client.arSet("ar-sparse", 0L, "a", "b", "c") - _ <- client.arMSet("ar-sparse", 10L -> "x", 20L -> "y") - got <- client.arMGet[String]("ar-sparse", 10L, 11L, 20L) - range <- client.arGetRange[String]("ar-sparse", 0L, 2L) - } yield { - assertEquals(got, Vector(Some("x"), None, Some("y"))) - assertEquals(range, Vector(Some("a"), Some("b"), Some("c"))) - } - } + clientTest("ARMSET/ARMGET/ARGETRANGE keep sparse gaps as None") { client => + client.arSet("ar-sparse", 0L, "a", "b", "c") >> + client.arMSet("ar-sparse", 10L -> "x", 20L -> "y") >> + client.arMGet[String]("ar-sparse", 10L, 11L, 20L).is(Vector(Some("x"), None, Some("y"))) >> + client.arGetRange[String]("ar-sparse", 0L, 2L).is(Vector(Some("a"), Some("b"), Some("c"))) } - test("ARRING wraps and ARLASTITEMS reads the most recent items") { - withClient { client => - for { - last <- client.arRing("ar-ring", 3L, "a", "b", "c", "d", "e") - len <- client.arLen("ar-ring") - items <- client.arLastItems[String]("ar-ring", 2L) - rev <- client.arLastItems[String]("ar-ring", 2L, rev = true) - } yield { - assertEquals(last, 1L) - assertEquals(len, 3L) - assertEquals(items, Vector("d", "e")) - assertEquals(rev, Vector("e", "d")) - } - } + clientTest("ARRING wraps and ARLASTITEMS reads the most recent items") { client => + client.arRing("ar-ring", 3L, "a", "b", "c", "d", "e").is(1L) >> + client.arLen("ar-ring").is(3L) >> + client.arLastItems[String]("ar-ring", 2L).is(Vector("d", "e")) >> + client.arLastItems[String]("ar-ring", 2L, rev = true).is(Vector("e", "d")) } - test("ARINSERT/ARNEXT/ARSEEK drive the write cursor") { - withClient { client => - for { - last <- client.arInsert("ar-cursor", "p", "q") - next <- client.arNext("ar-cursor") - seeked <- client.arSeek("ar-cursor", 100L) - absent <- client.arSeek("ar-cursor-absent", 5L) - } yield { - assertEquals(last, 1L) - assertEquals(next, Some(2L)) - assertEquals(seeked, true) - assertEquals(absent, false) - } - } + clientTest("ARINSERT/ARNEXT/ARSEEK drive the write cursor") { client => + client.arInsert("ar-cursor", "p", "q").is(1L) >> + client.arNext("ar-cursor").is(Some(2L)) >> + client.arSeek("ar-cursor", 100L).is(true) >> + client.arSeek("ar-cursor-absent", 5L).is(false) } - test("ARSCAN returns only existing index/value pairs") { - withClient { client => - for { - _ <- client.arSet("ar-scan", 0L, "a") - _ <- client.arSet("ar-scan", 5L, "f") - scan <- client.arScan[String]("ar-scan", 0L, 10L) - } yield assertEquals(scan, Vector(0L -> "a", 5L -> "f")) - } + clientTest("ARSCAN returns only existing index/value pairs") { client => + client.arSet("ar-scan", 0L, "a") >> + client.arSet("ar-scan", 5L, "f") >> + client.arScan[String]("ar-scan", 0L, 10L).is(Vector(0L -> "a", 5L -> "f")) } - test("ARDEL and ARDELRANGE delete by index and by ranges") { - withClient { client => - for { - _ <- client.arSet("ar-del", 0L, "a", "b", "c", "d", "e", "f", "g", "h") - deleted <- client.arDelRange("ar-del", 0L -> 1L, 4L -> 5L) - one <- client.arDel("ar-del", 2L) - scan <- client.arScan[String]("ar-del", 0L, 10L) - } yield { - assertEquals(deleted, 4L) - assertEquals(one, 1L) - assertEquals(scan, Vector(3L -> "d", 6L -> "g", 7L -> "h")) - } - } + clientTest("ARDEL and ARDELRANGE delete by index and by ranges") { client => + client.arSet("ar-del", 0L, "a", "b", "c", "d", "e", "f", "g", "h") >> + client.arDelRange("ar-del", 0L -> 1L, 4L -> 5L).is(4L) >> + client.arDel("ar-del", 2L).is(1L) >> + client.arScan[String]("ar-del", 0L, 10L).is(Vector(3L -> "d", 6L -> "g", 7L -> "h")) } - test("ARGREP matches indices and, WITHVALUES, index/value pairs") { - withClient { client => - for { - _ <- client.arSet("ar-grep", 0L, "apple", "banana", "apricot", "cherry") - glob <- client.arGrep("ar-grep", 0L, 10L)(ArMatch.Glob("ap*")) - withVal <- client.arGrepWithValues[String]("ar-grep", 0L, 10L)(ArMatch.Glob("ap*")) - anded <- client.arGrep("ar-grep", 0L, 10L, combine = ArGrepCombine.And)(ArMatch.Glob("a*"), ArMatch.Glob("*e")) - } yield { - assertEquals(glob, Vector(0L, 2L)) - assertEquals(withVal, Vector(0L -> "apple", 2L -> "apricot")) - assertEquals(anded, Vector(0L)) - } - } + clientTest("ARGREP matches indices and, WITHVALUES, index/value pairs") { client => + client.arSet("ar-grep", 0L, "apple", "banana", "apricot", "cherry") >> + client.arGrep("ar-grep", 0L, 10L)(ArMatch.Glob("ap*")).is(Vector(0L, 2L)) >> + client.arGrepWithValues[String]("ar-grep", 0L, 10L)(ArMatch.Glob("ap*")).is(Vector(0L -> "apple", 2L -> "apricot")) >> + client.arGrep("ar-grep", 0L, 10L, combine = ArGrepCombine.And)(ArMatch.Glob("a*"), ArMatch.Glob("*e")).is(Vector(0L)) } - test("AROP aggregates, bit-combines, and counts over a range") { - withClient { client => - for { - _ <- client.arSet("ar-op", 0L, "10", "20", "30") - sum <- client.arOpSum("ar-op", 0L, 2L) - min <- client.arOpMin("ar-op", 0L, 2L) - max <- client.arOpMax("ar-op", 0L, 2L) - and <- client.arOpAnd("ar-op", 0L, 2L) - or <- client.arOpOr("ar-op", 0L, 2L) - xor <- client.arOpXor("ar-op", 0L, 2L) - used <- client.arOpUsed("ar-op", 0L, 2L) - atMost <- client.arOpMatch("ar-op", 0L, 2L, "20") - } yield { - assertEquals(sum, Some(60.0)) - assertEquals(min, Some(10.0)) - assertEquals(max, Some(30.0)) - assertEquals(and, Some(0L)) - assertEquals(or, Some(30L)) - assertEquals(xor, Some(0L)) - assertEquals(used, 3L) - assertEquals(atMost, 1L) - } - } + clientTest("AROP aggregates, bit-combines, and counts over a range") { client => + client.arSet("ar-op", 0L, "10", "20", "30") >> + client.arOpSum("ar-op", 0L, 2L).is(Some(60.0)) >> + client.arOpMin("ar-op", 0L, 2L).is(Some(10.0)) >> + client.arOpMax("ar-op", 0L, 2L).is(Some(30.0)) >> + client.arOpAnd("ar-op", 0L, 2L).is(Some(0L)) >> + client.arOpOr("ar-op", 0L, 2L).is(Some(30L)) >> + client.arOpXor("ar-op", 0L, 2L).is(Some(0L)) >> + client.arOpUsed("ar-op", 0L, 2L).is(3L) >> + client.arOpMatch("ar-op", 0L, 2L, "20").is(1L) } - test("ARINFO and ARINFO FULL report metadata") { - withClient { client => - for { - _ <- client.arSet("ar-info", 0L, "a", "b", "c") - _ <- client.arMSet("ar-info", 100L -> "z") - info <- client.arInfo("ar-info") - full <- client.arInfoFull("ar-info") - } yield { - assertEquals(info.count, 4L) - assertEquals(info.len, 101L) - assertEquals(full.count, 4L) - assert(full.sparseSlices.forall(_ >= 0L)) - } + clientTest("ARINFO and ARINFO FULL report metadata") { client => + for { + _ <- client.arSet("ar-info", 0L, "a", "b", "c") + _ <- client.arMSet("ar-info", 100L -> "z") + info <- client.arInfo("ar-info") + full <- client.arInfoFull("ar-info") + } yield { + assertEquals(info.count, 4L) + assertEquals(info.len, 101L) + assertEquals(full.count, 4L) + assert(full.sparseSlices.forall(_ >= 0L)) } } } diff --git a/integration-tests/shared/src/test/scala/sage/integration/commands/RedisHashFieldExpirySuite.scala b/integration-tests/shared/src/test/scala/sage/integration/commands/RedisHashFieldExpirySuite.scala deleted file mode 100644 index 67c6ea13..00000000 --- a/integration-tests/shared/src/test/scala/sage/integration/commands/RedisHashFieldExpirySuite.scala +++ /dev/null @@ -1,120 +0,0 @@ -package sage.integration.commands - -import java.time.Instant - -import scala.concurrent.duration.* - -import kyo.compat.* - -import sage.commands.* -import sage.integration.{Images, ServerSuite} -import sage.integration.Ttls.{expires, expiresWithin} - -/** - * Hash field-expiration is a Redis-only family (absent in Valkey 8.1), so it has no cross-server counterpart. - */ -class RedisHashFieldExpirySuite extends ServerSuite(Images.redis) { - - test("HEXPIRE/HTTL/HPERSIST set, read, and clear per-field TTLs") { - withClient { client => - for { - _ <- client.hSet("hfe-ttl", ("a", "1"), ("b", "2")) - set <- client.hExpire("hfe-ttl", 100.seconds)("a", "missing") - ttl <- client.hTtl("hfe-ttl")("a", "b", "missing") - persist <- client.hPersist("hfe-ttl")("a", "b") - afterTtl <- client.hTtl("hfe-ttl")("a") - } yield { - assertEquals(set, Vector(FieldExpiry.Updated, FieldExpiry.NoField)) - assert(expiresWithin(ttl(0), 100.seconds)) - assertEquals(ttl(1), FieldTtl.NoExpiry) - assertEquals(ttl(2), FieldTtl.NoField) - assertEquals(persist, Vector(FieldPersist.Persisted, FieldPersist.NoExpiry)) - assertEquals(afterTtl, Vector(FieldTtl.NoExpiry)) - } - } - } - - test("a field-TTL command on a missing key reports NoField per field, not a null") { - withClient { client => - for { - result <- client.hExpire("hfe-missing", 100.seconds)("a", "b") - } yield assertEquals(result, Vector(FieldExpiry.NoField, FieldExpiry.NoField)) - } - } - - test("HEXPIREAT pins an absolute deadline that HEXPIRETIME reads back") { - withClient { client => - val at = Instant.ofEpochSecond(Instant.now().getEpochSecond + 3600) - for { - _ <- client.hSet("hfe-at", ("a", "1")) - set <- client.hExpireAt("hfe-at", at)("a") - time <- client.hExpireTime("hfe-at")("a") - } yield { - assertEquals(set, Vector(FieldExpiry.Updated)) - assertEquals(time, Vector(FieldExpiryTime.At(at))) - } - } - } - - test("HPTTL and HPEXPIRETIME read the millisecond-precision TTL and absolute deadline") { - withClient { client => - val at = Instant.ofEpochSecond(Instant.now().getEpochSecond + 3600) - for { - _ <- client.hSet("hfe-px", ("a", "1")) - _ <- client.hExpireAt("hfe-px", at)("a") - ttl <- client.hpTtl("hfe-px")("a", "missing") - time <- client.hpExpireTime("hfe-px")("a") - } yield { - assert(expires(ttl(0))) - assertEquals(ttl(1), FieldTtl.NoField) - assertEquals(time, Vector(FieldExpiryTime.At(at))) - } - } - } - - test("HGETDEL returns field values and removes them") { - withClient { client => - for { - _ <- client.hSet("hfe-getdel", ("a", "1"), ("b", "2")) - got <- client.hGetDel[String, String]("hfe-getdel")("a", "missing") - left <- client.hGetAll[String, String]("hfe-getdel") - } yield { - assertEquals(got, Vector(Some("1"), None)) - assertEquals(left, Map("b" -> "2")) - } - } - } - - test("HGETEX returns field values and sets their TTL") { - withClient { client => - for { - _ <- client.hSet("hfe-getex", ("a", "1")) - got <- client.hGetEx[String, String]("hfe-getex", GetExpiry.In(100.seconds))("a") - ttl <- client.hTtl("hfe-getex")("a") - } yield { - assertEquals(got, Vector(Some("1"))) - assert(expires(ttl(0))) - } - } - } - - test("HSETEX sets fields with a shared TTL and honors FNX/FXX") { - withClient { client => - for { - created <- client.hSetEx("hfe-setex", SetExpiry.In(100.seconds), HSetExCondition.IfNoneExist)(("a", "1"), ("b", "2")) - ttl <- client.hTtl("hfe-setex")("a") - blocked <- client.hSetEx("hfe-setex", condition = HSetExCondition.IfNoneExist)(("a", "9")) - aStill <- client.hGet[String, String]("hfe-setex", "a") - updated <- client.hSetEx("hfe-setex", SetExpiry.KeepTtl, HSetExCondition.IfAllExist)(("a", "10")) - aNew <- client.hGet[String, String]("hfe-setex", "a") - } yield { - assertEquals(created, true) - assert(expires(ttl(0))) - assertEquals(blocked, false) - assertEquals(aStill, Some("1")) - assertEquals(updated, true) - assertEquals(aNew, Some("10")) - } - } - } -} diff --git a/integration-tests/shared/src/test/scala/sage/integration/commands/RedisJsonExtrasSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/commands/RedisJsonExtrasSuite.scala deleted file mode 100644 index 11c463f6..00000000 --- a/integration-tests/shared/src/test/scala/sage/integration/commands/RedisJsonExtrasSuite.scala +++ /dev/null @@ -1,29 +0,0 @@ -package sage.integration.commands - -import kyo.compat.* - -import sage.commands.JsonPath -import sage.integration.{Images, ServerSuite} - -/** - * JSON commands present only on Redis (RedisJSON), not on Valkey Bundle 9.1.0: `JSON.MERGE` (RFC 7386). Mirrors the Redis-only extras suites - * of other families (ADR-0026); the shared [[JsonSuite]] stays backend-symmetric. - */ -class RedisJsonExtrasSuite extends ServerSuite(Images.redis) { - - test("JSON.MERGE updates existing and creates new members") { - withClient { client => - for { - _ <- client.jsonSet("merge", JsonPath.root, """{"a":1,"b":2}""") - _ <- client.jsonMerge("merge", JsonPath.root, """{"b":20,"c":3}""") - a <- client.jsonGet[String]("merge", JsonPath("$.a")) - b <- client.jsonGet[String]("merge", JsonPath("$.b")) - c <- client.jsonGet[String]("merge", JsonPath("$.c")) - } yield { - assert(a.exists(_.contains("1"))) - assert(b.exists(_.contains("20"))) - assert(c.exists(_.contains("3"))) - } - } - } -} diff --git a/integration-tests/shared/src/test/scala/sage/integration/commands/RedisStreamExtrasSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/commands/RedisStreamExtrasSuite.scala deleted file mode 100644 index 473f60fe..00000000 --- a/integration-tests/shared/src/test/scala/sage/integration/commands/RedisStreamExtrasSuite.scala +++ /dev/null @@ -1,73 +0,0 @@ -package sage.integration.commands - -import scala.concurrent.duration.* - -import kyo.compat.* - -import sage.commands.* -import sage.integration.{Images, ServerSuite} - -/** - * XDELEX/XACKDEL (8.2), XNACK (8.8) and XCFGSET are Redis-only stream commands absent in Valkey, so they have no cross-server counterpart. - */ -class RedisStreamExtrasSuite extends ServerSuite(Images.redis) { - - test("XCFGSET sets per-stream idempotent-message-processing config") { - withClient { client => - for { - _ <- client.xAdd("sx-cfg", XAddId.Explicit(StreamId(1L, 0L)))(("f", "1")) - both <- client.xCfgSet("sx-cfg", idmpDuration = Some(1.hour), idmpMaxSize = Some(100L)) - one <- client.xCfgSet("sx-cfg", idmpMaxSize = Some(50L)) - } yield { - assertEquals(both, ()) - assertEquals(one, ()) - } - } - } - - test("XDELEX reports per-id deletion status") { - withClient { client => - for { - _ <- client.xAdd("sx-delex", XAddId.Explicit(StreamId(1L, 0L)))(("f", "1")) - _ <- client.xAdd("sx-delex", XAddId.Explicit(StreamId(2L, 0L)))(("f", "2")) - status <- client.xDelEx("sx-delex")(StreamId(1L, 0L), StreamId(9L, 0L)) - len <- client.xLen("sx-delex") - } yield { - assertEquals(status, Vector(StreamEntryDeletion.Deleted, StreamEntryDeletion.NotFound)) - assertEquals(len, 1L) - } - } - } - - test("XACKDEL acknowledges and deletes in one step") { - withClient { client => - for { - _ <- client.xAdd("sx-ackdel", XAddId.Explicit(StreamId(1L, 0L)))(("f", "1")) - _ <- client.xGroupCreate("sx-ackdel", "g", GroupStartId.At(StreamId(0L, 0L))) - _ <- client.xReadGroup[String, String]("g", "c1")(("sx-ackdel", GroupReadId.New))() - status <- client.xAckDel("sx-ackdel", "g")(StreamId(1L, 0L)) - len <- client.xLen("sx-ackdel") - pend <- client.xPending("sx-ackdel", "g") - } yield { - assertEquals(status, Vector(StreamEntryDeletion.Deleted)) - assertEquals(len, 0L) - assertEquals(pend.total, 0L) - } - } - } - - test("XNACK releases a pending entry back to the group") { - withClient { client => - for { - _ <- client.xAdd("sx-nack", XAddId.Explicit(StreamId(1L, 0L)))(("f", "1")) - _ <- client.xGroupCreate("sx-nack", "g", GroupStartId.At(StreamId(0L, 0L))) - _ <- client.xReadGroup[String, String]("g", "c1")(("sx-nack", GroupReadId.New))() - released <- client.xNack("sx-nack", "g", NackMode.Fail)(StreamId(1L, 0L))() - pend <- client.xPending("sx-nack", "g") - } yield { - assertEquals(released, 1L) - assertEquals(pend.total, 1L) - } - } - } -} diff --git a/integration-tests/shared/src/test/scala/sage/integration/commands/RedisStringExtrasSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/commands/RedisStringExtrasSuite.scala deleted file mode 100644 index 5b7a964c..00000000 --- a/integration-tests/shared/src/test/scala/sage/integration/commands/RedisStringExtrasSuite.scala +++ /dev/null @@ -1,86 +0,0 @@ -package sage.integration.commands - -import scala.concurrent.duration.* - -import kyo.compat.* - -import sage.commands.* -import sage.integration.{Images, ServerSuite} -import sage.integration.Ttls.expires - -/** - * DIGEST/DELEX/MSETEX/INCREX are Redis-only string commands (absent in Valkey 8.1), so they have no cross-server counterpart. - */ -class RedisStringExtrasSuite extends ServerSuite(Images.redis) { - - test("DIGEST returns a stable hex digest, None for a missing key") { - withClient { client => - for { - _ <- client.set("sx-digest", "hello") - d1 <- client.digest("sx-digest") - d2 <- client.digest("sx-digest") - none <- client.digest("sx-digest-missing") - } yield { - assert(d1.exists(_.nonEmpty)) - assertEquals(d1, d2) - assertEquals(none, None) - } - } - } - - test("DELEX deletes only when the value or digest condition matches") { - withClient { client => - for { - _ <- client.set("sx-delex", "v1") - noMatch <- client.delex("sx-delex", DelexCondition.IfEq("other")) - present <- client.exists("sx-delex") - digest <- client.digest("sx-delex") - notNe <- client.delex[String]("sx-delex", DelexCondition.IfDigestNe(digest.get)) - matched <- client.delex("sx-delex", DelexCondition.IfEq("v1")) - gone <- client.exists("sx-delex") - } yield { - assertEquals(noMatch, false) - assertEquals(present, 1L) - assertEquals(notNe, false) - assertEquals(matched, true) - assertEquals(gone, 0L) - } - } - } - - test("MSETEX sets multiple keys with a shared TTL and respects NX") { - withClient { client => - for { - set <- client.msetEx(expiry = SetExpiry.In(100.seconds))(("sx-ms-a", "1"), ("sx-ms-b", "2")) - a <- client.get[String]("sx-ms-a") - ttl <- client.ttl("sx-ms-a") - blocked <- client.msetEx(condition = SetCondition.IfNotExists)(("sx-ms-a", "9")) - aStill <- client.get[String]("sx-ms-a") - } yield { - assertEquals(set, true) - assertEquals(a, Some("1")) - assert(expires(ttl)) - assertEquals(blocked, false) - assertEquals(aStill, Some("1")) - } - } - } - - test("INCREX increments with expiry, saturating bounds, and rejects when out of range") { - withClient { client => - for { - first <- client.increxBy("sx-incr", 5L, expiry = IncrExpiry.In(100.seconds)) - ttl <- client.ttl("sx-incr") - capped <- client.increxBy("sx-incr", 100L, saturate = true, upperBound = Some(10L)) - reject <- client.increxBy("sx-incr", 100L, upperBound = Some(10L)) - viaFlt <- client.increxByFloat("sx-incr-f", 1.5) - } yield { - assertEquals(first, IncrExResult(5L, 5L)) - assert(expires(ttl)) - assertEquals(capped, IncrExResult(10L, 5L)) - assertEquals(reject, IncrExResult(10L, 0L)) - assertEquals(viaFlt, IncrExResult(1.5, 1.5)) - } - } - } -} diff --git a/integration-tests/shared/src/test/scala/sage/integration/commands/ScriptingSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/commands/ScriptingSuite.scala index fa612d9c..1544b399 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/commands/ScriptingSuite.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/commands/ScriptingSuite.scala @@ -3,63 +3,37 @@ package sage.integration.commands import kyo.compat.* import sage.commands.FlushMode -import sage.integration.{Images, ServerSuite} +import sage.integration.BothServersSuite import sage.protocol.Frame -abstract class ScriptingSuite(image: String) extends ServerSuite(image) { +class ScriptingSuite extends BothServersSuite { private val absent = "0" * 40 - test("EVAL returns the raw reply and passes keys and args to the script") { - withClient { client => - for { - one <- client.eval("return 1") - keys <- client.eval("return #KEYS", Seq("a", "b")) - arg <- client.eval("return redis.call('set', KEYS[1], ARGV[1])", Seq("eval-k"), Seq("v")) - value <- client.get[String]("eval-k") - } yield { - assertEquals(one, Frame.Integer(1L)) - assertEquals(keys, Frame.Integer(2L)) - assertEquals(arg, Frame.SimpleString("OK")) - assertEquals(value, Some("v")) - } - } + clientTest("EVAL returns the raw reply and passes keys and args to the script") { client => + client.eval("return 1").is(Frame.Integer(1L)) >> + client.eval("return #KEYS", Seq("a", "b")).is(Frame.Integer(2L)) >> + client.eval("return redis.call('set', KEYS[1], ARGV[1])", Seq("eval-k"), Seq("v")).is(Frame.SimpleString("OK")) >> + client.get[String]("eval-k").is(Some("v")) } - test("EVAL_RO runs a read-only script") { - withClient(client => client.evalRo("return 42").map(assertEquals(_, Frame.Integer(42L)))) - } + clientTest("EVAL_RO runs a read-only script")(client => client.evalRo("return 42").is(Frame.Integer(42L))) - test("SCRIPT LOAD returns a sha that EVALSHA then runs; SCRIPT EXISTS reports per-sha presence") { - withClient { client => - for { - sha <- client.scriptLoad("return 7") - ran <- client.evalSha(sha) - ranRo <- client.evalShaRo(sha) - exists <- client.scriptExists(sha, absent) - } yield { - assertEquals(sha.length, 40) - assertEquals(ran, Frame.Integer(7L)) - assertEquals(ranRo, Frame.Integer(7L)) - assertEquals(exists, Vector(true, false)) - } - } + clientTest("SCRIPT LOAD returns a sha that EVALSHA then runs; SCRIPT EXISTS reports per-sha presence") { client => + for { + sha <- client.scriptLoad("return 7") + _ <- client.evalSha(sha).is(Frame.Integer(7L)) + _ <- client.evalShaRo(sha).is(Frame.Integer(7L)) + _ <- client.scriptExists(sha, absent).is(Vector(true, false)) + } yield assertEquals(sha.length, 40) } - test("SCRIPT FLUSH clears the cache") { - withClient { client => - for { - sha <- client.scriptLoad("return 1") - before <- client.scriptExists(sha) - _ <- client.scriptFlush(Some(FlushMode.Sync)) - after <- client.scriptExists(sha) - } yield { - assertEquals(before, Vector(true)) - assertEquals(after, Vector(false)) - } - } + clientTest("SCRIPT FLUSH clears the cache") { client => + for { + sha <- client.scriptLoad("return 1") + _ <- client.scriptExists(sha).is(Vector(true)) + _ <- client.scriptFlush(Some(FlushMode.Sync)) + _ <- client.scriptExists(sha).is(Vector(false)) + } yield () } } - -class RedisScriptingSuite extends ScriptingSuite(Images.redis) -class ValkeyScriptingSuite extends ScriptingSuite(Images.valkey) diff --git a/integration-tests/shared/src/test/scala/sage/integration/commands/ServerAdminSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/commands/ServerAdminSuite.scala index 8fdc8d45..60babdc4 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/commands/ServerAdminSuite.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/commands/ServerAdminSuite.scala @@ -4,106 +4,78 @@ import scala.concurrent.duration.* import kyo.compat.* -import sage.commands.Role -import sage.integration.{Images, ServerSuite} +import sage.commands.{CommandLogType, Role} +import sage.integration.BothServersSuite -abstract class ServerAdminSuite(image: String) extends ServerSuite(image) { +class ServerAdminSuite extends BothServersSuite { - test("CONFIG GET and SET read and write a parameter") { - withClient { client => - for { - before <- client.configGet("maxmemory") - _ <- client.configSet(("maxmemory", "100mb")) - after <- client.configGet("maxmemory") - _ <- client.configSet(("maxmemory", before.getOrElse("maxmemory", "0"))) - } yield { - assert(before.contains("maxmemory")) - assertEquals(after.get("maxmemory"), Some("104857600")) - } + clientTest("CONFIG GET and SET read and write a parameter") { client => + for { + before <- client.configGet("maxmemory") + _ <- client.configSet(("maxmemory", "100mb")) + after <- client.configGet("maxmemory") + _ <- client.configSet(("maxmemory", before.getOrElse("maxmemory", "0"))) + } yield { + assert(before.contains("maxmemory")) + assertEquals(after.get("maxmemory"), Some("104857600")) } } - test("DBSIZE, FLUSHDB, ECHO, TIME") { - withClient { client => - for { - _ <- client.set("admin-k", "v") - size <- client.dbSize - _ <- client.flushDb() - empty <- client.dbSize - echo <- client.echo("ping") - time <- client.time - } yield { - assert(size >= 1L) - assertEquals(empty, 0L) - assertEquals(echo, "ping") - assert(time.getEpochSecond > 1_000_000_000L) - } - } + clientTest("DBSIZE, FLUSHDB, ECHO, TIME") { client => + client.set("admin-k", "v") >> + client.dbSize.satisfies(_ >= 1L) >> + client.flushDb() >> + client.dbSize.is(0L) >> + client.echo("ping").is("ping") >> + client.time.satisfies(_.getEpochSecond > 1_000_000_000L) } - test("ROLE reports a standalone server as master") { - withClient { client => - client.role.map { - case Role.Master(_, _) => () - case other => fail(s"expected master, got $other") - } + clientTest("ROLE reports a standalone server as master") { client => + client.role.map { + case Role.Master(_, _) => () + case other => fail(s"expected master, got $other") } } - test("CLIENT ID/GETNAME/INFO/LIST and WAIT") { - withClient { client => - for { - id <- client.clientId - name <- client.clientGetName - info <- client.clientInfo - list <- client.clientList - waited <- client.waitReplicas(0L, 100.millis) - } yield { - assert(id > 0L) - assertEquals(name, "") - assert(info.contains("id=")) - assert(list.contains("addr=")) - assertEquals(waited, 0L) - } - } + clientTest("CLIENT ID/GETNAME/INFO/LIST and WAIT") { client => + client.clientId.satisfies(_ > 0L) >> + client.clientGetName.is("") >> + client.clientInfo.satisfies(_.contains("id=")) >> + client.clientList.satisfies(_.contains("addr=")) >> + client.waitReplicas(0L, 100.millis).is(0L) } - test("COMMAND COUNT/INFO/GETKEYS") { - withClient { client => - for { - count <- client.commandCount - infos <- client.commandInfo("get", "set") - keys <- client.commandGetKeys("SET", "k", "v") - } yield { - assert(count > 100L) - assertEquals(infos.map(_.name).toSet, Set("get", "set")) - assertEquals(keys, Vector("k")) - } - } + clientTest("COMMAND COUNT/INFO/GETKEYS") { client => + client.commandCount.satisfies(_ > 100L) >> + client.commandInfo("get", "set").map(_.map(_.name).toSet).is(Set("get", "set")) >> + client.commandGetKeys("SET", "k", "v").is(Vector("k")) } - test("MEMORY USAGE, SLOWLOG, LATENCY, ACL reads") { - withClient { client => - for { - _ <- client.set("mem-k", "value") - usage <- client.memoryUsage("mem-k") - _ <- client.slowLogReset - len <- client.slowLogLen - latest <- client.latencyLatest - who <- client.aclWhoAmI - users <- client.aclUsers - user <- client.aclGetUser("default") - } yield { - assert(usage.exists(_ > 0L)) - assertEquals(len, 0L) - assert(latest.isEmpty || latest.nonEmpty) - assertEquals(who, "default") - assert(users.contains("default")) - assert(user.exists(_.flags.nonEmpty)) - } - } + clientTest("MEMORY USAGE, SLOWLOG, LATENCY, ACL reads") { client => + client.set("mem-k", "value") >> + client.memoryUsage("mem-k").satisfies(_.exists(_ > 0L)) >> + client.slowLogReset >> + client.slowLogLen.is(0L) >> + client.latencyLatest >> + client.aclWhoAmI.is("default") >> + client.aclUsers.satisfies(_.contains("default")) >> + client.aclGetUser("default").satisfies(_.exists(_.flags.nonEmpty)) } -} -class RedisServerAdminSuite extends ServerAdminSuite(Images.redis) -class ValkeyServerAdminSuite extends ServerAdminSuite(Images.valkey) + // COMMANDLOG exists only on Valkey. + valkeyTest("COMMANDLOG GET/LEN/RESET over the slow log") { client => + client.configSet(("slowlog-log-slower-than", "0")) >> + client.commandLogReset(CommandLogType.Slow) >> + client.get[String]("cl-probe") >> + client.commandLogLen(CommandLogType.Slow).satisfies(_ > 0L) >> + client.commandLogGet(5L, CommandLogType.Slow).satisfies(recent => recent.nonEmpty && recent.forall(_.command.nonEmpty)) >> + client.configSet(("slowlog-log-slower-than", "10000")) >> + client.commandLogReset(CommandLogType.Slow) >> + client.commandLogLen(CommandLogType.Slow).is(0L) + } + + valkeyTest("COMMANDLOG LEN works for the large-request and large-reply types") { client => + client.commandLogLen(CommandLogType.LargeRequest).satisfies(_ >= 0L) >> + client.commandLogLen(CommandLogType.LargeReply).satisfies(_ >= 0L) + } +} diff --git a/integration-tests/shared/src/test/scala/sage/integration/commands/SetsSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/commands/SetsSuite.scala index 8a669cc8..40dcd2b9 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/commands/SetsSuite.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/commands/SetsSuite.scala @@ -3,112 +3,61 @@ package sage.integration.commands import kyo.compat.* import sage.commands.ScanCursor -import sage.integration.{Images, ServerSuite} - -abstract class SetsSuite(image: String) extends ServerSuite(image) { - - test("SADD SCARD SMEMBERS SISMEMBER SMISMEMBER and SREM manage membership") { - withClient { client => - for { - added <- client.sAdd("set-basic", "a", "b", "c") - dup <- client.sAdd("set-basic", "a") - card <- client.sCard("set-basic") - members <- client.sMembers[String]("set-basic") - isA <- client.sIsMember("set-basic", "a") - isZ <- client.sIsMember("set-basic", "z") - multi <- client.sMisMember("set-basic", "a", "z", "c") - removed <- client.sRem("set-basic", "a", "z") - after <- client.sMembers[String]("set-basic") - } yield { - assertEquals(added, 3L) - assertEquals(dup, 0L) - assertEquals(card, 3L) - assertEquals(members, Set("a", "b", "c")) - assertEquals(isA, true) - assertEquals(isZ, false) - assertEquals(multi, Vector(true, false, true)) - assertEquals(removed, 1L) - assertEquals(after, Set("b", "c")) - } - } +import sage.integration.BothServersSuite + +class SetsSuite extends BothServersSuite { + + clientTest("SADD SCARD SMEMBERS SISMEMBER SMISMEMBER and SREM manage membership") { client => + client.sAdd("set-basic", "a", "b", "c").is(3L) >> + client.sAdd("set-basic", "a").is(0L) >> + client.sCard("set-basic").is(3L) >> + client.sMembers[String]("set-basic").is(Set("a", "b", "c")) >> + client.sIsMember("set-basic", "a").is(true) >> + client.sIsMember("set-basic", "z").is(false) >> + client.sMisMember("set-basic", "a", "z", "c").is(Vector(true, false, true)) >> + client.sRem("set-basic", "a", "z").is(1L) >> + client.sMembers[String]("set-basic").is(Set("b", "c")) } - test("SPOP and SRANDMEMBER draw members, with and without a count") { - withClient { client => - for { - _ <- client.sAdd("set-draw", "a", "b", "c", "d") - popOne <- client.sPop[String]("set-draw") - popTwo <- client.sPopCount[String]("set-draw", 2L) - rndOne <- client.sRandMember[String]("set-draw") - rndDup <- client.sRandMemberCount[String]("set-draw", -5L) - emptyPop <- client.sPop[String]("set-missing") - } yield { - assert(popOne.exists(Set("a", "b", "c", "d"))) - assertEquals(popTwo.size, 2) - assert(popTwo.subsetOf(Set("a", "b", "c", "d"))) - assert(rndOne.isDefined) - assertEquals(rndDup.size, 5) - assertEquals(emptyPop, None) - } + clientTest("SPOP and SRANDMEMBER draw members, with and without a count") { client => + for { + _ <- client.sAdd("set-draw", "a", "b", "c", "d") + _ <- client.sPop[String]("set-draw").satisfies(_.exists(Set("a", "b", "c", "d"))) + popTwo <- client.sPopCount[String]("set-draw", 2L) + _ <- client.sRandMember[String]("set-draw").satisfies(_.isDefined) + _ <- client.sRandMemberCount[String]("set-draw", -5L).map(_.size).is(5) + _ <- client.sPop[String]("set-missing").is(None) + } yield { + assertEquals(popTwo.size, 2) + assert(popTwo.subsetOf(Set("a", "b", "c", "d"))) } } - test("SMOVE relocates a member between sets") { - withClient { client => - for { - _ <- client.sAdd("set-src", "x", "y") - _ <- client.sAdd("set-dst", "z") - moved <- client.sMove("set-src", "set-dst", "x") - absent <- client.sMove("set-src", "set-dst", "nope") - src <- client.sMembers[String]("set-src") - dst <- client.sMembers[String]("set-dst") - } yield { - assertEquals(moved, true) - assertEquals(absent, false) - assertEquals(src, Set("y")) - assertEquals(dst, Set("x", "z")) - } - } + clientTest("SMOVE relocates a member between sets") { client => + client.sAdd("set-src", "x", "y") >> + client.sAdd("set-dst", "z") >> + client.sMove("set-src", "set-dst", "x").is(true) >> + client.sMove("set-src", "set-dst", "nope").is(false) >> + client.sMembers[String]("set-src").is(Set("y")) >> + client.sMembers[String]("set-dst").is(Set("x", "z")) } - test("SDIFF SINTER SUNION and their STORE forms combine sets, SINTERCARD counts") { - withClient { client => - for { - _ <- client.sAdd("ops-a", "1", "2", "3") - _ <- client.sAdd("ops-b", "2", "3", "4") - diff <- client.sDiff[String]("ops-a", "ops-b") - inter <- client.sInter[String]("ops-a", "ops-b") - union <- client.sUnion[String]("ops-a", "ops-b") - card <- client.sInterCard("ops-a", "ops-b")() - cardLim <- client.sInterCard("ops-a", "ops-b")(limit = Some(1L)) - diffN <- client.sDiffStore("ops-diff", "ops-a", "ops-b") - interN <- client.sInterStore("ops-inter", "ops-a", "ops-b") - unionN <- client.sUnionStore("ops-union", "ops-a", "ops-b") - stored <- client.sMembers[String]("ops-union") - } yield { - assertEquals(diff, Set("1")) - assertEquals(inter, Set("2", "3")) - assertEquals(union, Set("1", "2", "3", "4")) - assertEquals(card, 2L) - assertEquals(cardLim, 1L) - assertEquals(diffN, 1L) - assertEquals(interN, 2L) - assertEquals(unionN, 4L) - assertEquals(stored, Set("1", "2", "3", "4")) - } - } + clientTest("SDIFF SINTER SUNION and their STORE forms combine sets, SINTERCARD counts") { client => + client.sAdd("ops-a", "1", "2", "3") >> + client.sAdd("ops-b", "2", "3", "4") >> + client.sDiff[String]("ops-a", "ops-b").is(Set("1")) >> + client.sInter[String]("ops-a", "ops-b").is(Set("2", "3")) >> + client.sUnion[String]("ops-a", "ops-b").is(Set("1", "2", "3", "4")) >> + client.sInterCard("ops-a", "ops-b")().is(2L) >> + client.sInterCard("ops-a", "ops-b")(limit = Some(1L)).is(1L) >> + client.sDiffStore("ops-diff", "ops-a", "ops-b").is(1L) >> + client.sInterStore("ops-inter", "ops-a", "ops-b").is(2L) >> + client.sUnionStore("ops-union", "ops-a", "ops-b").is(4L) >> + client.sMembers[String]("ops-union").is(Set("1", "2", "3", "4")) } - test("SSCAN streams members") { - withClient { client => - for { - _ <- client.sAdd("set-scan", "a", "b", "c") - page <- client.sScan[String]("set-scan", ScanCursor.start) - } yield assertEquals(page.items.toSet, Set("a", "b", "c")) - } + clientTest("SSCAN streams members") { client => + client.sAdd("set-scan", "a", "b", "c") >> + client.sScan[String]("set-scan", ScanCursor.start).map(_.items.toSet).is(Set("a", "b", "c")) } } - -class RedisSetsSuite extends SetsSuite(Images.redis) - -class ValkeySetsSuite extends SetsSuite(Images.valkey) diff --git a/integration-tests/shared/src/test/scala/sage/integration/commands/SortedSetsSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/commands/SortedSetsSuite.scala index 8fee6378..2ed75343 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/commands/SortedSetsSuite.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/commands/SortedSetsSuite.scala @@ -5,269 +5,138 @@ import scala.concurrent.duration.* import kyo.compat.* import sage.commands.{Aggregate, BlockTimeout, LexBoundary, Limit, MinMax, ScanCursor, ScoreBoundary, ZAddCondition, ZRange} -import sage.integration.{Images, ServerSuite} +import sage.integration.BothServersSuite -abstract class SortedSetsSuite(image: String) extends ServerSuite(image) { +class SortedSetsSuite extends BothServersSuite { - test("ZADD ZCARD ZSCORE ZMSCORE ZINCRBY and ZREM manage scored members") { - withClient { client => - for { - added <- client.zAdd("zset-basic")(("a", 1.0), ("b", 2.0), ("c", 3.0)) - card <- client.zCard("zset-basic") - score <- client.zScore[String]("zset-basic", "b") - scores <- client.zMScore[String]("zset-basic", "a", "missing", "c") - bumped <- client.zIncrBy("zset-basic", "a", 1.5) - removed <- client.zRem("zset-basic", "c", "missing") - } yield { - assertEquals(added, 3L) - assertEquals(card, 3L) - assertEquals(score, Some(2.0)) - assertEquals(scores, Vector(Some(1.0), None, Some(3.0))) - assertEquals(bumped, 2.5) - assertEquals(removed, 1L) - } - } + clientTest("ZADD ZCARD ZSCORE ZMSCORE ZINCRBY and ZREM manage scored members") { client => + client.zAdd("zset-basic")(("a", 1.0), ("b", 2.0), ("c", 3.0)).is(3L) >> + client.zCard("zset-basic").is(3L) >> + client.zScore[String]("zset-basic", "b").is(Some(2.0)) >> + client.zMScore[String]("zset-basic", "a", "missing", "c").is(Vector(Some(1.0), None, Some(3.0))) >> + client.zIncrBy("zset-basic", "a", 1.5).is(2.5) >> + client.zRem("zset-basic", "c", "missing").is(1L) } - test("ZADD conditions gate writes and ZADD INCR returns the new score or None") { - withClient { client => - for { - _ <- client.zAdd("zset-cond")(("a", 5.0)) - skipNew <- client.zAdd("zset-cond", ZAddCondition.IfExists)(("b", 1.0)) - addedNew <- client.zAdd("zset-cond", ZAddCondition.IfNotExists)(("b", 1.0)) - lower <- client.zAddIncr("zset-cond", ZAddCondition.IfExists)("a", -1.0) - skipped <- client.zAddIncr("zset-cond", ZAddCondition.IfExists)("c", 1.0) - gtNoBump <- client.zAddIncr("zset-cond", ZAddCondition.IfExistsAndGreater)("a", -1.0) - } yield { - assertEquals(skipNew, 0L) - assertEquals(addedNew, 1L) - assertEquals(lower, Some(4.0)) - assertEquals(skipped, None) - assertEquals(gtNoBump, None) - } - } + clientTest("ZADD conditions gate writes and ZADD INCR returns the new score or None") { client => + client.zAdd("zset-cond")(("a", 5.0)) >> + client.zAdd("zset-cond", ZAddCondition.IfExists)(("b", 1.0)).is(0L) >> + client.zAdd("zset-cond", ZAddCondition.IfNotExists)(("b", 1.0)).is(1L) >> + client.zAddIncr("zset-cond", ZAddCondition.IfExists)("a", -1.0).is(Some(4.0)) >> + client.zAddIncr("zset-cond", ZAddCondition.IfExists)("c", 1.0).is(None) >> + client.zAddIncr("zset-cond", ZAddCondition.IfExistsAndGreater)("a", -1.0).is(None) } - test("ZRANK and ZREVRANK report position, with the score on request") { - withClient { client => - for { - _ <- client.zAdd("zset-rank")(("a", 1.0), ("b", 2.0), ("c", 3.0)) - rank <- client.zRank[String]("zset-rank", "b") - rev <- client.zRevRank[String]("zset-rank", "b") - withScore <- client.zRankWithScore[String]("zset-rank", "c") - missing <- client.zRank[String]("zset-rank", "z") - } yield { - assertEquals(rank, Some(1L)) - assertEquals(rev, Some(1L)) - assertEquals(withScore, Some((2L, 3.0))) - assertEquals(missing, None) - } - } + clientTest("ZRANK and ZREVRANK report position, with the score on request") { client => + client.zAdd("zset-rank")(("a", 1.0), ("b", 2.0), ("c", 3.0)) >> + client.zRank[String]("zset-rank", "b").is(Some(1L)) >> + client.zRevRank[String]("zset-rank", "b").is(Some(1L)) >> + client.zRankWithScore[String]("zset-rank", "c").is(Some((2L, 3.0))) >> + client.zRank[String]("zset-rank", "z").is(None) } - test("ZRANGE reads by rank and score, with scores, reversal, and a limit") { - withClient { client => - for { - _ <- client.zAdd("zset-range")(("a", 1.0), ("b", 2.0), ("c", 3.0)) - byRank <- client.zRange[String]("zset-range", ZRange.ByRank(0L, -1L)) - revRank <- client.zRange[String]("zset-range", ZRange.ByRank(0L, -1L, rev = true)) - byScore <- client.zRange[String]("zset-range", ZRange.ByScore(ScoreBoundary.Inclusive(2.0), ScoreBoundary.PosInf)) - revScore <- - client.zRange[String]("zset-range", ZRange.ByScore(ScoreBoundary.Inclusive(1.0), ScoreBoundary.Inclusive(2.0), rev = true)) - withScores <- client.zRangeWithScores[String]("zset-range", ZRange.ByRank(0L, 1L)) - limited <- - client.zRange[String]("zset-range", ZRange.ByScore(ScoreBoundary.NegInf, ScoreBoundary.PosInf, limit = Some(Limit(1L, 1L)))) - } yield { - assertEquals(byRank, Vector("a", "b", "c")) - assertEquals(revRank, Vector("c", "b", "a")) - assertEquals(byScore, Vector("b", "c")) - assertEquals(revScore, Vector("b", "a")) - assertEquals(withScores, Vector("a" -> 1.0, "b" -> 2.0)) - assertEquals(limited, Vector("b")) - } - } + clientTest("ZRANGE reads by rank and score, with scores, reversal, and a limit") { client => + client.zAdd("zset-range")(("a", 1.0), ("b", 2.0), ("c", 3.0)) >> + client.zRange[String]("zset-range", ZRange.ByRank(0L, -1L)).is(Vector("a", "b", "c")) >> + client.zRange[String]("zset-range", ZRange.ByRank(0L, -1L, rev = true)).is(Vector("c", "b", "a")) >> + client.zRange[String]("zset-range", ZRange.ByScore(ScoreBoundary.Inclusive(2.0), ScoreBoundary.PosInf)).is(Vector("b", "c")) >> + client + .zRange[String]("zset-range", ZRange.ByScore(ScoreBoundary.Inclusive(1.0), ScoreBoundary.Inclusive(2.0), rev = true)) + .is(Vector("b", "a")) >> + client.zRangeWithScores[String]("zset-range", ZRange.ByRank(0L, 1L)).is(Vector("a" -> 1.0, "b" -> 2.0)) >> + client + .zRange[String]("zset-range", ZRange.ByScore(ScoreBoundary.NegInf, ScoreBoundary.PosInf, limit = Some(Limit(1L, 1L)))) + .is(Vector("b")) } - test("ZRANGE BYLEX and ZLEXCOUNT operate on equal-score members") { - withClient { client => - for { - _ <- client.zAdd("zset-lex")(("a", 0.0), ("b", 0.0), ("c", 0.0)) - all <- client.zRange[String]("zset-lex", ZRange.ByLex(LexBoundary.Min, LexBoundary.Max)) - range <- client.zRange[String]("zset-lex", ZRange.ByLex(LexBoundary.Inclusive("a"), LexBoundary.Exclusive("c"))) - count <- client.zLexCount[String]("zset-lex", LexBoundary.Min, LexBoundary.Max) - } yield { - assertEquals(all, Vector("a", "b", "c")) - assertEquals(range, Vector("a", "b")) - assertEquals(count, 3L) - } - } + clientTest("ZRANGE BYLEX and ZLEXCOUNT operate on equal-score members") { client => + client.zAdd("zset-lex")(("a", 0.0), ("b", 0.0), ("c", 0.0)) >> + client.zRange[String]("zset-lex", ZRange.ByLex(LexBoundary.Min, LexBoundary.Max)).is(Vector("a", "b", "c")) >> + client.zRange[String]("zset-lex", ZRange.ByLex(LexBoundary.Inclusive("a"), LexBoundary.Exclusive("c"))).is(Vector("a", "b")) >> + client.zLexCount[String]("zset-lex", LexBoundary.Min, LexBoundary.Max).is(3L) } - test("ZRANGESTORE copies a range into another key") { - withClient { client => - for { - _ <- client.zAdd("zrs-src")(("a", 1.0), ("b", 2.0), ("c", 3.0)) - n <- client.zRangeStore[String]("zrs-dst", "zrs-src", ZRange.ByScore(ScoreBoundary.Inclusive(2.0), ScoreBoundary.PosInf)) - stored <- client.zRange[String]("zrs-dst", ZRange.ByRank(0L, -1L)) - } yield { - assertEquals(n, 2L) - assertEquals(stored, Vector("b", "c")) - } - } + clientTest("ZRANGESTORE copies a range into another key") { client => + client.zAdd("zrs-src")(("a", 1.0), ("b", 2.0), ("c", 3.0)) >> + client.zRangeStore[String]("zrs-dst", "zrs-src", ZRange.ByScore(ScoreBoundary.Inclusive(2.0), ScoreBoundary.PosInf)).is(2L) >> + client.zRange[String]("zrs-dst", ZRange.ByRank(0L, -1L)).is(Vector("b", "c")) } - test("ZCOUNT counts members within a score band") { - withClient { client => - for { - _ <- client.zAdd("zset-count")(("a", 1.0), ("b", 2.0), ("c", 3.0)) - all <- client.zCount("zset-count", ScoreBoundary.NegInf, ScoreBoundary.PosInf) - band <- client.zCount("zset-count", ScoreBoundary.Exclusive(1.0), ScoreBoundary.Inclusive(3.0)) - } yield { - assertEquals(all, 3L) - assertEquals(band, 2L) - } - } + clientTest("ZCOUNT counts members within a score band") { client => + client.zAdd("zset-count")(("a", 1.0), ("b", 2.0), ("c", 3.0)) >> + client.zCount("zset-count", ScoreBoundary.NegInf, ScoreBoundary.PosInf).is(3L) >> + client.zCount("zset-count", ScoreBoundary.Exclusive(1.0), ScoreBoundary.Inclusive(3.0)).is(2L) } - test("ZPOPMIN ZPOPMAX and ZMPOP pop by score extreme") { - withClient { client => - for { - _ <- client.zAdd("zset-pop")(("a", 1.0), ("b", 2.0), ("c", 3.0), ("d", 4.0)) - min <- client.zPopMin[String]("zset-pop") - max <- client.zPopMax[String]("zset-pop") - twoMin <- client.zPopMinCount[String]("zset-pop", 2L) - emptied <- client.zPopMin[String]("zset-pop") - _ <- client.zAdd("zset-mpop")(("x", 1.0), ("y", 2.0)) - mpopped <- client.zMpop[String]("zset-empty", "zset-mpop")(MinMax.Min, count = Some(2L)) - } yield { - assertEquals(min, Some(("a", 1.0))) - assertEquals(max, Some(("d", 4.0))) - assertEquals(twoMin, Vector("b" -> 2.0, "c" -> 3.0)) - assertEquals(emptied, None) - assertEquals(mpopped, Some(("zset-mpop", Vector("x" -> 1.0, "y" -> 2.0)))) - } - } + clientTest("ZPOPMIN ZPOPMAX and ZMPOP pop by score extreme") { client => + client.zAdd("zset-pop")(("a", 1.0), ("b", 2.0), ("c", 3.0), ("d", 4.0)) >> + client.zPopMin[String]("zset-pop").is(Some(("a", 1.0))) >> + client.zPopMax[String]("zset-pop").is(Some(("d", 4.0))) >> + client.zPopMinCount[String]("zset-pop", 2L).is(Vector("b" -> 2.0, "c" -> 3.0)) >> + client.zPopMin[String]("zset-pop").is(None) >> + client.zAdd("zset-mpop")(("x", 1.0), ("y", 2.0)) >> + client.zMpop[String]("zset-empty", "zset-mpop")(MinMax.Min, count = Some(2L)).is(Some(("zset-mpop", Vector("x" -> 1.0, "y" -> 2.0)))) } - test("BZPOPMIN and BZMPOP pop present data and time out to None otherwise") { - withClient { client => - for { - _ <- client.zAdd("bz-data")(("a", 1.0), ("b", 2.0)) - popped <- client.bzPopMin[String]("bz-data")(BlockTimeout.After(1.second)) - mpop <- client.bzMpop[String]("bz-empty", "bz-data")(MinMax.Max, BlockTimeout.After(1.second)) - none <- client.bzPopMin[String]("bz-missing")(BlockTimeout.After(100.millis)) - } yield { - assertEquals(popped, Some(("bz-data", "a", 1.0))) - assertEquals(mpop, Some(("bz-data", Vector("b" -> 2.0)))) - assertEquals(none, None) - } - } + clientTest("BZPOPMIN and BZMPOP pop present data and time out to None otherwise") { client => + client.zAdd("bz-data")(("a", 1.0), ("b", 2.0)) >> + client.bzPopMin[String]("bz-data")(BlockTimeout.After(1.second)).is(Some(("bz-data", "a", 1.0))) >> + client.bzMpop[String]("bz-empty", "bz-data")(MinMax.Max, BlockTimeout.After(1.second)).is(Some(("bz-data", Vector("b" -> 2.0)))) >> + client.bzPopMin[String]("bz-missing")(BlockTimeout.After(100.millis)).is(None) } - test("ZRANDMEMBER draws a member, a count, and scored pairs") { - withClient { client => - for { - _ <- client.zAdd("zset-rand")(("a", 1.0), ("b", 2.0), ("c", 3.0)) - one <- client.zRandMember[String]("zset-rand") - dup <- client.zRandMemberCount[String]("zset-rand", -5L) - scored <- client.zRandMemberWithScores[String]("zset-rand", 3L) - } yield { - assert(one.exists(Set("a", "b", "c"))) - assertEquals(dup.size, 5) - assertEquals(scored.toMap, Map("a" -> 1.0, "b" -> 2.0, "c" -> 3.0)) - } - } + clientTest("ZRANDMEMBER draws a member, a count, and scored pairs") { client => + client.zAdd("zset-rand")(("a", 1.0), ("b", 2.0), ("c", 3.0)) >> + client.zRandMember[String]("zset-rand").satisfies(_.exists(Set("a", "b", "c"))) >> + client.zRandMemberCount[String]("zset-rand", -5L).map(_.size).is(5) >> + client.zRandMemberWithScores[String]("zset-rand", 3L).map(_.toMap).is(Map("a" -> 1.0, "b" -> 2.0, "c" -> 3.0)) } - test("ZUNION ZINTER ZDIFF combine sorted sets with weights and aggregation") { - withClient { client => - for { - _ <- client.zAdd("zops-a")(("x", 1.0), ("y", 2.0)) - _ <- client.zAdd("zops-b")(("y", 3.0), ("z", 4.0)) - union <- client.zUnionWithScores[String]("zops-a", "zops-b")() - weighted <- client.zUnionWithScores[String]("zops-a", "zops-b")(weights = Some(Vector(1.0, 2.0))) - maxAgg <- client.zUnionWithScores[String]("zops-a", "zops-b")(aggregate = Aggregate.Max) - inter <- client.zInter[String]("zops-a", "zops-b")() - diff <- client.zDiff[String]("zops-a", "zops-b") - card <- client.zInterCard("zops-a", "zops-b")() - unionN <- client.zUnionStore("zops-union", "zops-a", "zops-b")() - interN <- client.zInterStore("zops-inter", "zops-a", "zops-b")() - diffN <- client.zDiffStore("zops-diff", "zops-a", "zops-b") - } yield { - assertEquals(union.toMap, Map("x" -> 1.0, "y" -> 5.0, "z" -> 4.0)) - assertEquals(weighted.toMap, Map("x" -> 1.0, "y" -> 8.0, "z" -> 8.0)) - assertEquals(maxAgg.toMap, Map("x" -> 1.0, "y" -> 3.0, "z" -> 4.0)) - assertEquals(inter, Vector("y")) - assertEquals(diff, Vector("x")) - assertEquals(card, 1L) - assertEquals(unionN, 3L) - assertEquals(interN, 1L) - assertEquals(diffN, 1L) - } - } + clientTest("ZUNION ZINTER ZDIFF combine sorted sets with weights and aggregation") { client => + client.zAdd("zops-a")(("x", 1.0), ("y", 2.0)) >> + client.zAdd("zops-b")(("y", 3.0), ("z", 4.0)) >> + client.zUnionWithScores[String]("zops-a", "zops-b")().map(_.toMap).is(Map("x" -> 1.0, "y" -> 5.0, "z" -> 4.0)) >> + client + .zUnionWithScores[String]("zops-a", "zops-b")(weights = Some(Vector(1.0, 2.0))) + .map(_.toMap) + .is(Map("x" -> 1.0, "y" -> 8.0, "z" -> 8.0)) >> + client.zUnionWithScores[String]("zops-a", "zops-b")(aggregate = Aggregate.Max).map(_.toMap).is(Map("x" -> 1.0, "y" -> 3.0, "z" -> 4.0)) >> + client.zInter[String]("zops-a", "zops-b")().is(Vector("y")) >> + client.zDiff[String]("zops-a", "zops-b").is(Vector("x")) >> + client.zInterCard("zops-a", "zops-b")().is(1L) >> + client.zUnionStore("zops-union", "zops-a", "zops-b")().is(3L) >> + client.zInterStore("zops-inter", "zops-a", "zops-b")().is(1L) >> + client.zDiffStore("zops-diff", "zops-a", "zops-b").is(1L) } - test("ZREMRANGEBYRANK BYSCORE and BYLEX trim ranges") { - withClient { client => - for { - _ <- client.zAdd("zrem-rank")(("a", 1.0), ("b", 2.0), ("c", 3.0), ("d", 4.0)) - byRank <- client.zRemRangeByRank("zrem-rank", 0L, 0L) - _ <- client.zAdd("zrem-score")(("a", 1.0), ("b", 2.0), ("c", 3.0)) - byScore <- client.zRemRangeByScore("zrem-score", ScoreBoundary.Inclusive(2.0), ScoreBoundary.PosInf) - _ <- client.zAdd("zrem-lex")(("a", 0.0), ("b", 0.0), ("c", 0.0)) - byLex <- client.zRemRangeByLex[String]("zrem-lex", LexBoundary.Inclusive("a"), LexBoundary.Inclusive("b")) - } yield { - assertEquals(byRank, 1L) - assertEquals(byScore, 2L) - assertEquals(byLex, 2L) - } - } + clientTest("ZREMRANGEBYRANK BYSCORE and BYLEX trim ranges") { client => + client.zAdd("zrem-rank")(("a", 1.0), ("b", 2.0), ("c", 3.0), ("d", 4.0)) >> + client.zRemRangeByRank("zrem-rank", 0L, 0L).is(1L) >> + client.zAdd("zrem-score")(("a", 1.0), ("b", 2.0), ("c", 3.0)) >> + client.zRemRangeByScore("zrem-score", ScoreBoundary.Inclusive(2.0), ScoreBoundary.PosInf).is(2L) >> + client.zAdd("zrem-lex")(("a", 0.0), ("b", 0.0), ("c", 0.0)) >> + client.zRemRangeByLex[String]("zrem-lex", LexBoundary.Inclusive("a"), LexBoundary.Inclusive("b")).is(2L) } - test("ZSCAN streams member/score pairs") { - withClient { client => - for { - _ <- client.zAdd("zset-scan")(("a", 1.0), ("b", 2.0), ("c", 3.0)) - page <- client.zScan[String]("zset-scan", ScanCursor.start) - } yield assertEquals(page.items.toMap, Map("a" -> 1.0, "b" -> 2.0, "c" -> 3.0)) - } + clientTest("ZSCAN streams member/score pairs") { client => + client.zAdd("zset-scan")(("a", 1.0), ("b", 2.0), ("c", 3.0)) >> + client.zScan[String]("zset-scan", ScanCursor.start).map(_.items.toMap).is(Map("a" -> 1.0, "b" -> 2.0, "c" -> 3.0)) } - test("ZUNION ZINTER WITHSCORES, ZDIFF WITHSCORES, and ZREVRANK WITHSCORE return members with their scores") { - withClient { client => - for { - _ <- client.zAdd("zgap-a")(("x", 1.0), ("y", 2.0)) - _ <- client.zAdd("zgap-b")(("y", 3.0), ("z", 4.0)) - union <- client.zUnion[String]("zgap-a", "zgap-b")() - inter <- client.zInterWithScores[String]("zgap-a", "zgap-b")() - diff <- client.zDiffWithScores[String]("zgap-a", "zgap-b") - revRank <- client.zRevRankWithScore[String]("zgap-a", "x") - } yield { - assertEquals(union.toSet, Set("x", "y", "z")) - assertEquals(inter, Vector("y" -> 5.0)) - assertEquals(diff, Vector("x" -> 1.0)) - assertEquals(revRank, Some((1L, 1.0))) - } - } + clientTest("ZUNION ZINTER WITHSCORES, ZDIFF WITHSCORES, and ZREVRANK WITHSCORE return members with their scores") { client => + client.zAdd("zgap-a")(("x", 1.0), ("y", 2.0)) >> + client.zAdd("zgap-b")(("y", 3.0), ("z", 4.0)) >> + client.zUnion[String]("zgap-a", "zgap-b")().map(_.toSet).is(Set("x", "y", "z")) >> + client.zInterWithScores[String]("zgap-a", "zgap-b")().is(Vector("y" -> 5.0)) >> + client.zDiffWithScores[String]("zgap-a", "zgap-b").is(Vector("x" -> 1.0)) >> + client.zRevRankWithScore[String]("zgap-a", "x").is(Some((1L, 1.0))) } - test("ZPOPMAX with a count and BZPOPMAX pop the highest-scored members") { - withClient { client => - for { - _ <- client.zAdd("zgap-pop")(("a", 1.0), ("b", 2.0), ("c", 3.0)) - topTwo <- client.zPopMaxCount[String]("zgap-pop", 2L) - _ <- client.zAdd("zgap-bz")(("a", 1.0), ("b", 2.0)) - bzMax <- client.bzPopMax[String]("zgap-bz")(BlockTimeout.After(1.second)) - none <- client.bzPopMax[String]("zgap-bz-missing")(BlockTimeout.After(100.millis)) - } yield { - assertEquals(topTwo, Vector("c" -> 3.0, "b" -> 2.0)) - assertEquals(bzMax, Some(("zgap-bz", "b", 2.0))) - assertEquals(none, None) - } - } + clientTest("ZPOPMAX with a count and BZPOPMAX pop the highest-scored members") { client => + client.zAdd("zgap-pop")(("a", 1.0), ("b", 2.0), ("c", 3.0)) >> + client.zPopMaxCount[String]("zgap-pop", 2L).is(Vector("c" -> 3.0, "b" -> 2.0)) >> + client.zAdd("zgap-bz")(("a", 1.0), ("b", 2.0)) >> + client.bzPopMax[String]("zgap-bz")(BlockTimeout.After(1.second)).is(Some(("zgap-bz", "b", 2.0))) >> + client.bzPopMax[String]("zgap-bz-missing")(BlockTimeout.After(100.millis)).is(None) } } - -class RedisSortedSetsSuite extends SortedSetsSuite(Images.redis) - -class ValkeySortedSetsSuite extends SortedSetsSuite(Images.valkey) diff --git a/integration-tests/shared/src/test/scala/sage/integration/commands/StreamsSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/commands/StreamsSuite.scala index b89ccb63..1807fc82 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/commands/StreamsSuite.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/commands/StreamsSuite.scala @@ -5,194 +5,153 @@ import scala.concurrent.duration.* import kyo.compat.* import sage.commands.* -import sage.integration.{Images, ServerSuite} - -abstract class StreamsSuite(image: String) extends ServerSuite(image) { - - test("XADD XLEN XRANGE and XREVRANGE round-trip entries, preserving field order") { - withClient { client => - for { - id1 <- client.xAdd("s-range", XAddId.Explicit(StreamId(1L, 0L)))(("a", "1"), ("b", "2")) - id2 <- client.xAdd("s-range", XAddId.Explicit(StreamId(2L, 0L)))(("c", "3")) - len <- client.xLen("s-range") - all <- client.xRange[String, String]("s-range") - rev <- client.xRevRange[String, String]("s-range") - one <- client.xRange[String, String]("s-range", StreamRangeId.Exclusive(StreamId(1L, 0L))) - } yield { - assertEquals(id1, StreamId(1L, 0L)) - assertEquals(id2, StreamId(2L, 0L)) - assertEquals(len, 2L) - assertEquals(all, Vector(StreamEntry(StreamId(1L, 0L), Vector("a" -> "1", "b" -> "2")), StreamEntry(StreamId(2L, 0L), Vector("c" -> "3")))) - assertEquals(rev.map(_.id), Vector(StreamId(2L, 0L), StreamId(1L, 0L))) - assertEquals(one.map(_.id), Vector(StreamId(2L, 0L))) - } - } +import sage.integration.BothServersSuite + +class StreamsSuite extends BothServersSuite { + + clientTest("XADD XLEN XRANGE and XREVRANGE round-trip entries, preserving field order") { client => + client.xAdd("s-range", XAddId.Explicit(StreamId(1L, 0L)))(("a", "1"), ("b", "2")).is(StreamId(1L, 0L)) >> + client.xAdd("s-range", XAddId.Explicit(StreamId(2L, 0L)))(("c", "3")).is(StreamId(2L, 0L)) >> + client.xLen("s-range").is(2L) >> + client + .xRange[String, String]("s-range") + .is(Vector(StreamEntry(StreamId(1L, 0L), Vector("a" -> "1", "b" -> "2")), StreamEntry(StreamId(2L, 0L), Vector("c" -> "3")))) >> + client.xRevRange[String, String]("s-range").map(_.map(_.id)).is(Vector(StreamId(2L, 0L), StreamId(1L, 0L))) >> + client.xRange[String, String]("s-range", StreamRangeId.Exclusive(StreamId(1L, 0L))).map(_.map(_.id)).is(Vector(StreamId(2L, 0L))) } - test("XADD with * generates increasing ids, NOMKSTREAM declines a missing stream, and XDEL removes") { - withClient { client => - for { - absent <- client.xAddNoMkStream("s-del-missing")(("f", "v")) - a <- client.xAdd("s-del")(("f", "1")) - b <- client.xAdd("s-del")(("f", "2")) - del <- client.xDel("s-del")(a) - len <- client.xLen("s-del") - } yield { - assertEquals(absent, None) - assert(b > a) - assertEquals(del, 1L) - assertEquals(len, 1L) - } - } + clientTest("XADD with * generates increasing ids, NOMKSTREAM declines a missing stream, and XDEL removes") { client => + for { + _ <- client.xAddNoMkStream("s-del-missing")(("f", "v")).is(None) + a <- client.xAdd("s-del")(("f", "1")) + b <- client.xAdd("s-del")(("f", "2")) + _ <- client.xDel("s-del")(a).is(1L) + _ <- client.xLen("s-del").is(1L) + } yield assert(b > a) } - test("XTRIM MAXLEN caps the stream length") { - withClient { client => - for { - _ <- client.xAdd("s-trim", XAddId.Explicit(StreamId(1L, 0L)))(("f", "1")) - _ <- client.xAdd("s-trim", XAddId.Explicit(StreamId(2L, 0L)))(("f", "2")) - _ <- client.xAdd("s-trim", XAddId.Explicit(StreamId(3L, 0L)))(("f", "3")) - cut <- client.xTrim("s-trim", Trimming.Exact(TrimThreshold.MaxLen(1L))) - len <- client.xLen("s-trim") - } yield { - assertEquals(cut, 2L) - assertEquals(len, 1L) - } - } + clientTest("XTRIM MAXLEN caps the stream length") { client => + client.xAdd("s-trim", XAddId.Explicit(StreamId(1L, 0L)))(("f", "1")) >> + client.xAdd("s-trim", XAddId.Explicit(StreamId(2L, 0L)))(("f", "2")) >> + client.xAdd("s-trim", XAddId.Explicit(StreamId(3L, 0L)))(("f", "3")) >> + client.xTrim("s-trim", Trimming.Exact(TrimThreshold.MaxLen(1L))).is(2L) >> + client.xLen("s-trim").is(1L) } - test("XINFO STREAM reports length and the last generated id") { - withClient { client => - for { - _ <- client.xAdd("s-info", XAddId.Explicit(StreamId(1L, 0L)))(("f", "v")) - _ <- client.xAdd("s-info", XAddId.Explicit(StreamId(5L, 0L)))(("g", "w")) - info <- client.xInfoStream[String, String]("s-info") - } yield { - assertEquals(info.length, 2L) - assertEquals(info.lastGeneratedId, StreamId(5L, 0L)) - assertEquals(info.firstEntry, Some(StreamEntry(StreamId(1L, 0L), Vector("f" -> "v")))) - } + clientTest("XINFO STREAM reports length and the last generated id") { client => + for { + _ <- client.xAdd("s-info", XAddId.Explicit(StreamId(1L, 0L)))(("f", "v")) + _ <- client.xAdd("s-info", XAddId.Explicit(StreamId(5L, 0L)))(("g", "w")) + info <- client.xInfoStream[String, String]("s-info") + } yield { + assertEquals(info.length, 2L) + assertEquals(info.lastGeneratedId, StreamId(5L, 0L)) + assertEquals(info.firstEntry, Some(StreamEntry(StreamId(1L, 0L), Vector("f" -> "v")))) } } - test("a consumer group reads new entries, acknowledges them, and reports an empty PEL") { - withClient { client => - for { - _ <- client.xAdd("s-grp", XAddId.Explicit(StreamId(1L, 0L)))(("f", "1")) - _ <- client.xAdd("s-grp", XAddId.Explicit(StreamId(2L, 0L)))(("f", "2")) - _ <- client.xGroupCreate("s-grp", "g", GroupStartId.At(StreamId(0L, 0L))) - read <- client.xReadGroup[String, String]("g", "c1")(("s-grp", GroupReadId.New))(count = Some(10L)) - pending <- client.xPending("s-grp", "g") - acked <- client.xAck("s-grp", "g")(StreamId(1L, 0L), StreamId(2L, 0L)) - drained <- client.xPending("s-grp", "g") - groups <- client.xInfoGroups("s-grp") - } yield { - assertEquals( - read, - Vector("s-grp" -> Vector(StreamEntry(StreamId(1L, 0L), Vector("f" -> "1")), StreamEntry(StreamId(2L, 0L), Vector("f" -> "2")))) - ) - assertEquals(pending.total, 2L) - assertEquals(acked, 2L) - assertEquals(drained.total, 0L) - assertEquals(groups.map(_.name), Vector("g")) - } - } + clientTest("a consumer group reads new entries, acknowledges them, and reports an empty PEL") { client => + client.xAdd("s-grp", XAddId.Explicit(StreamId(1L, 0L)))(("f", "1")) >> + client.xAdd("s-grp", XAddId.Explicit(StreamId(2L, 0L)))(("f", "2")) >> + client.xGroupCreate("s-grp", "g", GroupStartId.At(StreamId(0L, 0L))) >> + client + .xReadGroup[String, String]("g", "c1")(("s-grp", GroupReadId.New))(count = Some(10L)) + .is(Vector("s-grp" -> Vector(StreamEntry(StreamId(1L, 0L), Vector("f" -> "1")), StreamEntry(StreamId(2L, 0L), Vector("f" -> "2"))))) >> + client.xPending("s-grp", "g").map(_.total).is(2L) >> + client.xAck("s-grp", "g")(StreamId(1L, 0L), StreamId(2L, 0L)).is(2L) >> + client.xPending("s-grp", "g").map(_.total).is(0L) >> + client.xInfoGroups("s-grp").map(_.map(_.name)).is(Vector("g")) } - test("XCLAIM and XAUTOCLAIM transfer pending entries to another consumer") { - withClient { client => - for { - _ <- client.xAdd("s-claim", XAddId.Explicit(StreamId(1L, 0L)))(("f", "1")) - _ <- client.xGroupCreate("s-claim", "g", GroupStartId.At(StreamId(0L, 0L))) - _ <- client.xReadGroup[String, String]("g", "c1")(("s-claim", GroupReadId.New))() - claimed <- client.xClaim[String, String]("s-claim", "g", "c2", Duration.Zero)(StreamId(1L, 0L))() - auto <- client.xAutoClaim[String, String]("s-claim", "g", "c3", Duration.Zero) - owners <- client.xPendingExtended("s-claim", "g") - } yield { - assertEquals(claimed, Vector(StreamEntry(StreamId(1L, 0L), Vector("f" -> "1")))) - assertEquals(auto.entries.map(_.id), Vector(StreamId(1L, 0L))) - assertEquals(owners.map(_.consumer), Vector("c3")) - } - } + clientTest("XCLAIM and XAUTOCLAIM transfer pending entries to another consumer") { client => + client.xAdd("s-claim", XAddId.Explicit(StreamId(1L, 0L)))(("f", "1")) >> + client.xGroupCreate("s-claim", "g", GroupStartId.At(StreamId(0L, 0L))) >> + client.xReadGroup[String, String]("g", "c1")(("s-claim", GroupReadId.New))() >> + client + .xClaim[String, String]("s-claim", "g", "c2", Duration.Zero)(StreamId(1L, 0L))() + .is(Vector(StreamEntry(StreamId(1L, 0L), Vector("f" -> "1")))) >> + client.xAutoClaim[String, String]("s-claim", "g", "c3", Duration.Zero).map(_.entries.map(_.id)).is(Vector(StreamId(1L, 0L))) >> + client.xPendingExtended("s-claim", "g").map(_.map(_.consumer)).is(Vector("c3")) } - test("XGROUP CREATECONSUMER/DELCONSUMER/SETID/DESTROY and XSETID manage a group and the stream's last id") { - withClient { client => - for { - _ <- client.xAdd("s-mgmt", XAddId.Explicit(StreamId(1L, 0L)))(("f", "1")) - _ <- client.xGroupCreate("s-mgmt", "g", GroupStartId.At(StreamId(0L, 0L))) - created <- client.xGroupCreateConsumer("s-mgmt", "g", "c1") - again <- client.xGroupCreateConsumer("s-mgmt", "g", "c1") - delPel <- client.xGroupDelConsumer("s-mgmt", "g", "c1") - _ <- client.xGroupSetId("s-mgmt", "g", GroupStartId.At(StreamId(1L, 0L))) - _ <- client.xSetId("s-mgmt", GroupStartId.At(StreamId(5L, 0L))) - info <- client.xInfoStream[String, String]("s-mgmt") - destroyed <- client.xGroupDestroy("s-mgmt", "g") - groups <- client.xInfoGroups("s-mgmt") - } yield { - assertEquals(created, true) - assertEquals(again, false) - assertEquals(delPel, 0L) - assertEquals(info.lastGeneratedId, StreamId(5L, 0L)) - assertEquals(destroyed, true) - assert(groups.isEmpty) - } - } + clientTest("XGROUP CREATECONSUMER/DELCONSUMER/SETID/DESTROY and XSETID manage a group and the stream's last id") { client => + client.xAdd("s-mgmt", XAddId.Explicit(StreamId(1L, 0L)))(("f", "1")) >> + client.xGroupCreate("s-mgmt", "g", GroupStartId.At(StreamId(0L, 0L))) >> + client.xGroupCreateConsumer("s-mgmt", "g", "c1").is(true) >> + client.xGroupCreateConsumer("s-mgmt", "g", "c1").is(false) >> + client.xGroupDelConsumer("s-mgmt", "g", "c1").is(0L) >> + client.xGroupSetId("s-mgmt", "g", GroupStartId.At(StreamId(1L, 0L))) >> + client.xSetId("s-mgmt", GroupStartId.At(StreamId(5L, 0L))) >> + client.xInfoStream[String, String]("s-mgmt").map(_.lastGeneratedId).is(StreamId(5L, 0L)) >> + client.xGroupDestroy("s-mgmt", "g").is(true) >> + client.xInfoGroups("s-mgmt").satisfies(_.isEmpty) } - test("XCLAIM and XAUTOCLAIM JUSTID transfer pending ids without the payload") { - withClient { client => - for { - _ <- client.xAdd("s-justid", XAddId.Explicit(StreamId(1L, 0L)))(("f", "1")) - _ <- client.xAdd("s-justid", XAddId.Explicit(StreamId(2L, 0L)))(("f", "2")) - _ <- client.xGroupCreate("s-justid", "g", GroupStartId.At(StreamId(0L, 0L))) - _ <- client.xReadGroup[String, String]("g", "c1")(("s-justid", GroupReadId.New))() - claimed <- client.xClaimJustId("s-justid", "g", "c2", Duration.Zero)(StreamId(1L, 0L))() - auto <- client.xAutoClaimJustId("s-justid", "g", "c3", Duration.Zero) - } yield { - assertEquals(claimed, Vector(StreamId(1L, 0L))) - assertEquals(auto.claimed, Vector(StreamId(1L, 0L), StreamId(2L, 0L))) - } - } + clientTest("XCLAIM and XAUTOCLAIM JUSTID transfer pending ids without the payload") { client => + client.xAdd("s-justid", XAddId.Explicit(StreamId(1L, 0L)))(("f", "1")) >> + client.xAdd("s-justid", XAddId.Explicit(StreamId(2L, 0L)))(("f", "2")) >> + client.xGroupCreate("s-justid", "g", GroupStartId.At(StreamId(0L, 0L))) >> + client.xReadGroup[String, String]("g", "c1")(("s-justid", GroupReadId.New))() >> + client.xClaimJustId("s-justid", "g", "c2", Duration.Zero)(StreamId(1L, 0L))().is(Vector(StreamId(1L, 0L))) >> + client.xAutoClaimJustId("s-justid", "g", "c3", Duration.Zero).map(_.claimed).is(Vector(StreamId(1L, 0L), StreamId(2L, 0L))) } - test("XINFO STREAM FULL decodes the group PEL after a consumer has read but not acked") { - withClient { client => - for { - _ <- client.xAdd("s-full", XAddId.Explicit(StreamId(1L, 0L)))(("f", "1")) - _ <- client.xGroupCreate("s-full", "g", GroupStartId.At(StreamId(0L, 0L))) - _ <- client.xReadGroup[String, String]("g", "c1")(("s-full", GroupReadId.New))() - full <- client.xInfoStreamFull[String, String]("s-full") - } yield { - assertEquals(full.length, 1L) - assertEquals(full.groups.map(_.name), Vector("g")) - assertEquals(full.groups.head.pending.map(_.id), Vector(StreamId(1L, 0L))) - assertEquals(full.groups.head.pending.head.consumer, Some("c1")) - assertEquals(full.groups.head.consumers.head.pending.map(_.id), Vector(StreamId(1L, 0L))) - } + clientTest("XINFO STREAM FULL decodes the group PEL after a consumer has read but not acked") { client => + for { + _ <- client.xAdd("s-full", XAddId.Explicit(StreamId(1L, 0L)))(("f", "1")) + _ <- client.xGroupCreate("s-full", "g", GroupStartId.At(StreamId(0L, 0L))) + _ <- client.xReadGroup[String, String]("g", "c1")(("s-full", GroupReadId.New))() + full <- client.xInfoStreamFull[String, String]("s-full") + } yield { + assertEquals(full.length, 1L) + assertEquals( + full.groups.map(g => (g.name, g.pending.map(p => (p.id, p.consumer)), g.consumers.map(c => (c.name, c.pending.map(_.id))))), + Vector(("g", Vector((StreamId(1L, 0L), Some("c1"))), Vector(("c1", Vector(StreamId(1L, 0L)))))) + ) } } - test("XREAD with BLOCK does not stall ordinary commands on the multiplexed connection") { - withClient { client => - for { - _ <- client.xAdd("s-block", XAddId.Explicit(StreamId(1L, 0L)))(("f", "0")) - out <- CIO.zip( - client.xRead[String, String](("s-block", ReadId.After(StreamId(1L, 0L))))(block = Some(BlockTimeout.After(5.seconds))), - for { - pong <- client.ping() - _ <- client.xAdd("s-block", XAddId.Explicit(StreamId(2L, 0L)))(("f", "1")) - } yield pong - ) - } yield { - val (read, pong) = out - assertEquals(pong, "PONG") - assertEquals(read, Vector("s-block" -> Vector(StreamEntry(StreamId(2L, 0L), Vector("f" -> "1"))))) - } - } + clientTest("XREAD with BLOCK does not stall ordinary commands on the multiplexed connection") { client => + client.xAdd("s-block", XAddId.Explicit(StreamId(1L, 0L)))(("f", "0")) >> + CIO + .zip( + client.xRead[String, String](("s-block", ReadId.After(StreamId(1L, 0L))))(block = Some(BlockTimeout.After(5.seconds))), + for { + pong <- client.ping() + _ <- client.xAdd("s-block", XAddId.Explicit(StreamId(2L, 0L)))(("f", "1")) + } yield pong + ) + .is((Vector("s-block" -> Vector(StreamEntry(StreamId(2L, 0L), Vector("f" -> "1")))), "PONG")) } -} -class RedisStreamsSuite extends StreamsSuite(Images.redis) + // XCFGSET, XDELEX, XACKDEL and XNACK exist only on Redis. + redisTest("XCFGSET sets per-stream idempotent-message-processing config") { client => + client.xAdd("sx-cfg", XAddId.Explicit(StreamId(1L, 0L)))(("f", "1")) >> + client.xCfgSet("sx-cfg", idmpDuration = Some(1.hour), idmpMaxSize = Some(100L)).is(()) >> + client.xCfgSet("sx-cfg", idmpMaxSize = Some(50L)).is(()) + } + + redisTest("XDELEX reports per-id deletion status") { client => + client.xAdd("sx-delex", XAddId.Explicit(StreamId(1L, 0L)))(("f", "1")) >> + client.xAdd("sx-delex", XAddId.Explicit(StreamId(2L, 0L)))(("f", "2")) >> + client.xDelEx("sx-delex")(StreamId(1L, 0L), StreamId(9L, 0L)).is(Vector(StreamEntryDeletion.Deleted, StreamEntryDeletion.NotFound)) >> + client.xLen("sx-delex").is(1L) + } -class ValkeyStreamsSuite extends StreamsSuite(Images.valkey) + redisTest("XACKDEL acknowledges and deletes in one step") { client => + client.xAdd("sx-ackdel", XAddId.Explicit(StreamId(1L, 0L)))(("f", "1")) >> + client.xGroupCreate("sx-ackdel", "g", GroupStartId.At(StreamId(0L, 0L))) >> + client.xReadGroup[String, String]("g", "c1")(("sx-ackdel", GroupReadId.New))() >> + client.xAckDel("sx-ackdel", "g")(StreamId(1L, 0L)).is(Vector(StreamEntryDeletion.Deleted)) >> + client.xLen("sx-ackdel").is(0L) >> + client.xPending("sx-ackdel", "g").map(_.total).is(0L) + } + + redisTest("XNACK releases a pending entry back to the group") { client => + client.xAdd("sx-nack", XAddId.Explicit(StreamId(1L, 0L)))(("f", "1")) >> + client.xGroupCreate("sx-nack", "g", GroupStartId.At(StreamId(0L, 0L))) >> + client.xReadGroup[String, String]("g", "c1")(("sx-nack", GroupReadId.New))() >> + client.xNack("sx-nack", "g", NackMode.Fail)(StreamId(1L, 0L))().is(1L) >> + client.xPending("sx-nack", "g").map(_.total).is(1L) + } +} diff --git a/integration-tests/shared/src/test/scala/sage/integration/commands/StringsSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/commands/StringsSuite.scala index 812c48b0..f4627243 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/commands/StringsSuite.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/commands/StringsSuite.scala @@ -6,203 +6,143 @@ import scala.concurrent.duration.* import kyo.compat.* -import sage.commands.{GetExpiry, LcsMatch, MatchRange, SetCondition, SetExpiry, Ttl} -import sage.integration.{Images, ServerSuite} +import sage.commands.* +import sage.integration.BothServersSuite import sage.integration.Ttls.expiresWithin -abstract class StringsSuite(image: String) extends ServerSuite(image) { - - test("APPEND grows the value and reports the new length") { - withClient { client => - for { - _ <- client.set("str-append", "abc") - length <- client.append("str-append", "def") - value <- client.get[String]("str-append") - } yield { - assertEquals(length, 6L) - assertEquals(value, Some("abcdef")) - } - } +class StringsSuite extends BothServersSuite { + + clientTest("APPEND grows the value and reports the new length") { client => + client.set("str-append", "abc") >> + client.append("str-append", "def").is(6L) >> + client.get[String]("str-append").is(Some("abcdef")) } - test("INCR DECR INCRBY DECRBY count atomically") { - withClient { client => - for { - _ <- client.set("str-counter", 10) - up <- client.incr("str-counter") - upBy <- client.incrBy("str-counter", 5L) - down <- client.decr("str-counter") - total <- client.decrBy("str-counter", 5L) - } yield { - assertEquals(up, 11L) - assertEquals(upBy, 16L) - assertEquals(down, 15L) - assertEquals(total, 10L) - } - } + clientTest("INCR DECR INCRBY DECRBY count atomically") { client => + client.set("str-counter", 10) >> + client.incr("str-counter").is(11L) >> + client.incrBy("str-counter", 5L).is(16L) >> + client.decr("str-counter").is(15L) >> + client.decrBy("str-counter", 5L).is(10L) } - test("INCRBYFLOAT increments with float precision") { - withClient { client => - for { - _ <- client.set("str-float", "10.5") - value <- client.incrByFloat("str-float", 0.25) - } yield assertEquals(value, 10.75) - } + clientTest("INCRBYFLOAT increments with float precision") { client => + client.set("str-float", "10.5") >> + client.incrByFloat("str-float", 0.25).is(10.75) } - test("GETDEL returns the value and removes the key") { - withClient { client => - for { - _ <- client.set("str-getdel", "gone") - value <- client.getDel[String]("str-getdel") - after <- client.get[String]("str-getdel") - missing <- client.getDel[String]("str-getdel-missing") - } yield { - assertEquals(value, Some("gone")) - assertEquals(after, None) - assertEquals(missing, None) - } - } + clientTest("GETDEL returns the value and removes the key") { client => + client.set("str-getdel", "gone") >> + client.getDel[String]("str-getdel").is(Some("gone")) >> + client.get[String]("str-getdel").is(None) >> + client.getDel[String]("str-getdel-missing").is(None) } - test("GETEX sets, keeps, and removes the ttl") { - withClient { client => - for { - _ <- client.set("str-getex", "v") - value <- client.getEx[String]("str-getex", GetExpiry.In(60.seconds)) - withTtl <- client.ttl("str-getex") - kept <- client.getEx[String]("str-getex") - stillTtl <- client.ttl("str-getex") - persisted <- client.getEx[String]("str-getex", GetExpiry.Persist) - noTtl <- client.ttl("str-getex") - } yield { - assertEquals(value, Some("v")) - assert(expiresWithin(withTtl, 60.seconds)) - assertEquals(kept, Some("v")) - assert(expiresWithin(stillTtl, 60.seconds)) - assertEquals(persisted, Some("v")) - assertEquals(noTtl, Ttl.NoExpiry) - } - } + clientTest("GETEX sets, keeps, and removes the ttl") { client => + client.set("str-getex", "v") >> + client.getEx[String]("str-getex", GetExpiry.In(60.seconds)).is(Some("v")) >> + client.ttl("str-getex").satisfies(expiresWithin(_, 60.seconds)) >> + client.getEx[String]("str-getex").is(Some("v")) >> + client.ttl("str-getex").satisfies(expiresWithin(_, 60.seconds)) >> + client.getEx[String]("str-getex", GetExpiry.Persist).is(Some("v")) >> + client.ttl("str-getex").is(Ttl.NoExpiry) } - test("GETRANGE SETRANGE STRLEN address the value by offset") { - withClient { client => - for { - _ <- client.set("str-range", "Hello World") - head <- client.getRange[String]("str-range", 0L, 4L) - length <- client.setRange("str-range", 6L, "Redis") - value <- client.get[String]("str-range") - len <- client.strLen("str-range") - } yield { - assertEquals(head, "Hello") - assertEquals(length, 11L) - assertEquals(value, Some("Hello Redis")) - assertEquals(len, 11L) - } - } + clientTest("GETRANGE SETRANGE STRLEN address the value by offset") { client => + client.set("str-range", "Hello World") >> + client.getRange[String]("str-range", 0L, 4L).is("Hello") >> + client.setRange("str-range", 6L, "Redis").is(11L) >> + client.get[String]("str-range").is(Some("Hello Redis")) >> + client.strLen("str-range").is(11L) } - test("MGET returns values positionally with None for missing keys") { - withClient { client => - for { - _ <- client.mSet(("str-mget-a", "1"), ("str-mget-b", "2")) - values <- client.mGet[String]("str-mget-a", "str-mget-missing", "str-mget-b") - } yield assertEquals(values, Vector(Some("1"), None, Some("2"))) - } + clientTest("MGET returns values positionally with None for missing keys") { client => + client.mSet(("str-mget-a", "1"), ("str-mget-b", "2")) >> + client.mGet[String]("str-mget-a", "str-mget-missing", "str-mget-b").is(Vector(Some("1"), None, Some("2"))) } - test("MSETNX writes all keys or none") { - withClient { client => - for { - first <- client.mSetNx(("str-msetnx-a", "1"), ("str-msetnx-b", "2")) - second <- client.mSetNx(("str-msetnx-b", "x"), ("str-msetnx-c", "3")) - c <- client.get[String]("str-msetnx-c") - } yield { - assertEquals(first, true) - assertEquals(second, false) - assertEquals(c, None) - } - } + clientTest("MSETNX writes all keys or none") { client => + client.mSetNx(("str-msetnx-a", "1"), ("str-msetnx-b", "2")).is(true) >> + client.mSetNx(("str-msetnx-b", "x"), ("str-msetnx-c", "3")).is(false) >> + client.get[String]("str-msetnx-c").is(None) } - test("SET honors the existence conditions") { - withClient { client => - for { - created <- client.set("str-cond", "one", condition = SetCondition.IfNotExists) - duplicate <- client.set("str-cond", "two", condition = SetCondition.IfNotExists) - updated <- client.set("str-cond", "three", condition = SetCondition.IfExists) - ghost <- client.set("str-cond-missing", "x", condition = SetCondition.IfExists) - value <- client.get[String]("str-cond") - } yield { - assertEquals(created, true) - assertEquals(duplicate, false) - assertEquals(updated, true) - assertEquals(ghost, false) - assertEquals(value, Some("three")) - } - } + clientTest("SET honors the existence conditions") { client => + client.set("str-cond", "one", condition = SetCondition.IfNotExists).is(true) >> + client.set("str-cond", "two", condition = SetCondition.IfNotExists).is(false) >> + client.set("str-cond", "three", condition = SetCondition.IfExists).is(true) >> + client.set("str-cond-missing", "x", condition = SetCondition.IfExists).is(false) >> + client.get[String]("str-cond").is(Some("three")) } - test("setGet returns the previous value") { - withClient { client => - for { - before <- client.setGet[String]("str-setget", "one") - after <- client.setGet[String]("str-setget", "two") - value <- client.get[String]("str-setget") - } yield { - assertEquals(before, None) - assertEquals(after, Some("one")) - assertEquals(value, Some("two")) - } - } + clientTest("setGet returns the previous value") { client => + client.setGet[String]("str-setget", "one").is(None) >> + client.setGet[String]("str-setget", "two").is(Some("one")) >> + client.get[String]("str-setget").is(Some("two")) } - test("SET expiry: In sets a ttl, KeepTtl preserves it, the default clears it") { - withClient { client => - for { - _ <- client.set("str-ttl", "v", expiry = SetExpiry.In(60.seconds)) - initial <- client.ttl("str-ttl") - _ <- client.set("str-ttl", "v2", expiry = SetExpiry.KeepTtl) - kept <- client.ttl("str-ttl") - _ <- client.set("str-ttl", "v3") - cleared <- client.ttl("str-ttl") - } yield { - assert(expiresWithin(initial, 60.seconds)) - assert(expiresWithin(kept, 60.seconds)) - assertEquals(cleared, Ttl.NoExpiry) - } - } + clientTest("SET expiry: In sets a ttl, KeepTtl preserves it, the default clears it") { client => + client.set("str-ttl", "v", expiry = SetExpiry.In(60.seconds)) >> + client.ttl("str-ttl").satisfies(expiresWithin(_, 60.seconds)) >> + client.set("str-ttl", "v2", expiry = SetExpiry.KeepTtl) >> + client.ttl("str-ttl").satisfies(expiresWithin(_, 60.seconds)) >> + client.set("str-ttl", "v3") >> + client.ttl("str-ttl").is(Ttl.NoExpiry) } - test("SET expiry: At pins an absolute deadline") { - withClient { client => - val deadline = Instant.ofEpochSecond(Instant.now().getEpochSecond + 3600) - for { - _ <- client.set("str-at", "v", expiry = SetExpiry.At(deadline)) - ttl <- client.ttl("str-at") - } yield assert(expiresWithin(ttl, 3600.seconds)) - } + clientTest("SET expiry: At pins an absolute deadline") { client => + val deadline = Instant.ofEpochSecond(Instant.now().getEpochSecond + 3600) + client.set("str-at", "v", expiry = SetExpiry.At(deadline)) >> + client.ttl("str-at").satisfies(expiresWithin(_, 3600.seconds)) } - test("LCS finds the subsequence, its length, and indexed matches") { - withClient { client => - for { - _ <- client.mSet(("str-lcs-1", "ohmytext"), ("str-lcs-2", "mynewtext")) - sequence <- client.lcs[String]("str-lcs-1", "str-lcs-2") - length <- client.lcsLen("str-lcs-1", "str-lcs-2") - idx <- client.lcsIdx("str-lcs-1", "str-lcs-2", minMatchLen = Some(4L), withMatchLen = true) - } yield { - assertEquals(sequence, "mytext") - assertEquals(length, 6L) - assertEquals(idx.length, 6L) - assertEquals(idx.matches, Vector(LcsMatch(MatchRange(4L, 7L), MatchRange(5L, 8L), Some(4L)))) - } + clientTest("LCS finds the subsequence, its length, and indexed matches") { client => + for { + _ <- client.mSet(("str-lcs-1", "ohmytext"), ("str-lcs-2", "mynewtext")) + _ <- client.lcs[String]("str-lcs-1", "str-lcs-2").is("mytext") + _ <- client.lcsLen("str-lcs-1", "str-lcs-2").is(6L) + idx <- client.lcsIdx("str-lcs-1", "str-lcs-2", minMatchLen = Some(4L), withMatchLen = true) + } yield { + assertEquals(idx.length, 6L) + assertEquals(idx.matches, Vector(LcsMatch(MatchRange(4L, 7L), MatchRange(5L, 8L), Some(4L)))) } } -} -class RedisStringsSuite extends StringsSuite(Images.redis) + // DIGEST, DELEX, MSETEX and INCREX exist only on Redis. + redisTest("DIGEST returns a stable hex digest, None for a missing key") { client => + for { + _ <- client.set("sx-digest", "hello") + d1 <- client.digest("sx-digest") + _ <- client.digest("sx-digest").is(d1) + _ <- client.digest("sx-digest-missing").is(None) + } yield assert(d1.exists(_.nonEmpty)) + } + + redisTest("DELEX deletes only when the value or digest condition matches") { client => + for { + _ <- client.set("sx-delex", "v1") + _ <- client.delex("sx-delex", DelexCondition.IfEq("other")).is(false) + _ <- client.exists("sx-delex").is(1L) + digest <- client.digest("sx-delex").flatMap(required("DIGEST", _)) + _ <- client.delex[String]("sx-delex", DelexCondition.IfDigestNe(digest)).is(false) + _ <- client.delex("sx-delex", DelexCondition.IfEq("v1")).is(true) + _ <- client.exists("sx-delex").is(0L) + } yield () + } + + redisTest("MSETEX sets multiple keys with a shared TTL and respects NX") { client => + client.msetEx(expiry = SetExpiry.In(100.seconds))(("sx-ms-a", "1"), ("sx-ms-b", "2")).is(true) >> + client.get[String]("sx-ms-a").is(Some("1")) >> + client.ttl("sx-ms-a").satisfies(expiresWithin(_, 100.seconds)) >> + client.msetEx(condition = SetCondition.IfNotExists)(("sx-ms-a", "9")).is(false) >> + client.get[String]("sx-ms-a").is(Some("1")) + } -class ValkeyStringsSuite extends StringsSuite(Images.valkey) + redisTest("INCREX increments with expiry, saturating bounds, and rejects when out of range") { client => + client.increxBy("sx-incr", 5L, expiry = IncrExpiry.In(100.seconds)).is(IncrExResult(5L, 5L)) >> + client.ttl("sx-incr").satisfies(expiresWithin(_, 100.seconds)) >> + client.increxBy("sx-incr", 100L, saturate = true, upperBound = Some(10L)).is(IncrExResult(10L, 5L)) >> + client.increxBy("sx-incr", 100L, upperBound = Some(10L)).is(IncrExResult(10L, 0L)) >> + client.increxByFloat("sx-incr-f", 1.5).is(IncrExResult(1.5, 1.5)) + } +} diff --git a/integration-tests/shared/src/test/scala/sage/integration/commands/ValkeyKeysExtrasSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/commands/ValkeyKeysExtrasSuite.scala deleted file mode 100644 index 1b9c84fe..00000000 --- a/integration-tests/shared/src/test/scala/sage/integration/commands/ValkeyKeysExtrasSuite.scala +++ /dev/null @@ -1,30 +0,0 @@ -package sage.integration.commands - -import kyo.compat.* - -import sage.integration.{Images, ServerSuite} - -/** - * DELIFEQ is a Valkey-only atomic compare-and-delete absent in Redis, so it has no cross-server counterpart. - */ -class ValkeyKeysExtrasSuite extends ServerSuite(Images.valkey) { - - test("DELIFEQ deletes only when the current value matches") { - withClient { client => - for { - _ <- client.set("vk-lock", "token-1") - mismatch <- client.delIfEq("vk-lock", "other") - present <- client.exists("vk-lock") - matched <- client.delIfEq("vk-lock", "token-1") - gone <- client.exists("vk-lock") - missing <- client.delIfEq("vk-lock-absent", "x") - } yield { - assertEquals(mismatch, false) - assertEquals(present, 1L) - assertEquals(matched, true) - assertEquals(gone, 0L) - assertEquals(missing, false) - } - } - } -} diff --git a/integration-tests/shared/src/test/scala/sage/integration/commands/ValkeyServerExtrasSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/commands/ValkeyServerExtrasSuite.scala deleted file mode 100644 index a41cafe9..00000000 --- a/integration-tests/shared/src/test/scala/sage/integration/commands/ValkeyServerExtrasSuite.scala +++ /dev/null @@ -1,44 +0,0 @@ -package sage.integration.commands - -import kyo.compat.* - -import sage.commands.CommandLogType -import sage.integration.{Images, ServerSuite} - -/** - * COMMANDLOG is Valkey's command-log observability (the SLOWLOG successor), absent in Redis, so it has no cross-server counterpart. - */ -class ValkeyServerExtrasSuite extends ServerSuite(Images.valkey) { - - test("COMMANDLOG GET/LEN/RESET over the slow log") { - withClient { client => - for { - _ <- client.configSet(("slowlog-log-slower-than", "0")) - _ <- client.commandLogReset(CommandLogType.Slow) - _ <- client.get[String]("cl-probe") - len <- client.commandLogLen(CommandLogType.Slow) - recent <- client.commandLogGet(5L, CommandLogType.Slow) - _ <- client.configSet(("slowlog-log-slower-than", "10000")) - _ <- client.commandLogReset(CommandLogType.Slow) - after <- client.commandLogLen(CommandLogType.Slow) - } yield { - assert(len > 0L) - assert(recent.nonEmpty) - assert(recent.forall(_.command.nonEmpty)) - assertEquals(after, 0L) - } - } - } - - test("COMMANDLOG LEN works for the large-request and large-reply types") { - withClient { client => - for { - req <- client.commandLogLen(CommandLogType.LargeRequest) - reply <- client.commandLogLen(CommandLogType.LargeReply) - } yield { - assert(req >= 0L) - assert(reply >= 0L) - } - } - } -} diff --git a/integration-tests/shared/src/test/scala/sage/integration/locking/LockSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/locking/LockSuite.scala index d32b7b3a..56dff296 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/locking/LockSuite.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/locking/LockSuite.scala @@ -1,190 +1,102 @@ package sage.integration.locking -import scala.concurrent.Future import scala.concurrent.duration.* import kyo.compat.* -import sage.Bytes import sage.SageException.LockLost -import sage.client.internal.Client -import sage.commands.Command -import sage.integration.{Eventually, Images, ServerSuite, Ttls} - -abstract class LockSuite(image: String) extends ServerSuite(image) { - private def killClient(id: Long): Command[Unit] = - Command( - "CLIENT", - Command.NoKeys, - Vector("KILL", "ID", id.toString).map(Bytes.utf8), - _ => Right(()) - ) - - private val pauseClients: Command[Unit] = - Command( - "CLIENT", - Command.NoKeys, - Vector("PAUSE", "600", "ALL").map(Bytes.utf8), - _ => Right(()) - ) - - private def withClients[A](body: (Client[CIO, String], Client[CIO, String]) => CIO[A]): Future[A] = - withContainers { server => - connectAndUse(configOf(server)) { first => - connectAndUse(configOf(server))(second => body(first, second)) - }.unsafeRun - } - - private def awaitRenewal(client: Client[CIO, String], key: String): CIO[Unit] = - Eventually.changes(30, 50.millis)(() => client.pTtl(key))((before, after) => - (Ttls.remaining(before), Ttls.remaining(after)) match { - case (Some(previous), Some(current)) => current > previous + 100.millis - case _ => false - } - )(ttl => s"$key was not renewed; its TTL was $ttl") - - private def commandCalls(info: String, command: String): Long = - info.linesIterator - .find(_.startsWith(s"cmdstat_${command.toLowerCase}:calls=")) - .fold(0L)(_.dropWhile(_ != '=').drop(1).takeWhile(_ != ',').toLong) - - test("independent clients contend on the same key and acquire after release") { - withClients { (first, second) => - val a = first.lock[String]() - val b = second.lock[String]() - for { - denied <- a.withLock("shared", 2.seconds) { - b.tryWithLock("shared")(CIO.fail(new AssertionError("contended body ran"))) - } - acquired <- b.tryWithLock("shared")(CIO.value(42)) - exists <- first.exists("4:lock:shared") - } yield { - assertEquals(denied, None) - assertEquals(acquired, Some(42)) - assertEquals(exists, 0L) - } - } - } - - test("waiting scopes protect a read-modify-write operation across clients") { - withClients { (first, second) => - first.set("counter", 0).flatMap { _ => - CIO - .foreach(1 to 12) { i => - val client = if (i % 2 == 0) first else second - client.lock[String]().withLock("counter", 5.seconds) { - client.get[Int]("counter").flatMap { current => - client.ping().flatMap(_ => client.set("counter", current.get + 1)) - } +import sage.integration.BothServersSuite + +class LockSuite extends BothServersSuite { + clientsTest("independent clients contend on the same key and acquire after release")((first, second) => contend(first.lock(), second, "shared")) + + clientsTest("waiting scopes protect a read-modify-write operation across clients") { (first, second) => + first.set("counter", 0).flatMap { _ => + CIO + .foreach(1 to 12) { i => + val client = if (i % 2 == 0) first else second + client.lock[String]().withLock("counter", 5.seconds) { + client.get[Int]("counter").flatMap(required("counter", _)).flatMap { current => + client.ping().flatMap(_ => client.set("counter", current + 1)) } } - .flatMap(_ => first.get[Int]("counter")) - .map(value => assertEquals(value, Some(12))) - } - } - } - - test("automatic renewal extends the lease and excludes another client") { - withClients { (first, second) => - first - .lock[String](leaseDuration = 900.millis) - .withLock("long", 2.seconds) { - awaitRenewal(second, "4:lock:long") - .flatMap(_ => second.lock[String]().tryWithLock("long")(CIO.value(42))) - .map(result => assertEquals(result, None)) } - .flatMap(_ => second.lock[String]().tryWithLock("long")(CIO.value(42))) - .map(result => assertEquals(result, Some(42))) + .flatMap(_ => first.get[Int]("counter")) + .is(Some(12)) } } - test("a standalone lock survives connection loss during renewal") { - withClients { (first, second) => - val holder = first.lock[String](leaseDuration = 1500.millis) - val contender = second.lock[String]() - first.clientId.flatMap { id => - holder - .withLock("reconnect", 2.seconds) { - second - .info("commandstats") - .flatMap { before => - second.pipeline((killClient(id), pauseClients)).flatMap { _ => - Eventually.converges(30, 50.millis)(() => second.info("commandstats"))( - commandCalls(_, "evalsha") > commandCalls(before, "evalsha") - )(info => s"the lock was not renewed after reconnecting: $info") - } - } - .flatMap(_ => contender.tryWithLock("reconnect")(CIO.value(1))) - .map(result => assertEquals(result, None)) - } - .flatMap(_ => contender.tryWithLock("reconnect")(CIO.value(42))) - .map(result => assertEquals(result, Some(42))) + clientsTest("automatic renewal extends the lease and excludes another client") { (first, second) => + first + .lock[String](leaseDuration = 900.millis) + .withLock("long", 2.seconds) { + awaitRenewal(second, "4:lock:long") + .flatMap(_ => second.lock[String]().tryWithLock("long")(CIO.value(42))) + .is(None) } - } + .flatMap(_ => second.lock[String]().tryWithLock("long")(CIO.value(42))) + .is(Some(42)) } - test("expired ownership cannot renew or remove a replacement owner's lock") { - withClients { (first, second) => - first - .lock[String](leaseDuration = 300.millis) - .tryWithLock("replaced") { - second.set("4:lock:replaced", "new-owner").flatMap(_ => CIO.never) - } - .liftToTry - .flatMap { result => - assert(result.failed.get.isInstanceOf[LockLost], result.toString) - second.get[String]("4:lock:replaced").map(value => assertEquals(value, Some("new-owner"))) + clientsTest("a standalone lock survives connection loss during renewal") { (first, second) => + val holder = first.lock[String](leaseDuration = 1500.millis) + val contender = second.lock[String]() + first.clientId.flatMap { id => + holder + .withLock("reconnect", 2.seconds) { + awaitCalls(second.info("commandstats"), "evalsha", 1)(renewed => + second.pipeline((admin("CLIENT", "KILL", "ID", id.toString), admin("CLIENT", "PAUSE", "600", "ALL"))) >> renewed + ) + .flatMap(_ => contender.tryWithLock("reconnect")(CIO.value(1))) + .is(None) } + .flatMap(_ => contender.tryWithLock("reconnect")(CIO.value(42))) + .is(Some(42)) } } - test("body failure releases the lease and preserves the error") { - withClient { client => - val failure = new IllegalStateException("body failed") - val lock = client.lock[String]() - lock.withLock("failure", 2.seconds)(CIO.fail(failure)).liftToTry.flatMap { result => - assert(result.failed.get eq failure) - lock.tryWithLock("failure")(CIO.value(42)).map(value => assertEquals(value, Some(42))) + clientsTest("expired ownership cannot renew or remove a replacement owner's lock") { (first, second) => + failsWith[LockLost]( + first.lock[String](leaseDuration = 300.millis).tryWithLock("replaced") { + second.set("4:lock:replaced", "new-owner").flatMap(_ => CIO.never) } - } + ).flatMap(_ => second.get[String]("4:lock:replaced").is(Some("new-owner"))) } - test("script cache flush during a scope is recovered by renewal and release") { - withClient { client => - client - .scriptFlush() - .flatMap { _ => - client.lock[String](leaseDuration = 900.millis).withLock("flush", 2.seconds) { - client - .scriptFlush() - .flatMap(_ => awaitRenewal(client, "4:lock:flush")) - .flatMap(_ => client.scriptFlush()) - } - } - .flatMap(_ => client.exists("4:lock:flush")) - .map(count => assertEquals(count, 0L)) + clientTest("body failure releases the lease and preserves the error") { client => + val failure = new IllegalStateException("body failed") + val lock = client.lock[String]() + failsWith[IllegalStateException](lock.withLock("failure", 2.seconds)(CIO.fail(failure))).flatMap { error => + assert(error eq failure) + lock.tryWithLock("failure")(CIO.value(42)).is(Some(42)) } } - test("different keys and namespaces remain independent, including through client.as") { - withClient { client => - client - .lock[String](namespace = "one") - .tryWithLock("key") { - for { - differentKey <- client.lock[String](namespace = "one").tryWithLock("other")(CIO.value(1)) - differentNamespace <- client.lock[String](namespace = "two").tryWithLock("key")(CIO.value(2)) - same <- client.as[Array[Byte]].lock[Array[Byte]](namespace = "one").tryWithLock("key".getBytes("UTF-8"))(CIO.value(3)) - } yield { - assertEquals(differentKey, Some(1)) - assertEquals(differentNamespace, Some(2)) - assertEquals(same, None) - } + clientTest("script cache flush during a scope is recovered by renewal and release") { client => + client + .scriptFlush() + .flatMap { _ => + client.lock[String](leaseDuration = 900.millis).withLock("flush", 2.seconds) { + client + .scriptFlush() + .flatMap(_ => awaitRenewal(client, "4:lock:flush")) + .flatMap(_ => client.scriptFlush()) } - .unit - } + } + .flatMap(_ => client.exists("4:lock:flush")) + .is(0L) } -} -class RedisLockSuite extends LockSuite(Images.redis) -class ValkeyLockSuite extends LockSuite(Images.valkey) + clientTest("different keys and namespaces remain independent, including through client.as") { client => + client + .lock[String](namespace = "one") + .tryWithLock("key") { + for { + differentKey <- client.lock[String](namespace = "one").tryWithLock("other")(CIO.value(1)) + differentNamespace <- client.lock[String](namespace = "two").tryWithLock("key")(CIO.value(2)) + same <- client.as[Array[Byte]].lock[Array[Byte]](namespace = "one").tryWithLock("key".getBytes("UTF-8"))(CIO.value(3)) + } yield (differentKey, differentNamespace, same) + } + .is(Some((Some(1), Some(2), None))) + } +} diff --git a/integration-tests/shared/src/test/scala/sage/integration/masterreplica/MasterReplicaSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/masterreplica/MasterReplicaSuite.scala index d5c8a6a0..190c477b 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/masterreplica/MasterReplicaSuite.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/masterreplica/MasterReplicaSuite.scala @@ -1,19 +1,18 @@ package sage.integration.masterreplica -import scala.concurrent.ExecutionContext import scala.concurrent.duration.* import com.dimafeng.testcontainers.GenericContainer -import com.dimafeng.testcontainers.munit.TestContainerForAll +import com.dimafeng.testcontainers.munit.{TestContainersForAll, TestContainersForEach} import kyo.compat.* +import munit.{Location, TestOptions} -import sage.{Bytes, Message} -import sage.SageException.{DecodeError, LockLost, TimedOut} +import sage.Message +import sage.SageException.{LockLost, TimedOut} import sage.client.{Endpoint, MasterReplicaConfig, ReadFrom, SageConfig, Topology} import sage.client.internal.Client -import sage.commands.{Command, Commands} -import sage.integration.{ContainerClient, Eventually, Images, Ttls} -import sage.protocol.Frame +import sage.commands.Commands +import sage.integration.{ContainerClient, Eventually, Images} /** * Shared setup for the master-replica suites. One container runs a master on port 6379 and a replica on port 6380. Like @@ -21,475 +20,208 @@ import sage.protocol.Frame * replica-pool endpoint, [[ReadFrom]] routing, and replication. The replica advertises the host and port mapped by Testcontainers so that * the test can reach the address from the master's `ROLE` reply. */ -abstract class MasterReplicaSuiteBase(image: String, serverBinary: String) extends munit.FunSuite with TestContainerForAll with ContainerClient { +abstract class MasterReplicaSuiteBase(image: String, serverBinary: String) extends ContainerClient { - // Run the process that survives each fault in the foreground so the container remains alive. The fault setup chooses whether that process - // is the master or replica. - protected def masterRunsInForeground: Boolean = false - - // --protected-mode no admits the testcontainers-mapped (non-loopback) connection; --save '' / --appendonly no keep the nodes in-memory - override val containerDef: GenericContainer.Def[GenericContainer] = { - val master = s"$serverBinary --port 6379 --save '' --appendonly no --protected-mode no --repl-diskless-sync-delay 0" - val replica = s"$serverBinary --port 6380 --save '' --appendonly no --protected-mode no --repl-diskless-sync-delay 0" - val (background, foreground) = if (masterRunsInForeground) (replica, master) else (master, replica) - GenericContainer.Def(image, exposedPorts = Seq(6379, 6380), command = Seq("sh", "-c", s"$background & exec $foreground")) - } - - given ExecutionContext = munitExecutionContext + override type Containers = GenericContainer protected val masterPort = 6379 protected val replicaPort = 6380 protected val marker = "mr:replica-only-marker" - protected def admin(name: String, args: String*): Command[Unit] = - Command(name, Command.NoKeys, args.toVector.map(Bytes.utf8), _ => Right(())) - - protected val infoReplication: Command[String] = - Command( - "INFO", - Command.NoKeys, - Vector(Bytes.utf8("replication")), - { - case Frame.BulkString(bytes) => Right(bytes.asUtf8String) - case Frame.VerbatimString(_, bytes) => Right(bytes.asUtf8String) - case other => Left(DecodeError("bulk or verbatim string", Frame.describe(other))) - } - ) - - protected def commandCalls(info: String, command: String): Long = - info.linesIterator - .find(_.startsWith(s"cmdstat_${command.toLowerCase}:calls=")) - .fold(0L)(_.dropWhile(_ != '=').drop(1).takeWhile(_ != ',').toLong) - - // Follow the master and announce the host-mapped port, then write a marker key the master never has, so a read's origin is observable. The - // marker goes in after the link is up, since REPLICAOF triggers a full resync that would wipe an earlier write. Idempotent across a suite's tests. - protected def ensureReplicating(replica: Client[CIO, String], announceHost: String, announcePort: Int): CIO[Unit] = - replica.run(infoReplication).flatMap { info => - if (info.contains("master_link_status:up")) CIO.value(()) - else - for { - _ <- replica.run(admin("CONFIG", "SET", "replica-announce-ip", announceHost)) - _ <- replica.run(admin("CONFIG", "SET", "replica-announce-port", announcePort.toString)) - _ <- replica.run(admin("REPLICAOF", "127.0.0.1", masterPort.toString)) - // Valkey's async full-sync takes a few seconds even for an empty dataset, so budget generously under CI load - _ <- awaitLinkUp(replica, 150) - _ <- replica.run(admin("CONFIG", "SET", "replica-read-only", "no")) - _ <- replica.run(admin("SET", marker, "from-replica")) - // restore read-only so a misrouted write to the replica fails loudly; replication still applies (it bypasses read-only) and the marker persists - _ <- replica.run(admin("CONFIG", "SET", "replica-read-only", "yes")) - } yield () - } - - protected def awaitLinkUp(replica: Client[CIO, String], attempts: Int): CIO[Unit] = - replica.run(infoReplication).flatMap { info => - if (info.contains("master_link_status:up")) CIO.value(()) - else if (attempts <= 0) CIO.fail(new RuntimeException(s"replica did not sync: $info")) - else CIO.sleep(100.millis).flatMap(_ => awaitLinkUp(replica, attempts - 1)) - } + // --protected-mode no admits the testcontainers-mapped (non-loopback) connection; --save '' / --appendonly no keep the nodes in-memory. + // Both servers run in the background, so a fault can shut down either one without stopping the container. + override def startContainers(): GenericContainer = { + def server(port: Int) = s"$serverBinary --port $port --save '' --appendonly no --protected-mode no --repl-diskless-sync-delay 0" + serverDef(image, Seq(masterPort, replicaPort), Seq("sh", "-c", s"${server(masterPort)} & ${server(replicaPort)} & exec tail -f /dev/null")) + .start() + } - // replication is async; poll until the replica catches up - protected def awaitValue(client: Client[CIO, String], key: String, expected: String, attempts: Int): CIO[Option[String]] = - client.get[String](key).flatMap { - case Some(v) if v == expected => CIO.value(Some(v)) - case _ if attempts <= 0 => CIO.value(None) - case _ => CIO.sleep(100.millis).flatMap(_ => awaitValue(client, key, expected, attempts - 1)) - } + final protected class Deployment(server: GenericContainer) { + private def at(port: Int) = Endpoint(server.host, server.mappedPort(port)) + val master: SageConfig = SageConfig(topology = Topology.Standalone(at(masterPort))) + val replica: SageConfig = SageConfig(topology = Topology.Standalone(at(replicaPort))) - // Ignore the result of commands such as SHUTDOWN, which report a connection failure because the server closes the socket. Later checks - // verify their effect. - protected def ignoreFailure(action: CIO[Unit]): CIO[Unit] = - action.fold(_ => CIO.value(()), _ => CIO.value(())) + def reading(policy: ReadFrom, minRefreshInterval: FiniteDuration = 5.seconds): SageConfig = + SageConfig( + topology = Topology.MasterReplica(Vector(at(masterPort)), MasterReplicaConfig(minRefreshInterval = minRefreshInterval)), + readFrom = policy + ) + } - protected def standalone(host: String, port: Int): SageConfig = - SageConfig(topology = Topology.Standalone(Endpoint(host, port))) + protected def replicationTest(options: TestOptions)(body: Deployment => CIO[Any])(using Location): Unit = + containerTest(options)(server => body(new Deployment(server))) - protected def masterReplica(host: String, masterMappedPort: Int, policy: ReadFrom, minRefreshInterval: FiniteDuration = 5.seconds): SageConfig = - SageConfig( - topology = Topology.MasterReplica(Vector(Endpoint(host, masterMappedPort)), MasterReplicaConfig(minRefreshInterval = minRefreshInterval)), - readFrom = policy - ) + // Follow the master and announce the host-mapped port, then write a marker key the master never has, so a read's origin is observable. The + // marker goes in after the link is up, since REPLICAOF triggers a full resync that would wipe an earlier write. + override def afterContainersStart(server: GenericContainer): Unit = + prepare(connectAndUse(new Deployment(server).replica) { replica => + replica.configSet("replica-announce-ip" -> server.host, "replica-announce-port" -> server.mappedPort(replicaPort).toString) >> + replica.run(admin("REPLICAOF", "127.0.0.1", masterPort.toString)) >> + // Valkey's async full-sync takes a few seconds even for an empty dataset, so budget generously under CI load + Eventually(150)(replica.role.satisfies(_.isConnectedReplica)) >> + replica.configSet("replica-read-only" -> "no") >> + replica.set(marker, "from-replica") >> + // restore read-only so a misrouted write to the replica fails loudly; replication still applies (it bypasses read-only) and the marker persists + replica.configSet("replica-read-only" -> "yes") + }) } /** * Read routing and master-pinned operations against a real master-replica deployment. Routing is proven with a marker key that lives only on * the replica: a read that sees it was served by the replica, and one that does not was served by the master, which never has it. */ -abstract class MasterReplicaSuite(image: String, serverBinary: String) extends MasterReplicaSuiteBase(image, serverBinary) { - - for (duringRenewal <- Vector(false, true)) - test(s"distributed locks fail when a replica stops acknowledging ${if (duringRenewal) "renewal" else "acquisition"}") { - withContainers { server => - val host = server.host - val pm = server.mappedPort(masterPort) - val pr = server.mappedPort(replicaPort) - connectAndUse(standalone(host, pr)) { replica => - ensureReplicating(replica, host, pr).flatMap { _ => - connectAndUse(masterReplica(host, pm, ReadFrom.Replica)) { client => - connectAndUse(standalone(host, pm)) { master => - val key = s"mr:unconfirmed-lock:$duringRenewal" - val bodyStarted = new java.util.concurrent.atomic.AtomicBoolean(false) - val bodyStopped = new java.util.concurrent.atomic.AtomicBoolean(false) - val pause = replica.run(admin("CLIENT", "PAUSE", "800", "ALL")) - val attempt = client.lock[String](600.millis).tryWithLock(key) { - CIO.defer(bodyStarted.set(true)).flatMap { _ => - if (duringRenewal) - CIO.ensure(CIO.defer(bodyStopped.set(true)))(pause.flatMap(_ => CIO.never)) - else CIO.unit - } - } - CIO.ensure(replica.run(admin("CLIENT", "UNPAUSE"))) { - for { - _ <- if (duringRenewal) CIO.unit else pause - result <- attempt.liftToTry - exists <- master.exists(s"4:lock:$key") - } yield { - if (duringRenewal) { - assert(result.failed.get.isInstanceOf[LockLost], result.toString) - assert(bodyStarted.get()) - assert(bodyStopped.get()) - } else { - assert(result.failed.get.isInstanceOf[TimedOut], result.toString) - assert(result.failed.get.getMessage.contains("replication confirmed by"), result.toString) - assert(!bodyStarted.get()) - } - assertEquals(exists, 0L) - } - } - } - } +abstract class MasterReplicaSuite(image: String, serverBinary: String) extends MasterReplicaSuiteBase(image, serverBinary) with TestContainersForAll { + + private def lockFailsWhenReplicaStops(phase: String)(scenario: (Client[CIO, String], String, CIO[Unit]) => CIO[Unit]): Unit = + replicationTest(s"distributed locks fail when a replica stops acknowledging $phase") { d => + connectAndUse(d.replica) { replica => + connectAndUse(d.reading(ReadFrom.Replica)) { client => + connectAndUse(d.master) { master => + val key = s"mr:unconfirmed-lock:$phase" + CIO + .ensure(replica.run(admin("CLIENT", "UNPAUSE")))(scenario(client, key, replica.run(admin("CLIENT", "PAUSE", "800", "ALL")))) + .flatMap(_ => master.exists(s"4:lock:$key").is(0L)) } - }.unsafeRun + } } } - test("distributed locks acquire, renew, and release on the master with Replica reads") { - withContainers { server => - val host = server.host - val pm = server.mappedPort(masterPort) - val pr = server.mappedPort(replicaPort) - - val program = - connectAndUse(standalone(host, pr))(ensureReplicating(_, host, pr)).flatMap { _ => - connectAndUse(masterReplica(host, pm, ReadFrom.Replica)) { client => - connectAndUse(standalone(host, pm)) { master => - val key = "mr:distributed-lock" - val holder = client.lock[String](leaseDuration = 900.millis) - val contender = master.lock[String]() - for { - fromReplica <- client.get[String](marker) - waitsBefore <- master.info("commandstats") - denied <- holder.withLock(key, 2.seconds) { - for { - before <- contender.tryWithLock(key)(CIO.value(1)) - _ <- Eventually.converges(30, 50.millis)(() => master.info("commandstats"))( - commandCalls(_, "wait") >= commandCalls(waitsBefore, "wait") + 2 - )(info => s"acquisition and renewal did not both wait for replication: $info") - after <- contender.tryWithLock(key)(CIO.value(2)) - } yield (before, after) - } - acquired <- contender.tryWithLock(key)(CIO.value(42)) - exists <- master.exists(s"4:lock:$key") - } yield { - assertEquals(fromReplica, Some("from-replica")) - assertEquals(denied, (None, None)) - assertEquals(acquired, Some(42)) - assertEquals(exists, 0L) - } - } - } - } - program.unsafeRun - } + lockFailsWhenReplicaStops("acquisition") { (client, key, pause) => + pause + .flatMap(_ => + failsWith[TimedOut](client.lock[String](600.millis).tryWithLock(key)(CIO.fail(new AssertionError("body ran without a confirmed lock")))) + ) + .map(e => assert(e.getMessage.contains("replication confirmed by"), e.getMessage)) } - test("distributed locks can skip replica acknowledgement") { - withContainers { server => - val host = server.host - val pm = server.mappedPort(masterPort) - val pr = server.mappedPort(replicaPort) - - connectAndUse(standalone(host, pr))(ensureReplicating(_, host, pr)).flatMap { _ => - connectAndUse(masterReplica(host, pm, ReadFrom.Replica)) { client => - connectAndUse(standalone(host, pm)) { master => - for { - before <- master.info("commandstats") - result <- client - .lock[String](leaseDuration = 900.millis, replicaAcknowledgement = false) - .tryWithLock("mr:no-replica-ack") { - Eventually - .changes(30, 50.millis)(() => master.pTtl("4:lock:mr:no-replica-ack"))((before, current) => - (Ttls.remaining(before), Ttls.remaining(current)) match { - case (Some(previous), Some(renewed)) => renewed > previous + 100.millis - case _ => false - } - )(ttl => s"lock was not renewed; its TTL was $ttl") - .map(_ => 42) - } - after <- master.info("commandstats") - } yield { - assertEquals(result, Some(42)) - assertEquals(commandCalls(after, "wait"), commandCalls(before, "wait")) - assert(commandCalls(after, "role") > commandCalls(before, "role")) - } - } - } - }.unsafeRun - } + lockFailsWhenReplicaStops("renewal") { (client, key, pause) => + val stopped = new java.util.concurrent.atomic.AtomicBoolean(false) + failsWith[LockLost]( + client.lock[String](600.millis).tryWithLock(key)(CIO.ensure(CIO.defer(stopped.set(true)))(pause.flatMap(_ => CIO.never))) + ) + .map(_ => assert(stopped.get(), "the body was not interrupted")) } - test("reads honor the ReadFrom policy and writes always reach the master") { - withContainers { server => - val host = server.host - val pm = server.mappedPort(masterPort) - val pr = server.mappedPort(replicaPort) - val replicaCfg = standalone(host, pr) - - val program = - connectAndUse(replicaCfg)(ensureReplicating(_, host, pr)) - .flatMap { _ => - connectAndUse(masterReplica(host, pm, ReadFrom.Replica)) { client => - for { - fromReplica <- client.get[String](marker) - // the write always goes to the master; reading it back off the replica proves both replication and replica routing - _ <- client.set("mr:k", "v") - replicated <- awaitValue(client, "mr:k", "v", 50) - } yield { - assertEquals(fromReplica, Some("from-replica")) - assertEquals(replicated, Some("v")) - } - } - } - .flatMap(_ => connectAndUse(masterReplica(host, pm, ReadFrom.Master))(_.get[String](marker).map(assertEquals(_, None)))) - .flatMap(_ => - connectAndUse(masterReplica(host, pm, ReadFrom.ReplicaPreferred))(_.get[String](marker).map(assertEquals(_, Some("from-replica")))) - ) - .flatMap(_ => connectAndUse(masterReplica(host, pm, ReadFrom.MasterPreferred))(_.get[String](marker).map(assertEquals(_, None)))) - program.unsafeRun + replicationTest("distributed locks acquire, renew, and release on the master with Replica reads") { d => + connectAndUse(d.reading(ReadFrom.Replica)) { client => + connectAndUse(d.master) { master => + client.get[String](marker).is(Some("from-replica")) >> + awaitCalls(master.info("commandstats"), "wait", 2)(contend(client.lock(leaseDuration = 900.millis), master, "mr:distributed-lock", _)) + } } } - test("transactions and pub/sub run on the master under the master-replica runtime") { - withContainers { server => - val host = server.host - val pm = server.mappedPort(masterPort) - val pr = server.mappedPort(replicaPort) - val replicaCfg = standalone(host, pr) - - val program = - connectAndUse(replicaCfg)(ensureReplicating(_, host, pr)).flatMap { _ => - connectAndUse(masterReplica(host, pm, ReadFrom.ReplicaPreferred)) { client => - for { - commit <- client.transaction(tx => tx.exec(Vector(Commands.incr[String]("mr:c"), Commands.incr[String]("mr:c")))) - sub <- client.subscribeChannels[String]("mr:news") - count <- client.publish("mr:news", "hello") - message <- sub.next - _ <- sub.close - } yield { - assertEquals(commit, Some(Vector(1L, 2L))) - assertEquals(count, 1L) - assertEquals(message, Some(Message("mr:news", "hello"))) - } - } + replicationTest("distributed locks can skip replica acknowledgement") { d => + connectAndUse(d.reading(ReadFrom.Replica)) { client => + connectAndUse(d.master) { master => + for { + before <- master.info("commandstats") + _ <- client + .lock[String](leaseDuration = 900.millis, replicaAcknowledgement = false) + .tryWithLock("mr:no-replica-ack")(awaitRenewal(master, "4:lock:mr:no-replica-ack").map(_ => 42)) + .is(Some(42)) + after <- master.info("commandstats") + } yield { + assertEquals(commandCalls(after, "wait"), commandCalls(before, "wait")) + assert(commandCalls(after, "role") > commandCalls(before, "role")) } - program.unsafeRun + } } } - test("pub/sub works when a subscription is the client's first operation") { - withContainers { server => - val host = server.host - val pm = server.mappedPort(masterPort) - val pr = server.mappedPort(replicaPort) - - val program = - connectAndUse(standalone(host, pr))(ensureReplicating(_, host, pr)).flatMap { _ => - connectAndUse(masterReplica(host, pm, ReadFrom.ReplicaPreferred)) { client => - for { - sub <- client.subscribeChannels[String]("mr:first") - count <- client.publish("mr:first", "hello") - message <- sub.next - _ <- sub.close - } yield { - assertEquals(count, 1L) - assertEquals(message, Some(Message("mr:first", "hello"))) - } - } - } - program.unsafeRun + replicationTest("reads honor the ReadFrom policy and writes always reach the master") { d => + connectAndUse(d.reading(ReadFrom.Replica)) { client => + client.get[String](marker).is(Some("from-replica")) >> + // the write always goes to the master; reading it back off the replica proves both replication and replica routing + client.set("mr:k", "v") >> + Eventually(50)(client.get[String]("mr:k").is(Some("v"))) } + .flatMap(_ => connectAndUse(d.reading(ReadFrom.Master))(_.get[String](marker).is(None))) + .flatMap(_ => connectAndUse(d.reading(ReadFrom.ReplicaPreferred))(_.get[String](marker).is(Some("from-replica")))) + .flatMap(_ => connectAndUse(d.reading(ReadFrom.MasterPreferred))(_.get[String](marker).is(None))) } -} -/** - * Tests recovery after promoting the replica and taking the old master out of write service. The first write to the old master triggers - * role discovery, and the caller retries against the promoted node. [[induceFailover]] runs the test once for a `READONLY` demotion and - * once for a connection loss. Each test uses its own container because promotion permanently changes the topology. - */ -abstract class MasterReplicaFailoverSuite(image: String, serverBinary: String, fault: String) extends MasterReplicaSuiteBase(image, serverBinary) { - - // promote the replica, then take the old master out of write service in a fault-specific way; the client is told nothing - protected def induceFailover(replicaCfg: SageConfig, masterCfg: SageConfig): CIO[Unit] + replicationTest("transactions and pub/sub run on the master under the master-replica runtime") { d => + connectAndUse(d.reading(ReadFrom.ReplicaPreferred)) { client => + client.transaction(tx => tx.exec(Vector(Commands.incr[String]("mr:c"), Commands.incr[String]("mr:c")))).is(Some(Vector(1L, 2L))) >> + withSubscription(client.subscribeChannels[String]("mr:news"))(sub => + client.publish("mr:news", "hello").is(1L) >> sub.next.is(Some(Message("mr:news", "hello"))) + ) + } + } - // retry the write until it succeeds; the first attempt reaches the old master and starts discovery, and a later attempt reaches the promoted master - private def writeUntilAccepted(client: Client[CIO, String], key: String, value: String, attempts: Int): CIO[Boolean] = - client - .set(key, value) - .fold( - _ => CIO.value(true), - _ => if (attempts <= 0) CIO.value(false) else CIO.sleep(100.millis).flatMap(_ => writeUntilAccepted(client, key, value, attempts - 1)) + replicationTest("pub/sub works when a subscription is the client's first operation") { d => + connectAndUse(d.reading(ReadFrom.ReplicaPreferred)) { client => + withSubscription(client.subscribeChannels[String]("mr:first"))(sub => + client.publish("mr:first", "hello").is(1L) >> sub.next.is(Some(Message("mr:first", "hello"))) ) - - test(s"the client recovers writes after the replica is promoted to master ($fault)") { - withContainers { server => - val host = server.host - val pm = server.mappedPort(masterPort) - val pr = server.mappedPort(replicaPort) - val replicaCfg = standalone(host, pr) - val masterCfg = standalone(host, pm) - // short refresh interval so the event-driven re-discovery is not throttled away during the retry window - val mrCfg = masterReplica(host, pm, ReadFrom.Master, minRefreshInterval = 100.millis) - - val program = - connectAndUse(replicaCfg)(ensureReplicating(_, host, pr)).flatMap { _ => - connectAndUse(mrCfg) { client => - for { - _ <- client.set("fo:before", "v1") - before <- client.get[String]("fo:before") - _ <- induceFailover(replicaCfg, masterCfg) - recovered <- writeUntilAccepted(client, "fo:after", "v2", 50) - after <- client.get[String]("fo:after") - } yield { - assertEquals(before, Some("v1")) - assert(recovered, "write never recovered after promotion") - assertEquals(after, Some("v2")) - } - } - } - program.unsafeRun } } } /** - * A failover that makes the old master follow the new master. The old master remains reachable but answers writes with `READONLY`, which - * tests topology discovery after an ownership failure. - */ -abstract class MasterReplicaDemotionFailoverSuite(image: String, serverBinary: String) - extends MasterReplicaFailoverSuite(image, serverBinary, "demoted master") { - - protected def induceFailover(replicaCfg: SageConfig, masterCfg: SageConfig): CIO[Unit] = - for { - _ <- connectAndUse(replicaCfg)(_.run(admin("REPLICAOF", "NO", "ONE"))) - _ <- connectAndUse(masterCfg)(_.run(admin("REPLICAOF", "127.0.0.1", replicaPort.toString))) - } yield () -} - -/** - * Failover where the old master stops with `SHUTDOWN`. The next write receives a connection-refused error, exercising role discovery after a - * connection loss. The master runs in the background, so the foreground replica keeps the container alive. + * Faults that permanently change the deployment, so each test starts its own container. */ -abstract class MasterReplicaConnectionLossFailoverSuite(image: String, serverBinary: String) - extends MasterReplicaFailoverSuite(image, serverBinary, "master down") { - - protected def induceFailover(replicaCfg: SageConfig, masterCfg: SageConfig): CIO[Unit] = - for { - _ <- connectAndUse(replicaCfg)(_.run(admin("REPLICAOF", "NO", "ONE"))) - _ <- ignoreFailure(connectAndUse(masterCfg)(_.run(admin("SHUTDOWN", "NOSAVE")))) - } yield () -} - -/** - * Tests a replica shutdown. The replica runs in the background, so shutting it down leaves the foreground master and the container alive. - * A strict `Replica` read then fails because no replica is available. `ReplicaPreferred` falls back to the master. Both clients connect - * before the shutdown, which makes the test use the failed replica instead of omitting it during discovery. - */ -abstract class MasterReplicaReplicaDownSuite(image: String, serverBinary: String) extends MasterReplicaSuiteBase(image, serverBinary) { - - override protected def masterRunsInForeground: Boolean = true - - test("a strict Replica read fails when the replica is down, while ReplicaPreferred falls back to the master") { - withContainers { server => - val host = server.host - val pm = server.mappedPort(masterPort) - val pr = server.mappedPort(replicaPort) - val replicaCfg = standalone(host, pr) +abstract class MasterReplicaFaultSuite(image: String, serverBinary: String) + extends MasterReplicaSuiteBase(image, serverBinary) + with TestContainersForEach { + + // Promotes the replica, then takes the old master out of write service without telling the client. The first write to the old master + // triggers role discovery, and the caller retries against the promoted node. + private def recoversAfterPromotion(fault: String)(retireOldMaster: Deployment => CIO[Unit]): Unit = + replicationTest(s"the client recovers writes after the replica is promoted to master ($fault)") { d => + // short refresh interval so the event-driven re-discovery is not throttled away during the retry window + connectAndUse(d.reading(ReadFrom.Master, minRefreshInterval = 100.millis)) { client => + client.set("fo:before", "v1") >> + client.get[String]("fo:before").is(Some("v1")) >> + connectAndUse(d.replica)(_.run(admin("REPLICAOF", "NO", "ONE"))) >> + retireOldMaster(d) >> + // the first attempt reaches the old master and starts discovery, and a later attempt reaches the promoted master + Eventually.succeeds(50)(client.set("fo:after", "v2")) >> + client.get[String]("fo:after").is(Some("v2")) + } + } - val program = - connectAndUse(replicaCfg)(ensureReplicating(_, host, pr)).flatMap { _ => - connectAndUse(masterReplica(host, pm, ReadFrom.Replica)) { strict => - connectAndUse(masterReplica(host, pm, ReadFrom.ReplicaPreferred)) { preferred => - for { - // a master-backed key for the fallback read; the replica-only marker for the warm-up reads - _ <- preferred.set("rd:k", "v") - warmStrict <- strict.get[String](marker) - warmPreferred <- preferred.get[String](marker) - _ <- ignoreFailure(connectAndUse(replicaCfg)(_.run(admin("SHUTDOWN", "NOSAVE")))) - strictFailed <- strict.get[String]("rd:k").fold(_ => CIO.value(false), _ => CIO.value(true)) - fallback <- awaitValue(preferred, "rd:k", "v", 50) - } yield { - assertEquals(warmStrict, Some("from-replica")) - assertEquals(warmPreferred, Some("from-replica")) - assert(strictFailed, "strict Replica read should fail when no replica is reachable") - assertEquals(fallback, Some("v")) - } - } - } - } - program.unsafeRun + // The old master follows the new master and answers writes with `READONLY`, which tests topology discovery after an ownership failure. + recoversAfterPromotion("demoted master")(d => connectAndUse(d.master)(_.run(admin("REPLICAOF", "127.0.0.1", replicaPort.toString)))) + + // The old master stops, so the next write gets a connection-refused error and discovery runs after a connection loss. + recoversAfterPromotion("master down")(d => shutdown(d.master)) + + // Both clients connect before the shutdown, so they use the failed replica instead of omitting it during discovery. + replicationTest("a strict Replica read fails when the replica is down, while ReplicaPreferred falls back to the master") { d => + connectAndUse(d.reading(ReadFrom.Replica)) { strict => + connectAndUse(d.reading(ReadFrom.ReplicaPreferred)) { preferred => + // a master-backed key for the fallback read; the replica-only marker for the warm-up reads + preferred.set("rd:k", "v") >> + strict.get[String](marker).is(Some("from-replica")) >> + preferred.get[String](marker).is(Some("from-replica")) >> + shutdown(d.replica) >> + failsWith[Throwable](strict.get[String]("rd:k")) >> + Eventually(50)(preferred.get[String]("rd:k").is(Some("v"))) + } } } -} -/** - * A replica configured with `replica-serve-stale-data no`. When its replication link breaks, it remains reachable but answers reads with - * `-MASTERDOWN`. The replica-preferred client must handle that response by retrying the read on the master. Both clients connect before the - * link breaks, ensuring that they first try the stale replica. - */ -abstract class MasterReplicaStaleReplicaSuite(image: String, serverBinary: String) extends MasterReplicaSuiteBase(image, serverBinary) { - - // 6399 listens to nothing, so the link stays down while the replica keeps answering + // With `replica-serve-stale-data no` and a broken link, the replica stays reachable but answers reads with `-MASTERDOWN`. 6399 listens to + // nothing, so the link stays down. private def refuseStaleReads(replica: Client[CIO, String]): CIO[Unit] = - for { - _ <- replica.run(admin("CONFIG", "SET", "replica-serve-stale-data", "no")) - _ <- replica.run(admin("REPLICAOF", "127.0.0.1", "6399")) - _ <- awaitLinkDown(replica, 100) - } yield () - - private def awaitLinkDown(replica: Client[CIO, String], attempts: Int): CIO[Unit] = - Eventually.converges(attempts)(() => replica.run(infoReplication))(_.contains("master_link_status:down"))(info => - s"replica link never went down: $info" - ) - - test("a replica answering MASTERDOWN falls through to the master under ReplicaPreferred, and fails a strict Replica read") { - withContainers { server => - val host = server.host - val pm = server.mappedPort(masterPort) - val pr = server.mappedPort(replicaPort) - val replicaCfg = standalone(host, pr) - - val program = - connectAndUse(replicaCfg)(ensureReplicating(_, host, pr)).flatMap { _ => - connectAndUse(masterReplica(host, pm, ReadFrom.Replica)) { strict => - connectAndUse(masterReplica(host, pm, ReadFrom.ReplicaPreferred)) { preferred => - for { - warmStrict <- strict.get[String](marker) - warmPreferred <- preferred.get[String](marker) - _ <- connectAndUse(replicaCfg)(refuseStaleReads) - _ <- preferred.set("sd:k", "v") - fallback <- preferred.get[String]("sd:k") - batchFallback <- preferred.pipeline(Seq.fill(2)(Commands.get[String, String]("sd:k"))) - strictFailed <- strict.get[String]("sd:k").fold(_ => CIO.value(false), _ => CIO.value(true)) - } yield { - assertEquals(warmStrict, Some("from-replica")) - assertEquals(warmPreferred, Some("from-replica")) - assertEquals(fallback, Some("v")) - assertEquals(batchFallback, Vector(Some("v"), Some("v"))) - assert(strictFailed, "strict Replica read should fail when the only replica refuses to serve stale data") - } - } - } - } - program.unsafeRun + replica.configSet("replica-serve-stale-data" -> "no") >> + replica.run(admin("REPLICAOF", "127.0.0.1", "6399")) >> + Eventually(100)(replica.role.satisfies(!_.isConnectedReplica)) + + // Both clients connect before the link breaks, so they first try the stale replica. + replicationTest("a replica answering MASTERDOWN falls through to the master under ReplicaPreferred, and fails a strict Replica read") { d => + connectAndUse(d.reading(ReadFrom.Replica)) { strict => + connectAndUse(d.reading(ReadFrom.ReplicaPreferred)) { preferred => + strict.get[String](marker).is(Some("from-replica")) >> + preferred.get[String](marker).is(Some("from-replica")) >> + connectAndUse(d.replica)(refuseStaleReads) >> + preferred.set("sd:k", "v") >> + preferred.get[String]("sd:k").is(Some("v")) >> + preferred.pipeline(Seq.fill(2)(Commands.get[String, String]("sd:k"))).is(Vector(Some("v"), Some("v"))) >> + failsWith[Throwable](strict.get[String]("sd:k")) + } } } } @@ -498,18 +230,6 @@ class RedisMasterReplicaSuite extends MasterReplicaSuite(Images.redis, "redis-se class ValkeyMasterReplicaSuite extends MasterReplicaSuite(Images.valkey, "valkey-server") -class RedisMasterReplicaDemotionFailoverSuite extends MasterReplicaDemotionFailoverSuite(Images.redis, "redis-server") - -class ValkeyMasterReplicaDemotionFailoverSuite extends MasterReplicaDemotionFailoverSuite(Images.valkey, "valkey-server") - -class RedisMasterReplicaConnectionLossFailoverSuite extends MasterReplicaConnectionLossFailoverSuite(Images.redis, "redis-server") - -class ValkeyMasterReplicaConnectionLossFailoverSuite extends MasterReplicaConnectionLossFailoverSuite(Images.valkey, "valkey-server") - -class RedisMasterReplicaReplicaDownSuite extends MasterReplicaReplicaDownSuite(Images.redis, "redis-server") - -class ValkeyMasterReplicaReplicaDownSuite extends MasterReplicaReplicaDownSuite(Images.valkey, "valkey-server") - -class RedisMasterReplicaStaleReplicaSuite extends MasterReplicaStaleReplicaSuite(Images.redis, "redis-server") +class RedisMasterReplicaFaultSuite extends MasterReplicaFaultSuite(Images.redis, "redis-server") -class ValkeyMasterReplicaStaleReplicaSuite extends MasterReplicaStaleReplicaSuite(Images.valkey, "valkey-server") +class ValkeyMasterReplicaFaultSuite extends MasterReplicaFaultSuite(Images.valkey, "valkey-server") diff --git a/integration-tests/shared/src/test/scala/sage/integration/ratelimit/RateLimiterSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/ratelimit/RateLimiterSuite.scala index 99bcb64c..1ab041fc 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/ratelimit/RateLimiterSuite.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/ratelimit/RateLimiterSuite.scala @@ -5,353 +5,203 @@ import scala.concurrent.duration.* import kyo.compat.* import sage.SageException.ServerError -import sage.integration.{Images, ServerSuite} -import sage.integration.Ttls.remaining +import sage.integration.BothServersSuite +import sage.integration.Ttls.expiresWithin import sage.ratelimit.{Decision, RateLimit, RateLimiter} +import sage.ratelimit.Decision.{Allowed, Denied} -abstract class RateLimiterSuite(image: String) extends ServerSuite(image) { +class RateLimiterSuite extends BothServersSuite { private def limit(capacity: Long) = RateLimit(capacity, refillTokens = 1, refillPeriod = 1.hour) private def bucketKey(subject: String) = s"9:ratelimit:$subject" - test("admits up to capacity, then denies (via client.rateLimiter)") { - withClient { client => - val rl = client.rateLimiter[String](limit(3)) - for { - d1 <- rl.tryAcquire("admit") - d2 <- rl.tryAcquire("admit") - d3 <- rl.tryAcquire("admit") - d4 <- rl.tryAcquire("admit") - } yield { - assert(d1.isAllowed, "1st allowed") - assert(d3.isAllowed && d3.remainingTokens == 0L, "3rd allowed, empties bucket") - assert(!d4.isAllowed, "4th denied") - d4 match { - case Decision.Denied(remaining, retryAfter) => - assertEquals(remaining, 0L) - assert(retryAfter > Duration.Zero) - case other => fail(s"expected Denied, got $other") - } - } - } + private def retryAfter(decision: Decision): FiniteDuration = decision match { + case Decision.Denied(_, wait) => wait + case other => fail(s"expected Denied, got $other") } - test("cost consumes multiple tokens") { - withClient { client => - val rl = client.rateLimiter[String](limit(5)) - for { - first <- rl.tryAcquire("cost", cost = 3) - second <- rl.tryAcquire("cost", cost = 3) - } yield { - assert(first.isAllowed && first.remainingTokens == 2L, "cost 3 of 5 leaves 2") - assert(!second.isAllowed && second.remainingTokens == 2L, "second cost-3 exceeds the 2 left, which stay untouched") - } - } + // Waits depend on the server clock between calls, so the wall-clock tests compare admission and remaining tokens only. + private def untimed(decision: Decision): Decision = decision match { + case Allowed(remaining, _) => Allowed(remaining, Duration.Zero) + case Denied(remaining, _) => Denied(remaining, Duration.Zero) } - - test("peek reports without consuming") { - withClient { client => - val rl = client.rateLimiter[String](limit(3)) - for { - _ <- rl.tryAcquire("peek") - peeked <- rl.peek("peek") - after <- rl.tryAcquire("peek") - } yield { - assert(peeked.isAllowed && peeked.remainingTokens == 2L) - assert(after.isAllowed && after.remainingTokens == 1L, "peek left the 2 tokens intact") - } + + clientTest("admits up to capacity, then denies (via client.rateLimiter)") { client => + val rl = client.rateLimiter[String](limit(3)) + for { + _ <- rl.tryAcquire("admit").map(untimed).is(Allowed(2, Duration.Zero)) + _ <- rl.tryAcquire("admit").map(untimed).is(Allowed(1, Duration.Zero)) + _ <- rl.tryAcquire("admit").map(untimed).is(Allowed(0, Duration.Zero)) + d4 <- rl.tryAcquire("admit") + } yield { + assertEquals(untimed(d4), Denied(0, Duration.Zero)) + assert(retryAfter(d4) > Duration.Zero, "4th denied") } } - test("peek on an empty bucket is denied, still consuming nothing") { - withClient { client => - val rl = client.rateLimiter[String](limit(1)) - for { - _ <- rl.tryAcquire("peek-empty") - peeked <- rl.peek("peek-empty") - } yield peeked match { - case Decision.Denied(remaining, retryAfter) => - assertEquals(remaining, 0L) - assert(retryAfter > Duration.Zero, "reports the wait until one token is available") - case other => fail(s"expected Denied, got $other") - } - } - } + clientTest("cost consumes multiple tokens") { client => + val rl = client.rateLimiter[String](limit(5)) + rl.tryAcquire("cost", cost = 3).map(untimed).is(Allowed(2, Duration.Zero)) >> + rl.tryAcquire("cost", cost = 3).map(untimed).is(Denied(2, Duration.Zero)) + } - test("reset refills the bucket") { - withClient { client => - val rl = client.rateLimiter[String](limit(1)) - for { - first <- rl.tryAcquire("reset") - denied <- rl.tryAcquire("reset") - _ <- rl.reset("reset") - again <- rl.tryAcquire("reset") - } yield { - assert(first.isAllowed) - assert(!denied.isAllowed) - assert(again.isAllowed, "reset restored capacity") - } - } - } + clientTest("peek reports without consuming") { client => + val rl = client.rateLimiter[String](limit(3)) + rl.tryAcquire("peek") >> + rl.peek("peek").map(untimed).is(Allowed(2, Duration.Zero)) >> + rl.tryAcquire("peek").map(untimed).is(Allowed(1, Duration.Zero)) + } + + clientTest("peek on an empty bucket is denied, still consuming nothing") { client => + val rl = client.rateLimiter[String](limit(1)) + for { + _ <- rl.tryAcquire("peek-empty") + peeked <- rl.peek("peek-empty") + } yield { + assertEquals(peeked.remainingTokens, 0L) + assert(retryAfter(peeked) > Duration.Zero, "reports the wait until one token is available") + } + } - test("the computed sha matches the server's, and EVALSHA runs once loaded") { - withClient { client => - val rl = RateLimiter[String](limit(2)) - for { - loaded <- client.scriptLoad(RateLimiter.script) - first <- client.run(rl.evalSha("es", 1)) - second <- client.run(rl.evalSha("es", 1)) - } yield { - assertEquals(loaded, RateLimiter.sha) - assert(first.isAllowed && first.remainingTokens == 1L) - assert(second.isAllowed && second.remainingTokens == 0L) - } - } - } + clientTest("reset refills the bucket") { client => + val rl = client.rateLimiter[String](limit(1)) + rl.tryAcquire("reset").map(untimed).is(Allowed(0, Duration.Zero)) >> + rl.tryAcquire("reset").map(untimed).is(Denied(0, Duration.Zero)) >> + rl.reset("reset") >> + rl.tryAcquire("reset").map(untimed).is(Allowed(0, Duration.Zero)) + } - test("the native EVALSHA path recovers from NOSCRIPT by re-running with the script body") { - withClient { client => - val rl = client.rateLimiter[String](limit(2), namespace = "noscript") - for { - _ <- client.scriptFlush() - first <- rl.tryAcquire("cold") - second <- rl.tryAcquire("cold") - } yield { - assert(first.isAllowed && first.remainingTokens == 1L, "a cold server is recovered by one EVAL of the body") - assert(second.isAllowed && second.remainingTokens == 0L, "that EVAL caches the script, so EVALSHA then serves directly") - } - } - } - - test("composable escape hatch: run a RateLimiter Command directly (plain EVAL)") { - withClient { client => - val rl = RateLimiter[String](limit(1)) - for { - allowed <- client.run(rl.tryAcquire("cmd")) - denied <- client.run(rl.tryAcquire("cmd")) - } yield { - assert(allowed.isAllowed) - assert(!denied.isAllowed) - } - } - } - - test("the composable EVAL guard and validate agree on every policy and cost rule") { - withClient { client => - val cases: List[(RateLimit, Long)] = List( - RateLimit(3, 1, 1.hour) -> 1L, - RateLimit(3, 1, 1.hour) -> 0L, - RateLimit(3, 1, 1.hour) -> 3L, - RateLimit(0, 1, 1.hour) -> 1L, - RateLimit(3, 0, 1.hour) -> 1L, - RateLimit(3, 1, 500.nanos) -> 1L, - RateLimit(9007199254740993L, 1, 1.hour) -> 1L, // 2^53 + 1 - RateLimit(3, 9007199254740993L, 1.hour) -> 1L, - RateLimit(3, 1, 9007199254740993L.micros) -> 1L, - RateLimit(1_000_000_000L, 1, 10000.seconds) -> 1L, - RateLimit(3, 1, 3002399751580331L.micros) -> 1L, // 3x = 2^53 + 1 - RateLimit(3, 1, 1.hour) -> 4L, - RateLimit(3, 1, 1.hour) -> -1L, - RateLimit(3, 1, 1.hour) -> 9007199254740993L - ) - def serverRejects(rl: RateLimiter[String], cost: Long): CIO[Boolean] = - client.run(rl.tryAcquire("k", cost)).map(_ => false).recover { - case ServerError("SAGE", _) => CIO.value(true) - case other => CIO.fail(other) - } - def loop(remaining: List[((RateLimit, Long), Int)]): CIO[Unit] = - remaining match { - case Nil => CIO.value(()) - case ((limit, cost), i) :: tail => - val rl = RateLimiter[String](limit, namespace = s"conform$i") - serverRejects(rl, cost).flatMap { rejected => - assertEquals(rejected, rl.validate(cost).isDefined, s"case $i: $limit cost=$cost") - loop(tail) - } - } - loop(cases.zipWithIndex) - } - } - - test("RateLimiterClient.command exposes the same composable check") { - withClient { client => - val rl = client.rateLimiter[String](limit(1)) - for { - allowed <- client.run(rl.command("cmd-client")) - denied <- client.run(rl.command("cmd-client")) - } yield { - assert(allowed.isAllowed) - assert(!denied.isAllowed) - } - } - } - - test("accrues at the exact rate even at a large capacity") { - withClient { client => - val rl = RateLimiter[String](RateLimit(capacity = 1_000_000_000L, refillTokens = 1000, refillPeriod = 1.second)) - val t0 = 7_000_000_000_000L - for { - _ <- client.run(rl.tryAcquireAt("big", 1_000_000_000L, t0)) - half <- client.run(rl.tryAcquireAt("big", 500, t0 + 500_000L)) - rest <- client.run(rl.tryAcquireAt("big", 500, t0 + 1_000_000L)) - over <- client.run(rl.tryAcquireAt("big", 1, t0 + 1_000_000L)) - } yield { - assert(half.isAllowed && half.remainingTokens == 0L, "0.5s accrues exactly 500 tokens") - assert(rest.isAllowed && rest.remainingTokens == 0L, "the next 0.5s accrues the other 500") - assert(!over.isAllowed, "and no more than the rate allows") - } - } - } - - test("refills continuously as injected time advances") { - withClient { client => - val rl = RateLimiter[String](RateLimit(capacity = 10, refillTokens = 10, refillPeriod = 1.second)) - val t0 = 2_000_000_000_000L - for { - drain <- client.run(rl.tryAcquireAt("refill", 10, t0)) - denied <- client.run(rl.tryAcquireAt("refill", 1, t0)) - partial <- client.run(rl.tryAcquireAt("refill", 3, t0 + 300_000L)) - empty <- client.run(rl.tryAcquireAt("refill", 1, t0 + 300_000L)) - } yield { - assert(drain.isAllowed && drain.remainingTokens == 0L, "cost 10 empties a 10-token bucket") - assert(!denied.isAllowed, "no tokens have refilled at t0") - assert(partial.isAllowed && partial.remainingTokens == 0L, "0.3s refills exactly 3 tokens") - assert(!empty.isAllowed, "and no more") - } - } - } - - test("a backward server clock does not double-credit refill") { - withClient { client => - val rl = RateLimiter[String](RateLimit(capacity = 10, refillTokens = 10, refillPeriod = 1.second)) - val t0 = 5_000_000_000_000L - for { - _ <- client.run(rl.tryAcquireAt("clock", 10, t0)) - backward <- client.run(rl.tryAcquireAt("clock", 1, t0 - 500_000L)) - recover <- client.run(rl.tryAcquireAt("clock", 1, t0 + 100_000L)) - extra <- client.run(rl.tryAcquireAt("clock", 1, t0 + 100_000L)) - } yield { - assert(!backward.isAllowed, "a regressed clock credits nothing") - assert(recover.isAllowed && recover.remainingTokens == 0L, "credit resumes from the high-water mark, not doubled") - assert(!extra.isAllowed) - } - } - } - - test("a non-integer refill rate credits whole tokens exactly and never re-counts elapsed time") { - withClient { client => - val rl = RateLimiter[String](RateLimit(capacity = 10, refillTokens = 3, refillPeriod = 1.second)) - val t0 = 3_000_000_000_000L - for { - _ <- client.run(rl.tryAcquireAt("frac", 10, t0)) - early <- client.run(rl.tryAcquireAt("frac", 1, t0 + 333_333L)) - repeat <- client.run(rl.tryAcquireAt("frac", 1, t0 + 333_333L)) - one <- client.run(rl.tryAcquireAt("frac", 1, t0 + 333_334L)) - } yield { - assert(!early.isAllowed, "just under 1/3 s accrues no whole token") - assert(!repeat.isAllowed, "the same timestamp does not mint from re-counted elapsed") - assert(one.isAllowed && one.remainingTokens == 0L, "one whole token at 333_334 micros, no more") - } - } - } - - test("a clock rollback folds the catch-up interval into retryAfter and the key's TTL") { - withClient { client => - val rl = RateLimiter[String](RateLimit(capacity = 1, refillTokens = 1, refillPeriod = 1.second)) - val t0 = 8_000_000_000_000L - for { - _ <- client.run(rl.tryAcquireAt("roll", 1, t0)) - rolled <- client.run(rl.tryAcquireAt("roll", 1, t0 - 4_000_000L)) - pttl <- client.pTtl(bucketKey("roll")) - } yield { - rolled match { - case Decision.Denied(_, retryAfter) => assertEquals(retryAfter, 5.seconds) - case other => fail(s"expected Denied, got $other") - } - assert(remaining(pttl).exists(left => left > 2.seconds && left <= 5.seconds), s"pttl was $pttl") - } - } - } - - test("a large refill period and odd rollback retain the exact reported wait past 2^53") { - withClient { client => - val period = 1L << 53 - val rl = RateLimiter[String](RateLimit(capacity = 1, refillTokens = 1, refillPeriod = period.micros)) - val t0 = 9_000_000_000_000L - for { - _ <- client.run(rl.tryAcquireAt("big-period", 1, t0)) - rolled <- client.run(rl.tryAcquireAt("big-period", 1, t0 - 9L)) - } yield rolled match { - case Decision.Denied(_, retryAfter) => assertEquals(retryAfter, (period + 9L).micros) - case other => fail(s"expected Denied, got $other") - } - } - } - - test("reusing a namespace with alternating policies never mints a fresh bucket") { - withClient { client => - val original = RateLimiter[String](RateLimit(capacity = 10, refillTokens = 10, refillPeriod = 1.second), namespace = "policy") - val tightened = RateLimiter[String](RateLimit(capacity = 1, refillTokens = 1, refillPeriod = 1.second), namespace = "policy") - val t0 = 6_000_000_000_000L - for { - old <- client.run(original.tryAcquireAt("same-subject", 1, t0)) - firstNew <- client.run(tightened.tryAcquireAt("same-subject", 1, t0)) - oldAgain <- client.run(original.tryAcquireAt("same-subject", 1, t0)) - secondNew <- client.run(tightened.tryAcquireAt("same-subject", 1, t0)) - rolled <- client.run(original.tryAcquireAt("same-subject", 1, t0 - 500_000L)) - recovered <- client.run(original.tryAcquireAt("same-subject", 1, t0)) - } yield { - assertEquals(old.remainingTokens, 9L) - assert(firstNew.isAllowed && firstNew.remainingTokens == 0L, "old credit is capped at the tightened capacity") - assert(!oldAgain.isAllowed, "the old policy cannot recreate a full bucket during an overlapping deployment") - assert(!secondNew.isAllowed, "alternating policy signatures cannot mint tokens") - assert(!rolled.isAllowed, "a policy switch during clock rollback credits nothing") - assert(!recovered.isAllowed, "clock recovery only reaches the preserved high-water mark") - } - } - } - - test("malformed bucket fields are rejected instead of granting tokens") { - withClient { client => - val rl = RateLimiter[String](limit(1)) - def rejected(subject: String): CIO[Boolean] = - client.run(rl.tryAcquireAt(subject, 1, 4_000_000_000_000L)).map(_ => false).recover { - case ServerError("SAGE", message) if message.contains("invalid rate-limit state") => CIO.value(true) - case other => CIO.fail(other) - } - for { - _ <- client.run(rl.tryAcquireAt("missing", 1, 4_000_000_000_000L)) - _ <- client.hDel(bucketKey("missing"), "ts") - missingField <- rejected("missing") - _ <- client.run(rl.tryAcquireAt("out-of-range", 1, 4_000_000_000_000L)) - _ <- client.hSet(bucketKey("out-of-range"), ("t", "2")) - outOfRange <- rejected("out-of-range") - _ <- client.run(rl.tryAcquireAt("full-with-fraction", 1, 4_000_000_000_000L)) - _ <- client.hSet(bucketKey("full-with-fraction"), ("t", "1"), ("f", "1")) - fullFraction <- rejected("full-with-fraction") - } yield { - assert(missingField, "a missing field must not restore capacity") - assert(outOfRange, "out-of-range tokens must not exceed capacity") - assert(fullFraction, "a full bucket must not carry fractional refill credit") - } - } - } - - test("an already-full bucket reports zero time to full") { - withClient { client => - val rl = RateLimiter[String](RateLimit(capacity = 5, refillTokens = 5, refillPeriod = 1.second)) - for { - peeked <- client.run(rl.peek("full")) - } yield peeked match { - case Decision.Allowed(remaining, resetAfter) => - assertEquals(remaining, 5L) - assertEquals(resetAfter, Duration.Zero) - case other => fail(s"expected Allowed, got $other") - } - } + clientTest("the computed sha matches the server's, and EVALSHA runs once loaded") { client => + val rl = RateLimiter[String](limit(2)) + client.scriptLoad(RateLimiter.script).is(RateLimiter.compiled.sha) >> + client.run(rl.eval(cached = true, "es", 1, peek = false)).map(untimed).is(Allowed(1, Duration.Zero)) >> + client.run(rl.eval(cached = true, "es", 1, peek = false)).map(untimed).is(Allowed(0, Duration.Zero)) + } + + clientTest("the native EVALSHA path recovers from NOSCRIPT by re-running with the script body") { client => + val rl = client.rateLimiter[String](limit(2), namespace = "noscript") + client.scriptFlush() >> + rl.tryAcquire("cold").map(untimed).is(Allowed(1, Duration.Zero)) >> + rl.tryAcquire("cold").map(untimed).is(Allowed(0, Duration.Zero)) + } + + clientTest("the composable EVAL guard and validate agree on every policy and cost rule") { client => + val cases: List[(RateLimit, Long)] = List( + RateLimit(3, 1, 1.hour) -> 1L, + RateLimit(3, 1, 1.hour) -> 0L, + RateLimit(3, 1, 1.hour) -> 3L, + RateLimit(0, 1, 1.hour) -> 1L, + RateLimit(3, 0, 1.hour) -> 1L, + RateLimit(3, 1, 500.nanos) -> 1L, + RateLimit(9007199254740993L, 1, 1.hour) -> 1L, // 2^53 + 1 + RateLimit(3, 9007199254740993L, 1.hour) -> 1L, + RateLimit(3, 1, 9007199254740993L.micros) -> 1L, + RateLimit(1_000_000_000L, 1, 10000.seconds) -> 1L, + RateLimit(3, 1, 3002399751580331L.micros) -> 1L, // 3x = 2^53 + 1 + RateLimit(3, 1, 1.hour) -> 4L, + RateLimit(3, 1, 1.hour) -> -1L, + RateLimit(3, 1, 1.hour) -> 9007199254740993L + ) + def serverRejects(rl: RateLimiter[String], cost: Long): CIO[Boolean] = + client.run(rl.tryAcquire("k", cost)).map(_ => false).recover { + case ServerError("SAGE", _) => CIO.value(true) + case other => CIO.fail(other) + } + CIO.foreachDiscard(cases.zipWithIndex) { case ((limit, cost), i) => + val rl = RateLimiter[String](limit, namespace = s"conform$i") + serverRejects(rl, cost).map(rejected => assertEquals(rejected, rl.validate(cost).isDefined, s"case $i: $limit cost=$cost")) + } + } + + clientTest("RateLimiterClient.command exposes the same composable check") { client => + val rl = client.rateLimiter[String](limit(1)) + client.run(rl.command("cmd-client")).map(untimed).is(Allowed(0, Duration.Zero)) >> + client.run(rl.command("cmd-client")).map(untimed).is(Denied(0, Duration.Zero)) + } + + clientTest("accrues at the exact rate even at a large capacity") { client => + val rl = RateLimiter[String](RateLimit(capacity = 1_000_000_000L, refillTokens = 1000, refillPeriod = 1.second)) + val t0 = 7_000_000_000_000L + client.run(rl.tryAcquireAt("big", 1_000_000_000L, t0)) >> + client.run(rl.tryAcquireAt("big", 500, t0 + 500_000L)).is(Allowed(0, 1_000_000.seconds)) >> + client.run(rl.tryAcquireAt("big", 500, t0 + 1_000_000L)).is(Allowed(0, 1_000_000.seconds)) >> + client.run(rl.tryAcquireAt("big", 1, t0 + 1_000_000L)).is(Denied(0, 1.millis)) + } + + clientTest("refills continuously as injected time advances") { client => + val rl = RateLimiter[String](RateLimit(capacity = 10, refillTokens = 10, refillPeriod = 1.second)) + val t0 = 2_000_000_000_000L + client.run(rl.tryAcquireAt("refill", 10, t0)).is(Allowed(0, 1.second)) >> + client.run(rl.tryAcquireAt("refill", 1, t0)).is(Denied(0, 100.millis)) >> + client.run(rl.tryAcquireAt("refill", 3, t0 + 300_000L)).is(Allowed(0, 1.second)) >> + client.run(rl.tryAcquireAt("refill", 1, t0 + 300_000L)).is(Denied(0, 100.millis)) + } + + clientTest("a backward server clock does not double-credit refill") { client => + val rl = RateLimiter[String](RateLimit(capacity = 10, refillTokens = 10, refillPeriod = 1.second)) + val t0 = 5_000_000_000_000L + client.run(rl.tryAcquireAt("clock", 10, t0)) >> + client.run(rl.tryAcquireAt("clock", 1, t0 - 500_000L)).is(Denied(0, 600.millis)) >> + client.run(rl.tryAcquireAt("clock", 1, t0 + 100_000L)).is(Allowed(0, 1.second)) >> + client.run(rl.tryAcquireAt("clock", 1, t0 + 100_000L)).is(Denied(0, 100.millis)) + } + + clientTest("a non-integer refill rate credits whole tokens exactly and never re-counts elapsed time") { client => + val rl = RateLimiter[String](RateLimit(capacity = 10, refillTokens = 3, refillPeriod = 1.second)) + val t0 = 3_000_000_000_000L + client.run(rl.tryAcquireAt("frac", 10, t0)) >> + client.run(rl.tryAcquireAt("frac", 1, t0 + 333_333L)).is(Denied(0, 1.micro)) >> + client.run(rl.tryAcquireAt("frac", 1, t0 + 333_333L)).is(Denied(0, 1.micro)) >> + client.run(rl.tryAcquireAt("frac", 1, t0 + 333_334L)).is(Allowed(0, 3333333.micros)) + } + + clientTest("a clock rollback folds the catch-up interval into retryAfter and the key's TTL") { client => + val rl = RateLimiter[String](RateLimit(capacity = 1, refillTokens = 1, refillPeriod = 1.second)) + val t0 = 8_000_000_000_000L + client.run(rl.tryAcquireAt("roll", 1, t0)) >> + client.run(rl.tryAcquireAt("roll", 1, t0 - 4_000_000L)).map(retryAfter).is(5.seconds) >> + client.pTtl(bucketKey("roll")).satisfies(expiresWithin(_, 5.seconds, above = 2.seconds)) + } + + clientTest("a large refill period and odd rollback retain the exact reported wait past 2^53") { client => + val period = 1L << 53 + val rl = RateLimiter[String](RateLimit(capacity = 1, refillTokens = 1, refillPeriod = period.micros)) + val t0 = 9_000_000_000_000L + client.run(rl.tryAcquireAt("big-period", 1, t0)) >> + client.run(rl.tryAcquireAt("big-period", 1, t0 - 9L)).map(retryAfter).is((period + 9L).micros) + } + + clientTest("reusing a namespace with alternating policies never mints a fresh bucket") { client => + val original = RateLimiter[String](RateLimit(capacity = 10, refillTokens = 10, refillPeriod = 1.second), namespace = "policy") + val tightened = RateLimiter[String](RateLimit(capacity = 1, refillTokens = 1, refillPeriod = 1.second), namespace = "policy") + val t0 = 6_000_000_000_000L + client.run(original.tryAcquireAt("same-subject", 1, t0)).is(Allowed(9, 100.millis)) >> + client.run(tightened.tryAcquireAt("same-subject", 1, t0)).is(Allowed(0, 1.second)) >> + client.run(original.tryAcquireAt("same-subject", 1, t0)).is(Denied(0, 100.millis)) >> + client.run(tightened.tryAcquireAt("same-subject", 1, t0)).is(Denied(0, 1.second)) >> + client.run(original.tryAcquireAt("same-subject", 1, t0 - 500_000L)).is(Denied(0, 600.millis)) >> + client.run(original.tryAcquireAt("same-subject", 1, t0)).is(Denied(0, 100.millis)) + } + + clientTest("malformed bucket fields are rejected instead of granting tokens") { client => + val rl = RateLimiter[String](limit(1)) + def rejected(subject: String): CIO[Unit] = + failsWith[ServerError](client.run(rl.tryAcquireAt(subject, 1, 4_000_000_000_000L))) + .map(assertEquals(_, ServerError("SAGE", "invalid rate-limit state"), subject)) + client.run(rl.tryAcquireAt("missing", 1, 4_000_000_000_000L)) >> + client.hDel(bucketKey("missing"), "ts") >> + rejected("missing") >> + client.run(rl.tryAcquireAt("out-of-range", 1, 4_000_000_000_000L)) >> + client.hSet(bucketKey("out-of-range"), ("t", "2")) >> + rejected("out-of-range") >> + client.run(rl.tryAcquireAt("full-with-fraction", 1, 4_000_000_000_000L)) >> + client.hSet(bucketKey("full-with-fraction"), ("t", "1"), ("f", "1")) >> + rejected("full-with-fraction") + } + + clientTest("an already-full bucket reports zero time to full") { client => + val rl = RateLimiter[String](RateLimit(capacity = 5, refillTokens = 5, refillPeriod = 1.second)) + client.run(rl.peek("full")).is(Decision.Allowed(5L, Duration.Zero)) } } - -class RedisRateLimiterSuite extends RateLimiterSuite(Images.redis) -class ValkeyRateLimiterSuite extends RateLimiterSuite(Images.valkey) diff --git a/integration-tests/shared/src/test/scala/sage/integration/security/TlsAuthSuite.scala b/integration-tests/shared/src/test/scala/sage/integration/security/TlsAuthSuite.scala index a8764176..54b621f3 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/security/TlsAuthSuite.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/security/TlsAuthSuite.scala @@ -1,27 +1,27 @@ package sage.integration.security -import scala.concurrent.{ExecutionContext, Future} - import com.dimafeng.testcontainers.GenericContainer -import com.dimafeng.testcontainers.munit.TestContainerForAll import kyo.compat.* import org.testcontainers.containers.wait.strategy.Wait import org.testcontainers.images.builder.Transferable import sage.SageException.{ServerError, TlsError} import sage.client.{AuthConfig, SageConfig, TlsConfig, TrustSource} -import sage.integration.{ContainerClient, Images} +import sage.integration.{BothServersSuite, Images} /** - * One TLS+ACL server hosts every case. The args start with `-`, so the redis and valkey entrypoints each prepend their own server + * One TLS+ACL server per image hosts every case. The args start with `-`, so the redis and valkey entrypoints each prepend their own server * binary; `--port 0` makes the single exposed port speak only TLS. The cert is generated per run with the Docker host in its SAN * ([[TlsFixture]]), so hostname verification passes whatever host Testcontainers reports. */ -abstract class TlsAuthSuite(image: String) extends munit.FunSuite with TestContainerForAll with ContainerClient { +class TlsAuthSuite extends BothServersSuite { + + override protected def redisDef: GenericContainer.Def[GenericContainer] = tlsServer(Images.redis) + override protected def valkeyDef: GenericContainer.Def[GenericContainer] = tlsServer(Images.valkey) // Copy certificates through the Docker API so the suite also works with a remote daemon. A bind mount would look for the path on the daemon // host instead of the test runner. - override val containerDef: GenericContainer.Def[GenericContainer] = + private def tlsServer(image: String): GenericContainer.Def[GenericContainer] = new GenericContainer.Def[GenericContainer]({ val container = GenericContainer( image, @@ -55,58 +55,34 @@ abstract class TlsAuthSuite(image: String) extends munit.FunSuite with TestConta waitStrategy = Wait.forLogMessage(".*Ready to accept connections.*", 1) ) container.underlyingUnsafeContainer - .withCopyToContainer(Transferable.of(TlsFixture.serverCertPem), "/tls/server.crt") - .withCopyToContainer(Transferable.of(TlsFixture.serverKeyPem), "/tls/server.key") + .withCopyToContainer(Transferable.of(TlsFixture.material.certPem), "/tls/server.crt") + .withCopyToContainer(Transferable.of(TlsFixture.material.keyPem), "/tls/server.key") container }) {} - given ExecutionContext = munitExecutionContext - - private val caPath = TlsFixture.serverCert + private val caPath = TlsFixture.material.certFile private val app = AuthConfig(username = "app", password = "apppass") private def configWith(server: GenericContainer, trust: TrustSource = TrustSource.System, auth: AuthConfig = app): SageConfig = configOf(server).copy(tls = Some(TlsConfig(trust)), auth = Some(auth)) - private def connectAndPing(config: SageConfig): Future[String] = connectAndUse(config)(_.ping()).unsafeRun - - test("rejects the server certificate by default: the private CA is not in the system trust store") { - withContainers { server => - connectAndPing(configWith(server)).failed.map(error => assert(error.isInstanceOf[TlsError], error)) - } - } - - test("connects when the private CA is supplied as PEM trust material") { - withContainers { server => - connectAndPing(configWith(server, TrustSource.Pem(caPath))).map(pong => assertEquals(pong, "PONG")) - } + serverTest("rejects the server certificate by default: the private CA is not in the system trust store") { server => + failsWith[TlsError](connectAndUse(configWith(server))(_.ping())) } - test("connects with verification disabled") { - withContainers { server => - connectAndPing(configWith(server, TrustSource.Insecure)).map(pong => assertEquals(pong, "PONG")) - } + serverTest("connects with verification disabled") { server => + connectAndUse(configWith(server, TrustSource.Insecure))(_.ping().is("PONG")) } - test("ACL auth over TLS succeeds for a named user via HELLO, and round-trips a command") { - withContainers { server => - connectAndUse(configWith(server, TrustSource.Pem(caPath))) { client => - for { - _ <- client.set("tls:key", "value") - value <- client.get[String]("tls:key") - } yield value - }.unsafeRun.map(value => assertEquals(value, Some("value"))) + serverTest("ACL auth over TLS succeeds for a named user via HELLO, and round-trips a command") { server => + connectAndUse(configWith(server, TrustSource.Pem(caPath))) { client => + client.set("tls:key", "value").flatMap(_ => client.get[String]("tls:key").is(Some("value"))) } } - test("bad credentials fail with a server error") { - withContainers { server => - val config = configWith(server, TrustSource.Pem(caPath), AuthConfig(username = "app", password = "wrong")) - connectAndPing(config).failed.map(error => assert(error.isInstanceOf[ServerError], error)) - } + serverTest("bad credentials fail with a server error") { server => + failsWith[ServerError]( + connectAndUse(configWith(server, TrustSource.Pem(caPath), AuthConfig(username = "app", password = "wrong")))(_.ping()) + ) } } - -class RedisTlsAuthSuite extends TlsAuthSuite(Images.redis) - -class ValkeyTlsAuthSuite extends TlsAuthSuite(Images.valkey) diff --git a/integration-tests/shared/src/test/scala/sage/integration/security/TlsFixture.scala b/integration-tests/shared/src/test/scala/sage/integration/security/TlsFixture.scala index e2a887de..d39898db 100644 --- a/integration-tests/shared/src/test/scala/sage/integration/security/TlsFixture.scala +++ b/integration-tests/shared/src/test/scala/sage/integration/security/TlsFixture.scala @@ -17,24 +17,12 @@ import org.testcontainers.DockerClientFactory */ object TlsFixture { - private lazy val material = generate() - - /** - * A local file holding the server certificate, for the client's PEM trust material. - */ - def serverCert: Path = material.certFile - - /** - * The server certificate in PEM, to copy into the container as `--tls-cert-file`. - */ - def serverCertPem: Array[Byte] = material.certPem - /** - * The server private key in PEM, to copy into the container as `--tls-key-file`. + * The server certificate and key in PEM, copied into the container, plus a local file holding the certificate for the client's trust. */ - def serverKeyPem: Array[Byte] = material.keyPem + final case class Material(certFile: Path, certPem: Array[Byte], keyPem: Array[Byte]) - final private case class Material(certFile: Path, certPem: Array[Byte], keyPem: Array[Byte]) + lazy val material: Material = generate() private def generate(): Material = { val dir = Files.createTempDirectory("sage-tls") diff --git a/integration-tests/zio/src/test/scala/sage/integration/ZioSmokeSuite.scala b/integration-tests/zio/src/test/scala/sage/integration/ZioSmokeSuite.scala index e82b0a4e..9e85fb04 100644 --- a/integration-tests/zio/src/test/scala/sage/integration/ZioSmokeSuite.scala +++ b/integration-tests/zio/src/test/scala/sage/integration/ZioSmokeSuite.scala @@ -4,198 +4,88 @@ import java.util.concurrent.TimeUnit import scala.concurrent.duration.FiniteDuration +import kyo.compat.* +import munit.{Location, TestOptions} import zio.* import sage.* import sage.backend.* import sage.client.{DedicatedPoolConfig, SageConfig} -class ZioSmokeSuite extends ServerSuite(Images.redis) { +class ZioSmokeSuite extends SmokeSuite { - private def withNativeClient(body: SageClient => RIO[Scope, Unit]): Unit = withTunedClient(identity)(body) - - private def withTunedClient(tune: SageConfig => SageConfig)(body: SageClient => RIO[Scope, Unit]): Unit = - withContainers { server => + private def nativeTest(options: TestOptions, tune: SageConfig => SageConfig = identity)( + body: SageClient => RIO[Scope, Unit] + )(using Location): Unit = + test(options)(withContainers { server => val program: Task[Unit] = ZIO.scoped(SageClient.scoped(tune(configOf(server))).flatMap(body)) Unsafe.unsafe(implicit u => Runtime.default.unsafe.run(program).getOrThrowFiberFailure()) - } + }) - test("a distributed lock scopes native effects and skips contended bodies") { - withNativeClient { client => - val locks = client.lock[String]() - var evaluated = false - for { - busy <- locks.withLock("native-lock", FiniteDuration(2L, TimeUnit.SECONDS)) { - locks.tryWithLock("native-lock") { - evaluated = true - client.ping() - } - } - acquired <- locks.tryWithLock("native-lock")(client.ping()) - } yield { - assertEquals(busy, None) - assertEquals(acquired, Some("PONG")) - assert(!evaluated) - } - } - } + private val lift: [A] => IO[SageException, A] => CIO[A] = [A] => (io: IO[SageException, A]) => CIO.lift(io) - test("an end user connects and round-trips with native ZIO") { - withNativeClient { client => - for { - pong <- client.ping() - _ <- ZIO.foreachParDiscard(1 to 50)(i => client.set(s"key-$i", s"value-$i")) - values <- ZIO.foreachPar((1 to 50).toList)(i => client.get[String](s"key-$i")) - } yield { - assertEquals(pong, "PONG") - assertEquals(values, (1 to 50).toList.map(i => Some(s"value-$i"))) - } - } - } + nativeTest("a distributed lock scopes native effects and skips contended bodies")(lockScopesNativeEffects(_)(lift).lower) - test("a pipeline returns a typed tuple natively, surfacing failures per position") { - withNativeClient { client => - for { - _ <- client.set("pipe:a", "x") - _ <- client.set("pipe:n", 10) - out <- client.pipeline((Commands.get[String, String]("pipe:a"), Commands.incrBy[String]("pipe:n", 5))) - _ <- client.set("pipe:str", "hello") - attempt <- client.pipelineAttempt((Commands.get[String, String]("pipe:str"), Commands.incr[String]("pipe:str"))) - } yield { - assertEquals(out, (Some("x"), 15L)) - assert(attempt._1 == Right(Some("hello")), attempt._1) - assert(attempt._2.isLeft, attempt._2) - } - } - } + private val pool = DedicatedPoolConfig(maxConnections = 2, acquireTimeout = FiniteDuration(1L, TimeUnit.SECONDS)) - test("a transaction commits atomically with native ZIO, guarded by WATCH") { - withNativeClient { client => - for { - _ <- client.set("tx:n", 1) - out <- client.transaction { tx => - for { - _ <- tx.watch("tx:n") - _ <- tx.get[Int]("tx:n") - res <- tx.exec((Commands.incr[String]("tx:n"), Commands.incrBy[String]("tx:n", 4))) - } yield res - } - } yield assertEquals(out, Some((2L, 6L))) - } + nativeTest("an interrupted blocking command releases its pooled slot instead of leaking it", _.copy(dedicatedPool = pool)) { client => + for { + _ <- ZIO.foreachDiscard(1 to 4)(_ => client.blPop[String]("leak:empty")(BlockTimeout.Forever).timeout(Duration.fromMillis(150))) + none <- client.blPop[String]("leak:empty")(BlockTimeout.After(FiniteDuration(100L, TimeUnit.MILLISECONDS))) + } yield assertEquals(none, None) } - test("an interrupted blocking command releases its pooled slot instead of leaking it") { - val pool = DedicatedPoolConfig(maxConnections = 2, acquireTimeout = FiniteDuration(1L, TimeUnit.SECONDS)) - withTunedClient(_.copy(dedicatedPool = pool)) { client => - for { - _ <- ZIO.foreachDiscard(1 to 4)(_ => client.blPop[String]("leak:empty")(BlockTimeout.Forever).timeout(Duration.fromMillis(150))) - none <- client.blPop[String]("leak:empty")(BlockTimeout.After(FiniteDuration(100L, TimeUnit.MILLISECONDS))) - } yield assertEquals(none, None) + nativeTest("re-running the same blocking command value succeeds each round instead of hanging after the first") { client => + val blPop = client.blPop[String]("h1:reuse")(BlockTimeout.After(FiniteDuration(100L, TimeUnit.MILLISECONDS))) + for { + first <- blPop + second <- blPop.timeout(Duration.fromSeconds(5)) + } yield { + assertEquals(first, None) + assertEquals(second, Some(None), "re-running the same blocking effect hung: its lease was captured and single-shot") } } - test("re-running the same blocking command value succeeds each round instead of hanging after the first") { - withNativeClient { client => - val blPop = client.blPop[String]("h1:reuse")(BlockTimeout.After(FiniteDuration(100L, TimeUnit.MILLISECONDS))) - for { - first <- blPop - second <- blPop.timeout(Duration.fromSeconds(5)) - } yield { - assertEquals(first, None) - assert(second.isDefined, "re-running the same blocking effect hung: its lease was captured and single-shot") - assertEquals(second.flatten, None) - } - } + nativeTest("scanAll streams every key as a native ZStream") { client => + scanAllFindsEveryKey(client)(lift)(client.scanAll(pattern = Some("scan-*"), count = Some(10L)).runCollect).lower } - test("scanAll streams every key as a native ZStream") { - withNativeClient { client => - for { - _ <- ZIO.foreachParDiscard(1 to 50)(i => client.set(s"scan-$i", "v")) - keys <- client.scanAll(pattern = Some("scan-*"), count = Some(10L)).runCollect - } yield assertEquals(keys.toSet, (1 to 50).map(i => s"scan-$i").toSet) - } + nativeTest("subscribe delivers published messages as a native ZStream") { client => + for { + stream <- client.subscribeScoped[String]("smoke") + _ <- ZIO.foreachDiscard(1 to 3)(i => client.publish("smoke", s"m$i")) + messages <- stream.take(3).runCollect + } yield assertEquals(messages.toList, List("m1", "m2", "m3").map(Message("smoke", _))) } - test("subscribe delivers published messages as a native ZStream") { - withNativeClient { client => - for { - stream <- client.subscribeScoped[String]("smoke") - _ <- ZIO.foreachDiscard(1 to 3)(i => client.publish("smoke", s"m$i")) - messages <- stream.take(3).runCollect - } yield { - assertEquals(messages.map(_.channel).toSet, Set("smoke")) - assertEquals(messages.map(_.payload).toList, List("m1", "m2", "m3")) - } - } + nativeTest("hScanAll streams every field/value pair as a native ZStream") { client => + for { + _ <- ZIO.foreachParDiscard(1 to 50)(i => client.hSet("hscan", (s"f$i", s"v$i"))) + pairs <- client.hScanAll[String, String]("hscan", count = Some(10L)).runCollect + } yield assertEquals(pairs.toMap, (1 to 50).map(i => s"f$i" -> s"v$i").toMap) } - test("hScanAll streams every field/value pair as a native ZStream") { - withNativeClient { client => - for { - _ <- ZIO.foreachParDiscard(1 to 50)(i => client.hSet("hscan", (s"f$i", s"v$i"))) - pairs <- client.hScanAll[String, String]("hscan", count = Some(10L)).runCollect - } yield assertEquals(pairs.toMap, (1 to 50).map(i => s"f$i" -> s"v$i").toMap) - } - } - - test("sScanAll streams every member as a native ZStream") { - withNativeClient { client => - for { - _ <- ZIO.foreachParDiscard(1 to 50)(i => client.sAdd("sscan", s"m$i")) - members <- client.sScanAll[String]("sscan", count = Some(10L)).runCollect - } yield assertEquals(members.toSet, (1 to 50).map(i => s"m$i").toSet) - } - } - - test("zScanAll streams every member/score pair as a native ZStream") { - withNativeClient { client => - for { - _ <- ZIO.foreachParDiscard(1 to 50)(i => client.zAdd("zscan")((s"m$i", i.toDouble))) - pairs <- client.zScanAll[String]("zscan", count = Some(10L)).runCollect - } yield assertEquals(pairs.toMap, (1 to 50).map(i => s"m$i" -> i.toDouble).toMap) - } - } - - test("xRangeAll pages every entry as a native ZStream") { - withNativeClient { client => - for { - _ <- ZIO.foreachDiscard(1 to 50)(i => client.xAdd("xrangeall", XAddId.Explicit(StreamId(i.toLong, 0L)))(("f", s"v$i"))) - entries <- client.xRangeAll[String, String]("xrangeall", batch = 10L).runCollect - } yield assertEquals(entries.map(_.id).toList, (1 to 50).map(i => StreamId(i.toLong, 0L)).toList) - } - } - - test("xConsume tails a group and auto-acks each entry after the handler succeeds") { - withNativeClient { client => - val block = BlockTimeout.After(FiniteDuration(200L, TimeUnit.MILLISECONDS)) - for { - _ <- ZIO.foreachDiscard(1 to 3)(i => client.xAdd("xconsume", XAddId.Explicit(StreamId(i.toLong, 0L)))(("f", s"v$i"))) - _ <- client.xGroupCreate("xconsume", "g", GroupStartId.At(StreamId(0L, 0L))) - seen <- Ref.make(Vector.empty[String]) - fiber <- client.xConsume[String, String]("g", "c", "xconsume", block = block)(entry => seen.update(_ :+ entry.fields.head._2)).fork - _ <- seen.get.repeatUntil(_.size >= 3).timeoutFail(new RuntimeException("xConsume did not deliver"))(Duration.fromSeconds(10)) - _ <- fiber.interrupt - got <- seen.get - pend <- client.xPending("xconsume", "g") - } yield { - assertEquals(got.sorted, Vector("v1", "v2", "v3")) - assertEquals(pend.total, 0L) - } - } + nativeTest("xRangeAll pages every entry as a native ZStream") { client => + for { + _ <- ZIO.foreachDiscard(1 to 50)(i => client.xAdd("xrangeall", XAddId.Explicit(StreamId(i.toLong, 0L)))(("f", s"v$i"))) + entries <- client.xRangeAll[String, String]("xrangeall", batch = 10L).runCollect + } yield assertEquals(entries.map(_.id).toList, (1 to 50).map(i => StreamId(i.toLong, 0L)).toList) } - test("client.rateLimiter admits up to capacity then denies") { - withNativeClient { client => - val rl = client.rateLimiter[String](RateLimit(capacity = 2, refillTokens = 1, refillPeriod = FiniteDuration(1L, TimeUnit.HOURS))) - for { - first <- rl.tryAcquire("smoke") - second <- rl.tryAcquire("smoke") - denied <- rl.tryAcquire("smoke") - } yield { - assert(first.isAllowed && second.isAllowed, "the first two are admitted") - assert(!denied.isAllowed, "the third is denied once the bucket empties") - } + nativeTest("xConsume tails a group and auto-acks each entry after the handler succeeds") { client => + val block = BlockTimeout.After(FiniteDuration(200L, TimeUnit.MILLISECONDS)) + for { + _ <- ZIO.foreachDiscard(1 to 3)(i => client.xAdd("xconsume", XAddId.Explicit(StreamId(i.toLong, 0L)))(("f", s"v$i"))) + _ <- client.xGroupCreate("xconsume", "g", GroupStartId.At(StreamId(0L, 0L))) + seen <- Ref.make(Vector.empty[String]) + fiber <- client.xConsume[String, String]("g", "c", "xconsume", block = block)(entry => seen.update(_ ++ entry.fields.map(_._2))).fork + _ <- seen.get.repeatUntil(_.size >= 3).timeoutFail(new RuntimeException("xConsume did not deliver"))(Duration.fromSeconds(10)) + _ <- fiber.interrupt + got <- seen.get + pend <- client.xPending("xconsume", "g") + } yield { + assertEquals(got.sorted, Vector("v1", "v2", "v3")) + assertEquals(pend.total, 0L) } } } diff --git a/sage-client/ce/src/main/scala/sage/backend/SageClient.scala b/sage-client/ce/src/main/scala/sage/backend/SageClient.scala index 46dd227d..f395fead 100644 --- a/sage-client/ce/src/main/scala/sage/backend/SageClient.scala +++ b/sage-client/ce/src/main/scala/sage/backend/SageClient.scala @@ -8,7 +8,7 @@ import kyo.compat.* import sage.{Message, PatternMessage} import sage.client.SageConfig -import sage.client.internal.{Client, LoweredClient, Paged, ScanStep, ScanTarget, Subscription} +import sage.client.internal.{Client, LoweredClient, Paged, Subscription} import sage.codec.{KeyCodec, ValueCodec} import sage.commands.* @@ -28,7 +28,7 @@ extension [K](client: Client[IO, K])(using @unused ev: KeyCodec[K]) { count: Option[Long] = None, ofType: Option[RedisType] = None ): fs2.Stream[IO, K] = - scanStreamAll(target => cursor => client.runOn(target, Keys.scan[K](cursor, pattern, count, ofType))) + paged(Paged.scanAll[K](client.runner, pattern, count, ofType)) /** * Iterates over all HSCAN field/value pairs until the server returns a zero cursor. An empty page with a non-zero cursor continues the scan. @@ -38,7 +38,7 @@ extension [K](client: Client[IO, K])(using @unused ev: KeyCodec[K]) { pattern: Option[String] = None, count: Option[Long] = None ): fs2.Stream[IO, (F, V)] = - scanStream(cursor => client.run(Hashes.hScan[K, F, V](key, cursor, pattern, count))) + paged(Paged.scanKey(client.runner)(cursor => Hashes.hScan[K, F, V](key, cursor, pattern, count))) /** * Iterates over all SSCAN members until the server returns a zero cursor. An empty page with a non-zero cursor continues the scan. @@ -48,7 +48,7 @@ extension [K](client: Client[IO, K])(using @unused ev: KeyCodec[K]) { pattern: Option[String] = None, count: Option[Long] = None ): fs2.Stream[IO, V] = - scanStream(cursor => client.run(Sets.sScan[K, V](key, cursor, pattern, count))) + paged(Paged.scanKey(client.runner)(cursor => Sets.sScan[K, V](key, cursor, pattern, count))) /** * Iterates over all ZSCAN member/score pairs until the server returns a zero cursor. An empty page with a non-zero cursor continues the scan. @@ -58,18 +58,11 @@ extension [K](client: Client[IO, K])(using @unused ev: KeyCodec[K]) { pattern: Option[String] = None, count: Option[Long] = None ): fs2.Stream[IO, (V, Double)] = - scanStream(cursor => client.run(SortedSets.zScan[K, V](key, cursor, pattern, count))) + paged(Paged.scanKey(client.runner)(cursor => SortedSets.zScan[K, V](key, cursor, pattern, count))) // convert pages from the shared Paged helper into individual fs2 Stream elements - private def paged[S, A](init: S)(step: Paged.Step[S, A]): fs2.Stream[IO, A] = - CStream.unfold[S, Vector[A]](init)(step).flatMap(items => CStream.init(items)).lower - - private def scanStream[A](fetch: ScanCursor => IO[ScanPage[A]]): fs2.Stream[IO, A] = - paged[Option[ScanCursor], A](Some(ScanCursor.start))(Paged.byCursor(cursor => CIO.lift(fetch(cursor)))) - - // scan each target in sequence with its own node-local cursor. A cluster scan visits every master that owns slots. - private def scanStreamAll[A](fetch: ScanTarget => ScanCursor => IO[ScanPage[A]]): fs2.Stream[IO, A] = - paged[ScanStep, A](ScanStep.Begin)(Paged.acrossTargets(CIO.lift(client.scanTargets))(target => cursor => CIO.lift(fetch(target)(cursor)))) + private def paged[S, A](pages: Paged.Pages[S, A]): fs2.Stream[IO, A] = + CStream.unfold[S, Vector[A]](pages.init)(pages.step).flatMap(items => CStream.init(items)).lower /** * Lazily pages an entire stream by range, batching `XRANGE` and advancing past the last id each page. Stops when a page comes back empty. @@ -80,9 +73,7 @@ extension [K](client: Client[IO, K])(using @unused ev: KeyCodec[K]) { end: StreamRangeId = StreamRangeId.Max, batch: Long = 100L ): fs2.Stream[IO, StreamEntry[F, V]] = - paged[Option[StreamRangeId], StreamEntry[F, V]](Some(start))( - Paged.byRange(batch)(from => CIO.lift(client.run(Streams.xRange[K, F, V](key, from, end, Some(batch))))) - ) + paged(Paged.xRangeAll[K, F, V](client.runner, key, start, end, batch)) /** * Auto-claims idle pending entries for `consumer`, advancing the `XAUTOCLAIM` cursor until it returns to the start. Entries whose data @@ -96,9 +87,7 @@ extension [K](client: Client[IO, K])(using @unused ev: KeyCodec[K]) { start: StreamId = StreamId.Zero, count: Option[Long] = None ): fs2.Stream[IO, StreamEntry[F, V]] = - paged[Option[StreamId], StreamEntry[F, V]](Some(start))( - Paged.byAutoClaim(from => CIO.lift(client.run(Streams.xAutoClaim[K, F, V](key, group, consumer, minIdle, from, count)))) - ) + paged(Paged.xAutoClaimAll[K, F, V](client.runner, key, group, consumer, minIdle, start, count)) /** * Follows a stream without a consumer group. It first reads every entry after `from`, then waits for new entries. The explicit entry ID @@ -111,11 +100,7 @@ extension [K](client: Client[IO, K])(using @unused ev: KeyCodec[K]) { count: Option[Long] = None, block: BlockTimeout = Paged.defaultPoll ): fs2.Stream[IO, StreamEntry[F, V]] = - paged[StreamId, StreamEntry[F, V]](from)( - Paged.tail(last => - CIO.lift(client.run(Streams.xRead[K, F, V]((key, ReadId.After(last)))(count = count, block = Some(block)))).map(_.flatMap(_._2)) - ) - ) + paged(Paged.xTail[K, F, V](client.runner, key, from, count, block)) /** * Follows a stream as part of a consumer group. It processes this consumer's pending entries first, then waits for new entries. Each @@ -128,28 +113,11 @@ extension [K](client: Client[IO, K])(using @unused ev: KeyCodec[K]) { count: Option[Long] = None, block: BlockTimeout = Paged.defaultPoll )(handle: StreamEntry[F, V] => IO[Unit]): IO[Unit] = - consumeStream[F, V](group, consumer, key, count, block) + paged(Paged.xConsume[K, F, V](client.runner, group, consumer, key, count, block)) .evalMap(entry => handle(entry) >> client.run(Streams.xAck(key, group)(entry.id)).void) .compile .drain - private def consumeStream[F: KeyCodec, V: ValueCodec]( - group: String, - consumer: String, - key: K, - count: Option[Long], - block: BlockTimeout - ): fs2.Stream[IO, StreamEntry[F, V]] = - paged[Either[StreamId, Unit], StreamEntry[F, V]](Left(StreamId.Zero))( - Paged.consume( - drainPending = after => - CIO.lift(client.run(Streams.xReadGroup[K, F, V](group, consumer)((key, GroupReadId.After(after)))(count = count))).map(_.flatMap(_._2)), - tailNew = CIO - .lift(client.run(Streams.xReadGroup[K, F, V](group, consumer)((key, GroupReadId.New))(count = count, block = Some(block)))) - .map(_.flatMap(_._2)) - ) - ) - /** * Subscribes to one or more channels. Closing the stream's scope unsubscribes. Sage resubscribes after reconnecting, but messages * published while the connection is down are lost. diff --git a/sage-client/ce/src/test/scala/sage/client/internal/CeLockCancellationSpec.scala b/sage-client/ce/src/test/scala/sage/client/internal/CeLockCancellationSpec.scala index 1714d9e1..86a97532 100644 --- a/sage-client/ce/src/test/scala/sage/client/internal/CeLockCancellationSpec.scala +++ b/sage-client/ce/src/test/scala/sage/client/internal/CeLockCancellationSpec.scala @@ -13,15 +13,15 @@ import sage.commands.Command import sage.protocol.Frame class CeLockCancellationSpec extends LockCancellationSpec { - override protected def tryWithLock[A](commands: CommandRunner[CIO, String], lease: FiniteDuration)(body: CIO[A]): CIO[Option[A]] = + override protected def tryWithLock[A](commands: SharedRunner, lease: FiniteDuration)(body: CIO[A]): CIO[Option[A]] = CIO.lift(new SageClient.Lowered(new LockTestClient(commands)).lock[String](lease).tryWithLock("key")(body.lower)) - override protected def withLock[A](commands: CommandRunner[CIO, String], lease: FiniteDuration, wait: FiniteDuration)(body: CIO[A]): CIO[A] = + override protected def withLock[A](commands: SharedRunner, lease: FiniteDuration, wait: FiniteDuration)(body: CIO[A]): CIO[A] = CIO.lift(new SageClient.Lowered(new LockTestClient(commands)).lock[String](lease).withLock("key", wait)(body.lower)) List("acquire", "renew", "release").foreach { stalled => test(s"a delayed $stalled callback does not hold up the lock deadline") { - val commands = new CommandRunner[CIO, String] { + val commands = new SharedRunner { def run[A](command: Command[A]): CIO[A] = CIO.async { callback => val result = command.decode(Frame.Integer(1)).toTry if (command.args(4).asUtf8String == stalled) Scheduler.real.after(5.seconds)(callback(result)) diff --git a/sage-client/kyo/src/main/scala/sage/backend/SageClient.scala b/sage-client/kyo/src/main/scala/sage/backend/SageClient.scala index d4f19535..f860c891 100644 --- a/sage-client/kyo/src/main/scala/sage/backend/SageClient.scala +++ b/sage-client/kyo/src/main/scala/sage/backend/SageClient.scala @@ -9,7 +9,7 @@ import _root_.kyo.compat.* import sage.{Message, PatternMessage, SageException} import sage.client.SageConfig -import sage.client.internal.{Client, LoweredClient, Paged, ScanStep, ScanTarget, Subscription} +import sage.client.internal.{Client, LoweredClient, Paged, Subscription} import sage.codec.{KeyCodec, ValueCodec} import sage.commands.* @@ -45,7 +45,7 @@ extension [K](client: Client[[A] =>> A < (Abort[SageException] & Async), K])(usi count: Option[Long] = None, ofType: Option[RedisType] = None )(using Tag[K]): Stream[K, Abort[SageException] & Async] = - scanStreamAll(target => cursor => client.runOn(target, Keys.scan[K](cursor, pattern, count, ofType))) + paged(Paged.scanAll[K](client.runner, pattern, count, ofType)) /** * Iterates over all HSCAN field/value pairs until the server returns a zero cursor. An empty page with a non-zero cursor continues the scan. @@ -55,7 +55,7 @@ extension [K](client: Client[[A] =>> A < (Abort[SageException] & Async), K])(usi pattern: Option[String] = None, count: Option[Long] = None )(using Tag[F], Tag[V]): Stream[(F, V), Abort[SageException] & Async] = - scanStream(cursor => client.run(Hashes.hScan[K, F, V](key, cursor, pattern, count))) + paged(Paged.scanKey(client.runner)(cursor => Hashes.hScan[K, F, V](key, cursor, pattern, count))) /** * Iterates over all SSCAN members until the server returns a zero cursor. An empty page with a non-zero cursor continues the scan. @@ -65,7 +65,7 @@ extension [K](client: Client[[A] =>> A < (Abort[SageException] & Async), K])(usi pattern: Option[String] = None, count: Option[Long] = None )(using Tag[V]): Stream[V, Abort[SageException] & Async] = - scanStream(cursor => client.run(Sets.sScan[K, V](key, cursor, pattern, count))) + paged(Paged.scanKey(client.runner)(cursor => Sets.sScan[K, V](key, cursor, pattern, count))) /** * Iterates over all ZSCAN member/score pairs until the server returns a zero cursor. An empty page with a non-zero cursor continues the scan. @@ -75,26 +75,15 @@ extension [K](client: Client[[A] =>> A < (Abort[SageException] & Async), K])(usi pattern: Option[String] = None, count: Option[Long] = None )(using Tag[V]): Stream[(V, Double), Abort[SageException] & Async] = - scanStream(cursor => client.run(SortedSets.zScan[K, V](key, cursor, pattern, count))) + paged(Paged.scanKey(client.runner)(cursor => SortedSets.zScan[K, V](key, cursor, pattern, count))) // A chunk size of 1 emits each page immediately. Kyo's default chunk size of 4096 would delay an unbounded stream such as xTail or // xConsume. Paged provides the iteration logic; this adapter converts its CIO and Option result to Kyo types. - private def paged[S, A](start: S)(step: Paged.Step[S, A])(using Tag[A]): Stream[A, Abort[SageException] & Async] = + private def paged[S, A](pages: Paged.Pages[S, A])(using Tag[A]): Stream[A, Abort[SageException] & Async] = Stream - .unfold[S, Vector[A], Abort[SageException] & Async](start, chunkSize = 1)(s => refine(step(s).lower).map(Maybe.fromOption)) + .unfold[S, Vector[A], Abort[SageException] & Async](pages.init, chunkSize = 1)(s => refine(pages.step(s).lower).map(Maybe.fromOption)) .flatMap(items => Stream.init(items)) - private def scanStream[A](fetch: ScanCursor => ScanPage[A] < (Abort[SageException] & Async))( - using Tag[A] - ): Stream[A, Abort[SageException] & Async] = - paged[Option[ScanCursor], A](Some(ScanCursor.start))(Paged.byCursor(cursor => CIO.lift(fetch(cursor)))) - - // scan each target with its own node-local cursor. A cluster has one target for every slot-owning master. - private def scanStreamAll[A]( - fetch: ScanTarget => ScanCursor => ScanPage[A] < (Abort[SageException] & Async) - )(using Tag[A]): Stream[A, Abort[SageException] & Async] = - paged[ScanStep, A](ScanStep.Begin)(Paged.acrossTargets(CIO.lift(client.scanTargets))(target => cursor => CIO.lift(fetch(target)(cursor)))) - /** * Lazily pages an entire stream by range, batching `XRANGE` and advancing past the last id each page. Stops when a page comes back empty. */ @@ -104,9 +93,7 @@ extension [K](client: Client[[A] =>> A < (Abort[SageException] & Async), K])(usi end: StreamRangeId = StreamRangeId.Max, batch: Long = 100L )(using Tag[F], Tag[V]): Stream[StreamEntry[F, V], Abort[SageException] & Async] = - paged[Option[StreamRangeId], StreamEntry[F, V]](Some(start))( - Paged.byRange(batch)(from => CIO.lift(client.run(Streams.xRange[K, F, V](key, from, end, Some(batch))))) - ) + paged(Paged.xRangeAll[K, F, V](client.runner, key, start, end, batch)) /** * Auto-claims idle pending entries for `consumer`, advancing the `XAUTOCLAIM` cursor until it returns to the start. Entries whose data @@ -120,9 +107,7 @@ extension [K](client: Client[[A] =>> A < (Abort[SageException] & Async), K])(usi start: StreamId = StreamId.Zero, count: Option[Long] = None )(using Tag[F], Tag[V]): Stream[StreamEntry[F, V], Abort[SageException] & Async] = - paged[Option[StreamId], StreamEntry[F, V]](Some(start))( - Paged.byAutoClaim(from => CIO.lift(client.run(Streams.xAutoClaim[K, F, V](key, group, consumer, minIdle, from, count)))) - ) + paged(Paged.xAutoClaimAll[K, F, V](client.runner, key, group, consumer, minIdle, start, count)) /** * Follows a stream without a consumer group. It first reads every entry after `from`, then waits for new entries. The explicit entry ID @@ -135,11 +120,7 @@ extension [K](client: Client[[A] =>> A < (Abort[SageException] & Async), K])(usi count: Option[Long] = None, block: BlockTimeout = Paged.defaultPoll )(using Tag[F], Tag[V]): Stream[StreamEntry[F, V], Abort[SageException] & Async] = - paged[StreamId, StreamEntry[F, V]](from)( - Paged.tail(last => - CIO.lift(client.run(Streams.xRead[K, F, V]((key, ReadId.After(last)))(count = count, block = Some(block)))).map(_.flatMap(_._2)) - ) - ) + paged(Paged.xTail[K, F, V](client.runner, key, from, count, block)) /** * Follows a stream as part of a consumer group. It processes this consumer's pending entries first, then waits for new entries. Each @@ -152,26 +133,9 @@ extension [K](client: Client[[A] =>> A < (Abort[SageException] & Async), K])(usi count: Option[Long] = None, block: BlockTimeout = Paged.defaultPoll )(handle: StreamEntry[F, V] => Unit < (Abort[SageException] & Async))(using Tag[F], Tag[V], Frame): Unit < (Abort[SageException] & Async) = - consumeStream[F, V](group, consumer, key, count, block) + paged(Paged.xConsume[K, F, V](client.runner, group, consumer, key, count, block)) .foreach(entry => handle(entry).flatMap(_ => client.run(Streams.xAck(key, group)(entry.id)).map(_ => ()))) - private def consumeStream[F: KeyCodec, V: ValueCodec]( - group: String, - consumer: String, - key: K, - count: Option[Long], - block: BlockTimeout - )(using Tag[F], Tag[V]): Stream[StreamEntry[F, V], Abort[SageException] & Async] = - paged[Either[StreamId, Unit], StreamEntry[F, V]](Left(StreamId.Zero))( - Paged.consume( - drainPending = after => - CIO.lift(client.run(Streams.xReadGroup[K, F, V](group, consumer)((key, GroupReadId.After(after)))(count = count))).map(_.flatMap(_._2)), - tailNew = CIO - .lift(client.run(Streams.xReadGroup[K, F, V](group, consumer)((key, GroupReadId.New))(count = count, block = Some(block)))) - .map(_.flatMap(_._2)) - ) - ) - /** * Subscribes to one or more channels. Closing the enclosing `Scope` unsubscribes. Sage resubscribes after reconnecting, but messages * published while the connection is down are lost. diff --git a/sage-client/ox/src/main/scala/sage/backend/SageClient.scala b/sage-client/ox/src/main/scala/sage/backend/SageClient.scala index f15301c3..9dd9c1f1 100644 --- a/sage-client/ox/src/main/scala/sage/backend/SageClient.scala +++ b/sage-client/ox/src/main/scala/sage/backend/SageClient.scala @@ -1,7 +1,5 @@ package sage.backend -import java.util.concurrent.atomic.AtomicBoolean - import scala.annotation.unused import scala.concurrent.duration.FiniteDuration @@ -11,7 +9,7 @@ import kyo.compat.* import sage.{Message, PatternMessage} import sage.client.SageConfig -import sage.client.internal.{Client, LoweredClient, Paged, ScanStep, ScanTarget, Subscription} +import sage.client.internal.{Client, LoweredClient, Paged, Subscription} import sage.codec.{KeyCodec, ValueCodec} import sage.commands.* @@ -31,7 +29,7 @@ extension [K](client: Client[[A] =>> Ox ?=> A, K])(using @unused ev: KeyCodec[K] count: Option[Long] = None, ofType: Option[RedisType] = None ): Ox ?=> Flow[K] = - scanStreamAll(target => cursor => client.runOn(target, Keys.scan[K](cursor, pattern, count, ofType))) + paged(Paged.scanAll[K](client.runner, pattern, count, ofType)) /** * Iterates over all HSCAN field/value pairs until the server returns a zero cursor. An empty page with a non-zero cursor continues the scan. @@ -41,7 +39,7 @@ extension [K](client: Client[[A] =>> Ox ?=> A, K])(using @unused ev: KeyCodec[K] pattern: Option[String] = None, count: Option[Long] = None ): Ox ?=> Flow[(F, V)] = - scanStream(cursor => client.run(Hashes.hScan[K, F, V](key, cursor, pattern, count))) + paged(Paged.scanKey(client.runner)(cursor => Hashes.hScan[K, F, V](key, cursor, pattern, count))) /** * Iterates over all SSCAN members until the server returns a zero cursor. An empty page with a non-zero cursor continues the scan. @@ -51,7 +49,7 @@ extension [K](client: Client[[A] =>> Ox ?=> A, K])(using @unused ev: KeyCodec[K] pattern: Option[String] = None, count: Option[Long] = None ): Ox ?=> Flow[V] = - scanStream(cursor => client.run(Sets.sScan[K, V](key, cursor, pattern, count))) + paged(Paged.scanKey(client.runner)(cursor => Sets.sScan[K, V](key, cursor, pattern, count))) /** * Iterates over all ZSCAN member/score pairs until the server returns a zero cursor. An empty page with a non-zero cursor continues the scan. @@ -61,18 +59,11 @@ extension [K](client: Client[[A] =>> Ox ?=> A, K])(using @unused ev: KeyCodec[K] pattern: Option[String] = None, count: Option[Long] = None ): Ox ?=> Flow[(V, Double)] = - scanStream(cursor => client.run(SortedSets.zScan[K, V](key, cursor, pattern, count))) + paged(Paged.scanKey(client.runner)(cursor => SortedSets.zScan[K, V](key, cursor, pattern, count))) // convert pages from the shared Paged helper into individual Flow elements - private def paged[S, A](init: S)(step: Paged.Step[S, A]): Ox ?=> Flow[A] = - CStream.unfold[S, Vector[A]](init)(step).flatMap(items => CStream.init(items)).lower - - private def scanStream[A](fetch: ScanCursor => (Ox ?=> ScanPage[A])): Ox ?=> Flow[A] = - paged[Option[ScanCursor], A](Some(ScanCursor.start))(Paged.byCursor(cursor => CIO.lift(fetch(cursor)))) - - // scan each target in sequence with its own node-local cursor. A cluster scan visits every master that owns slots. - private def scanStreamAll[A](fetch: ScanTarget => ScanCursor => (Ox ?=> ScanPage[A])): Ox ?=> Flow[A] = - paged[ScanStep, A](ScanStep.Begin)(Paged.acrossTargets(CIO.lift(client.scanTargets))(target => cursor => CIO.lift(fetch(target)(cursor)))) + private def paged[S, A](pages: Paged.Pages[S, A]): Ox ?=> Flow[A] = + CStream.unfold[S, Vector[A]](pages.init)(pages.step).flatMap(items => CStream.init(items)).lower /** * Lazily pages an entire stream by range, batching `XRANGE` and advancing past the last id each page. Stops when a page comes back empty. @@ -83,9 +74,7 @@ extension [K](client: Client[[A] =>> Ox ?=> A, K])(using @unused ev: KeyCodec[K] end: StreamRangeId = StreamRangeId.Max, batch: Long = 100L ): Ox ?=> Flow[StreamEntry[F, V]] = - paged[Option[StreamRangeId], StreamEntry[F, V]](Some(start))( - Paged.byRange(batch)(from => CIO.lift(client.run(Streams.xRange[K, F, V](key, from, end, Some(batch))))) - ) + paged(Paged.xRangeAll[K, F, V](client.runner, key, start, end, batch)) /** * Auto-claims idle pending entries for `consumer`, advancing the `XAUTOCLAIM` cursor until it returns to the start. Entries whose data @@ -99,9 +88,7 @@ extension [K](client: Client[[A] =>> Ox ?=> A, K])(using @unused ev: KeyCodec[K] start: StreamId = StreamId.Zero, count: Option[Long] = None ): Ox ?=> Flow[StreamEntry[F, V]] = - paged[Option[StreamId], StreamEntry[F, V]](Some(start))( - Paged.byAutoClaim(from => CIO.lift(client.run(Streams.xAutoClaim[K, F, V](key, group, consumer, minIdle, from, count)))) - ) + paged(Paged.xAutoClaimAll[K, F, V](client.runner, key, group, consumer, minIdle, start, count)) /** * Follows a stream without a consumer group. It first reads every entry after `from`, then waits for new entries. The explicit entry ID @@ -114,11 +101,7 @@ extension [K](client: Client[[A] =>> Ox ?=> A, K])(using @unused ev: KeyCodec[K] count: Option[Long] = None, block: BlockTimeout = Paged.defaultPoll ): Ox ?=> Flow[StreamEntry[F, V]] = - paged[StreamId, StreamEntry[F, V]](from)( - Paged.tail(last => - CIO.lift(client.run(Streams.xRead[K, F, V]((key, ReadId.After(last)))(count = count, block = Some(block)))).map(_.flatMap(_._2)) - ) - ) + paged(Paged.xTail[K, F, V](client.runner, key, from, count, block)) /** * Follows a stream as part of a consumer group. It processes this consumer's pending entries first, then waits for new entries. Each @@ -131,29 +114,12 @@ extension [K](client: Client[[A] =>> Ox ?=> A, K])(using @unused ev: KeyCodec[K] count: Option[Long] = None, block: BlockTimeout = Paged.defaultPoll )(handle: StreamEntry[F, V] => (Ox ?=> Unit)): Ox ?=> Unit = - consumeStream[F, V](group, consumer, key, count, block).runForeach { entry => + paged(Paged.xConsume[K, F, V](client.runner, group, consumer, key, count, block)).runForeach { entry => handle(entry) client.run(Streams.xAck(key, group)(entry.id)) () } - private def consumeStream[F: KeyCodec, V: ValueCodec]( - group: String, - consumer: String, - key: K, - count: Option[Long], - block: BlockTimeout - ): Ox ?=> Flow[StreamEntry[F, V]] = - paged[Either[StreamId, Unit], StreamEntry[F, V]](Left(StreamId.Zero))( - Paged.consume( - drainPending = after => - CIO.lift(client.run(Streams.xReadGroup[K, F, V](group, consumer)((key, GroupReadId.After(after)))(count = count))).map(_.flatMap(_._2)), - tailNew = CIO - .lift(client.run(Streams.xReadGroup[K, F, V](group, consumer)((key, GroupReadId.New))(count = count, block = Some(block)))) - .map(_.flatMap(_._2)) - ) - ) - /** * Subscribes to one or more channels each time the returned `Flow` runs. Ending the flow unsubscribes. Sage resubscribes after * reconnecting, but messages published while the connection is down are lost. With standalone and master-replica clients, use @@ -212,10 +178,7 @@ extension [K](client: Client[[A] =>> Ox ?=> A, K])(using @unused ev: KeyCodec[K] } private def scopedStreamOf[A](open: => Subscription[[X] =>> Ox ?=> X, A]): Ox ?=> Flow[A] = { - val sub = open - val closed = new AtomicBoolean(false) - def unsubscribe(): Unit = if (closed.compareAndSet(false, true)) sub.close - useInScope(sub)(_ => unsubscribe()) + val sub = useInScope(open)(_.close) Flow.usingEmit { emit => var continue = true while (continue) diff --git a/sage-client/ox/src/test/scala/sage/client/internal/OxLockCancellationSpec.scala b/sage-client/ox/src/test/scala/sage/client/internal/OxLockCancellationSpec.scala index c8648176..ce73e5a8 100644 --- a/sage-client/ox/src/test/scala/sage/client/internal/OxLockCancellationSpec.scala +++ b/sage-client/ox/src/test/scala/sage/client/internal/OxLockCancellationSpec.scala @@ -13,10 +13,10 @@ import sage.backend.SageClient class OxLockCancellationSpec extends LockCancellationSpec { private given ExecutionContext = munitExecutionContext - override protected def tryWithLock[A](commands: CommandRunner[CIO, String], lease: FiniteDuration)(body: CIO[A]): CIO[Option[A]] = + override protected def tryWithLock[A](commands: SharedRunner, lease: FiniteDuration)(body: CIO[A]): CIO[Option[A]] = CIO.deferLift(new SageClient.Lowered(new LockTestClient(commands)).lock[String](lease).tryWithLock("key")(body.lower)) - override protected def withLock[A](commands: CommandRunner[CIO, String], lease: FiniteDuration, wait: FiniteDuration)(body: CIO[A]): CIO[A] = + override protected def withLock[A](commands: SharedRunner, lease: FiniteDuration, wait: FiniteDuration)(body: CIO[A]): CIO[A] = CIO.deferLift(new SageClient.Lowered(new LockTestClient(commands)).lock[String](lease).withLock("key", wait)(body.lower)) test("a body control exception ends the scope and releases ownership") { diff --git a/sage-client/pekko/src/main/scala/sage/backend/SageClient.scala b/sage-client/pekko/src/main/scala/sage/backend/SageClient.scala index b039ac0e..246727a7 100644 --- a/sage-client/pekko/src/main/scala/sage/backend/SageClient.scala +++ b/sage-client/pekko/src/main/scala/sage/backend/SageClient.scala @@ -13,7 +13,7 @@ import org.apache.pekko.stream.scaladsl.{Keep, Sink, Source} import sage.{Message, PatternMessage, SageException} import sage.client.SageConfig -import sage.client.internal.{Client, LoweredClient, Paged, ScanStep, ScanTarget, Subscription} +import sage.client.internal.{Client, LoweredClient, Paged, Subscription} import sage.codec.{KeyCodec, ValueCodec} import sage.commands.* @@ -33,7 +33,7 @@ extension [K](client: Client[Future, K])(using @unused ev: KeyCodec[K]) { count: Option[Long] = None, ofType: Option[RedisType] = None ): Source[K, NotUsed] = - scanSourceAll(target => cursor => client.runOn(target, Keys.scan[K](cursor, pattern, count, ofType))) + pagedSource(Paged.scanAll[K](client.runner, pattern, count, ofType)) /** * Iterates over all HSCAN field/value pairs until the server returns a zero cursor. An empty page with a non-zero cursor continues the scan. @@ -43,7 +43,7 @@ extension [K](client: Client[Future, K])(using @unused ev: KeyCodec[K]) { pattern: Option[String] = None, count: Option[Long] = None ): Source[(F, V), NotUsed] = - scanSource(cursor => client.run(Hashes.hScan[K, F, V](key, cursor, pattern, count))) + pagedSource(Paged.scanKey(client.runner)(cursor => Hashes.hScan[K, F, V](key, cursor, pattern, count))) /** * Iterates over all SSCAN members until the server returns a zero cursor. An empty page with a non-zero cursor continues the scan. @@ -53,7 +53,7 @@ extension [K](client: Client[Future, K])(using @unused ev: KeyCodec[K]) { pattern: Option[String] = None, count: Option[Long] = None ): Source[V, NotUsed] = - scanSource(cursor => client.run(Sets.sScan[K, V](key, cursor, pattern, count))) + pagedSource(Paged.scanKey(client.runner)(cursor => Sets.sScan[K, V](key, cursor, pattern, count))) /** * Iterates over all ZSCAN member/score pairs until the server returns a zero cursor. An empty page with a non-zero cursor continues the scan. @@ -63,7 +63,7 @@ extension [K](client: Client[Future, K])(using @unused ev: KeyCodec[K]) { pattern: Option[String] = None, count: Option[Long] = None ): Source[(V, Double), NotUsed] = - scanSource(cursor => client.run(SortedSets.zScan[K, V](key, cursor, pattern, count))) + pagedSource(Paged.scanKey(client.runner)(cursor => SortedSets.zScan[K, V](key, cursor, pattern, count))) /** * Lazily pages an entire stream by range, batching `XRANGE` and advancing past the last id each page. Stops when a page comes back empty. @@ -74,9 +74,7 @@ extension [K](client: Client[Future, K])(using @unused ev: KeyCodec[K]) { end: StreamRangeId = StreamRangeId.Max, batch: Long = 100L ): Source[StreamEntry[F, V], NotUsed] = - pagedSource[Option[StreamRangeId], StreamEntry[F, V]](Some(start))( - Paged.byRange(batch)(from => CIO.lift(client.run(Streams.xRange[K, F, V](key, from, end, Some(batch))))) - ) + pagedSource(Paged.xRangeAll[K, F, V](client.runner, key, start, end, batch)) /** * Auto-claims idle pending entries for `consumer`, advancing the `XAUTOCLAIM` cursor until it returns to the start. Entries whose data @@ -90,9 +88,7 @@ extension [K](client: Client[Future, K])(using @unused ev: KeyCodec[K]) { start: StreamId = StreamId.Zero, count: Option[Long] = None ): Source[StreamEntry[F, V], NotUsed] = - pagedSource[Option[StreamId], StreamEntry[F, V]](Some(start))( - Paged.byAutoClaim(from => CIO.lift(client.run(Streams.xAutoClaim[K, F, V](key, group, consumer, minIdle, from, count)))) - ) + pagedSource(Paged.xAutoClaimAll[K, F, V](client.runner, key, group, consumer, minIdle, start, count)) /** * Follows a stream without a consumer group: replays every entry after `from`, then blocks for new entries forever. Cancel the `Source` @@ -107,13 +103,7 @@ extension [K](client: Client[Future, K])(using @unused ev: KeyCodec[K]) { SageClient.boundedPoll(block, "xTail") match { case Left(e) => Source.failed(e) case Right(poll) => - pagedSource[StreamId, StreamEntry[F, V]](from)( - Paged.tail(last => - CIO - .lift(client.run(Streams.xRead[K, F, V]((key, ReadId.After(last)))(count = count, block = Some(poll)))) - .map(_.flatMap(_._2)) - ) - ) + pagedSource(Paged.xTail[K, F, V](client.runner, key, from, count, poll)) } /** @@ -133,7 +123,7 @@ extension [K](client: Client[Future, K])(using @unused ev: KeyCodec[K]) { case Right(poll) => given Materializer = SystemMaterializer(system).materializer given ExecutionContext = system.executionContext - val (killSwitch, done) = consumeSource[F, V](group, consumer, key, count, poll) + val (killSwitch, done) = pagedSource(Paged.xConsume[K, F, V](client.runner, group, consumer, key, count, poll)) .viaMat(KillSwitches.single)(Keep.right) .mapAsync(1)(entry => handle(entry).flatMap(_ => client.run(Streams.xAck(key, group)(entry.id)).map(_ => ()))) .toMat(Sink.ignore)(Keep.both) @@ -163,36 +153,12 @@ extension [K](client: Client[Future, K])(using @unused ev: KeyCodec[K]) { def sSubscribe[V: ValueCodec](channel: String, rest: String*): Source[Message[V], Future[Done]] = subscriptionSource(client.subscribeShardChannels[V](channel, rest*)) - private def scanSource[A](fetch: ScanCursor => Future[ScanPage[A]]): Source[A, NotUsed] = - pagedSource[Option[ScanCursor], A](Some(ScanCursor.start))(Paged.byCursor(cursor => CIO.lift(fetch(cursor)))) - - // scan each target with its own node-local cursor. A cluster has one target for every slot-owning master. - private def scanSourceAll[A](fetch: ScanTarget => ScanCursor => Future[ScanPage[A]]): Source[A, NotUsed] = - pagedSource[ScanStep, A](ScanStep.Begin)( - Paged.acrossTargets(CIO.lift(client.scanTargets))(target => cursor => CIO.lift(fetch(target)(cursor))) - ) - - private def consumeSource[F: KeyCodec, V: ValueCodec]( - group: String, - consumer: String, - key: K, - count: Option[Long], - block: BlockTimeout - ): Source[StreamEntry[F, V], NotUsed] = - pagedSource[Either[StreamId, Unit], StreamEntry[F, V]](Left(StreamId.Zero))( - Paged.consume( - drainPending = after => - CIO.lift(client.run(Streams.xReadGroup[K, F, V](group, consumer)((key, GroupReadId.After(after)))(count = count))).map(_.flatMap(_._2)), - tailNew = CIO - .lift(client.run(Streams.xReadGroup[K, F, V](group, consumer)((key, GroupReadId.New))(count = count, block = Some(block)))) - .map(_.flatMap(_._2)) - ) - ) - // convert pages from the shared Paged helper into individual Source elements - private def pagedSource[S, A](init: S)(step: Paged.Step[S, A]): Source[A, NotUsed] = + private def pagedSource[S, A](pages: Paged.Pages[S, A]): Source[A, NotUsed] = Source - .unfoldAsync[S, Vector[A]](init)(s => step(s).unsafeRun.map(_.map { case (items, next) => (next, items) })(using ExecutionContext.parasitic)) + .unfoldAsync[S, Vector[A]](pages.init)(s => + pages.step(s).unsafeRun.map(_.map { case (items, next) => (next, items) })(using ExecutionContext.parasitic) + ) .mapConcat(identity) // Open on materialization and close on cancellation, completion, or failure. Complete Future[Done] after the subscription open call finishes; diff --git a/sage-client/pekko/src/test/scala/sage/client/internal/LockFutureSpec.scala b/sage-client/pekko/src/test/scala/sage/client/internal/LockFutureSpec.scala index 9f4ceaf6..4680531d 100644 --- a/sage-client/pekko/src/test/scala/sage/client/internal/LockFutureSpec.scala +++ b/sage-client/pekko/src/test/scala/sage/client/internal/LockFutureSpec.scala @@ -39,12 +39,11 @@ class LockFutureSpec extends munit.FunSuite { } ), Scheduler.real, - Vector(Connection.hello()), SageConfig(dedicatedPool = DedicatedPoolConfig(maxConnections = 1, acquireTimeout = 200.millis), closeTimeout = Duration.Zero), Vector(master), MasterReplicaConfig() ) - live.bootstrapRoles() + live.start() val client = new SageClient.Lowered(live) val checked = for { error <- client @@ -67,7 +66,7 @@ class LockFutureSpec extends munit.FunSuite { val continued = Promise[Int]() val bodyStarted = new AtomicBoolean(false) val released = new AtomicBoolean(false) - val commands = new CommandRunner[CIO, String] { + val commands = new SharedRunner { def run[A](command: Command[A]): CIO[A] = CIO.defer(()).flatMap { _ => val reply = command.args(4).asUtf8String match { case "renew" => if (bodyStarted.get()) 0L else 1L diff --git a/sage-client/shared/src/main/scala/sage/client/SageConfig.scala b/sage-client/shared/src/main/scala/sage/client/SageConfig.scala index 34142403..c95fe7aa 100644 --- a/sage-client/shared/src/main/scala/sage/client/SageConfig.scala +++ b/sage-client/shared/src/main/scala/sage/client/SageConfig.scala @@ -202,7 +202,7 @@ object SageConfig { uri.split("://", 2) match { case Array(scheme, rest) => for { - tls <- scheme.toLowerCase match { + tls <- scheme.toLowerCase(java.util.Locale.ROOT) match { case "redis" => Right(None) case "rediss" => Right(Some(TlsConfig())) case other => fail[Option[TlsConfig]](s"unsupported scheme '$other' (expected redis or rediss)") diff --git a/sage-client/shared/src/main/scala/sage/client/internal/Backoff.scala b/sage-client/shared/src/main/scala/sage/client/internal/Backoff.scala deleted file mode 100644 index 7901a8fe..00000000 --- a/sage-client/shared/src/main/scala/sage/client/internal/Backoff.scala +++ /dev/null @@ -1,14 +0,0 @@ -package sage.client.internal - -import sage.client.BackoffConfig - -private[internal] object Backoff { - - // Reconnects and cluster-redirect retries share this formula: exponential backoff capped at maxDelay, then full jitter in [0, base]. - def jitteredMillis(config: BackoffConfig, attempt: Int, scheduler: Scheduler): Long = { - val capped = config.maxDelay.toMillis - val raw = config.initialDelay.toMillis.toDouble * math.pow(config.multiplier, attempt.toDouble) - val base = if (raw.isInfinite || raw >= capped.toDouble) capped else math.max(0L, raw.toLong) - scheduler.jitterMillis(base + 1) - } -} diff --git a/sage-client/shared/src/main/scala/sage/client/internal/Bootstrap.scala b/sage-client/shared/src/main/scala/sage/client/internal/Bootstrap.scala index e02af4e9..e5c814b7 100644 --- a/sage-client/shared/src/main/scala/sage/client/internal/Bootstrap.scala +++ b/sage-client/shared/src/main/scala/sage/client/internal/Bootstrap.scala @@ -3,10 +3,9 @@ package sage.client.internal import java.util.concurrent.{CountDownLatch, TimeUnit} import java.util.concurrent.atomic.AtomicReference -import scala.util.{Failure, Success, Try} +import scala.util.{Failure, Try} -import sage.SageException.{ConnectionLost, ServerError} -import sage.client.{AuthConfig, BuildInfo} +import sage.client.{BuildInfo, SageConfig} import sage.commands.{Command, Connection} private[client] object Bootstrap { @@ -16,63 +15,33 @@ private[client] object Bootstrap { * share this list, keeping connection identification consistent across topologies. `SELECT` lives here rather than as a runtime command * because it would move the database under every fiber sharing the connection. */ - def commands(auth: Option[AuthConfig], database: Int, clientName: Option[String]): Vector[Command[?]] = { + def commands(config: SageConfig): Vector[Command[?]] = { val identification = Vector( Connection.clientSetInfo("LIB-NAME", "sage"), Connection.clientSetInfo("LIB-VER", BuildInfo.version) - ) ++ clientName.map(Connection.clientSetName).toVector - val selectDb = if (database > 0) Vector(Connection.select(database)) else Vector.empty - (Connection.hello(auth.map(a => a.username -> a.password)) +: identification) ++ selectDb + ) ++ config.clientName.map(Connection.clientSetName).toVector + val selectDb = if (config.database > 0) Vector(Connection.select(config.database)) else Vector.empty + (Connection.hello(config.auth.map(a => a.username -> a.password)) +: identification) ++ selectDb } /** - * Runs the connection-setup handshake on a freshly opened connection: submits each command in turn and blocks for its reply up to - * `connectTimeoutMillis`, then closes the half-built connection and throws on a timeout or a failed reply so the caller discards it. - * `submit` enqueues one command and delivers its decoded reply; replies are FIFO, so each command is awaited before the next is sent. - * A [[bestEffort]] command whose reply is a `ServerError` is tolerated rather than fatal (and reported to `onTolerated`); a - * `ConnectionLost`/`DecodeError` stays fatal even for it. Used by the Multiplexed, Dedicated, and Subscription connections, which differ - * only in how a reply is obtained. + * Submits one command and blocks the calling thread for its reply, failing with `timedOut` when none arrives in time. */ - def run( - commands: Vector[Command[?]], - connectTimeoutMillis: Long, - submit: (Command[?], Try[Any] => Unit) => Unit, - close: () => Unit, - onTolerated: Command[?] => Unit = _ => () - ): Unit = - commands.foreach { command => - awaitReply[Any](connectTimeoutMillis)(callback => submit(command, callback)) match { - case None => - close() - throw ConnectionLost(mayHaveExecuted = false) - case Some(Failure(_: ServerError)) if bestEffort(command) => onTolerated(command) - case Some(Failure(error)) => - close() - throw error - case Some(Success(_)) => () - } - } - - /** - * Submits one command and blocks the calling thread for its reply; `None` is the timeout. - */ - def awaitReply[A](timeoutMillis: Long)(submit: (Try[A] => Unit) => Unit): Option[Try[A]] = { + def awaitReply[A](timeoutMillis: Long, timedOut: => Throwable)(submit: (Try[A] => Unit) => Unit): Try[A] = { val latch = new CountDownLatch(1) val outcome = new AtomicReference[Try[A]]() submit { result => outcome.set(result) latch.countDown() } - if (latch.await(timeoutMillis, TimeUnit.MILLISECONDS)) Some(outcome.get()) else None + if (latch.await(timeoutMillis, TimeUnit.MILLISECONDS)) outcome.get() else Failure(timedOut) } /** * Whether a server-error reply to this command may be tolerated during bootstrap. `CLIENT SETINFO` qualifies because it is library - * identification added in Redis 7.2, so an older server rejects it with an error every client ignores. `CLIENT TRACKING` qualifies because - * a server that permits `HELLO` but denies tracking (an ACL restriction, a proxy) should still connect and serve cached reads uncached - * rather than fail the connection (ADR-0045). Every other bootstrap command is load-bearing, so its failure stays fatal. + * identification added in Redis 7.2, so an older server rejects it with an error every client ignores. Every other bootstrap command is + * load-bearing, so its failure stays fatal. */ - private def bestEffort(command: Command[?]): Boolean = - Connection.isClientTracking(command) || - (command.name == "CLIENT" && command.args.headOption.exists(_.asUtf8String == "SETINFO")) + def bestEffort(command: Command[?]): Boolean = + command.name == "CLIENT" && command.args.headOption.exists(_.asUtf8String == "SETINFO") } diff --git a/sage-client/shared/src/main/scala/sage/client/internal/Client.scala b/sage-client/shared/src/main/scala/sage/client/internal/Client.scala index f5d9f506..7c64ccfa 100644 --- a/sage-client/shared/src/main/scala/sage/client/internal/Client.scala +++ b/sage-client/shared/src/main/scala/sage/client/internal/Client.scala @@ -1,9 +1,7 @@ package sage.client.internal import java.time.Instant -import java.util.concurrent.atomic.{AtomicBoolean, AtomicReference} -import java.util.concurrent.locks.ReentrantLock -import javax.net.ssl.SSLException +import java.util.concurrent.atomic.AtomicReference import scala.concurrent.duration.{Duration, FiniteDuration} import scala.util.{Failure, Success, Try} @@ -31,12 +29,6 @@ trait CommandRunner[F[_], K](using KeyCodec[K]) { */ def run[A](command: Command[A]): F[A] - private[sage] def lockWrite( - command: Command[Boolean], - @scala.annotation.unused timeout: FiniteDuration, - @scala.annotation.unused replicaAcknowledgement: Boolean - ): F[Boolean] = run(command) - /** * Returns a view that uses another key type and reuses the same connection. Command builders encode keys before calling `run`, so the * runner itself does not depend on the key type. For example, `client.as[Array[Byte]]` accepts binary keys, including inside a @@ -47,11 +39,6 @@ trait CommandRunner[F[_], K](using KeyCodec[K]) { val self = this new CommandRunner[F, K2] { def run[A](command: Command[A]): F[A] = self.run(command) - override private[sage] def lockWrite( - command: Command[Boolean], - timeout: FiniteDuration, - replicaAcknowledgement: Boolean - ): F[Boolean] = self.lockWrite(command, timeout, replicaAcknowledgement) } } @@ -2162,18 +2149,15 @@ trait CommandRunner[F[_], K](using KeyCodec[K]) { final def arInfoFull(key: K): F[ArrayInfoFull] = run(Arrays.arInfoFull(key)) } -// One independent keyspace visited by SCAN: the standalone server or one slot-owning cluster master. A missing node means that normal -// routing should be used, either for a standalone server or before cluster topology discovery finishes. -final private[sage] case class ScanTarget(node: Option[Node]) - -private[sage] object ScanTarget { - val any: ScanTarget = ScanTarget(None) +// One independent keyspace visited by SCAN: the standalone server or one slot-owning cluster master. `run` sends a page to that keyspace. +private[sage] trait ScanTarget { + def run[A](command: Command[A]): CIO[A] } // tracks a cluster-wide SCAN while each target is scanned to its node-local zero cursor in turn. private[sage] enum ScanStep { case Begin - case Visit(cursor: ScanCursor, remaining: Vector[ScanTarget]) + case Visit(cursor: ScanCursor, target: ScanTarget, rest: Vector[ScanTarget]) case End } @@ -2188,16 +2172,14 @@ trait Client[F[_], K] extends CommandRunner[F, K] { /** * Runs a read with client-side caching. A cached value remains until the server invalidates it, its `ttl` expires, or the cache evicts it * to stay within `maxBytes`. A command qualifies when it is a cacheable read with at least one key and its result changes only when data - * at those keys changes. A write, a keyless read, or a time-varying or non-deterministic read (`TTL`, `SRANDMEMBER`) fails with - * [[sage.SageException.NotCacheable]]. Cached reads use the master regardless of the `ReadFrom` policy. In a cluster, each slot-owning + * at those keys changes. A write, a keyless read, a blocking command, or a time-varying or non-deterministic read (`TTL`, `SRANDMEMBER`) + * fails with [[sage.SageException.NotCacheable]]. Cached reads use the master regardless of the `ReadFrom` policy. In a cluster, each slot-owning * master has an independent cache. The same call works with every topology. When caching is disabled or the server rejects tracking, it * runs uncached. */ def cached[A](command: Command[A], ttl: FiniteDuration): F[A] - private[sage] def pipeline[Out, R](p: Pipeline[Out, R]): F[Out] - - private[sage] def pipelineAttempt[Out, R](p: Pipeline[Out, R]): F[R] + private[sage] def pipeline[R](p: Pipeline[R]): F[R] /** * Runs a fixed-arity batch of commands in one round-trip, yielding a result tuple that mirrors the argument tuple element-for-element @@ -2217,13 +2199,13 @@ trait Client[F[_], K] extends CommandRunner[F, K] { * Like the tuple [[pipeline]], but yields the per-position results, each slot a `Right`/`Left`, instead of failing on the first error. */ def pipelineAttempt[T <: NonEmptyTuple](commands: T)(using Tuple.IsMappedBy[Command][T]): F[Tuple.Map[Tuple.InverseMap[T, Command], Attempt]] = - pipelineAttempt(Pipeline.fromTuple(commands)) + pipeline(Pipeline.fromTupleAttempt(commands)) /** * Like the `Seq` [[pipeline]], but yields the per-position results, each slot a `Right`/`Left`, instead of failing on the first error. */ def pipelineAttempt[A](commands: Seq[Command[A]]): F[Vector[Attempt[A]]] = - pipelineAttempt(Pipeline.sequence(commands)) + pipeline(Pipeline.sequenceAttempt(commands)) /** * Opens a [[TransactionScope]] on a leased Dedicated Connection for `MULTI`/`EXEC`, optionally guarded by `WATCH`. @@ -2273,10 +2255,8 @@ trait Client[F[_], K] extends CommandRunner[F, K] { */ def subscribeShardChannels[V: ValueCodec](channel: String, rest: String*): F[Subscription[F, Message[V]]] - // return every keyspace SCAN must visit. runOn sends the next page to the node that issued its cursor. - private[sage] def scanTargets: F[Vector[ScanTarget]] - - private[sage] def runOn[A](target: ScanTarget, command: Command[A]): F[A] + // the shared CIO client behind this one. Concrete only because MiMa reports new abstract members; every Sage client overrides it. + private[sage] def runner: SharedRunner = SharedRunner.unavailable /** * Releases all connections and the client's resources. @@ -2291,20 +2271,13 @@ trait Client[F[_], K] extends CommandRunner[F, K] { val self = this new Client[F, K2] { def run[A](command: Command[A]): F[A] = self.run(command) - override private[sage] def lockWrite( - command: Command[Boolean], - timeout: FiniteDuration, - replicaAcknowledgement: Boolean - ): F[Boolean] = self.lockWrite(command, timeout, replicaAcknowledgement) def cached[A](command: Command[A], ttl: FiniteDuration): F[A] = self.cached(command, ttl) - private[sage] def pipeline[Out, R](p: Pipeline[Out, R]): F[Out] = self.pipeline(p) - private[sage] def pipelineAttempt[Out, R](p: Pipeline[Out, R]): F[R] = self.pipelineAttempt(p) + private[sage] def pipeline[R](p: Pipeline[R]): F[R] = self.pipeline(p) def transaction[A](body: TransactionScope[F, K2] => F[A]): F[A] = self.transaction(scope => body(scope.as[K2])) def subscribeChannels[V: ValueCodec](channel: String, rest: String*) = self.subscribeChannels(channel, rest*) def subscribePatterns[V: ValueCodec](pattern: String, rest: String*) = self.subscribePatterns(pattern, rest*) def subscribeShardChannels[V: ValueCodec](channel: String, rest: String*) = self.subscribeShardChannels(channel, rest*) - private[sage] def scanTargets: F[Vector[ScanTarget]] = self.scanTargets - private[sage] def runOn[A](target: ScanTarget, command: Command[A]): F[A] = self.runOn(target, command) + override private[sage] def runner: SharedRunner = self.runner private[sage] def rateLimitAcquire[RK](executor: RateLimitExecutor[RK], subject: RK, cost: Long, peek: Boolean): F[Decision] = self.rateLimitAcquire(executor, subject, cost, peek) private[sage] def lockTryWith[LK, A](executor: LockExecutor[LK], key: LK)(body: => F[A]): F[Option[A]] = @@ -2321,16 +2294,25 @@ object Client { private val defaults = SageConfig() // Server invalidations can keep a cached read current only when the command names at least one key. Reject keyless reads even if they are - // otherwise deterministic. - private[internal] def cacheable(command: Command[?]): Boolean = command.cacheable && command.keyIndices.nonEmpty + // otherwise deterministic, and reject key positions outside the arguments, since the cache cannot read those keys. + private[internal] def cacheable(command: Command[?]): Boolean = + command.cacheable && !command.isBlocking && command.keyIndices.nonEmpty && !command.hasMalformedKeys private[internal] def notCacheable(command: Command[?]): NotCacheable = - NotCacheable(s"${command.name} is not cacheable: cached requires a cacheable command with at least one key") + NotCacheable(s"${command.name} is not cacheable: cached requires a cacheable command with at least one key within its arguments") - // a closed transport may throw before registering its callback. Complete CIO.async with that synchronous failure. + // A submit can throw before it registers its callback, for example when the cluster refresh it runs is interrupted. + // Complete CIO.async with that failure. private[internal] def completing[A](complete: Try[A] => Unit)(submit: => Unit): Unit = try submit - catch { case NonFatal(error) => complete(Failure(error)) } + catch { case error: Throwable => complete(Failure(error)) } + + // run a script by digest, and send its body once if the server does not have it cached + private[internal] def withScriptFallback[A](run: Boolean => CIO[A]): CIO[A] = + run(true).recover { + case ServerError("NOSCRIPT", _) => run(false) + case other => CIO.fail(other) + } // create a lease for each execution. Leases cannot be reused after cancellation, and interruption uses the lease to release the pool slot. private[internal] def withLeaseIfBlocking[A](command: Command[?])(body: DedicatedPool.Lease => CIO[A]): CIO[A] = @@ -2350,13 +2332,8 @@ object Client { } } - // attribute the node at completion. A batch that never reaches the wire leaves its callbacks unattributed. - private def attributeOnComplete(cb: Try[Any] => Unit, node: Node): Try[Any] => Unit = - result => { - Events.attributeNode(cb, node) - cb(result) - } - + // A pipeline batched onto a single connection: scatter each reply into its submission-order slot, and on a disconnect report a failed + // completion per position before failing the effect once. Shared by standalone/master-replica; the cluster runtime splits per node instead. final private[internal] class TrackedBatch( events: Events, commands: Vector[Command[?]], @@ -2372,19 +2349,16 @@ object Client { def callbacks(node: Option[Node]): Vector[Try[Any] => Unit] = node match { - case Some(n) => tracked.map(attributeOnComplete(_, n)) + // attribute the node at completion. A batch that never reaches the wire leaves its callbacks unattributed. + case Some(n) => tracked.map(Events.completeAt(_, n)) case None => tracked } def settleAll(node: Node, results: Vector[Try[Any]]): Unit = - results.indices.foreach { i => - Events.attributeNode(tracked(i), node) - tracked(i)(results(i)) - } + results.indices.foreach(i => Events.completeAt(tracked(i), node)(results(i))) - def failUnsent(onUnsent: () => Unit): Unit = { + def failUnsent(): Unit = { val error = NotConnected() - onUnsent() tracked.foreach(Events.abandonSpan(_, error)) if (events.emitsEvents) commands.foreach(c => events.emit(SageEvent.CommandCompleted(c.name, None, Duration.Zero, Outcome.Failed(error)))) @@ -2392,37 +2366,47 @@ object Client { } } - // A pipeline batched onto a single connection: scatter each reply into its submission-order slot, and on a disconnect report a failed - // completion per position before failing the effect once. Shared by standalone/master-replica; the cluster runtime splits per node instead. - private[internal] def submitBatchOnOne( - events: Events, - commands: Vector[Command[?]], - spans: Vector[CommandSpan], - submitAll: (Vector[Command[?]], Vector[Try[Any] => Unit]) => Boolean, - complete: Try[Vector[Either[SageException, Any]]] => Unit, - onUnsent: () => Unit, - node: Option[Node] = None - ): Unit = { - val batch = new TrackedBatch(events, commands, spans, complete) - if (!submitAll(commands, batch.callbacks(node))) batch.failUnsent(onUnsent) - } - /** - * The construction entry point each backend's `connect`/`scoped` builds on: validates `config`, then connects per its [[Topology]]. + * The construction entry point each backend's `connect`/`scoped` builds on: validates `config`, then connects per its [[Topology]]. It + * fails with [[sage.SageException.ConnectionFailed]] when the server does not answer `HELLO` within `connectTimeout` or closes the + * socket during setup, and a master-replica connect fails with [[sage.SageException.TimedOut]] when its `ROLE` probes time out. */ def connect(config: SageConfig): CIO[Client[CIO, String]] = validate(config) match { case Some(problem) => CIO.fail(InvalidArgument(problem)) case None => - config.topology match { - case Topology.Standalone(endpoint) => connectStandalone(config, endpoint) - case Topology.Cluster(seeds, clusterConfig) => - ClusterLive.connect(config, seeds.map(e => Node(e.host, e.port)), clusterConfig, Scheduler.real, translateHandshake) - case Topology.MasterReplica(seeds, masterReplica) => - MasterReplicaLive.connect(config, seeds.map(e => Node(e.host, e.port)), masterReplica, Scheduler.real, translateHandshake) + start { + val factory = transports(config) + def events(server: Option[Node]) = Events(config.listeners, config.tracer, server) + def nodes(seeds: Vector[Endpoint]) = seeds.map(e => Node(e.host, e.port)) + config.topology match { + case Topology.Standalone(endpoint) => + val node = Node(endpoint.host, endpoint.port) + new Live(factory(node), Scheduler.real, config, events(Some(node))) + case Topology.Cluster(seeds, clusterConfig) => new ClusterLive(factory, Scheduler.real, config, clusterConfig, nodes(seeds), events(None)) + case Topology.MasterReplica(seeds, masterReplica) => + new MasterReplicaLive(factory, Scheduler.real, config, nodes(seeds), masterReplica, events(None)) + } } } + // connect reports every ordinary start failure as a SageException + private def start(build: => LiveClient): CIO[Client[CIO, String]] = + CIO.blocking { + val live = build + try live.start() + catch { case NonFatal(error) => throw translateHandshake(error) } + live + } + + // Each call builds the node's TLS context, which throws a TlsError for unusable trust material. The returned factory shares the context + // with every connection it creates (the node's reconnects and dedicated connections), and creating a transport does no I/O. + private[internal] def transports(config: SageConfig): Node => MultiplexedConnection.TransportFactory = + node => { + val upgrade = Tls.buildUpgrade(config.tls, node.host, node.port) + (onFrame, onClosed) => SocketTransport.connect(node.host, node.port, config.connectTimeout, upgrade, onFrame, onClosed) + } + // report invalid configuration through the connect effect instead of throwing during construction. private def validate(config: SageConfig): Option[String] = { // skip ping interval and timeout validation when the watchdog is disabled because those values are not used @@ -2476,168 +2460,67 @@ object Client { private def atLeastOneMilliOrInfinite(value: Duration, label: String): Option[String] = cond(value == Duration.Inf || (value.isFinite && value.toMillis >= 1L), s"$label must be at least 1ms (or Inf)") - private def connectStandalone(config: SageConfig, endpoint: Endpoint): CIO[Client[CIO, String]] = - // Build the TLS context once so invalid trust material fails during client creation. Capture it in the reconnect factory to apply the - // same upgrade to the multiplexed connection and every dedicated connection. - CIO.blocking(Tls.buildUpgrade(config.tls, endpoint.host, endpoint.port)).flatMap { upgrade => - connectWith( - (onFrame, onClosed) => SocketTransport.connect(endpoint.host, endpoint.port, config.connectTimeout, upgrade, onFrame, onClosed), - Scheduler.real, - config, - Events(config.listeners, config.tracer, serverNode = Some(Node(endpoint.host, endpoint.port))) - ) - } - // run the HELLO 3 bootstrap for every connection. The initial connection reports a handshake failure; reconnects retry it. private[client] def connectWith( factory: MultiplexedConnection.TransportFactory, scheduler: Scheduler = Scheduler.real, config: SageConfig = defaults, events: Events = Events.disabled - ): CIO[Client[CIO, String]] = { - val cachingEnabled = config.clientCache.enabled - val connectTimeout = config.connectTimeout - val bootstrap = Bootstrap.commands(config.auth, config.database, config.clientName) - // Enable tracking on the multiplexed connection, where cached reads run. Dedicated and subscription connections use the plain bootstrap. - // When caching is disabled, omit tracking from every connection so servers that deny CLIENT TRACKING can still connect. - val multiplexedBootstrap = if (cachingEnabled) bootstrap :+ Connection.clientTrackingOnOptin else bootstrap - CIO - .blocking( - MultiplexedConnection.connect( - factory, - scheduler, - multiplexedBootstrap, - config.reconnect, - config.watchdog, - connectTimeout, - config.closeTimeout, - config.clientCache.maxBytes, - None, - events - ) - ) - .map { connection => - val pool = DedicatedPool.forConnection(factory, bootstrap, scheduler, connection, config.dedicatedPool, connectTimeout.toMillis) - // open the subscription socket on first use, after the Multiplexed Connection becomes live - val subscriptions = new SubscriptionConnection( - factory, - bootstrap, - scheduler, - config.reconnect, - config.watchdog, - connectTimeout.toMillis, - config.pubsub.bufferSize, - () => connection.isLive, - events = events - ) - new Live(new NodeClient(connection, pool), subscriptions, cachingEnabled, events) - } - .mapError { error => - events.close() - translateHandshake(error) - } - } + ): CIO[Client[CIO, String]] = + start(new Live(factory, scheduler, config, events)) - // Redis versions before 6.0 report HELLO as an unknown command. Newer servers use NOPROTO for unsupported protocol versions. Convert TLS - // certificate and hostname failures to TlsError so all expected connection failures remain SageException values. - private def translateHandshake(error: Throwable): Throwable = + // Redis versions before 6.0 report HELLO as an unknown command. Newer servers use NOPROTO for unsupported protocol versions. A setup reply + // that times out or loses its socket is a failed connect, and a raw network error would otherwise escape the sealed hierarchy. + private[internal] def translateHandshake(error: Throwable): Throwable = error match { - case e: ServerError if e.code == "NOPROTO" || e.getMessage.toLowerCase.contains("unknown command") => + case e: ServerError if e.code == "NOPROTO" || e.getMessage.toLowerCase(java.util.Locale.ROOT).contains("unknown command") => UnsupportedServer(s"sage requires RESP3 (Redis 6.0+ or any Valkey); server rejected HELLO 3: ${e.getMessage}") - case e: SSLException => - TlsError(s"TLS handshake failed: ${e.getMessage}") - case e: SageException => e - // a raw network error would otherwise escape the sealed hierarchy - case other => + case e: SageException if !e.isInstanceOf[ConnectionLost] => e + case other => val failed = ConnectionFailed(s"could not connect: $other") failed.initCause(other) failed } - final private class Live( - nodeClient: NodeClient, - subscriptions: SubscriptionConnection, - cachingEnabled: Boolean, - events: Events - ) extends Client[CIO, String] { - - def run[A](command: Command[A]): CIO[A] = - Client.withLeaseIfBlocking(command) { lease => - CIO.async { callback => - val tracked = Events.trackCommand(events, command, callback) - Client.completing(tracked)(nodeClient.submit(command, asking = false, tracked, lease)) - } - } - - def cached[A](command: Command[A], ttl: FiniteDuration): CIO[A] = - if (!Client.cacheable(command)) CIO.fail(Client.notCacheable(command)) - else if (!cachingEnabled) run(command) // run uncached because connection setup did not enable CLIENT TRACKING - else CIO.async(callback => Client.completing(callback)(nodeClient.cachedSubmit(command, ttl.toMillis, callback))) - - def scanTargets: CIO[Vector[ScanTarget]] = CIO.value(Vector(ScanTarget.any)) - - def runOn[A](target: ScanTarget, command: Command[A]): CIO[A] = run(command) + // the same factory serves the multiplexed connection and every dedicated and subscription connection + final private class Live(factory: MultiplexedConnection.TransportFactory, scheduler: Scheduler, config: SageConfig, events: Events) + extends LiveClient(events) { - private[sage] def rateLimitAcquire[RK](executor: RateLimitExecutor[RK], subject: RK, cost: Long, peek: Boolean): CIO[Decision] = - executor.evalSha(this, subject, cost, peek) + // only the multiplexed connection enables tracking; dedicated and subscription connections never cache + private val nodeClient = new MultiplexedConnection(factory, scheduler, config, MultiplexedConnection.NodeRole.Master, None, events) + // open the subscription socket on first use, after the Multiplexed Connection becomes live + private val subscriptions = + new SubscriptionConnection(factory, scheduler, config, () => nodeClient.isLive, SubscriptionConnection.OnLoss.Reconnect(() => (), events)) - private[sage] def lockTryWith[LK, A](executor: LockExecutor[LK], key: LK)(body: => CIO[A]): CIO[Option[A]] = - executor.tryWithLock(this, key)(body) - - private[sage] def lockWith[LK, A](executor: LockExecutor[LK], key: LK, waitTimeout: FiniteDuration)(body: => CIO[A]): CIO[A] = - executor.withLock(this, key, waitTimeout)(body) + def run[A](command: Command[A]): CIO[A] = + Client.withLeaseIfBlocking(command)(lease => tracked(command)(nodeClient.submit(command, _, lease = lease))) - private[sage] def pipeline[Out, R](p: Pipeline[Out, R]): CIO[Out] = - submitPipeline(p).flatMap(TxSupport.collapseStrict(_, p.toOut)) + private val fetchTracking = Events.fetchTracking(events) - private[sage] def pipelineAttempt[Out, R](p: Pipeline[Out, R]): CIO[R] = - submitPipeline(p).map(p.toResults) + protected def cachedChecked[A](command: Command[A], ttl: FiniteDuration): CIO[A] = + CIO.async[A](complete => Client.completing(complete)(nodeClient.cachedSubmit(command, ttl.toMillis, complete, fetchTracking))) - // Release the transaction connection after success, failure, or interruption. Return it to the pool only after EXEC or UNWATCH has - // cleared WATCH/MULTI state and no replies remain pending. Discard it when watched keys or commands may still be active. - def transaction[A](body: TransactionScope[CIO, String] => CIO[A]): CIO[A] = - CIO.acquireReleaseWith(acquireScope)(releaseScope)(scope => CIO.unit.flatMap(_ => body(scope))) + // a standalone server uses one subscription connection for all shard channels. + protected def pubsub: SubscriptionConnection.PubSub = subscriptions - private def acquireScope: CIO[TxScope] = + protected def openTransaction: CIO[LiveTransactionScope] = CIO.blocking { - try new TxScope(nodeClient.acquireForTransaction(), events = events) + try new TxScope(nodeClient.pool.acquireForTransaction(), nodeClient.pool.releaseTransaction, events = events) catch { case e: SageException => throw e case NonFatal(_) => throw ConnectionLost(mayHaveExecuted = false) } } - private def releaseScope(scope: TxScope): CIO[Unit] = - CIO.blocking(nodeClient.releaseTransaction(scope.conn, scope.sealAndReusable())) - - private def submitPipeline[Out, R](p: Pipeline[Out, R]): CIO[Vector[Either[SageException, Any]]] = - if (p.commands.isEmpty) - CIO.value(Vector.empty) - else if (p.commands.exists(_.isBlocking)) - CIO.fail(InvalidArgument("a Pipeline cannot carry blocking commands; run them individually on the client")) - else - CIO.async { complete => - Client.submitBatchOnOne( - events, - p.commands, - Events.startSpans(events, p.commands), - nodeClient.submitAll, - complete, - onUnsent = () => () - ) - } - - def subscribeChannels[V: ValueCodec](channel: String, rest: String*): CIO[Subscription[CIO, Message[V]]] = - CIO.blocking(channelMessages(subscriptions.subscribeChannels(channel +: rest.toVector))) - - def subscribePatterns[V: ValueCodec](pattern: String, rest: String*): CIO[Subscription[CIO, PatternMessage[V]]] = - CIO.blocking(patternMessages(subscriptions.subscribePatterns(pattern +: rest.toVector))) + protected def submitPipeline[R](p: Pipeline[R]): CIO[Vector[Either[SageException, Any]]] = + CIO.async { complete => + val batch = new TrackedBatch(events, p.commands, Events.startSpans(events, p.commands), complete) + if (!nodeClient.submitAll(p.commands, batch.callbacks(None))) batch.failUnsent() + } - // a standalone server uses one subscription connection for all shard channels. - def subscribeShardChannels[V: ValueCodec](channel: String, rest: String*): CIO[Subscription[CIO, Message[V]]] = - CIO.blocking(channelMessages(subscriptions.subscribeShard(channel +: rest.toVector))) + protected def establish(): Unit = nodeClient.start(): Unit - def close: CIO[Unit] = CIO.blocking { + protected def shutdown(): Unit = { subscriptions.close() nodeClient.close() events.close() @@ -2659,9 +2542,10 @@ object Client { } CIO.ensure(deregister) { CIO.async { complete => + // an ended subscription yields None, or the server's error when it refused a name val cb: Option[SubscriptionConnection.Delivery] => Unit = { case Some(delivery) => complete(Try(build(delivery))) - case None => complete(Success(None)) + case None => complete(raw.failure.fold[Try[Option[M]]](Success(None))(Failure(_))) } registered.set(cb) raw.next(cb) @@ -2674,14 +2558,14 @@ object Client { // a channel/shard delivery is a Message private[internal] def channelMessages[V](raw: SubscriptionConnection.RawSubscription)(using ValueCodec[V]): Subscription[CIO, Message[V]] = messages(raw) { - case SubscriptionConnection.Delivery.Channel(ch, payload) => Some(Message(ch, decodeOrThrow[V](payload))) - case _ => None + case Message(ch, payload) => Some(Message(ch, decodeOrThrow[V](payload))) + case _ => None } private[internal] def patternMessages[V](raw: SubscriptionConnection.RawSubscription)(using ValueCodec[V]): Subscription[CIO, PatternMessage[V]] = messages(raw) { - case SubscriptionConnection.Delivery.Pattern(pat, ch, payload) => Some(PatternMessage(pat, ch, decodeOrThrow[V](payload))) - case _ => None + case PatternMessage(pat, ch, payload) => Some(PatternMessage(pat, ch, decodeOrThrow[V](payload))) + case _ => None } // fail the stream on a bad payload rather than dropping it @@ -2696,105 +2580,27 @@ object Client { case NonFatal(e) => throw DecodeError.fromThrowable(e) } - final private[internal] class TxScope(val conn: DedicatedConnection, onFault: Throwable => Unit = _ => (), events: Events = Events.disabled) - extends TransactionScope[CIO, String] { + final private[internal] class TxScope( + conn: DedicatedConnection, + returnConn: (DedicatedConnection, Boolean) => Unit, + refresh: RefreshPolicy => Unit = _ => (), + events: Events = Events.disabled + ) extends LiveTransactionScope(events, refresh) { - // true after WATCH is attempted and false after EXEC or UNWATCH; prevents reuse while the server may still track watched keys - val armed = new AtomicBoolean(false) + protected def leasedConn: DedicatedConnection = conn - private def faulting[A](complete: Try[A] => Unit): Try[A] => Unit = { - case failure @ Failure(error) => - onFault(error) - complete(failure) - case success => complete(success) - } + protected def giveBack(conn: DedicatedConnection, reusable: Boolean): Unit = returnConn(conn, reusable) - // Coordinate submission with release under one lock. A command accepted before release is recorded as in flight before - // [[sealAndReusable]] checks the connection, preventing the finalizer from recycling a busy connection. Commands submitted after release - // are rejected, preventing an old transaction handle from using a connection that another transaction has borrowed. - private val lock = new ReentrantLock() - private var released = false + protected type Target = Unit + protected def targetOf(command: Command[?]): Unit = () + protected def targetOf(commands: Vector[Command[?]]): Unit = () - private def submitting[A](complete: Try[A] => Unit)(submit: => Unit): Unit = { + protected def withConn[A](target: Unit, complete: Try[A] => Unit)(use: DedicatedConnection => Unit): Unit = { lock.lock() try if (released) complete(Failure(TxSupport.scopeReleasedError)) - else Client.completing(complete)(submit) - finally lock.unlock() - } - - // run once by the lease finalizer: seals the scope against further operations and reports whether the connection may be recycled - private[internal] def sealAndReusable(): Boolean = { - lock.lock() - try { - released = true - conn.isHealthy && conn.isQuiescent && !armed.get - } finally lock.unlock() - } - - private def isReleased: Boolean = { - lock.lock() - try released + else Client.completing(complete)(use(conn)) finally lock.unlock() } - - def watch[K: KeyCodec](key: K, rest: K*): CIO[Unit] = - CIO.async[Unit] { complete => - val watchCmd = Connection.watch(key, rest*) - val tracked = Events.trackSpan(events, watchCmd, complete) - submitting(tracked) { - armed.set(true) - conn.submit(watchCmd, faulting(tracked)) - } - } - - def run[A](command: Command[A]): CIO[A] = - if (isReleased) - CIO.fail(TxSupport.scopeReleasedError) - else if (command.isBlocking) - CIO.fail(InvalidArgument("a Transaction cannot run blocking commands; run them individually on the client")) - else - CIO.async[A] { complete => - val tracked = Events.trackSpan(events, command, complete) - submitting(tracked)(conn.submit(command, faulting(tracked))) - } - - def discard: CIO[Unit] = - CIO.async[Unit] { complete => - submitting(complete) { - armed.set(false) - conn.submit(Connection.unwatch, faulting(complete)) - } - } - - private[sage] def exec[Out, R](p: Pipeline[Out, R]): CIO[Option[Out]] = - runExec(p).flatMap { - case None => CIO.value(None) - case Some(results) => TxSupport.collapseStrict(results, p.toOut).map(Some(_)) - } - - private[sage] def execAttempt[Out, R](p: Pipeline[Out, R]): CIO[Option[R]] = - runExec(p).map(_.map(p.toResults)) - - // return None when EXEC reports a WATCH abort and Some with one decoded result per command; a queueing error fails the effect before execution - private def runExec[Out, R](p: Pipeline[Out, R]): CIO[Option[Vector[Either[SageException, Any]]]] = - if (isReleased) - CIO.fail(TxSupport.scopeReleasedError) - // skip MULTI/EXEC for an empty pipeline only when WATCH is inactive; watched keys still require EXEC to detect concurrent changes - else if (p.commands.isEmpty && !armed.get) - CIO.value(Some(Vector.empty)) - else if (p.commands.exists(_.isBlocking)) - CIO.fail(InvalidArgument("a Transaction cannot carry blocking commands; run them individually on the client")) - else - CIO - .async[Vector[Frame]] { complete => - val tracked = Events.trackSpan(events, Connection.multi, complete) - submitting(tracked)(conn.submitRaw(Connection.multi +: p.commands :+ Connection.exec, faulting(tracked))) - } - .flatMap { frames => - armed.set(false) // EXEC clears WATCH/MULTI state server-side whether it committed or aborted - TxSupport.execErrors(frames).foreach(onFault) - TxSupport.interpretExec(p.commands, frames) - } } } diff --git a/sage-client/shared/src/main/scala/sage/client/internal/ClientCache.scala b/sage-client/shared/src/main/scala/sage/client/internal/ClientCache.scala index 0e243508..47af62f0 100644 --- a/sage-client/shared/src/main/scala/sage/client/internal/ClientCache.scala +++ b/sage-client/shared/src/main/scala/sage/client/internal/ClientCache.scala @@ -19,18 +19,18 @@ final private[client] class ClientCache(maxBytes: Long) { import ClientCache.* import ClientCache.Acquire.* - private val lock = new ReentrantLock() + private val lock = new ReentrantLock() // accessOrder = true moves a read entry to the end of the map. Eviction then removes the least recently used entry first. - private val entries = new java.util.LinkedHashMap[Key, Entry](16, 0.75f, true) - private val reverse = mutable.HashMap.empty[Key, mutable.HashSet[Key]] - private val pending = mutable.HashMap.empty[Key, InFlight] - private var bytesUsed: Long = 0L - @volatile private var epoch: CacheEpoch = CacheEpoch.initial - @volatile private var rerouteWatermark: CacheEpoch = CacheEpoch.initial + private val entries = new java.util.LinkedHashMap[Key, Entry](16, 0.75f, true) + private val reverse = mutable.HashMap.empty[Key, mutable.HashSet[Key]] + private val pending = mutable.HashMap.empty[Key, Fetching] + private var bytesUsed: Long = 0L + // a flush retires every hit handed out before it + @volatile private var epoch = 0L /** * Tries to serve `commandBytes` from the cache. [[Hit]] returns the stored frame. Decode it and complete the caller. [[Fetch]] means that - * this caller is the first to miss. Read from the server, then call [[store]] or [[fail]]. [[Wait]] means that another fetch is in flight + * this caller is the first to miss. Read from the server, then pass its ticket to [[store]] or [[fail]]. [[Wait]] means that another fetch is in flight * and now owns `waiter`. Do nothing for this caller. [[Fetch]] and [[Wait]] enqueue `waiter`; [[Hit]] does not. */ def acquire(commandBytes: Bytes, trackedKeys: Vector[Bytes], now: Long, waiter: Try[Frame] => Unit): Acquire = { @@ -43,42 +43,35 @@ final private[client] class ClientCache(maxBytes: Long) { removeEntry(key, entry) } pending.get(key) match { - case Some(inFlight) => - inFlight.waiters += waiter + case Some(fetching) => + fetching.waiters += waiter Wait case None => - val inFlight = new InFlight(trackedKeys.map(new Key(_))) - inFlight.waiters += waiter - pending.update(key, inFlight) - Fetch + val fetching = new Fetching(key, trackedKeys.map(new Key(_))) + fetching.waiters += waiter + pending.update(key, fetching) + Fetch(fetching) } } finally lock.unlock() } - def store(commandBytes: Bytes, trackedKeys: Vector[Bytes], frame: Frame, now: Long, ttlMillis: Long): Unit = { - val key = new Key(commandBytes) - val size = frameSize(frame) // walked outside the lock so a large reply can't stall acquire/invalidate - var waiters: mutable.ArrayBuffer[Try[Frame] => Unit] = null + // Waiters are read after the unlock: once the fetch leaves `pending`, no acquire can add to them. + def store(fetching: Fetching, frame: Frame, now: Long, ttlMillis: Long): Unit = { + val size = frameSize(frame) // walked outside the lock so a large reply can't stall acquire/invalidate lock.lock() try { - val inFlight = pending.remove(key) - waiters = inFlight.map(_.waiters).orNull - val dirty = inFlight.exists(_.dirty) - // Reuse the Key objects created by the matching acquire. Create them here only when no matching fetch is recorded. An entry larger than - // the cache limit cannot be stored, so return its reply without caching it. - if (!dirty && size <= maxBytes) - insert(key, new Entry(frame, size, now + ttlMillis, inFlight.map(_.keys).getOrElse(trackedKeys.map(new Key(_))))) + pending.remove(fetching.key) + // An entry larger than the cache limit cannot be stored, so return its reply without caching it. + if (!fetching.dirty && size <= maxBytes) insert(fetching.key, new Entry(frame, size, now + ttlMillis, fetching.keys)) } finally lock.unlock() - if (waiters != null) waiters.foreach(_.apply(Success(frame))) + fetching.waiters.foreach(_.apply(Success(frame))) } - def fail(commandBytes: Bytes, error: Throwable): Unit = { - val key = new Key(commandBytes) - var waiters: mutable.ArrayBuffer[Try[Frame] => Unit] = null + def fail(fetching: Fetching, error: Throwable): Unit = { lock.lock() - try waiters = pending.remove(key).map(_.waiters).orNull + try pending.remove(fetching.key) finally lock.unlock() - if (waiters != null) waiters.foreach(_.apply(Failure(error))) + fetching.waiters.foreach(_.apply(Failure(error))) } def invalidate(redisKey: Bytes): Unit = { @@ -91,39 +84,23 @@ final private[client] class ClientCache(maxBytes: Long) { if (entry != null) removeEntry(ck, entry) } } - pending.valuesIterator.foreach(inFlight => if (inFlight.keys.contains(tracked)) inFlight.dirty = true) + pending.valuesIterator.foreach(fetching => if (fetching.keys.contains(tracked)) fetching.dirty = true) } finally lock.unlock() } def flush(): Unit = { lock.lock() try { - clearEntries() - epoch = epoch.next + entries.clear() + reverse.clear() + bytesUsed = 0L + pending.valuesIterator.foreach(_.dirty = true) + epoch += 1 } finally lock.unlock() } - def flushForReroute(): Unit = { - lock.lock() - try { - clearEntries() - val retired = epoch.next - // publish the watermark before the epoch, which readers check first. This prevents a reader from pairing the new epoch with the old watermark. - rerouteWatermark = retired - epoch = retired - } finally lock.unlock() - } - - private def clearEntries(): Unit = { - entries.clear() - reverse.clear() - bytesUsed = 0L - pending.valuesIterator.foreach(_.dirty = true) - } - - def isCurrent(stamped: CacheEpoch): Boolean = epoch == stamped - - def rerouteRetired(stamped: CacheEpoch): Boolean = rerouteWatermark.isAfter(stamped) + // false once a flush has retired the hit; the caller looks the command up again + def isCurrent(hit: Hit): Boolean = hit.epoch == epoch private def insert(key: Key, entry: Entry): Unit = { val previous = entries.put(key, entry) @@ -170,26 +147,18 @@ private[client] object ClientCache { } } - opaque type CacheEpoch = Long - object CacheEpoch { - val initial: CacheEpoch = 0L - } - extension (e: CacheEpoch) { - def next: CacheEpoch = e + 1L - def isAfter(other: CacheEpoch): Boolean = e > other - } - enum Acquire { - case Hit(frame: Frame, epoch: CacheEpoch) - case Fetch + case Hit(frame: Frame, epoch: Long) + case Fetch(ticket: Fetching) case Wait } final private class Entry(val frame: Frame, val sizeBytes: Long, val expiresAt: Long, val keys: Vector[Key]) - final private class InFlight(val keys: Vector[Key]) { - val waiters = mutable.ArrayBuffer.empty[Try[Frame] => Unit] - var dirty: Boolean = false + // the in-flight server read that the first missing caller owns; guarded by the cache lock until it leaves `pending` + final class Fetching private[ClientCache] (private[ClientCache] val key: Key, private[ClientCache] val keys: Vector[Key]) { + private[ClientCache] val waiters = mutable.ArrayBuffer.empty[Try[Frame] => Unit] + private[ClientCache] var dirty: Boolean = false } // approximate retained size: payload bytes plus a flat per-node overhead, enough to bound memory without walking object headers exactly diff --git a/sage-client/shared/src/main/scala/sage/client/internal/ClusterLive.scala b/sage-client/shared/src/main/scala/sage/client/internal/ClusterLive.scala index afca2e42..a54d39fc 100644 --- a/sage-client/shared/src/main/scala/sage/client/internal/ClusterLive.scala +++ b/sage-client/shared/src/main/scala/sage/client/internal/ClusterLive.scala @@ -2,27 +2,27 @@ package sage.client.internal import java.util.Locale import java.util.concurrent.CountDownLatch -import java.util.concurrent.atomic.{AtomicBoolean, AtomicReference} +import java.util.concurrent.atomic.AtomicReference import java.util.concurrent.locks.ReentrantLock import scala.collection.mutable -import scala.concurrent.duration.* import scala.util.{Failure, Success, Try} import scala.util.control.NonFatal +import RoutedClient.DispatchMode +import RoutedClient.DispatchMode.* import kyo.compat.* -import sage.{Bytes, CommandSpan, Message, PatternMessage, SageEvent, SageException} +import sage.{Bytes, CommandSpan, SageEvent, SageException} import sage.SageException.{ConnectionLost, CrossSlot, DecodeError, InvalidArgument, NotConnected, ServerError, TimedOut, UnsupportedServer} -import sage.client.{BackoffConfig, ClusterConfig, DedicatedPoolConfig, ReadFrom, SageConfig, WatchdogConfig} -import sage.cluster.{ClusterTopology, Node, NodeGroup, Redirect, RedirectKind, Rejected, Route, Shard, Slot, SplitPlan} -import sage.codec.{KeyCodec, ValueCodec} -import sage.commands.{BroadcastReduce, Cluster, Command, Connection, Pipeline, Reply} +import sage.client.{ClusterConfig, ReadFrom, SageConfig} +import sage.cluster.{ClusterTopology, Node, NodeGroup, Redirect, RedirectKind, Route, Shard, Slot, SlotRange, SplitPlan} +import sage.codec.KeyCodec +import sage.commands.{Cluster, Command, Pipeline, Reply} import sage.protocol.Frame -import sage.ratelimit.Decision /** - * Implements the cluster client with one [[NodeClient]] per master and a [[ClusterTopology]] that can be refreshed. The topology identifies + * Implements the cluster client with one [[MultiplexedConnection]] per master and a [[ClusterTopology]] that can be refreshed. The topology identifies * the node for each command, and this class handles connections, redirects, and failover. Configuration selects this implementation without * changing the `Client` type. * @@ -40,257 +40,96 @@ import sage.ratelimit.Decision final private[client] class ClusterLive( nodeFactory: Node => MultiplexedConnection.TransportFactory, scheduler: Scheduler, - bootstrap: Vector[Command[?]], - reconnect: BackoffConfig, - watchdog: WatchdogConfig, - connectTimeout: FiniteDuration, - closeTimeout: FiniteDuration, - dedicatedPool: DedicatedPoolConfig, + config: SageConfig, cluster: ClusterConfig, - pubsubBufferSize: Int, seeds: Vector[Node], - readFrom: ReadFrom = ReadFrom.Master, - events: Events = Events.disabled, - cachingEnabled: Boolean = false, - cacheMaxBytes: Long = 0L -) extends Client[CIO, String] { + events: Events = Events.disabled +) extends RoutedClient( + nodeFactory, + scheduler, + // replica connections remain separate from the master registry used for command routing and redirects + MultiplexedConnection.NodeRole.ClusterReplica, + config, + cluster.minRefreshInterval, + cluster.topologyRefreshInterval, + events + ) { private val topologyRef = new AtomicReference[ClusterTopology](ClusterTopology.from(Vector.empty)) - // Send READONLY during setup to let replicas serve reads for their master's slots. Replica connections remain separate from the master - // registry used for command routing and redirects. - private val replicaPool = new NodePool( - nodeFactory, - scheduler, - bootstrap :+ Connection.readonly, - reconnect, - watchdog, - connectTimeout, - closeTimeout, - dedicatedPool, - events = events - ) private val keylessCursor = new java.util.concurrent.atomic.AtomicInteger() private val subscriptions = new ClusterSubscriptions( nodeFactory, - bootstrap, scheduler, - reconnect, - watchdog, - connectTimeout.toMillis, - pubsubBufferSize, + config, () => topologyRef.get(), - () => refresh(force = true), - () => pickNode(topologyRef.get()) - ) - - // Store master connections separately from replica connections. Master failures affect redirects and topology refresh, and their setup - // omits READONLY. - private val masterPool = new NodePool( - nodeFactory, - scheduler, - if (cachingEnabled) bootstrap :+ Connection.clientTrackingOnOptin else bootstrap, - reconnect, - watchdog, - connectTimeout, - closeTimeout, - dedicatedPool, - cacheMaxBytes = if (cachingEnabled) cacheMaxBytes else 0L, - events = events, - dedicatedBootstrap = Some(bootstrap) + () => refreshThrottle(force = true), + () => pickNode(topologyRef.get()), + events ) - private val reads = new ReadRouting(masterPool, replicaPool, scheduler, readFrom, () => triggerRefresh()) - // set once by close; routing refuses afterwards, so close is terminal like the standalone client's - @volatile private var closed = false - - private val refreshThrottle = new RefreshThrottle(scheduler, cluster.minRefreshInterval.toMillis) // if every seed fails, report the final connection or handshake error to the caller. - private[client] def bootstrapTopology(): Unit = { - var lastError: Throwable = NotConnected() - val candidates = seeds.iterator - while (candidates.hasNext) { - val node = candidates.next() - try - querySlotsVia(node) match { - case Right(shards) => - adopt(node, shards) - startRefreshPoll() - return - case Left(error) => lastError = error - } - catch { case NonFatal(error) => lastError = error } - } - closeAll() - throw lastError - } - - def run[A](command: Command[A]): CIO[A] = { - def body(lease: DedicatedPool.Lease): CIO[A] = - CIO.async[A] { complete => - val tracked = Events.trackCommand(events, command, complete) - Client.completing(tracked)(dispatch(command, cluster.maxRedirects, tracked, context = DispatchContext.withLease(lease))) - } - Client.withLeaseIfBlocking(command)(body) - } - - override private[sage] def lockWrite( - command: Command[Boolean], - timeout: FiniteDuration, - replicaAcknowledgement: Boolean - ): CIO[Boolean] = - Client.withLockLease(timeout, scheduler) { (lease, deadlineMillis) => - CIO.async { complete => - val tracked = Events.trackCommand(events, command, complete) - Client.completing(tracked) { - dispatch( - command, - cluster.maxRedirects, - tracked, - allowReplica = false, - context = DispatchContext(lease, DispatchMode.Confirmed(deadlineMillis, replicaAcknowledgement)) - ) - } - } - } - - def cached[A](command: Command[A], ttl: FiniteDuration): CIO[A] = - if (!Client.cacheable(command)) CIO.fail(Client.notCacheable(command)) - else if (!cachingEnabled) - CIO.async[A] { complete => - val tracked = Events.trackCommand(events, command, complete) - Client.completing(tracked)(dispatch(command, cluster.maxRedirects, tracked, allowReplica = false)) - } - else - CIO.async[A] { complete => - val deferred = Events.deferSpan(events, command) - Client.completing(complete)( - dispatch( - command, - cluster.maxRedirects, - complete, - allowReplica = false, - context = DispatchContext(null, DispatchMode.Cached(ttl.toMillis, deferred)) - ) - ) - } - - private[sage] def rateLimitAcquire[RK](executor: RateLimitExecutor[RK], subject: RK, cost: Long, peek: Boolean): CIO[Decision] = - executor.evalSha(this, subject, cost, peek) + protected def discover(): Either[Throwable, Unit] = + seeds.foldLeft[Either[Throwable, Unit]](Left(NotConnected()))((found, node) => + if (found.isRight) found else Try(querySlotsVia(node).map(adopt)).toEither.flatten + ) - private[sage] def lockTryWith[LK, A](executor: LockExecutor[LK], key: LK)(body: => CIO[A]): CIO[Option[A]] = - executor.tryWithLock(this, key)(body) + protected def route[A](command: Command[A], complete: Try[A] => Unit, lease: DedicatedPool.Lease, mode: DispatchMode): Unit = + dispatch(Request(command, complete, lease, mode), cluster.maxRedirects) - private[sage] def lockWith[LK, A](executor: LockExecutor[LK], key: LK, waitTimeout: FiniteDuration)(body: => CIO[A]): CIO[A] = - executor.withLock(this, key, waitTimeout)(body) + protected def replicaCount(master: Node): Int = topologyRef.get().replicasForMaster(master).size // SCAN cursors are node-local. A full scan visits every master that owns slots. Resharding during the scan can still miss or duplicate keys. - def scanTargets: CIO[Vector[ScanTarget]] = + override def scanTargets: CIO[Vector[ScanTarget]] = CIO.blocking { - val masters = slotOwningMasters(topologyRef.get()) - if (masters.isEmpty) Vector(ScanTarget.any) else masters.map(node => ScanTarget(Some(node))) + val masters = topologyRef.get().masters + if (masters.isEmpty) Vector(this) else masters.map(pinnedTo) } - private def slotOwningMasters(topology: ClusterTopology): Vector[Node] = - topology.shards.collect { case shard if shard.slots.nonEmpty => shard.master }.distinct - // Resume a SCAN page on the node that issued its cursor. If that node is unavailable, fail the scan because another master would interpret // the node-local cursor against a different keyspace. redirectsLeft = 0 disables rerouting. - def runOn[A](target: ScanTarget, command: Command[A]): CIO[A] = - target.node match { - case Some(node) => - def body(lease: DedicatedPool.Lease): CIO[A] = - CIO.async[A] { complete => - val tracked = Events.trackCommand(events, command, complete) - Client.completing(tracked)(sendTo(node, command, asking = false, redirectsLeft = 0, tracked, DispatchContext.withLease(lease))) - } - Client.withLeaseIfBlocking(command)(body) - case None => run(command) + private def pinnedTo(node: Node): ScanTarget = + new ScanTarget { + def run[A](command: Command[A]): CIO[A] = + Client.withLeaseIfBlocking(command) { lease => + tracked(command)(t => sendTo(node, Request(command, t, lease, readMode(command)), asking = false, redirectsLeft = 0)) + } } - private[sage] def pipeline[Out, R](p: Pipeline[Out, R]): CIO[Out] = submitPipeline(p).flatMap(TxSupport.collapseStrict(_, p.toOut)) - private[sage] def pipelineAttempt[Out, R](p: Pipeline[Out, R]): CIO[R] = submitPipeline(p).map(p.toResults) - - def transaction[A](body: TransactionScope[CIO, String] => CIO[A]): CIO[A] = - CIO.acquireReleaseWith(acquireScope)(releaseScope)(scope => CIO.unit.flatMap(_ => body(scope))) - - // classic subscriptions share a connection to an arbitrary master because PUBLISH broadcasts across the cluster. - def subscribeChannels[V: ValueCodec](channel: String, rest: String*): CIO[Subscription[CIO, Message[V]]] = - CIO.blocking(Client.channelMessages(subscriptions.subscribeChannels(channel +: rest.toVector))) - - def subscribePatterns[V: ValueCodec](pattern: String, rest: String*): CIO[Subscription[CIO, PatternMessage[V]]] = - CIO.blocking(Client.patternMessages(subscriptions.subscribePatterns(pattern +: rest.toVector))) - - // route each shard channel to a sharded subscription connection for its slot's node. Update the subscription when ownership changes. - def subscribeShardChannels[V: ValueCodec](channel: String, rest: String*): CIO[Subscription[CIO, Message[V]]] = - CIO.blocking(Client.channelMessages(subscriptions.subscribeShard(channel +: rest.toVector))) - - def close: CIO[Unit] = CIO.blocking(closeAll()) + // Classic subscriptions share a connection to an arbitrary master because PUBLISH broadcasts across the cluster. Each shard channel uses a + // sharded subscription connection for its slot's node, and the subscription follows ownership changes. + protected def pubsub: SubscriptionConnection.PubSub = subscriptions // --- routing ------------------------------------------------------------------------------------------------------------------------- - private enum DispatchMode { - case Ordinary - case Cached(ttlMillis: Long, deferred: () => CommandSpan) - case Confirmed(deadlineMillis: Long, replicaAcknowledgement: Boolean) - } - import DispatchMode.* - - final private case class DispatchContext(lease: DedicatedPool.Lease, mode: DispatchMode) - private object DispatchContext { - val Default: DispatchContext = DispatchContext(null, Ordinary) - def withLease(lease: DedicatedPool.Lease): DispatchContext = DispatchContext(lease, Ordinary) - } + final private case class Request[A](command: Command[A], complete: Try[A] => Unit, lease: DedicatedPool.Lease, mode: DispatchMode) - // Both cached reads and confirmed writes require the master. - private def replicaAllowed(mode: DispatchMode): Boolean = mode == Ordinary - - private def dispatch[A]( - command: Command[A], - redirectsLeft: Int, - complete: Try[A] => Unit, - allowReplica: Boolean = true, - context: DispatchContext = DispatchContext.Default - ): Unit = - if (closed) complete(Failure(NotConnected())) + private def dispatch[A](req: Request[A], redirectsLeft: Int): Unit = + if (closed) req.complete(Failure(NotConnected())) else { val topology = topologyRef.get() - if (command.allMasters) - if (slotOwningMasters(topology).forall(node => masterPool.existing(node) != null)) - broadcast(topology, command, redirectsLeft, complete, masterPool.existing) - else scheduler.offload(broadcast(topology, command, redirectsLeft, complete, masterPool.getOrEstablishOrNull)) + if (req.command.allMasters) broadcast(topology, req, redirectsLeft) else - topology.route(command) match { - case Route.ToNode(node, slot) => sendOwned(command, node, slot, redirectsLeft, complete, allowReplica, context) - case Route.Keyless => - if (servesFromReplica(command, allowReplica)) sendKeylessRead(topology, command, redirectsLeft, complete) - else sendToAny(topology, command, redirectsLeft, complete, context) - case Route.Unowned(_) => scheduler.offload(onUnowned(command, redirectsLeft, complete, context)) - case Route.CrossSlot(slots) => - multiSlotPolicy(command) match { - case Some(policy) => scatterMultiSlot(command, policy, redirectsLeft, complete, allowReplica, context) - case None => complete(Failure(crossSlot(command.name, slots))) + topology.route(req.command) match { + case Route.ToNode(shard, _) => sendOwned(req, shard, redirectsLeft) + case Route.Keyless => + if (req.mode == ReplicaRead) sendKeylessRead(topology, req, redirectsLeft) + else sendToAny(topology, req, redirectsLeft) + case Route.Unowned(_) => scheduler.offload(onUnowned(req, redirectsLeft)) + case Route.CrossSlot => + multiSlotPolicy(req.command) match { + case Some(policy) => scatterMultiSlot(req, policy, redirectsLeft) + case None => req.complete(Failure(crossSlot(req.command))) } - case Route.Malformed => - complete(Failure(malformedKeys(command.name))) + case Route.Malformed => + req.complete(Failure(malformedKeys(req.command.name))) } } - private def sendOwned[A]( - command: Command[A], - node: Node, - slot: Slot, - redirectsLeft: Int, - complete: Try[A] => Unit, - allowReplica: Boolean, - context: DispatchContext - ): Unit = - if (servesFromReplica(command, allowReplica)) sendRead(command, node, slot, redirectsLeft, complete) - else sendTo(node, command, asking = false, redirectsLeft, complete, context) - - private def servesFromReplica(command: Command[?], allowReplica: Boolean): Boolean = - allowReplica && readFrom != ReadFrom.Master && ReadRouting.replicaEligible(command) + private def sendOwned[A](req: Request[A], shard: Shard, redirectsLeft: Int): Unit = + if (req.mode == ReplicaRead) walkRead(req, reads.candidatesFor(shard.master, shard.replicas), shard.master, redirectsLeft) + else sendTo(shard.master, req, asking = false, redirectsLeft) private enum MultiSlotMerge { case Positional, Sum, AllSucceeded @@ -304,64 +143,47 @@ final private[client] class ClusterLive( // Split a supported cross-slot command into one command for each slot. Send every group through normal dispatch to apply replica policy, // topology refresh, and MOVED or ASK handling. After every group completes, decode the combined result with the original command. MGET // restores values to their original positions, integer commands add the per-slot counts, and MSET requires an OK reply from every slot. - private def scatterMultiSlot[A]( - command: Command[A], - policy: MultiSlotPolicy, - redirectsLeft: Int, - complete: Try[A] => Unit, - allowReplica: Boolean, - context: DispatchContext - ): Unit = { - val bySlot = mutable.LinkedHashMap.empty[Slot, mutable.ArrayBuffer[MultiSlotEntry]] + private def scatterMultiSlot[A](req: Request[A], policy: MultiSlotPolicy, redirectsLeft: Int): Unit = { + val command = req.command + val bySlot = mutable.LinkedHashMap.empty[Slot, mutable.ArrayBuffer[MultiSlotEntry]] command.keyIndices.iterator.zipWithIndex.foreach { case (argIndex, resultIndex) => val key = command.args(argIndex) val entry = MultiSlotEntry(resultIndex, argIndex) val group = bySlot.getOrElseUpdate(Slot.of(key), mutable.ArrayBuffer.empty) group += entry } - val groups = bySlot.valuesIterator.map(_.toVector).toVector - - lazy val values = new java.util.concurrent.atomic.AtomicReferenceArray[Frame](command.keyIndices.size) - lazy val total = new java.util.concurrent.atomic.AtomicLong(0L) - val remaining = new java.util.concurrent.atomic.AtomicInteger(groups.size) - val firstError = new java.util.concurrent.atomic.AtomicReference[Throwable](null) - - def settle(group: Vector[MultiSlotEntry], result: Try[Frame]): Unit = { - (policy.merge, result) match { - case (MultiSlotMerge.Positional, Success(Frame.Array(elements))) if elements.size == group.size => - group.iterator.zip(elements).foreach { case (entry, frame) => values.set(entry.resultIndex, frame) } - case (MultiSlotMerge.Positional, Success(Frame.Array(elements))) => - firstError.compareAndSet( - null, - DecodeError(s"an array of ${group.size} MGET values", s"an array of ${elements.size} values") - ) - case (MultiSlotMerge.Positional, Success(other)) => - firstError.compareAndSet(null, DecodeError(s"an array of ${group.size} MGET values", Frame.describe(other))) - case (MultiSlotMerge.Sum, Success(Frame.Integer(value))) => total.addAndGet(value) - case (MultiSlotMerge.Sum, Success(other)) => - firstError.compareAndSet(null, DecodeError("an integer count", Frame.describe(other))) - case (MultiSlotMerge.AllSucceeded, Success(Frame.SimpleString("OK"))) => () - case (MultiSlotMerge.AllSucceeded, Success(other)) => - firstError.compareAndSet(null, DecodeError("simple string 'OK'", Frame.describe(other))) - case (_, Failure(error)) => - firstError.compareAndSet(null, error) - } - - if (remaining.decrementAndGet() == 0) - Option(firstError.get()) match { - case Some(error) => complete(Failure(error)) - case None => - val merged = policy.merge match { - case MultiSlotMerge.Positional => Frame.Array(Vector.tabulate(command.keyIndices.size)(values.get)) - case MultiSlotMerge.Sum => Frame.Integer(total.get()) - case MultiSlotMerge.AllSucceeded => Frame.SimpleString("OK") + val groups = bySlot.valuesIterator.map(_.toVector).toVector + + val collector = gather(command, groups.size, req.complete) { frames => + policy.merge match { + case MultiSlotMerge.Positional => + val values = new Array[Frame](command.keyIndices.size) + groups.lazyZip(frames).foreach { (group, frame) => + frame match { + case Frame.Array(elements) if elements.size == group.size => + group.lazyZip(elements).foreach((entry, value) => values(entry.resultIndex) = value) + case Frame.Array(elements) => + throw DecodeError(s"an array of ${group.size} MGET values", s"an array of ${elements.size} values") + case other => throw DecodeError(s"an array of ${group.size} MGET values", Frame.describe(other)) } - complete(Reply.decode(command, merged)) - } + } + Frame.Array(values.toVector) + case MultiSlotMerge.Sum => + Frame.Integer(frames.foldLeft(0L) { + case (total, Frame.Integer(value)) => total + value + case (_, other) => throw DecodeError("an integer count", Frame.describe(other)) + }) + case MultiSlotMerge.AllSucceeded => + frames.foreach { + case Frame.SimpleString("OK") => () + case other => throw DecodeError("simple string 'OK'", Frame.describe(other)) + } + Frame.SimpleString("OK") + } } val raw = command.rawFrame - groups.foreach { group => + groups.iterator.zipWithIndex.foreach { case (group, index) => val args = Vector.newBuilder[Bytes] args.sizeHint(group.size * policy.argsPerKey) group.foreach { entry => @@ -376,7 +198,7 @@ final private[client] class ClusterLive( keyIndices = Vector.tabulate(group.size)(_ * policy.argsPerKey), args = args.result() ) - dispatch(sub, redirectsLeft, result => settle(group, result), allowReplica = allowReplica, context = context) + dispatch(Request(sub, collector.set(index, _), req.lease, req.mode), redirectsLeft) } } @@ -401,401 +223,210 @@ final private[client] class ClusterLive( command.args.size > suffixArgs && command.keyIndices == Vector.tabulate(command.args.size - suffixArgs)(identity) - private def sendRead[A](command: Command[A], master: Node, slot: Slot, redirectsLeft: Int, complete: Try[A] => Unit): Unit = { - val replicas = topologyRef.get().shardForSlot(slot).map(_.replicas).getOrElse(Vector.empty) - walkRead(command, reads.candidatesFor(master, replicas), master, redirectsLeft, complete) - } - // RANDOMKEY has no slot. Try replicas across the cluster in round-robin order, then apply the configured fallback policy. // ReadFrom.Replica uses only replica candidates. - private def sendKeylessRead[A](topology: ClusterTopology, command: Command[A], redirectsLeft: Int, complete: Try[A] => Unit): Unit = + private def sendKeylessRead[A](topology: ClusterTopology, req: Request[A], redirectsLeft: Int): Unit = pickNode(topology) match { case Some(master) => val replicas = topology.shards.iterator.flatMap(_.replicas).toVector.distinct - walkRead(command, reads.candidatesFor(master, replicas, keylessCursor.getAndIncrement()), master, redirectsLeft, complete) - case None => complete(Failure(NotConnected())) + walkRead(req, reads.candidatesFor(master, replicas, keylessCursor.getAndIncrement()), master, redirectsLeft) + case None => req.complete(Failure(NotConnected())) } - private def walkRead[A](command: Command[A], candidates: Vector[Node], master: Node, redirectsLeft: Int, complete: Try[A] => Unit): Unit = - reads.walk(command, candidates, master, complete)((node, error, rest) => - onReadFailure(node, command, error, rest, master, redirectsLeft, complete) - ) + private def walkRead[A](req: Request[A], candidates: Vector[Node], master: Node, redirectsLeft: Int): Unit = + reads.walk(req.command, candidates, master, req.complete)((node, error, rest) => onReadFailure(node, req, error, rest, master, redirectsLeft)) - private def onReadFailure[A]( - node: Node, - command: Command[A], - error: Throwable, - rest: Vector[Node], - master: Node, - redirectsLeft: Int, - complete: Try[A] => Unit - ): Unit = + // Lost or Unavailable walks the remaining candidates. MOVED refreshes the topology and re-dispatches so the read keeps its replica policy. + private def onReadFailure[A](node: Node, req: Request[A], error: Throwable, rest: Vector[Node], master: Node, redirectsLeft: Int): Unit = Fault.categorize(error) match { - case Fault.Redirected(redirect) => - redirect.kind match { - // ReadFrom.Replica cannot follow ASK because the importing master holds the key during migration. MOVED refreshes the topology and - // re-dispatches the command. - case RedirectKind.Ask => - if (readFrom == ReadFrom.Replica) complete(Failure(NotConnected())) - else onRedirect(node, redirect, command, redirectsLeft, complete) - case RedirectKind.Moved if redirectsLeft <= 0 => - refreshBeforeFailing() - complete(Failure(ServerError("ERR", s"exceeded ${cluster.maxRedirects} cluster redirects for ${command.name}"))) - case RedirectKind.Moved => - refresh(force = true) - scheduler.offload(dispatch(command, redirectsLeft - 1, complete)) - } - // re-dispatch only when the failure reports that the command was not executed, as required by onUnreachable - case Fault.Lost(executed) => - if (rest.nonEmpty) walkRead(command, rest, master, redirectsLeft, complete) - else if (executed) { - triggerRefresh() - Events.attributeNode(complete, node) - complete(Failure(error)) - } else onUnreachable(command, redirectsLeft, complete) - case Fault.Unavailable(clusterWide) => - if (rest.nonEmpty) walkRead(command, rest, master, redirectsLeft, complete) - else onRetryable(command, error, clusterWide, redirectsLeft, complete) - case Fault.TryAgain => onRetryable(command, error, refreshFirst = false, redirectsLeft, complete) - case Fault.Demoted => - triggerRefresh() - Events.attributeNode(complete, node) - complete(Failure(error)) - case Fault.Fatal => - Events.attributeNode(complete, node) - complete(Failure(error)) - } - - private def broadcast[A]( - topology: ClusterTopology, - command: Command[A], - redirectsLeft: Int, - complete: Try[A] => Unit, - resolve: Node => NodeClient - ): Unit = - command.broadcast match { - case BroadcastReduce.First => sendToAllMasters(topology, command, redirectsLeft, complete, resolve) - case BroadcastReduce.Concat => broadcastCombine(topology, command, concatFrames, redirectsLeft, complete, resolve) - case BroadcastReduce.Fold(fold) => broadcastCombine(topology, command, _.reduce(fold), redirectsLeft, complete, resolve) + case Fault.Lost(_) | Fault.Unavailable(_) if rest.nonEmpty => walkRead(req, rest, master, redirectsLeft) + case Fault.Redirected(redirect) if redirect.kind == RedirectKind.Moved && redirectsLeft > 0 => + refreshThrottle(force = true) + scheduler.offload(dispatch(req, redirectsLeft - 1)) + case _ => onFailure(node, req, error, redirectsLeft) } // A broadcast command (SCRIPT LOAD, FUNCTION LOAD, …) runs on every slot-owning master, since a cluster replicates no script/function - // cache; any node failing terminally fails the command - private def sendToAllMasters[A]( - topology: ClusterTopology, - command: Command[A], - redirectsLeft: Int, - complete: Try[A] => Unit, - resolve: Node => NodeClient - ): Unit = { - val masters = slotOwningMasters(topology) - if (masters.isEmpty) sendToAny(topology, command, cluster.maxRedirects, complete) + // cache; any node failing terminally fails the command. Replies are combined before decoding: KEYS concatenates the keys returned by each + // node, and WAIT and WAITAOF use the lowest acknowledgement counts returned by any shard. + private def broadcast[A](topology: ClusterTopology, req: Request[A], redirectsLeft: Int): Unit = { + val masters = topology.masters + val command = req.command + if (masters.isEmpty) sendToAny(topology, req, cluster.maxRedirects) else { - val remaining = new java.util.concurrent.atomic.AtomicInteger(masters.size) - val firstError = new java.util.concurrent.atomic.AtomicReference[Throwable](null) - val firstValue = new java.util.concurrent.atomic.AtomicReference[Try[A]](null) - def settle(result: Try[A]): Unit = { - result match { - case Success(_) => firstValue.compareAndSet(null, result) - case Failure(e) => firstError.compareAndSet(null, e) - } - if (remaining.decrementAndGet() == 0) - complete(Option(firstError.get()).map(Failure(_)).getOrElse(firstValue.get())) - } - masters.foreach(node => submitBroadcast(node, command, resolve, redirectsLeft, settle)) - } - } - - // Combine replies from an all-masters command before decoding them. KEYS concatenates the keys returned by each node. WAIT and WAITAOF use - // the lowest acknowledgement counts returned by any shard. If a node cannot complete the command, fail the combined result. - private def broadcastCombine[A]( - topology: ClusterTopology, - command: Command[A], - combine: Vector[Frame] => Frame, - redirectsLeft: Int, - complete: Try[A] => Unit, - resolve: Node => NodeClient - ): Unit = { - val masters = slotOwningMasters(topology) - if (masters.isEmpty) sendToAny(topology, command, cluster.maxRedirects, complete) - else { - val raw = command.rawFrame - val frames = new java.util.concurrent.atomic.AtomicReferenceArray[Frame](masters.size) - val remaining = new java.util.concurrent.atomic.AtomicInteger(masters.size) - val firstError = new java.util.concurrent.atomic.AtomicReference[Throwable](null) - def settle(index: Int, result: Try[Frame]): Unit = { - result match { - case Success(frame) => frames.set(index, frame) - case Failure(e) => firstError.compareAndSet(null, e) - } - if (remaining.decrementAndGet() == 0) - Option(firstError.get()) match { - case Some(e) => complete(Failure(e)) - case None => complete(Try(combine(Vector.tabulate(masters.size)(frames.get))).flatMap(Reply.decode(command, _))) - } - } + val raw = command.rawFrame + val collector = gather(command, masters.size, req.complete)(frames => command.reduceReplies(frames(0), frames.drop(1))) masters.iterator.zipWithIndex.foreach { case (node, index) => - submitBroadcast(node, raw, resolve, redirectsLeft, result => settle(index, result)) + submitBroadcast(node, raw, redirectsLeft, collector.set(index, _)) } } } - private def submitBroadcast[B](node: Node, command: Command[B], resolve: Node => NodeClient, attemptsLeft: Int, settle: Try[B] => Unit): Unit = { - val nc = resolve(node) - // resolve on the caller's thread; when no connection is available, offload the refresh and retry because refresh may block - if (nc == null) scheduler.offload(retryBroadcast(node, command, NotConnected(), refreshFirst = true, attemptsLeft, settle)) - else - nc.submit[B]( + // after every part replies, completes with the first failure by part index or with the decoded combination of all frames + private def gather[A](command: Command[A], parts: Int, complete: Try[A] => Unit)(combine: Vector[Frame] => Frame) = + new TxSupport.IndexedCollector[Try[Frame]]( + parts, + results => + complete( + results + .collectFirst { case Failure(e) => Failure(e) } + .getOrElse(Try(combine(results.collect { case Success(f) => f }))) + .flatMap(Reply.decode(command, _)) + ) + ) + + // withClient runs the unreachable branch only after offloading the connect, so the retry's blocking refresh never runs on the caller's thread + private def submitBroadcast[B](node: Node, command: Command[B], attemptsLeft: Int, settle: Try[B] => Unit): Unit = + masterPool.withClient(node)(retryBroadcast(node, command, NotConnected(), attemptsLeft, settle)) { + _.submit[B]( command, - asking = false, { case Success(value) => settle(Success(value)) case Failure(error) => scheduler.offload(onBroadcastFailure(node, command, error, attemptsLeft, settle)) } ) - } + } // retry only the node whose connection was lost or whose request was temporarily refused; retrying WAIT starts its timeout again there private def onBroadcastFailure[B](node: Node, command: Command[B], error: Throwable, attemptsLeft: Int, settle: Try[B] => Unit): Unit = Fault.categorize(error) match { - case Fault.Lost(false) => retryBroadcast(node, command, error, refreshFirst = true, attemptsLeft, settle) - case Fault.TryAgain | Fault.Unavailable(false) => retryBroadcast(node, command, error, refreshFirst = false, attemptsLeft, settle) + case Fault.Lost(false) | Fault.TryAgain | Fault.Unavailable(false) => retryBroadcast(node, command, error, attemptsLeft, settle) // a cluster-wide refusal may mean the selected masters are stale; refresh the topology, then return the error - case Fault.Unavailable(true) => + case Fault.Unavailable(true) => refreshBeforeFailing() settle(Failure(error)) - case fault => + case fault => if (fault.refreshPolicy != RefreshPolicy.Skip) triggerRefresh() settle(Failure(error)) } // retry this node only while it remains a slot-owning master - private def retryBroadcast[B]( - node: Node, - command: Command[B], - error: Throwable, - refreshFirst: Boolean, - attemptsLeft: Int, - settle: Try[B] => Unit - ): Unit = + private def retryBroadcast[B](node: Node, command: Command[B], error: Throwable, attemptsLeft: Int, settle: Try[B] => Unit): Unit = { + val refreshFirst = refreshesFirst(error) if (attemptsLeft <= 0) { if (refreshFirst) refreshBeforeFailing() settle(Failure(error)) } else afterBackoff(attemptsLeft) { - if (refreshFirst) refresh(force = true) - if (closed || !slotOwningMasters(topologyRef.get()).contains(node)) settle(Failure(error)) - else submitBroadcast(node, command, masterPool.getOrEstablishOrNull, attemptsLeft - 1, settle) + if (refreshFirst) refreshThrottle(force = true) + if (closed || !topologyRef.get().masters.contains(node)) settle(Failure(error)) + else submitBroadcast(node, command, attemptsLeft - 1, settle) } + } - private def concatFrames(frames: Vector[Frame]): Frame = - Frame.Array(frames.flatMap { - case Frame.Array(elements) => elements - case Frame.Set(elements) => elements - case other => Vector(other) - }) - - private def sendToAny[A]( - topology: ClusterTopology, - command: Command[A], - redirectsLeft: Int, - complete: Try[A] => Unit, - context: DispatchContext = DispatchContext.Default - ): Unit = + private def sendToAny[A](topology: ClusterTopology, req: Request[A], redirectsLeft: Int): Unit = pickNode(topology) match { - case Some(node) => sendTo(node, command, asking = false, redirectsLeft, complete, context) - case None => complete(Failure(NotConnected())) + case Some(node) => sendTo(node, req, asking = false, redirectsLeft) + case None => req.complete(Failure(NotConnected())) } - private def sendTo[A]( - node: Node, - command: Command[A], - asking: Boolean, - redirectsLeft: Int, - complete: Try[A] => Unit, - context: DispatchContext - ): Unit = { - val existing = masterPool.existing(node) - if (existing != null) submitTo(existing, node, command, asking, redirectsLeft, complete, context) - else - scheduler.offload { - val nc = masterPool.getOrEstablishOrNull(node) - if (nc == null) onUnreachable(command, redirectsLeft, complete, context) - else submitTo(nc, node, command, asking, redirectsLeft, complete, context) - } - } + private def sendTo[A](node: Node, req: Request[A], asking: Boolean, redirectsLeft: Int): Unit = + masterPool.withClient(node)(onUnreachable(req, redirectsLeft))(submitTo(_, node, req, asking, redirectsLeft)) - private def submitTo[A]( - nc: NodeClient, - node: Node, - command: Command[A], - asking: Boolean, - redirectsLeft: Int, - complete: Try[A] => Unit, - context: DispatchContext - ): Unit = { + private def submitTo[A](nc: MultiplexedConnection, node: Node, req: Request[A], asking: Boolean, redirectsLeft: Int): Unit = { val onReply: Try[A] => Unit = { - case Success(value) => - Events.attributeNode(complete, node) - complete(Success(value)) - case Failure(error) => scheduler.offload(onFailure(node, command, error, redirectsLeft, complete, context)) - } - context.mode match { - case Cached(ttlMillis, deferred) if !asking => nc.cachedSubmit[A](command, ttlMillis, onReply, deferred) - case Confirmed(deadlineMillis, replicaAcknowledgement) => - val replication = new LockReplication( - scheduler, - topologyRef.get().replicasForMaster(node).size, - deadlineMillis, - () => refreshThrottle.request(refreshWork), - replicaAcknowledgement - ) - nc.submitLockWrite( - command, - asking, - onReply, - context.lease, - replication - ) - case Ordinary | Cached(_, _) => nc.submit[A](command, asking, onReply, context.lease) + case success @ Success(_) => Events.completeAt(req.complete, node)(success) + case Failure(error) => scheduler.offload(onFailure(node, req, error, redirectsLeft)) } + submitOn(nc, node, req.command, req.lease, req.mode, asking, onReply) } - private def onFailure[A]( - node: Node, - command: Command[A], - error: Throwable, - redirectsLeft: Int, - complete: Try[A] => Unit, - context: DispatchContext - ): Unit = + private def onFailure[A](node: Node, req: Request[A], error: Throwable, redirectsLeft: Int): Unit = Fault.categorize(error) match { - case Fault.Redirected(redirect) => onRedirect(node, redirect, command, redirectsLeft, complete, context, Some(error)) - case Fault.Lost(false) => onUnreachable(command, redirectsLeft, complete, context) - case Fault.TryAgain => onRetryable(command, error, refreshFirst = false, redirectsLeft, complete, context) - case Fault.Unavailable(clusterWide) => onRetryable(command, error, clusterWide, redirectsLeft, complete, context) - case Fault.Demoted | Fault.Lost(true) => - triggerRefresh() - Events.attributeNode(complete, node) - complete(Failure(error)) - case Fault.Fatal => - Events.attributeNode(complete, node) - complete(Failure(error)) + case Fault.Lost(false) => onUnreachable(req, redirectsLeft) + case fault => + // the node received the command, so an exhausted retry reports it unless a later attempt reaches another node + Events.attributeNode(req.complete, node) + fault match { + case Fault.Redirected(redirect) => onRedirect(node, redirect, req, error, redirectsLeft) + case Fault.TryAgain | Fault.Unavailable(_) => onRetryable(req, error, redirectsLeft) + case Fault.Demoted | Fault.Lost(true) => + triggerRefresh() + req.complete(Failure(error)) + case _ => req.complete(Failure(error)) + } } - private def onRedirect[A]( - from: Node, - redirect: Redirect, - command: Command[A], - redirectsLeft: Int, - complete: Try[A] => Unit, - context: DispatchContext = DispatchContext.Default, - exhaustedFailure: Option[Throwable] = None - ): Unit = { + private def onRedirect[A](from: Node, redirect: Redirect, req: Request[A], error: Throwable, redirectsLeft: Int): Unit = { // a MOVED proves `from` lost the slot; retire its cache even if the retry budget is now exhausted if (redirect.kind == RedirectKind.Moved) flushNode(from) - if (redirectsLeft <= 0) { + // ReadFrom.Replica cannot follow ASK because the importing master holds the key during migration + if (redirect.kind == RedirectKind.Ask && req.mode == ReplicaRead && config.readFrom == ReadFrom.Replica) { + req.complete(Failure(NotConnected())) + } else if (redirectsLeft <= 0) { if (redirect.kind == RedirectKind.Moved) refreshBeforeFailing() - val limitFailure = ServerError("ERR", s"exceeded ${cluster.maxRedirects} cluster redirects for ${command.name}") // Lock writes retry redirect faults within their lock budget after the topology refresh. Ordinary commands report the redirect limit. - val failure = context.mode match { - case Confirmed(_, _) => exhaustedFailure.getOrElse(limitFailure) - case _ => limitFailure + val failure = req.mode match { + case Confirmed(_, _) => error + case _ => ServerError("ERR", s"exceeded ${cluster.maxRedirects} cluster redirects for ${req.command.name}") } - complete(Failure(failure)) + req.complete(Failure(failure)) } else { - val target = resolve(redirect.target, from) + val target = redirect.target(from) redirect.kind match { case RedirectKind.Moved => triggerRefresh() - sendTo(target, command, asking = false, redirectsLeft - 1, complete, context) - case RedirectKind.Ask => sendTo(target, command, asking = true, redirectsLeft - 1, complete, context) + sendTo(target, req, asking = false, redirectsLeft - 1) + case RedirectKind.Ask => sendTo(target, req, asking = true, redirectsLeft - 1) } } } // The command was not sent; refresh the topology before routing it again, delay retries with jitter while failover completes, and use // redirectsLeft to limit the number of attempts - private def onUnreachable[A]( - command: Command[A], - redirectsLeft: Int, - complete: Try[A] => Unit, - context: DispatchContext = DispatchContext.Default - ): Unit = onRetryable(command, NotConnected(), refreshFirst = true, redirectsLeft, complete, context) + private def onUnreachable[A](req: Request[A], redirectsLeft: Int): Unit = onRetryable(req, NotConnected(), redirectsLeft) // retry temporary refusals such as TRYAGAIN, LOADING, MASTERDOWN, and CLUSTERDOWN with bounded jitter - private def onRetryable[A]( - command: Command[A], - error: Throwable, - refreshFirst: Boolean, - redirectsLeft: Int, - complete: Try[A] => Unit, - context: DispatchContext = DispatchContext.Default - ): Unit = + private def onRetryable[A](req: Request[A], error: Throwable, redirectsLeft: Int): Unit = { + val refreshFirst = refreshesFirst(error) if (redirectsLeft <= 0) { if (refreshFirst) refreshBeforeFailing() - complete(Failure(error)) + req.complete(Failure(error)) } else { - if (refreshFirst) refresh(force = true) - afterBackoff(redirectsLeft)( - dispatch(command, redirectsLeft - 1, complete, allowReplica = replicaAllowed(context.mode), context = context) - ) + if (refreshFirst) refreshThrottle(force = true) + afterBackoff(redirectsLeft)(dispatch(req, redirectsLeft - 1)) } + } + + private def refreshesFirst(error: Throwable): Boolean = Fault.categorize(error).refreshPolicy == RefreshPolicy.Forced // increase the jittered delay with each attempt to reduce request load during failover or migration private def afterBackoff(attemptsLeft: Int)(retry: => Unit): Unit = - scheduler.after(Backoff.jitteredMillis(reconnect, (cluster.maxRedirects - attemptsLeft).max(0), scheduler).millis)(retry) + scheduler.afterBackoff(config.reconnect, (cluster.maxRedirects - attemptsLeft).max(0))(retry) // refresh immediately before returning an error on paths where no later retry can trigger another refresh - private def refreshBeforeFailing(): Unit = refresh(force = true) + private def refreshBeforeFailing(): Unit = refreshThrottle(force = true) - private def onUnowned[A]( - command: Command[A], - redirectsLeft: Int, - complete: Try[A] => Unit, - context: DispatchContext - ): Unit = { - refresh(force = false) - val topology = topologyRef.get() - val allowReplica = replicaAllowed(context.mode) - topology.route(command) match { + private def onUnowned[A](req: Request[A], redirectsLeft: Int): Unit = { + refreshThrottle(force = false) + val topology = topologyRef.get() + topology.route(req.command) match { // apply the read policy after the slot resolves. Eligible reads still use replica routing. - case Route.ToNode(node, slot) => - sendOwned(command, node, slot, redirectsLeft, complete, allowReplica, context) + case Route.ToNode(shard, _) => sendOwned(req, shard, redirectsLeft) // ReadFrom.Replica has no master fallback. Refresh and retry within the configured limit. - case _ if allowReplica && readFrom == ReadFrom.Replica && ReadRouting.replicaEligible(command) => - onUnreachable(command, redirectsLeft, complete, context) + case _ if req.mode == ReplicaRead && config.readFrom == ReadFrom.Replica => onUnreachable(req, redirectsLeft) // if the refreshed topology still has no owner, send to any master and handle its MOVED or CLUSTERDOWN reply - case _ => sendToAny(topology, command, redirectsLeft, complete, context) + case _ => sendToAny(topology, req, redirectsLeft) } } - // an empty redirect host means "the node I just talked to" (e.g. `MOVED 3999 :6381`) - private def resolve(target: Node, from: Node): Node = if (target.host.isEmpty) Node(from.host, target.port) else target - private def pickNode(topology: ClusterTopology): Option[Node] = masterPool.firstLiveNode.orElse(topology.shards.headOption.map(_.master)) - private def crossSlot(name: String, slots: Set[Slot]): CrossSlot = - CrossSlot(s"$name: keys span ${slots.size} slots; a single command must touch exactly one") + private def crossSlot(command: Command[?]): CrossSlot = { + val slots = command.keyIndices.iterator.flatMap(command.args.lift).map(Slot.of).distinct.size + CrossSlot(s"${command.name}: keys span $slots slots; a single command must touch exactly one") + } private def malformedKeys(name: String): InvalidArgument = InvalidArgument(s"$name: declared key positions fall outside its arguments") - // transactions do not follow redirects. Refresh without throttling so the caller's retry uses the latest ownership information. - private def forceRefresh(): Unit = scheduler.offload(refresh(force = true)) - // --- pipelines (split per node, batch each, merge in submission order) ---------------------------------------------------------------- - private def submitPipeline[Out, R](p: Pipeline[Out, R]): CIO[Vector[Either[SageException, Any]]] = - if (p.commands.isEmpty) - CIO.value(Vector.empty) - // reject the whole pipeline before submission when it contains a blocking command - else if (p.commands.exists(_.isBlocking)) - CIO.fail(InvalidArgument("a Pipeline cannot carry blocking commands; run them individually on the client")) + protected def submitPipeline[R](p: Pipeline[R]): CIO[Vector[Either[SageException, Any]]] = // A pipeline batches commands per node, but all-masters commands must run on every master. Reject them before submission because running // one on a single node could break a later key-routed EVALSHA or FCALL, or return only part of the keyspace. - else if (p.commands.exists(_.allMasters)) + if (p.commands.exists(_.allMasters)) CIO.fail( InvalidArgument("a Pipeline cannot carry an all-masters command (e.g. SCRIPT LOAD, FUNCTION LOAD, KEYS); run it individually on the client") ) @@ -805,153 +436,122 @@ final private[client] class ClusterLive( } // use per-command dispatch for positions that the current topology cannot resolve. Complete after every position has succeeded or failed. - private def runPipeline[Out, R]( - p: Pipeline[Out, R], + private def runPipeline[R]( + p: Pipeline[R], complete: Try[Vector[Either[SageException, Any]]] => Unit, deferred: Vector[() => CommandSpan] ): Unit = { - val plan = topologyRef.get().split(p) + val plan = topologyRef.get().split(p.commands) // reject a malformed command before starting spans or submitting any part of the pipeline - plan.rejected.iterator.collectFirst { case (index, Rejected.Malformed) => index } match { - case Some(index) => complete(Failure(malformedKeys(p.commands(index).name))) - case None => dispatchPipeline(p, complete, deferred, plan) + plan.routes.indexOf(Route.Malformed) match { + case -1 => new PipelineRun(p.commands, plan, deferred, complete).start() + case index => complete(Failure(malformedKeys(p.commands(index).name))) } } - private def dispatchPipeline[Out, R]( - p: Pipeline[Out, R], - complete: Try[Vector[Either[SageException, Any]]] => Unit, + final private class PipelineRun( + commands: Vector[Command[?]], + plan: SplitPlan, deferred: Vector[() => CommandSpan], - plan: SplitPlan - ): Unit = { - val n = p.commands.length - val collector = - new TxSupport.IndexedCollector[Either[SageException, Any]](n, results => complete(Success(results))) - // settling a command releases its latch; a retry waits for the previous command on the same slot to preserve write order - val gates = Vector.fill(n)(new CountDownLatch(1)) - val slotAt = p.commands.map(slotOf) - val emits = Vector.tabulate(n) { i => - val span = if (deferred.isEmpty) CommandSpan.noop else Events.startDeferred(deferred(i)) - val settle: Try[Any] => Unit = result => { - gates(i).countDown() - collector.set(i, TxSupport.toEither(result)) + complete: Try[Vector[Either[SageException, Any]]] => Unit + ) { + private val collector = + new TxSupport.IndexedCollector[Either[SageException, Any]](commands.length, results => complete(Success(results))) + // reroutes and retries keep the original choice, preventing a slot from being split across a master and replica. + private val routing = pipelineMode(commands) + private val positions = Vector.tabulate(commands.length)(new Position(_)) + + final private class Position(index: Int) { + val command = commands(index) + // settling a command releases its latch; a retry waits for the previous command on the same slot to preserve write order + private val settled = new CountDownLatch(1) + // the slot a retry orders on; -1 for a keyless or cross-slot position, which orders against nothing + val slot: Int = plan.routes(index) match { + case Route.ToNode(_, slot) => slot.value + case Route.Unowned(slot) => slot.value + case _ => -1 } - Events.trackCommand[Any](events, p.commands(i), settle, span) - } - def awaitTurn(index: Int): Unit = { - val previous = if (slotAt(index) < 0) -1 else slotAt.lastIndexOf(slotAt(index), index - 1) - if (previous >= 0) gates(previous).await() - } - // reroutes keep the original choice, preventing a slot from being split across a master and replica. - val useReplica = readFrom != ReadFrom.Master && p.commands.forall(ReadRouting.replicaEligible) - // run rerouting on the scheduler because awaitTurn may block - def reroute(index: Int): Unit = scheduler.offload { - awaitTurn(index) - dispatch(p.commands(index), cluster.maxRedirects, emits(index), allowReplica = useReplica) - } - - plan.rejected.foreach { - case (index, Rejected.CrossSlot(slots)) => - if (multiSlotPolicy(p.commands(index)).nonEmpty) reroute(index) - else emits(index)(Failure(crossSlot(p.commands(index).name, slots))) - case (index, Rejected.Unowned(_)) => reroute(index) // dispatch refreshes then re-routes - case (_, Rejected.Malformed) => () // unreachable: the guard above returned - } - // add keyless commands to the first node batch. If there is no keyed batch, route each one independently. - if (plan.perNode.isEmpty) plan.keyless.foreach(reroute) - plan.perNode.zipWithIndex.foreach { case (NodeGroup(node, positions), groupIndex) => - // sort positions to preserve submission order within each node's batch, including keyless commands added to the first group - sendBatch(node, if (groupIndex == 0) (positions ++ plan.keyless).sorted else positions, p, emits, reroute, awaitTurn, useReplica) - } - } + val emit: Try[Any] => Unit = Events.trackCommand[Any]( + events, + command, + result => { + settled.countDown() + collector.set(index, TxSupport.toEither(result)) + }, + if (deferred.isEmpty) CommandSpan.noop else Events.startDeferred(deferred(index)) + ) - // the slot a retry orders on; -1 for a keyless or cross-slot position, which orders against nothing - private def slotOf(command: Command[?]): Int = - topologyRef.get().route(command) match { - case Route.ToNode(_, slot) => slot.value - case Route.Unowned(slot) => slot.value - case _ => -1 - } + private def awaitTurn(): Unit = + if (slot >= 0) { + val previous = positions.lastIndexWhere(_.slot == slot, index - 1) + if (previous >= 0) positions(previous).settled.await() + } - private def sendBatch[Out, R]( - node: Node, - indices: Vector[Int], - p: Pipeline[Out, R], - emits: Vector[Try[Any] => Unit], - reroute: Int => Unit, - awaitTurn: Int => Unit, - useReplica: Boolean - ): Unit = - // attribute the batch to the node that handles it, which is a replica when useReplica is true - if (useReplica) { - val replicas = topologyRef.get().shards.collectFirst { case s if s.master == node => s.replicas }.getOrElse(Vector.empty) - reads.pickOne(reads.candidatesFor(node, replicas), node) { - case Some(picked) => submitBatch(picked.node, picked.client, indices, p, emits, reroute, awaitTurn, useReplica) - case None => indices.foreach(reroute) + // run rerouting on the scheduler because awaitTurn may block + def reroute(): Unit = scheduler.offload { + awaitTurn() + dispatch(Request(command, emit, null, routing), cluster.maxRedirects) } - } else { - val existing = masterPool.existing(node) - if (existing != null) submitBatch(node, existing, indices, p, emits, reroute, awaitTurn, useReplica) - else - scheduler.offload { - val nc = masterPool.getOrEstablishOrNull(node) - submitBatch(node, nc, indices, p, emits, reroute, awaitTurn, useReplica) - } - } - private def submitBatch[Out, R]( - target: Node, - nc: NodeClient, - indices: Vector[Int], - p: Pipeline[Out, R], - emits: Vector[Try[Any] => Unit], - reroute: Int => Unit, - awaitTurn: Int => Unit, - useReplica: Boolean - ): Unit = { - def settle(index: Int, result: Try[Any]): Unit = { - Events.attributeNode(emits(index), target) - emits(index)(result) - } - val callbacks: Vector[Try[Any] => Unit] = indices.map { index => (result: Try[Any]) => - result match { - case Success(_) => settle(index, result) + def onBatchReply(target: Node): Try[Any] => Unit = { + case success @ Success(_) => Events.completeAt(emit, target)(success) // a fault's disposition can block on CLUSTER SLOTS, whose reply needs this very reader thread - case Failure(error) => + case Failure(error) => scheduler.offload { - awaitTurn(index) + awaitTurn() Fault.categorize(error) match { - // ASK keeps the exporting node as the slot owner in the topology. Send the command directly to the importing node with ASKING - // instead of routing it back to the exporter. MOVED and connection loss use normal routing. - case Fault.Redirected(redirect) => - redirect.kind match { - // ReadFrom.Replica rejects ASK because the importing node is a master, matching the single-read path in onReadFailure. - case RedirectKind.Ask if useReplica && readFrom == ReadFrom.Replica => settle(index, Failure(NotConnected())) - case RedirectKind.Ask => - onRedirect(target, redirect, p.commands(index), cluster.maxRedirects, emits(index)) - case RedirectKind.Moved => reroute(index) - } - case Fault.Lost(false) => reroute(index) - case Fault.TryAgain => onRetryable(p.commands(index), error, refreshFirst = false, cluster.maxRedirects, emits(index)) - case Fault.Unavailable(clusterWide) => onRetryable(p.commands(index), error, clusterWide, cluster.maxRedirects, emits(index)) - case Fault.Demoted | Fault.Lost(true) => - triggerRefresh() - settle(index, result) - case Fault.Fatal => settle(index, result) + // MOVED and connection loss use normal routing. ASK keeps the exporting node as the slot owner in the topology, so onFailure + // sends the command directly to the importing node with ASKING instead of routing it back to the exporter. + case Fault.Redirected(redirect) if redirect.kind == RedirectKind.Moved => reroute() + case Fault.Lost(false) => reroute() + case _ => + onFailure(target, Request(command, emit, null, routing), error, cluster.maxRedirects) } } } } + + def start(): Unit = { + plan.routes.iterator.zip(positions).foreach { + case (Route.CrossSlot, position) => + if (multiSlotPolicy(position.command).nonEmpty) position.reroute() + else position.emit(Failure(crossSlot(position.command))) + case (Route.Unowned(_), position) => position.reroute() // dispatch refreshes then re-routes + case (Route.Keyless, position) if plan.perNode.isEmpty => position.reroute() + case _ => () + } + plan.perNode.foreach { case NodeGroup(shard, indices) => send(shard, indices) } + } + + // attribute the batch to the node that handles it, which is a replica when the routing allows one + private def send(shard: Shard, indices: Vector[Int]): Unit = + if (routing == ReplicaRead) + reads.pickOne(reads.candidatesFor(shard.master, shard.replicas), shard.master) { + case Some(picked) => submit(picked.node, picked.client, indices) + case None => indices.foreach(positions(_).reroute()) + } + else masterPool.withClient(shard.master)(indices.foreach(positions(_).reroute()))(submit(shard.master, _, indices)) + // if the node is unavailable before the batch is submitted, route each command in the batch again - if (nc == null || !nc.submitAll(indices.map(p.commands), callbacks)) indices.foreach(reroute) + private def submit(target: Node, nc: MultiplexedConnection, indices: Vector[Int]): Unit = + if (!nc.submitAll(indices.map(positions(_).command), indices.map(positions(_).onBatchReply(target)))) + indices.foreach(positions(_).reroute()) } // --- transactions (one leased connection, optionally pinned to a key's slot) --------------------------------------------------------- - private def acquireScope: CIO[ClusterTxScope] = + protected def openTransaction: CIO[LiveTransactionScope] = if (closed) CIO.fail(NotConnected()) else CIO.value(new ClusterTxScope) - private def releaseScope(scope: ClusterTxScope): CIO[Unit] = CIO.blocking(scope.release()) + // A transaction cannot follow a redirect without breaking MULTI/EXEC atomicity. After an ownership or connection failure, refresh the + // topology in the background. A later transaction attempt then selects a connection using the updated topology. Data errors do not refresh. + // A forced refresh skips the throttle so the caller's retry uses the latest ownership information. + private def refreshFor(policy: RefreshPolicy): Unit = + policy match { + case RefreshPolicy.Forced => scheduler.offload(refreshThrottle(force = true)) + case RefreshPolicy.Throttled => triggerRefresh() + case RefreshPolicy.Skip => () + } /** * A cluster transaction scope. It leases a dedicated connection when the first command is submitted. A keyed command pins the transaction @@ -959,399 +559,200 @@ final private[client] class ClusterLive( * Later keys must use the same slot or fail with [[CrossSlot]]. The transaction does not follow redirects or reconnect after a connection * loss. These failures trigger a background topology refresh, and the caller can retry the full transaction. */ - final private class ClusterTxScope extends TransactionScope[CIO, String] { - - private val lock = new ReentrantLock() - private var released = false - private var nodeClient: NodeClient = null - private var conn: DedicatedConnection = null - private var pinnedNode: Node = null - private var pinnedSlot: Option[Slot] = None - private val armed = new AtomicBoolean(false) - - def watch[K: KeyCodec](key: K, rest: K*): CIO[Unit] = { - val command = Connection.watch(key, rest*) - CIO.async[Unit] { complete => - val tracked = Events.trackSpan(events, command, complete) - scheduler.offload(withConn(command, tracked) { c => - armed.set(true) - c.submit(command, faulting(tracked)) - }) - } - } + final private class ClusterTxScope extends LiveTransactionScope(events, refreshFor) { - def run[A](command: Command[A]): CIO[A] = - if (command.isBlocking) - CIO.fail(InvalidArgument("a Transaction cannot run blocking commands; run them individually on the client")) - else if (command.requiresClusterWideTxResult) + // guarded by `lock` + private var leased: Leased = null + private val acquiring = new ReentrantLock() + + protected type Target = Either[Throwable, Option[Slot]] + + override def run[A](command: Command[A]): CIO[A] = + if (command.requiresClusterWideTxResult) CIO.fail( InvalidArgument( s"${command.name} returns a cluster-wide result that a single-node Transaction cannot produce; run it individually on the client" ) ) - else - CIO.async[A] { complete => - val tracked = Events.trackSpan(events, command, complete) - scheduler.offload(withConn(command, tracked)(c => c.submit(command, faulting(tracked)))) - } + else super.run(command) - def discard: CIO[Unit] = - CIO.async[Unit] { complete => - scheduler.offload { - lock.lock() - try - if (released) complete(Failure(TxSupport.scopeReleasedError)) - else if (conn == null) complete(Success(())) // the transaction has not leased a connection or sent WATCH - else - Client.completing(complete) { - armed.set(false) - conn.submit(Connection.unwatch, faulting(complete)) - } - finally lock.unlock() - } - } - - // A transaction cannot follow a redirect without breaking MULTI/EXEC atomicity. After an ownership or connection failure, refresh the - // topology in the background. A later transaction attempt then selects a connection using the updated topology. Data errors do not refresh. - private def refreshOnFault(error: Throwable): Unit = refreshFor(Vector(Fault.categorize(error))) - - private def refreshFor(faults: Vector[Fault]): Unit = - faults.iterator.map(_.refreshPolicy).maxByOption(_.ordinal) match { - case Some(RefreshPolicy.Forced) => forceRefresh() - case Some(RefreshPolicy.Throttled) => triggerRefresh() - case _ => () - } - - private def faulting[A](complete: Try[A] => Unit): Try[A] => Unit = { - case failure @ Failure(error) => - refreshOnFault(error) - complete(failure) - case success => complete(success) - } - - private def refreshOnExecFault(frames: Vector[Frame]): Unit = - refreshFor(TxSupport.execErrors(frames).map(Fault.categorize).toVector) - - private[sage] def exec[Out, R](p: Pipeline[Out, R]): CIO[Option[Out]] = - runExec(p).flatMap { - case None => CIO.value(None) - case Some(results) => TxSupport.collapseStrict(results, p.toOut).map(Some(_)) - } - - private[sage] def execAttempt[Out, R](p: Pipeline[Out, R]): CIO[Option[R]] = - runExec(p).map(_.map(p.toResults)) - - private def runExec[Out, R](p: Pipeline[Out, R]): CIO[Option[Vector[Either[SageException, Any]]]] = - if (isReleased) - CIO.fail(TxSupport.scopeReleasedError) - else if (p.commands.isEmpty && !armed.get) - CIO.value(Some(Vector.empty)) - else if (p.commands.exists(_.isBlocking)) - CIO.fail(InvalidArgument("a Transaction cannot carry blocking commands; run them individually on the client")) - else if (p.commands.exists(_.requiresClusterWideTxResult)) + override protected def sendMultiExec[R](p: Pipeline[R]): CIO[TxSupport.ExecReplies] = + if (p.commands.exists(_.requiresClusterWideTxResult)) CIO.fail( InvalidArgument( "a Transaction cannot carry a command that returns a cluster-wide result; run it individually on the client" ) ) - else - CIO - .async[Vector[Frame]] { complete => - val tracked = Events.trackSpan(events, Connection.multi, complete) - scheduler.offload(submitExec(p, tracked)) - } - .flatMap { frames => - armed.set(false) // EXEC clears WATCH/MULTI state server-side whether it committed or aborted - refreshOnExecFault(frames) - TxSupport.interpretExec(p.commands, frames) - } + else super.sendMultiExec(p) // validate every pipeline slot before sending MULTI. Reject a cross-slot transaction before submitting any commands. - private def submitExec[Out, R](p: Pipeline[Out, R], complete: Try[Vector[Frame]] => Unit): Unit = - onConn(pipelineSlot(p), complete)(c => c.submitRaw(Connection.multi +: p.commands :+ Connection.exec, faulting(complete))) - - private def withConn[A](command: Command[?], complete: Try[A] => Unit)(use: DedicatedConnection => Unit): Unit = - onConn(commandSlot(command), complete)(use) + protected def withConn[A](target: Target, complete: Try[A] => Unit)(use: DedicatedConnection => Unit): Unit = + scheduler.offload(onConn(target, complete)(use)) // Check the released state and submit while holding `lock` so release cannot race with a submission. Acquire outside the lock so release() // can finish while a connection is being opened. - private def onConn[A](slotResult: Either[Throwable, Option[Slot]], complete: Try[A] => Unit)(use: DedicatedConnection => Unit): Unit = { - var retry = true - while (retry) { - retry = false - var fault: Throwable = null - var acquire = false - var slot: Option[Slot] = None - lock.lock() - try - if (released) complete(Failure(TxSupport.scopeReleasedError)) - else - slotResult match { - case Left(error) => - fault = error - complete(Failure(error)) - case Right(s) => - if (conn == null) { - acquire = true - slot = s - } else - checkPin(s) match { - case Left(error) => - fault = error - complete(Failure(error)) - case Right(()) => Client.completing(complete)(use(conn)) - } - } - finally lock.unlock() - if (fault != null) refreshOnFault(fault) - else if (acquire) - acquireConn(slot) match { - case Left(error) => - refreshOnFault(error) + private def onConn[A](target: Target, complete: Try[A] => Unit)(use: DedicatedConnection => Unit): Unit = { + val leasing = target.flatMap(ensureLeased) + var fault: Throwable = null + lock.lock() + try + if (released) complete(Failure(TxSupport.scopeReleasedError)) + else + leasing.flatMap(checkPin) match { + case Left(error) => + fault = error complete(Failure(error)) - case Right((nc, c, node, pin)) => - var giveBack = false - lock.lock() - try - if (released) { - giveBack = true - complete(Failure(TxSupport.scopeReleasedError)) - } - // lost the acquire race; re-validate against the winner's pin - else if (conn != null) { - giveBack = true - retry = true - } else { - nodeClient = nc - conn = c - pinnedNode = node - pinnedSlot = pin - Client.completing(complete)(use(conn)) - } - finally lock.unlock() - if (giveBack) nc.releaseTransaction(c, reusable = true) + case Right(()) => Client.completing(complete)(use(leased.conn)) } - } + finally lock.unlock() + if (fault != null) onFault(fault) + } + + // `acquiring` serializes the first lease so concurrent first commands share one connection; release() never takes it + private def ensureLeased(slot: Option[Slot]): Either[Throwable, Option[Slot]] = { + acquiring.lock() + try + if (underLock(released || leased != null)) Right(slot) + else + acquireConn(slot).map { acquired => + if (!underLock { if (!released) leased = acquired; !released }) + acquired.nc.pool.releaseTransaction(acquired.conn, reusable = true) + slot + } + finally acquiring.unlock() } - // must hold `lock` with conn != null. If a keyless command acquired the connection, accept the first keyed slot only when its node owns it. + private inline def underLock[A](inline body: A): A = { + lock.lock() + try body + finally lock.unlock() + } + + // must hold `lock` with leased != null. If a keyless command acquired the connection, accept the first keyed slot only when its node owns it. private def checkPin(slot: Option[Slot]): Either[Throwable, Unit] = slot match { case None => Right(()) case Some(s) => - pinnedSlot match { + leased.slot match { case Some(ps) if ps == s => Right(()) case Some(ps) => Left(CrossSlot(s"transaction touches slot ${s.value} but is pinned to slot ${ps.value}; MULTI/EXEC requires a single slot")) case None => - if (topologyRef.get().nodeForSlot(s).contains(pinnedNode)) { - pinnedSlot = Some(s) + if (topologyRef.get().nodeForSlot(s).contains(leased.node)) { + leased = leased.copy(slot = Some(s)) Right(()) } else Left(CrossSlot(s"transaction touches slot ${s.value} on a node other than its pinned one; MULTI/EXEC requires a single slot")) } } // runs outside `lock`: may force a topology refresh, connect, or wait for a pool slot - private def acquireConn(slot: Option[Slot]): Either[Throwable, (NodeClient, DedicatedConnection, Node, Option[Slot])] = { - val target = slot match { - case Some(s) => nodeForSlotRefreshing(s).map(_ -> slot) - case None => pickNode(topologyRef.get()).map(_ -> None) - } - target match { - case None => Left(NotConnected()) - case Some((node, pin)) => + private def acquireConn(slot: Option[Slot]): Either[Throwable, Leased] = + slot.fold(pickNode(topologyRef.get()))(nodeForSlotRefreshing) match { + case None => Left(NotConnected()) + case Some(node) => try { val nc = masterPool.getOrEstablish(node) - Right((nc, nc.acquireForTransaction(), node, pin)) + Right(Leased(nc, nc.pool.acquireForTransaction(), node, slot)) } catch { case error: SageException => Left(error) case NonFatal(_) => Left(ConnectionLost(mayHaveExecuted = false)) } } - } private def nodeForSlotRefreshing(slot: Slot): Option[Node] = topologyRef.get().nodeForSlot(slot).orElse { - refresh(force = true) + refreshThrottle(force = true) topologyRef.get().nodeForSlot(slot) } // select the transaction connection by key slot. Keep the slot even when the topology does not currently identify its owner. - private def commandSlot(command: Command[?]): Either[Throwable, Option[Slot]] = + protected def targetOf(command: Command[?]): Either[Throwable, Option[Slot]] = topologyRef.get().route(command) match { - case Route.Malformed => Left(malformedKeys(command.name)) - case Route.Keyless => Right(None) - case Route.ToNode(_, slot) => Right(Some(slot)) - case Route.Unowned(slot) => Right(Some(slot)) - case Route.CrossSlot(slots) => Left(crossSlot(command.name, slots)) + case Route.Malformed => Left(malformedKeys(command.name)) + case Route.Keyless => Right(None) + case Route.ToNode(_, slot) => Right(Some(slot)) + case Route.Unowned(slot) => Right(Some(slot)) + case Route.CrossSlot => Left(crossSlot(command)) } - private def pipelineSlot[Out, R](p: Pipeline[Out, R]): Either[Throwable, Option[Slot]] = { - var acc = Option.empty[Slot] - val it = p.commands.iterator - while (it.hasNext) - commandSlot(it.next()) match { - case Left(error) => return Left(error) - case Right(None) => () - case Right(Some(slot)) => - acc match { - case None => acc = Some(slot) - case Some(prev) => - if (prev != slot) return Left(CrossSlot("transaction keys span multiple slots; MULTI/EXEC requires a single slot")) - } - } - Right(acc) - } + protected def targetOf(commands: Vector[Command[?]]): Either[Throwable, Option[Slot]] = + commands.foldLeft[Either[Throwable, Option[Slot]]](Right(None)) { (acc, command) => + acc.flatMap(pinned => + targetOf(command).flatMap { + case Some(slot) if pinned.exists(_ != slot) => + Left(CrossSlot("transaction keys span multiple slots; MULTI/EXEC requires a single slot")) + case slot => Right(pinned.orElse(slot)) + } + ) + } - private def isReleased: Boolean = { - lock.lock() - try released - finally lock.unlock() - } + protected def leasedConn: DedicatedConnection = if (leased == null) null else leased.conn - // Reject further operations, then release the transaction connection. Reuse it only when healthy, with no pending commands or watched - // keys. A transaction that did not submit any commands has no connection to release. - private[internal] def release(): Unit = { - lock.lock() - val (nc, c, reusable) = - try { - released = true - (nodeClient, conn, conn != null && conn.isHealthy && conn.isQuiescent && !armed.get) - } finally lock.unlock() - if (nc != null) nc.releaseTransaction(c, reusable) - } + // runs after release() sealed the scope, so `leased` no longer changes + protected def giveBack(conn: DedicatedConnection, reusable: Boolean): Unit = leased.nc.pool.releaseTransaction(conn, reusable) } - private def flushNode(node: Node): Unit = - if (cachingEnabled) { - val nc = masterPool.existing(node) - if (nc != null) nc.flushCache() - } - - // --- topology refresh (single-flight, throttled) ------------------------------------------------------------------------------------- - - private val refreshWork: () => Unit = () => runRefresh() + // a transaction's leased connection, its node, and the slot it is pinned to (None until a keyed command arrives) + final private case class Leased(nc: MultiplexedConnection, conn: DedicatedConnection, node: Node, slot: Option[Slot]) - private def triggerRefresh(): Unit = refreshThrottle.trigger(refreshWork) - - private def startRefreshPoll(): Unit = refreshThrottle.startPolling(cluster.topologyRefreshInterval)(triggerRefresh()) + private def flushNode(node: Node): Unit = { + val nc = masterPool.existing(node) + if (nc != null) nc.flushCache() + } - // wait for any current refresh to finish before callers read `topologyRef` - private def refresh(force: Boolean): Unit = if (!closed) refreshThrottle(force)(runRefresh()) + // --- topology refresh (single-flight, throttled) ------------------------------------------------------------------------------------- - // skip a queued refresh after close so it does not open new connections - private def runRefresh(): Unit = - if (!closed) - querySlots(refreshCandidates()) match { - case Some((from, shards)) => adopt(from, shards) - // if no candidate answers CLUSTER SLOTS, slot ownership is unknown. Clear every client-side cache. - case None => if (cachingEnabled) masterPool.foreachEstablished(_.flushCache()) - } + protected def rediscover(): Unit = + querySlots(refreshCandidates()) match { + case Some(ranges) => adopt(ranges) + // if no candidate answers CLUSTER SLOTS, slot ownership is unknown. Clear every client-side cache. + case None => masterPool.foreachEstablished(_.flushCache()) + } private def refreshCandidates(): Vector[Node] = (masterPool.candidatesByLiveness ++ seeds).distinct - private def querySlots(candidates: Vector[Node]): Option[(Node, Vector[Shard])] = + private def querySlots(candidates: Vector[Node]): Option[Vector[SlotRange]] = candidates.iterator.flatMap(trySlots).nextOption() - private def trySlots(node: Node): Option[(Node, Vector[Shard])] = - try querySlotsVia(node).toOption.map(node -> _) + private def trySlots(node: Node): Option[Vector[SlotRange]] = + try querySlotsVia(node).toOption catch { case NonFatal(_) => None } // treat an empty CLUSTER SLOTS reply as unavailable topology information. A node can return it before joining a formed cluster. - private def querySlotsVia(node: Node): Either[Throwable, Vector[Shard]] = { - val nc = masterPool.getOrEstablish(node) - Bootstrap.awaitReply[Vector[Shard]](connectTimeout.toMillis)(callback => nc.submit(Cluster.slots, asking = false, callback)) match { - case None => - Left(TimedOut(s"CLUSTER SLOTS on ${node.host}:${node.port} timed out after ${connectTimeout.toMillis}ms")) - case Some(Success(shards)) if shards.nonEmpty => - Right(shards) - case Some(Success(_)) => - Left(UnsupportedServer(s"${node.host}:${node.port} owns no slots: it is not part of a formed cluster")) - case Some(Failure(error: ServerError)) => - Left(UnsupportedServer(s"${node.host}:${node.port} rejected CLUSTER SLOTS: ${error.getMessage}")) - case Some(Failure(error)) => - Left(error) + private def querySlotsVia(node: Node): Either[Throwable, Vector[SlotRange]] = { + val nc = masterPool.getOrEstablish(node) + def timedOut = TimedOut(s"CLUSTER SLOTS on ${node.host}:${node.port} timed out after ${config.connectTimeout.toMillis}ms") + Bootstrap.awaitReply[Vector[SlotRange]](config.connectTimeout.toMillis, timedOut)(nc.submit(Cluster.slots(node), _)) match { + case Success(ranges) if ranges.nonEmpty => Right(ranges) + case Success(_) => Left(UnsupportedServer(s"${node.host}:${node.port} owns no slots: it is not part of a formed cluster")) + case Failure(error: ServerError) => Left(UnsupportedServer(s"${node.host}:${node.port} rejected CLUSTER SLOTS: ${error.getMessage}")) + case Failure(error) => Left(error) } } - // Prune bundles for masters that are no longer listed. This stops reconnect loops for nodes that have left. An empty announce-IP from CLUSTER SLOTS - // means "the node I queried", so substitute `from` as redirects do - private def adopt(from: Node, shards: Vector[Shard]): Unit = { - val resolved = shards.map(shard => shard.copy(master = resolve(shard.master, from), replicas = shard.replicas.map(resolve(_, from)))) + // Prune bundles for masters that are no longer listed. This stops reconnect loops for nodes that have left. + private def adopt(ranges: Vector[SlotRange]): Unit = { val oldTopology = topologyRef.get() - val previous = if (events.emitsEvents) slotOwningMasters(oldTopology).toSet else Set.empty[Node] - val newTopology = ClusterTopology.from(resolved) + val previous = if (events.emitsEvents) oldTopology.masters.toSet else Set.empty[Node] + val newTopology = ClusterTopology.from(ranges) // retire losing masters' caches before the new topology is published - if (cachingEnabled) newTopology.mastersLosingSlots(oldTopology).foreach(flushNode) + newTopology.mastersLosingSlots(oldTopology).foreach(flushNode) topologyRef.set(newTopology) // skip the empty -> populated bootstrap transition: discovering the topology at connect is not a change if (events.emitsEvents && previous.nonEmpty) { - val current = slotOwningMasters(newTopology) + val current = newTopology.masters if (current.toSet != previous) events.emit(SageEvent.TopologyChanged(current)) } - val masters = resolved.map(_.master).toSet + val masters = ranges.map(_.master).toSet masterPool.retain(masters.contains) // prune replica connections and their cursors for replicas the new topology no longer lists, mirroring the master prune - val replicaNodes = resolved.iterator.flatMap(_.replicas).toSet + val replicaNodes = ranges.iterator.flatMap(_.replicas).toSet replicaPool.retain(replicaNodes.contains) + // replicas also receive PUBLISH, so a master demoted in place keeps the classic subscriptions + subscriptions.retain(node => masters(node) || replicaNodes(node)) reads.retain(masters.contains) // Reassign shard subscriptions only when slot ownership changes. Doing this for every forced refresh during failover would create a - // refresh and reconciliation loop. Classic subscriptions need reassignment only when their connection closes. + // refresh and reconciliation loop. if (!newTopology.sameOwnership(oldTopology)) subscriptions.onTopologyChanged() } - - private def closeAll(): Unit = { - closed = true - refreshThrottle.stopPolling() - masterPool.close() - subscriptions.close() - replicaPool.close() - events.close() - } -} - -private[client] object ClusterLive { - - def connect( - config: SageConfig, - seeds: Vector[Node], - cluster: ClusterConfig, - scheduler: Scheduler, - translate: Throwable => Throwable - ): CIO[Client[CIO, String]] = - CIO.blocking[Client[CIO, String]] { - val bootstrap = Bootstrap.commands(config.auth, config.database, config.clientName) - val factory: Node => MultiplexedConnection.TransportFactory = node => { - val upgrade = Tls.buildUpgrade(config.tls, node.host, node.port) - (onFrame, onClosed) => SocketTransport.connect(node.host, node.port, config.connectTimeout, upgrade, onFrame, onClosed) - } - val events = Events(config.listeners, config.tracer) - val live = new ClusterLive( - factory, - scheduler, - bootstrap, - config.reconnect, - config.watchdog, - config.connectTimeout, - config.closeTimeout, - config.dedicatedPool, - cluster, - config.pubsub.bufferSize, - seeds, - config.readFrom, - events, - config.clientCache.enabled, - config.clientCache.maxBytes - ) - // Translate discovery's handshake/TLS failures here rather than via mapError, which the per-backend CIO alias does not reconcile through - // `Client`'s invariant type parameter - try { - live.bootstrapTopology() - live - } catch { - case NonFatal(error) => - events.close() - throw translate(error) - } - } } diff --git a/sage-client/shared/src/main/scala/sage/client/internal/ClusterSubscriptions.scala b/sage-client/shared/src/main/scala/sage/client/internal/ClusterSubscriptions.scala index a8e2adbe..38c1cd13 100644 --- a/sage-client/shared/src/main/scala/sage/client/internal/ClusterSubscriptions.scala +++ b/sage-client/shared/src/main/scala/sage/client/internal/ClusterSubscriptions.scala @@ -3,286 +3,165 @@ package sage.client.internal import java.util.concurrent.locks.ReentrantLock import scala.collection.mutable -import scala.util.control.NonFatal +import scala.util.Try import SubscriptionConnection.{Kind, RawSubscription, Sink} import sage.Bytes import sage.SageException.NotConnected -import sage.client.{BackoffConfig, WatchdogConfig} +import sage.client.SageConfig import sage.cluster.{ClusterTopology, Node, Slot} -import sage.commands.Command /** * Manages pub/sub subscriptions in a cluster. Classic channel and pattern subscriptions share one connection to an arbitrary master - * because `PUBLISH` broadcasts across the cluster. If that node becomes unavailable, the manager chooses another master. Shard channel - * subscriptions use one connection per owning node, created when first needed and closed after its last subscription ends. + * because `PUBLISH` broadcasts across the cluster. If that node becomes unavailable or leaves the cluster, the connection reconnects to + * another master. Shard channel subscriptions use one connection per owning node, created when first needed and closed after its last + * subscription ends. * - * Cluster subscription connections do not reconnect themselves. When one closes, the manager refreshes the topology and assigns its - * subscribers to the current owner. It also performs this reconciliation after topology changes discovered by commands. Each - * `SSUBSCRIBE` contains channels from one slot to avoid `CROSSSLOT` errors. + * Shard connections do not reconnect themselves. When one closes, the manager refreshes the topology and assigns its subscribers to the + * current owner. It also performs this reconciliation after topology changes discovered by commands. */ final private[client] class ClusterSubscriptions( nodeFactory: Node => MultiplexedConnection.TransportFactory, - bootstrap: Vector[Command[?]], scheduler: Scheduler, - reconnect: BackoffConfig, - watchdog: WatchdogConfig, - connectTimeoutMillis: Long, - bufferSize: Int, + config: SageConfig, topologyOf: () => ClusterTopology, refresh: () => Unit, - pickMaster: () => Option[Node] -) { + pickMaster: () => Option[Node], + events: Events +) extends SubscriptionConnection.PubSub { private val lock = new ReentrantLock() - // --- classic state (guarded by lock) --- - private var classicConn: SubscriptionConnection = null - private val classicSubs = mutable.LinkedHashSet.empty[ClassicSub] - // --- sharded state (guarded by lock) --- private val shardConns = mutable.HashMap.empty[Node, SubscriptionConnection] private val shardSubs = mutable.LinkedHashSet.empty[ShardSub] private var closed = false + // Retries wait at most four initial delays, so a subscription resumes soon after a long failover ends. They refresh the topology at most + // once per that delay. + private val retryBackoff = config.reconnect.copy(maxDelay = config.reconnect.maxDelay.min(config.reconnect.initialDelay * 4)) + private val retryRefresh = new RefreshThrottle(scheduler, retryBackoff.maxDelay.toMillis, refresh) + + // one retry reconciles every subscription, including those placed while it waits + private val retries = new Reconnects(scheduler, retryBackoff, lock) + private inline def locked[A](inline body: A): A = { lock.lock() try body finally lock.unlock() } - // run one pass at a time while using the outer lock to update its state. If schedule is called during a pass, run one more pass afterward. - final private class CoalescedPass(body: () => Unit) { - private var running = false - private var queued = false - - def schedule(): Unit = { - val go = locked { - if (closed) false - else if (running) { - queued = true - false - } else { - running = true - true - } - } - if (go) scheduler.offload(run()) - } - - private def run(): Unit = - try body() - finally { - val again = locked { - if (!closed && queued) { - queued = false - true - } else { - running = false - false - } - } - if (again) scheduler.offload(run()) - } - } - - // refresh first so a replacement for the failed master is chosen from the current topology. - private val classicRehome = new CoalescedPass(() => { - refresh() - rehomeClassic() - }) - private val shardReconcile = new CoalescedPass(() => reconcileShard()) + // one pass at a time; a request during a pass runs one more pass afterward + private val shardReconcile = new RefreshThrottle(scheduler, 0L, () => reconcileShard()) // --- classic (channels / patterns) ------------------------------------------------------------------------------------------------------- - def subscribeChannels(channels: Vector[String]): RawSubscription = classic(channels, Kind.Channel) - - def subscribePatterns(patterns: Vector[String]): RawSubscription = classic(patterns, Kind.Pattern) - - private def classic(names: Vector[String], kind: Kind): RawSubscription = { - val sink = new Sink(names, kind, bufferSize) - val sub = ClassicSub(sink, names, kind) - try { - val conn = - locked { - if (closed) throw NotConnected() - classicSubs += sub - ensureClassicConn() - } - conn.attach(sink, names, kind) - } catch { - case e: Throwable => - locked(classicSubs -= sub) - sink.terminate() - throw e - } - new RawSubscription(sink, () => closeClassic(sub)) - } - - // must hold lock - private def ensureClassicConn(): SubscriptionConnection = { - if (classicConn == null) - pickMaster() match { - case Some(node) => classicConn = newConnection(node, () => onClassicTerminated()) - case None => throw NotConnected() - } - classicConn - } + private val classicOn = new SubscriptionConnection.Following(nodeFactory, pickMaster) - private def closeClassic(sub: ClassicSub): Unit = { - val (conn, teardown) = - locked { - classicSubs -= sub - val c = classicConn - val teardown = classicSubs.isEmpty - if (teardown) classicConn = null - (c, teardown) - } - if (conn != null) { - conn.detach(sub.sink, sub.names, sub.kind) - sub.sink.terminate() - if (teardown) conn.close() - } else sub.sink.terminate() - } + // Each attempt connects to the master picked at that time. After a loss, every reconnect attempt refreshes the topology first. + private val classic = new SubscriptionConnection( + classicOn.factory, + scheduler, + config.copy(reconnect = retryBackoff), + isLive = () => true, + SubscriptionConnection.OnLoss.Reconnect(() => retryRefresh(force = false), events, () => classicOn.node, immediately = true) + ) - private def onClassicTerminated(): Unit = { - locked { classicConn = null } - classicRehome.schedule() - } + def subscribeChannels(channels: Vector[String]): RawSubscription = classic.owned(channels, Kind.Channel, failIfUnconfirmed = false) - private def rehomeClassic(): Unit = { - val subs = locked(if (closed || classicSubs.isEmpty) Vector.empty[ClassicSub] else classicSubs.toVector) - if (subs.nonEmpty) { - val conn = locked { - if (closed || classicSubs.isEmpty) null - else { - if (classicConn == null) pickMaster().foreach(node => classicConn = newConnection(node, () => onClassicTerminated())) - classicConn - } - } - if (conn == null) scheduleRehomeRetry() // a master is not available yet; retry when the topology may contain one - else { - // When attachment fails during establishment, it resets the connection to Idle and rethrows without calling onTerminated. Ignoring - // that failure would leave every classic subscription on the dead connection. Drop the connection and retry here. - var failed = false - subs.foreach { sub => - try { - conn.attach(sub.sink, sub.names, sub.kind) - // if the subscription closed while attach was running, closeClassic could not detach it yet. Detach it here after attach finishes. - if (!locked(classicSubs.contains(sub))) { conn.detach(sub.sink, sub.names, sub.kind): Unit } - } catch { case NonFatal(_) => failed = true } - } - if (failed) { - // Attachment may have restored some subscriptions on this connection. Close it before retrying to avoid duplicate delivery and an - // unused open socket. - locked(if (classicConn eq conn) classicConn = null) - conn.shutdown() - scheduleRehomeRetry() - } - } - } - } + def subscribePatterns(patterns: Vector[String]): RawSubscription = classic.owned(patterns, Kind.Pattern, failIfUnconfirmed = false) - private def scheduleRehomeRetry(): Unit = - if (!locked(closed)) scheduler.after(reconnect.initialDelay)(classicRehome.schedule()) + // `redis-cli --cluster del-node` resets a removed node without stopping it, so PUBLISH stops reaching it while the connection stays open + def retain(listed: Node => Boolean): Unit = classicOn.retain(listed) // --- sharded (shard channels) ------------------------------------------------------------------------------------------------------------ - // ensure/get both yield nothing once closed, so no connection is created during teardown - private val conns: Placement.Conns = new Placement.Conns { - def ensure(node: Node): Option[Placement.ShardConn] = locked(if (closed) None else Some(ensureShardConn(node))) - def get(node: Node): Option[Placement.ShardConn] = locked(shardConns.get(node)) - } - def subscribeShard(channels: Vector[String]): RawSubscription = { - val sink = new Sink(channels, Kind.Shard, bufferSize) - val sub = ShardSub(sink, channels) + val sink = new Sink(channels, Kind.Shard, config.pubsub.bufferSize) + val sub = ShardSub(sink) locked { if (closed) throw NotConnected() shardSubs += sub } - // if placement fails, closeShard detaches any completed subscriptions and terminates the sink - try place(sub) - catch { - case e: Throwable => - closeShard(sub) - throw e - } + // Keep channels with an unowned slot or an unreachable owner pending and retry them. If the owner refuses a channel, closeShard detaches + // any completed subscriptions and terminates the sink. + onThrow { + // a slot with no owner is usually mid-failover, and a refresh may already name its new owner + val topo = topologyOf() + if (channels.exists(channel => topo.nodeForSlot(Slot.of(Bytes.utf8(channel))).isEmpty)) refresh() + if (!reconcile(sub, planFor(sub.channels))) { sink.failure.foreach(throw _); scheduleRetry() } + }(_ => closeShard(sub)) new RawSubscription(sink, () => closeShard(sub)) } - // if the initial connection attempt fails, cancel the whole subscription. Keep channels with an unowned slot pending and retry them. - private def place(sub: ShardSub): Unit = { - if (hasUnownedSlot(sub.channels)) refresh() - sub.placement.place(planFor(sub.channels), conns) - if (!sub.placement.fullyPlaced) scheduleRetry() - } - - // Group channels by owning node and then by slot, with one SSUBSCRIBE for each group. Omit an unowned slot from this attempt; the caller - // refreshes the topology before each retry. - private def planFor(channels: Vector[String]): Placement.Plan = { - val topo = topologyOf() - val byNode = mutable.HashMap.empty[Node, mutable.HashMap[Slot, mutable.ArrayBuffer[String]]] - channels.foreach { channel => - val slot = Slot.of(Bytes.utf8(channel)) - topo.nodeForSlot(slot).foreach { node => - byNode.getOrElseUpdate(node, mutable.HashMap.empty).getOrElseUpdate(slot, mutable.ArrayBuffer.empty) += channel - } - } - byNode.iterator.map { case (node, slots) => node -> slots.valuesIterator.map(_.toVector).toVector }.toMap + // Evaluates the plan under the subscription's lock, so a plan from an older topology cannot replace a newer one. No connection is created + // once the manager is closed. + private def reconcile(sub: ShardSub, planned: => ClusterSubscriptions.Plan): Boolean = { + sub.lock.lock() + try + ClusterSubscriptions.reconcile(sub.sink, planned, locked(shardConns.toVector), node => locked(Option.unless(closed)(ensureShardConn(node)))) + finally sub.lock.unlock() } - private def hasUnownedSlot(channels: Vector[String]): Boolean = { + // Group channels by owning node. Omit a channel whose slot is unowned from this attempt; the caller refreshes the topology before each retry. + private def planFor(channels: Vector[String]): ClusterSubscriptions.Plan = { val topo = topologyOf() - channels.exists(channel => topo.nodeForSlot(Slot.of(Bytes.utf8(channel))).isEmpty) + channels.groupBy(channel => topo.nodeForSlot(Slot.of(Bytes.utf8(channel)))).collect { case (Some(node), names) => node -> names } } // must hold lock private def ensureShardConn(node: Node): SubscriptionConnection = - shardConns.getOrElseUpdate(node, newConnection(node, () => onShardConnTerminated(node))) - - private def onShardConnTerminated(node: Node): Unit = { - locked(shardConns.remove(node)) - // A dropped connection can mean that the slot migrated. The server sends sunsubscribe and then disconnects. Refresh because the stale - // topology still names the disconnected owner, which planFor would not consider unowned. Reconcile after the refresh finds the new owner. - scheduler.offload { - refresh() - shardReconcile.schedule() - } + shardConns.getOrElseUpdate( + node, { + // a dropped channel is not a failure, so it re-homes without backoff + val report = + SubscriptionConnection.OnLoss.Report(onShardConnTerminated(node, _), () => scheduler.offload(refreshAndReconcile()), () => scheduleRetry()) + new SubscriptionConnection(nodeFactory(node), scheduler, config, () => true, report) + } + ) + + // a connection that reports its loss late must not remove the connection that replaced it + private def forget(node: Node, conn: SubscriptionConnection): Unit = locked(if (shardConns.get(node).contains(conn)) shardConns -= node) + + private def onShardConnTerminated(node: Node, conn: SubscriptionConnection): Unit = { + forget(node, conn) + scheduleRetry(immediately = true) } - def onTopologyChanged(): Unit = shardReconcile.schedule() + def onTopologyChanged(): Unit = shardReconcile.request() - // Assign each subscription to the current owners of its channels. Refresh at most once per pass and retry incomplete work after transient - // failover errors. + // Assign each subscription to the current owners of its channels, and retry incomplete work after transient failover errors. private def reconcileShard(): Unit = { val subs = locked(if (closed) Vector.empty else shardSubs.toVector) if (subs.nonEmpty) { - if (subs.exists(sub => hasUnownedSlot(sub.channels))) refresh() var incomplete = false subs.foreach { sub => - val failed = sub.placement.reconcile(planFor(sub.channels), conns) + val placed = sub.sink.failure.isEmpty && Try(reconcile(sub, planFor(sub.channels))).getOrElse(false) // If the subscription closed during this pass, closeShard may have detached it before reconcile attached it again. Reconcile with an - // empty plan to remove those attachments. - if (!locked(shardSubs.contains(sub))) { sub.placement.reconcile(Map.empty, conns): Unit } - // The plan omits an unowned slot, and reconcile does not treat that omission as a failure. Check fullyPlaced and retry while a requested - // channel remains unattached. - else if (failed || !sub.placement.fullyPlaced) incomplete = true + // empty plan to remove those attachments, as for a subscription the owner refused. + if (!locked(shardSubs.contains(sub)) || sub.sink.failure.nonEmpty) { reconcile(sub, Map.empty): Unit } + else if (!placed) incomplete = true } evictEmptyShardConns() - if (incomplete) scheduleRetry() + if (incomplete) scheduleRetry() else locked(retries.live()) } } - // an incomplete placement (owner unreachable, or a Slot still unowned mid-failover) retries after a short delay until it converges - private def scheduleRetry(): Unit = - if (!locked(closed)) scheduler.after(reconnect.initialDelay)(shardReconcile.schedule()) + // A lost shard connection or an incomplete placement (owner unreachable, or a slot still unowned mid-failover) retries with backoff. A + // retry refreshes first unless another refreshed within the retry delay, because the topology still names the old owner until a refresh + // or failover replaces it. With `immediately`, the first retry after a stable period does not wait. + private def scheduleRetry(immediately: Boolean = false): Unit = locked(retries.schedule(!closed, _ => (), immediately)(refreshAndReconcile())) + + private def refreshAndReconcile(): Unit = { + retryRefresh(force = false) + shardReconcile.request() + } private def closeShard(sub: ShardSub): Unit = { locked(shardSubs -= sub) - sub.placement.reconcile(Map.empty, conns) // detach every placement; the empty plan leaves the ledger empty + reconcile(sub, Map.empty): Unit // detach every placement sub.sink.terminate() evictEmptyShardConns() } @@ -290,47 +169,58 @@ final private[client] class ClusterSubscriptions( private def evictEmptyShardConns(): Unit = { val candidates = locked(shardConns.iterator.collect { case (node, conn) if conn.isEmpty => node -> conn }.toVector) candidates.foreach { case (node, conn) => - if (conn.closeIfEmpty()) locked(if (shardConns.get(node).contains(conn)) shardConns -= node) + if (conn.closeIfEmpty()) forget(node, conn) } } // --- shared ------------------------------------------------------------------------------------------------------------------------------ - private def newConnection(node: Node, onTerminated: () => Unit): SubscriptionConnection = - new SubscriptionConnection( - nodeFactory(node), - bootstrap, - scheduler, - reconnect, - watchdog, - connectTimeoutMillis, - bufferSize, - isLive = () => true, - cluster = true, - onTerminated = onTerminated - ) - def close(): Unit = { - val (classic, shard) = + val shard = locked { closed = true - val c = classicConn - classicConn = null val s = shardConns.values.toVector shardConns.clear() // terminate sinks before closing connections. This releases any reader waiting because of backpressure before close waits for it. - (classicSubs.toVector.map(_.sink) ++ shardSubs.toVector.map(_.sink)).foreach(_.terminate()) - classicSubs.clear() + shardSubs.foreach(_.sink.terminate()) shardSubs.clear() - (c, s) + s } - if (classic != null) classic.close() + shardReconcile.stop() + classic.close() shard.foreach(_.close()) } - final private case class ClassicSub(sink: Sink, names: Vector[String], kind: Kind) + final private class ShardSub(val sink: Sink) { + def channels: Vector[String] = sink.names + val lock = new ReentrantLock() + } +} - final private class ShardSub(val sink: Sink, val channels: Vector[String]) { - val placement = new Placement(sink, channels) +private[internal] object ClusterSubscriptions { + + type Plan = Map[Node, Vector[String]] + + // what reconciliation needs of a shard connection; [[SubscriptionConnection]] is the production one + trait ShardConn { + def attach(sink: Sink, names: Vector[String]): Unit + def detach(sink: Sink, names: Vector[String]): Unit + def namesOf(sink: Sink): Vector[String] + } + + // Detaches what `sink` holds outside the plan on the `current` connections, attaches every planned node's channels, and returns true when + // every requested channel is attached. The caller omits unowned slots from the plan and retries until this holds. + def reconcile(sink: Sink, plan: Plan, current: Vector[(Node, ShardConn)], ensure: Node => Option[ShardConn]): Boolean = { + current.foreach { case (node, conn) => + val keep = plan.getOrElse(node, Vector.empty).toSet + val gone = conn.namesOf(sink).filterNot(keep) + if (gone.nonEmpty) conn.detach(sink, gone) + } + // a concurrent eviction, a refused channel or a failure to create the connection leaves the channels pending for a retry + val attached = plan.iterator.flatMap { case (node, names) => + Try(ensure(node)).toOption.flatten.filter(conn => Try(conn.attach(sink, names)).isSuccess).iterator.flatMap(_.namesOf(sink)) + }.toSet + // count distinct attached channels across all nodes. Counting each node separately could hide a missing channel when another is recorded twice. + attached.size >= sink.names.distinct.size } } diff --git a/sage-client/shared/src/main/scala/sage/client/internal/DedicatedConnection.scala b/sage-client/shared/src/main/scala/sage/client/internal/DedicatedConnection.scala deleted file mode 100644 index d16f3726..00000000 --- a/sage-client/shared/src/main/scala/sage/client/internal/DedicatedConnection.scala +++ /dev/null @@ -1,187 +0,0 @@ -package sage.client.internal - -import java.util.concurrent.ConcurrentLinkedQueue -import java.util.concurrent.atomic.{AtomicBoolean, AtomicInteger, AtomicReference, AtomicReferenceArray} - -import scala.util.{Failure, Success, Try} - -import sage.Bytes -import sage.SageException -import sage.SageException.ConnectionLost -import sage.commands.{Command, Reply} -import sage.protocol.Frame - -/** - * A connection borrowed exclusively from the [[DedicatedPool]]. It does not reconnect or run a watchdog. If the connection is lost, the - * pool discards it and in-flight work fails with `ConnectionLost(mayHaveExecuted = true)`. Replies are matched in order. This allows the - * synchronous `HELLO` setup to finish before the connection is borrowed and keeps a transaction's `MULTI`, queued-command, and `EXEC` - * replies aligned with their commands. - */ -final private[client] class DedicatedConnection private ( - factory: MultiplexedConnection.TransportFactory, - connectTimeoutMillis: Long -) { - - // the generation recorded when this connection joins the pool; see [[DedicatedPool.establishOutsideLock]] - @volatile private var stampedEpoch: MultiplexedConnection.Generation = MultiplexedConnection.Generation.initial - - private val pending = new ConcurrentLinkedQueue[DedicatedConnection.Waiter]() - private val transportRef = new AtomicReference[Transport]() - @volatile private var dead: Boolean = false - // Count a command when submit accepts it. The asynchronous transport may hold it before writeAttempted adds it to pending. Checking only - // pending.isEmpty could therefore return a connection to the pool while it still has a queued batch. - private val inFlight = new AtomicInteger(0) - - def epoch: MultiplexedConnection.Generation = stampedEpoch - - def stampEpoch(generation: MultiplexedConnection.Generation): Unit = stampedEpoch = generation - - def isHealthy: Boolean = !dead - - def isQuiescent: Boolean = !dead && inFlight.get() == 0 - - def submit[A](command: Command[A], callback: Try[A] => Unit): Unit = { - inFlight.incrementAndGet() - transportRef.get().send(new Entry(command, callback)) - } - - /** - * Sends `commands` as a single pipelined write and returns their raw reply frames in order. Used for a transaction's `MULTI` … `EXEC` - * batch, whose `EXEC` array the caller decodes per-position against the original commands; the frames are returned undecoded. - */ - def submitRaw(commands: Vector[Command[?]], callback: Try[Vector[Frame]] => Unit): Unit = { - inFlight.incrementAndGet() - transportRef.get().send(new RawBatch(commands, callback)) - } - - def close(): Unit = { - dead = true - val transport = transportRef.get() - if (transport != null) transport.close() - } - - /** - * Opens the socket and runs the bootstrap synchronously; throws (no retry) if the connect or handshake fails. - */ - def establish(bootstrap: Vector[Command[?]]): Unit = { - start() - runBootstrap(bootstrap) - } - - private def start(): Unit = { - val transport = factory(onFrame, onClosed) - transportRef.set(transport) - if (dead) transport.close() - else transport.start() - } - - private def runBootstrap(bootstrap: Vector[Command[?]]): Unit = - Bootstrap.run(bootstrap, connectTimeoutMillis, (c, cb) => submit(c, cb), () => close()) - - private def onFrame(frame: Frame): Unit = - frame match { - case _: Frame.Push => () - case reply => - val waiter = pending.poll() - if (waiter == null) close() // a reply with nothing pending means the stream desynced; discard - else { - // close before delivering a READONLY reply, ensuring the pool discards this connection when it is released - if (Poison.isReadonly(reply)) close() - waiter.complete(reply) - } - } - - private def onClosed(): Unit = { - dead = true - var waiter = pending.poll() - while (waiter != null) { - waiter.fail(ConnectionLost(mayHaveExecuted = true)) - waiter = pending.poll() - } - } - - final private class Entry[A](command: Command[A], callback: Try[A] => Unit) extends Transport.Item with DedicatedConnection.Waiter { - - var payload: Bytes = command.encode - - override def clearPayload(): Unit = payload = Bytes.empty - - def writeAttempted(): Unit = - pending.add(this): Unit - - // the transport calls dropped during teardown before onClosed. Mark the connection dead first to prevent the pool from reusing it during close. - def dropped(): Unit = { - dead = true - inFlight.decrementAndGet() - callback(Failure(ConnectionLost(mayHaveExecuted = false))) - } - - def complete(frame: Frame): Unit = { - val result = Reply.decode(command, frame) - inFlight.decrementAndGet() - callback(result) - } - - def fail(error: SageException): Unit = { - inFlight.decrementAndGet() - callback(Failure(error)) - } - } - - // Write the batch as one transport item and store each reply at its matching index. Invoke the callback after all replies arrive, or after - // the first failure. - final private class RawBatch(commands: Vector[Command[?]], callback: Try[Vector[Frame]] => Unit) extends Transport.Item { - - private val n = commands.length - private val frames = new AtomicReferenceArray[Frame](n) - private val remaining = new AtomicInteger(n) - private val done = new AtomicBoolean(false) - - var payload: Bytes = Bytes.concatBy(commands)(_.encode) - - override def clearPayload(): Unit = payload = Bytes.empty - - def writeAttempted(): Unit = { - var i = 0 - while (i < n) { - pending.add(new Slot(i)) - i += 1 - } - } - - def dropped(): Unit = { - dead = true - finish(Failure(ConnectionLost(mayHaveExecuted = false))) - } - - private def finish(result: Try[Vector[Frame]]): Unit = - if (done.compareAndSet(false, true)) { - inFlight.decrementAndGet() - callback(result) - } - - final private class Slot(index: Int) extends DedicatedConnection.Waiter { - - def complete(frame: Frame): Unit = { - frames.set(index, frame) - if (remaining.decrementAndGet() == 0) finish(Success(Vector.tabulate(n)(frames.get))) - } - - def fail(error: SageException): Unit = finish(Failure(error)) - } - } -} - -private[client] object DedicatedConnection { - - private trait Waiter { - def complete(frame: Frame): Unit - def fail(error: SageException): Unit - } - - /** - * Builds an unconnected connection; the pool runs the blocking [[DedicatedConnection.establish]] separately. - */ - def create(factory: MultiplexedConnection.TransportFactory, connectTimeoutMillis: Long): DedicatedConnection = - new DedicatedConnection(factory, connectTimeoutMillis) -} diff --git a/sage-client/shared/src/main/scala/sage/client/internal/DedicatedPool.scala b/sage-client/shared/src/main/scala/sage/client/internal/DedicatedPool.scala index e8f4b8cc..7939412a 100644 --- a/sage-client/shared/src/main/scala/sage/client/internal/DedicatedPool.scala +++ b/sage-client/shared/src/main/scala/sage/client/internal/DedicatedPool.scala @@ -1,17 +1,19 @@ package sage.client.internal -import java.util.concurrent.atomic.AtomicReference +import java.util.concurrent.atomic.{AtomicReference, AtomicReferenceArray} import java.util.concurrent.locks.ReentrantLock +import scala.annotation.tailrec import scala.collection.mutable import scala.concurrent.duration.* -import scala.util.{Failure, Try} +import scala.util.{Failure, Success, Try} import scala.util.control.NonFatal import sage.SageException import sage.SageException.{ConnectionLost, NotConnected, TimedOut} import sage.client.DedicatedPoolConfig import sage.commands.{Command, Connection} +import sage.protocol.Frame /** * A pool of dedicated connections for blocking commands, transactions, and lock replication checks. Connections are created when needed, @@ -27,8 +29,6 @@ final private[client] class DedicatedPool( bootstrap: Vector[Command[?]], scheduler: Scheduler, isLive: () => Boolean, - liveGeneration: () => Option[MultiplexedConnection.Generation], - isCurrent: MultiplexedConnection.Generation => Boolean, config: DedicatedPoolConfig, connectTimeoutMillis: Long ) { @@ -40,7 +40,6 @@ final private[client] class DedicatedPool( private val live = mutable.Set.empty[DedicatedConnection] // connections whose socket is being opened outside the lock, so close() can abort one still connecting private val establishing = mutable.Set.empty[DedicatedConnection] - private var reserved = 0 private var closing = false private val sweepHandle: Scheduler.Cancelable = config.idleTimeout match { @@ -50,23 +49,11 @@ final private[client] class DedicatedPool( /** * Runs a blocking command on a borrowed connection and releases it after the reply or failure. Acquisition runs on another thread - * because it may wait for a pool slot or open a socket. + * because it may wait for a pool slot or open a socket. After an `ASK` redirect, `asking` writes `ASKING` and the command consecutively on + * the same leased connection. The `ASKING` reply is discarded, and the command's reply releases the connection. */ - def use[A](command: Command[A], callback: Try[A] => Unit, lease: DedicatedPool.Lease = new DedicatedPool.Lease): Unit = - leaseAndSubmit(command, asking = false, callback, lease) - - /** - * Runs a blocking command after an `ASK` redirect. `ASKING` and the command are written consecutively on the same leased connection. - * The `ASKING` reply is discarded, and the command's reply releases the connection. - */ - def useAsking[A](command: Command[A], callback: Try[A] => Unit, lease: DedicatedPool.Lease): Unit = - leaseAndSubmit(command, asking = true, callback, lease) - - private def leaseAndSubmit[A](command: Command[A], asking: Boolean, callback: Try[A] => Unit, lease: DedicatedPool.Lease): Unit = - useConnection(callback, lease) { (conn, complete) => - if (asking) conn.submit[Unit](Connection.asking, _ => ()) - conn.submit(command, complete) - } + def use[A](command: Command[A], callback: Try[A] => Unit, lease: DedicatedPool.Lease = new DedicatedPool.Lease, asking: Boolean = false): Unit = + useConnection(callback, lease, asking)(_.submit(command, _)) // WAIT blocks its socket. Keep the write and confirmation on one leased connection so other commands can proceed independently. def useLockWrite[A]( @@ -76,12 +63,12 @@ final private[client] class DedicatedPool( lease: DedicatedPool.Lease, replication: LockReplication ): Unit = - useConnection(callback, lease, replication.cancelled, Some(replication.deadlineMillis))(replication.submit(_, command, asking, _)) + useConnection(callback, lease, asking, Some(replication.deadlineMillis))(replication.submit(_, command, _)) private def useConnection[A]( callback: Try[A] => Unit, lease: DedicatedPool.Lease, - onCancel: () => Unit = () => (), + asking: Boolean, deadlineMillis: Option[Long] = None )(submit: (DedicatedConnection, Try[A] => Unit) => Unit): Unit = if (!isLive()) callback(Failure(NotConnected())) @@ -100,11 +87,10 @@ final private[client] class DedicatedPool( callback(Failure(error)) case Right(conn) => // Cancellation can precede attachment while acquisition waits for a socket or pool slot. - val onInterrupt = () => { - onCancel() - callback(Failure(ConnectionLost(mayHaveExecuted = true))) - } + val onInterrupt = () => callback(Failure(ConnectionLost(mayHaveExecuted = true))) if (lease.attach(this, conn, onInterrupt)) { + // the leased connection is exclusive, so ASKING stays adjacent to the command + if (asking) conn.submit[Unit](Connection.asking, _ => ()) submit( conn, result => @@ -140,6 +126,13 @@ final private[client] class DedicatedPool( private[internal] def wakeWaiters(): Unit = locked(available.signalAll()) + // a connection admitted before the loss may point to the previous server after a failover, so it is discarded when released + private[internal] def onLivenessLost(): Unit = + locked { + live.foreach(_.markDead()) + available.signalAll() + } + def close(): Unit = { val toClose = locked { closing = true @@ -154,69 +147,52 @@ final private[client] class DedicatedPool( } private def acquire(lease: Option[DedicatedPool.Lease] = None, deadlineMillis: Option[Long] = None): DedicatedConnection = { - val budgetNanos = deadlineMillis.fold(config.acquireTimeout.toNanos) { deadline => + val budgetNanos = deadlineMillis.fold(config.acquireTimeout.toNanos) { deadline => math.min(config.acquireTimeout.toNanos, (deadline - scheduler.nowMillis).millis.toNanos) } // Condition waits use real nanoseconds, so only the remaining duration crosses from the scheduler's monotonic clock. - val deadlineNanos = System.nanoTime() + budgetNanos - locked { - while (true) { - if (lease.exists(_.isCancelled)) throw ConnectionLost(mayHaveExecuted = true) - if (closing) throw NotConnected() - // reject acquisition while the shared connection is reconnecting; repeat the liveness check after each wake-up - if (!isLive()) throw NotConnected() - val remaining = deadlineNanos - System.nanoTime() - // A lock deadline limits the write itself. An expired write must not take a slot even when one is immediately available. - if (deadlineMillis.isDefined && remaining <= 0L) throw acquireTimedOut(budgetNanos.nanos) - val reused = takeIdleLocked() - if (reused != null) return reused - if (live.size + reserved < config.maxConnections) { - reserved += 1 - return establishOutsideLock() - } - if (remaining <= 0L) throw acquireTimedOut(budgetNanos.nanos) + val deadlineNanos = System.nanoTime() + budgetNanos + @tailrec def loop(): DedicatedConnection = { + if (lease.exists(_.isCancelled)) throw ConnectionLost(mayHaveExecuted = true) + if (closing) throw NotConnected() + // reject acquisition while the shared connection is reconnecting; repeat the liveness check after each wake-up + if (!isLive()) throw NotConnected() + val remaining = deadlineNanos - System.nanoTime() + // A lock deadline limits the write itself. An expired write must not take a slot even when one is immediately available. + if (deadlineMillis.isDefined && remaining <= 0L) throw acquireTimedOut(budgetNanos.nanos) + val reused = takeIdleLocked() + if (reused != null) reused + else if (live.size + establishing.size < config.maxConnections) establishOutsideLock() + else if (remaining <= 0L) throw acquireTimedOut(budgetNanos.nanos) + else { available.awaitNanos(remaining): Unit + loop() } - throw new IllegalStateException("unreachable") } + locked(loop()) } - // Entered holding the lock with `reserved` already incremented; registers the connection, drops the lock for the blocking establish, then - // re-accounts under it. + // Entered holding the lock; registers the connection, drops the lock for the blocking establish, then re-accounts under it. private def establishOutsideLock(): DedicatedConnection = { - val connection = DedicatedConnection.create(factory, connectTimeoutMillis) + val connection = new DedicatedConnection(factory, scheduler) establishing += connection lock.unlock() - try connection.establish(bootstrap) - catch { - case e: Throwable => - locked { - establishing -= connection - reserved -= 1 - available.signal() - } - lock.lock() // re-take so acquire()'s `locked` block unlocks exactly once on exit - throw e + onThrow(connection.handshake(bootstrap, connectTimeoutMillis)) { _ => + locked { + establishing -= connection + available.signal() + } + lock.lock() // re-take so acquire()'s `locked` block unlocks exactly once on exit } lock.lock() establishing -= connection - reserved -= 1 - // record the current generation only after the connection is ready to join the pool. This handles reconnects during establishment. - if (closing) { + if (closing || !isLive()) { available.signal() scheduleClose(connection) throw NotConnected() } - liveGeneration() match { - case None => - available.signal() - scheduleClose(connection) - throw NotConnected() - case Some(gen) => - connection.stampEpoch(gen) - live += connection - connection - } + live += connection + connection } private def acquireTimedOut(budget: FiniteDuration): TimedOut = @@ -234,7 +210,7 @@ final private[client] class DedicatedPool( private def release(connection: DedicatedConnection): Unit = locked { - if (closing || !healthyAndCurrent(connection)) discardLocked(connection) + if (closing || !healthy(connection)) discardLocked(connection) else idle.append(DedicatedPool.Idle(connection, scheduler.nowMillis)) available.signal() } @@ -254,12 +230,9 @@ final private[client] class DedicatedPool( toClose.foreach(scheduleClose) // never close on the timer thread: close() joins I/O threads } - // a connection from an older generation may still point to the previous server after a reconnect or DNS failover. - private def healthyAndCurrent(connection: DedicatedConnection): Boolean = - connection.isHealthy && isCurrent(connection.epoch) + private def healthy(connection: DedicatedConnection): Boolean = !connection.isDead && isLive() - private def reusable(entry: DedicatedPool.Idle): Boolean = - healthyAndCurrent(entry.connection) && !expired(entry) + private def reusable(entry: DedicatedPool.Idle): Boolean = healthy(entry.connection) && !expired(entry) private def expired(entry: DedicatedPool.Idle): Boolean = config.idleTimeout.isFinite && scheduler.nowMillis - entry.idleSinceMillis >= config.idleTimeout.toMillis @@ -281,29 +254,37 @@ final private[client] class DedicatedPool( } } -private[client] object DedicatedPool { +/** + * A connection borrowed exclusively from the [[DedicatedPool]]. It does not reconnect or run a watchdog. If the connection is lost, the + * pool discards it and in-flight work fails with `ConnectionLost(mayHaveExecuted = true)`. Replies are matched in order. This allows the + * synchronous `HELLO` setup to finish before the connection is borrowed and keeps a transaction's `MULTI`, queued-command, and `EXEC` + * replies aligned with their commands. + */ +final private[client] class DedicatedConnection(factory: MultiplexedConnection.TransportFactory, scheduler: Scheduler) + extends Pipe(factory, scheduler) { - def forConnection( - factory: MultiplexedConnection.TransportFactory, - bootstrap: Vector[Command[?]], - scheduler: Scheduler, - connection: MultiplexedConnection, - config: DedicatedPoolConfig, - connectTimeoutMillis: Long - ): DedicatedPool = { - val pool = new DedicatedPool( - factory, - bootstrap, - scheduler, - () => connection.isLive, - () => connection.liveGeneration(), - connection.isCurrent, - config, - connectTimeoutMillis - ) - connection.setOnLivenessLost(() => pool.wakeWaiters()) - pool + // The pool discards a dead connection on release, and closing here would fail the rest of a MULTI/EXEC batch the server discards anyway. + override protected def onReadOnly(): Unit = () + + /** + * Sends `MULTI`, `commands` and `EXEC` as a single pipelined write and returns their raw reply frames. The caller decodes the `EXEC` + * array per position against the original commands. + */ + def submitExec(commands: Vector[Command[?]], callback: Try[TxSupport.ExecReplies] => Unit): Unit = { + val queued = (Connection.multi +: commands).map(_.rawFrame) + val replies = new AtomicReferenceArray[Try[Frame]](queued.length) + // replies arrive in write order, so EXEC's reply comes after every queued reply is stored + val onExec: Try[Frame] => Unit = { + case Failure(lost: ConnectionLost) => callback(Failure(lost)) + case exec => callback(Success(TxSupport.ExecReplies(Vector.tabulate(queued.length)(replies.get), exec))) + } + reserve(queued.length + 1) + val entries = Vector.tabulate(queued.length)(i => new Entry[Frame](queued(i), replies.set(i, _))) + sendAll(entries :+ new Entry(Connection.exec.rawFrame, onExec)) } +} + +private[client] object DedicatedPool { final case class Idle(connection: DedicatedConnection, idleSinceMillis: Long) diff --git a/sage-client/shared/src/main/scala/sage/client/internal/Events.scala b/sage-client/shared/src/main/scala/sage/client/internal/Events.scala index 387a3a8e..9c62aaed 100644 --- a/sage-client/shared/src/main/scala/sage/client/internal/Events.scala +++ b/sage-client/shared/src/main/scala/sage/client/internal/Events.scala @@ -125,7 +125,7 @@ private[client] object Events { private val noSpanFactory: () => CommandSpan = () => CommandSpan.noop - // Capture tracing context now and return a function that starts the span later. Cached reads call it only when a cache miss reaches the server. + // Capture tracing context now and return a function that starts the span later. def deferSpan(events: Events, command: Command[?]): () => CommandSpan = events.tracer match { case Some(t) => @@ -138,9 +138,6 @@ private[client] object Events { try factory() catch { case NonFatal(_) => CommandSpan.noop } - def startOrDefer(events: Events, command: Command[?], deferred: () => CommandSpan): CommandSpan = - if (deferred == null) startSpan(events, command) else startDeferred(deferred) - def startSpans(events: Events, commands: Vector[Command[?]]): Vector[CommandSpan] = if (events.tracer.isEmpty) Vector.empty else commands.map(c => startSpan(events, c)) @@ -159,6 +156,37 @@ private[client] object Events { if (events.tracer.isEmpty) callback else new CommandEmit[A](command.name, System.nanoTime(), events, callback, startSpan(events, command), emitsEvent = false) + // Traces a cached read once it is sent to the server or fails, and returns the callback that completes it from then on. + trait Fetching { + def fetching[A](command: Command[?], callback: Try[A] => Unit): Try[A] => Unit + + // Called before the cache lookup. Unless the lookup leads to `fetching`, the read reports only Cache.Hit, even when it fails. + def lookingUp(): Unit = () + } + + val untraced: Fetching = new Fetching { + def fetching[A](command: Command[?], callback: Try[A] => Unit): Try[A] => Unit = callback + } + + // A standalone read is sent at most once, so it is tracked from its send like any command, and a local hit allocates nothing. + def fetchTracking(events: Events): Fetching = + if (!events.enabled) untraced + else + new Fetching { + def fetching[A](command: Command[?], callback: Try[A] => Unit): Try[A] => Unit = trackCommand(events, command, callback) + } + + // Traces a cached read that redirects and retries may send several times. The tracing context is captured on the caller's thread. + inline def trackCached[A](events: Events, command: Command[?], callback: Try[A] => Unit)(inline use: (Try[A] => Unit, Fetching) => Unit): Unit = + if (!events.enabled) use(callback, untraced) + else { + val emit = cachedEmit(events, command, callback) + use(emit, emit) + } + + private def cachedEmit[A](events: Events, command: Command[?], callback: Try[A] => Unit): (Try[A] => Unit) & Fetching = + new CachedEmit[A](command, events, callback, deferSpan(events, command)) + // Record the final routed node before the command completes. Ignore callbacks that do not track command events. def attributeNode(callback: AnyRef, node: Node): Unit = callback match { @@ -166,6 +194,11 @@ private[client] object Events { case _ => () } + def completeAt[A](callback: Try[A] => Unit, node: Node)(result: Try[A]): Unit = { + attributeNode(callback, node) + callback(result) + } + def abandonSpan(callback: AnyRef, error: Throwable): Unit = callback match { case emit: CommandEmit[?] => emit.abandon(error) @@ -186,7 +219,7 @@ private[client] object Events { try span.settled(outcome) catch { case NonFatal(_) => () } - final private class CommandEmit[A]( + private class CommandEmit[A]( name: String, startNanos: Long, events: Events, @@ -195,21 +228,50 @@ private[client] object Events { emitsEvent: Boolean = true ) extends (Try[A] => Unit) { - @volatile private var node: Option[Node] = None + @volatile protected var node: Option[Node] = None - def at(n: Node): Unit = { - node = Some(n) - routeSpan(span, n) - } + protected def current: CommandSpan = span + + // a later attempt replaces the node; the span is routed once, when it settles + def at(n: Node): Unit = node = Some(n) - def abandon(error: Throwable): Unit = settleSpan(span, Outcome.Failed(error)) + def abandon(error: Throwable): Unit = settleSpan(current, Outcome.Failed(error)) def apply(result: Try[A]): Unit = { val outcome = Outcome.of(result) - settleSpan(span, outcome) + node.foreach(routeSpan(current, _)) + settleSpan(current, outcome) if (emitsEvent && events.emitsEvents) events.emit(SageEvent.CommandCompleted(name, node, FiniteDuration(System.nanoTime() - startNanos, NANOSECONDS), outcome)) callback(result) } } + + // The span starts when the read is first sent or fails, so a read served locally reports nothing. Redirects and retries share the span. + final private class CachedEmit[A](command: Command[?], events: Events, callback: Try[A] => Unit, deferred: () => CommandSpan) + extends CommandEmit[A](command.name, System.nanoTime(), events, callback, CommandSpan.noop) + with Fetching { + + @volatile private var started: CommandSpan = null + @volatile private var local = false + + override protected def current: CommandSpan = if (started == null) CommandSpan.noop else started + + // the parts of a cross-slot read can start the span from several threads + private def start(): Unit = if (started == null) synchronized(if (started == null) started = startDeferred(deferred)) + + override def lookingUp(): Unit = local = true + + def fetching[B](command: Command[?], callback: Try[B] => Unit): Try[B] => Unit = { + start() + callback + } + + override def apply(result: Try[A]): Unit = + if (started == null && (result.isSuccess || local)) callback(result) + else { + start() + super.apply(result) + } + } } diff --git a/sage-client/shared/src/main/scala/sage/client/internal/Fault.scala b/sage-client/shared/src/main/scala/sage/client/internal/Fault.scala index c0728511..5a8a050d 100644 --- a/sage-client/shared/src/main/scala/sage/client/internal/Fault.scala +++ b/sage-client/shared/src/main/scala/sage/client/internal/Fault.scala @@ -42,17 +42,20 @@ private[client] object Fault { def categorize(error: Throwable): Fault = error match { - case e: ServerError => - Redirect.parse(e.getMessage) match { - case Some(redirect) => Fault.Redirected(redirect) - case None if e.code == "READONLY" => Fault.Demoted - case None if e.code == "TRYAGAIN" => Fault.TryAgain - case None if e.code == "CLUSTERDOWN" => Fault.Unavailable(clusterWide = true) - case None if e.code == "LOADING" || e.code == "MASTERDOWN" => Fault.Unavailable(clusterWide = false) - case None => Fault.Fatal - } - case NotConnected() => Fault.Lost(mayHaveExecuted = false) - case ConnectionLost(executed) => Fault.Lost(executed) - case _ => Fault.Fatal + case ServerError("MOVED", detail) => redirected(RedirectKind.Moved, detail) + case ServerError("ASK", detail) => redirected(RedirectKind.Ask, detail) + case ServerError("READONLY", _) => Fault.Demoted + case ServerError("TRYAGAIN", _) => Fault.TryAgain + case ServerError("CLUSTERDOWN", _) => Fault.Unavailable(clusterWide = true) + case ServerError("LOADING" | "MASTERDOWN", _) => Fault.Unavailable(clusterWide = false) + case NotConnected() => Fault.Lost(mayHaveExecuted = false) + case ConnectionLost(executed) => Fault.Lost(executed) + case _ => Fault.Fatal + } + + private def redirected(kind: RedirectKind, detail: String): Fault = + Redirect.parse(kind, detail) match { + case Some(redirect) => Fault.Redirected(redirect) + case None => Fault.Fatal } } diff --git a/sage-client/shared/src/main/scala/sage/client/internal/LiveClient.scala b/sage-client/shared/src/main/scala/sage/client/internal/LiveClient.scala new file mode 100644 index 00000000..486276e6 --- /dev/null +++ b/sage-client/shared/src/main/scala/sage/client/internal/LiveClient.scala @@ -0,0 +1,203 @@ +package sage.client.internal + +import scala.annotation.unused +import scala.concurrent.duration.FiniteDuration +import scala.util.Try + +import kyo.compat.* + +import sage.{Message, PatternMessage, SageException} +import sage.SageException.InvalidArgument +import sage.client.{ReadFrom, SageConfig} +import sage.cluster.Node +import sage.codec.ValueCodec +import sage.commands.{Command, Pipeline} +import sage.ratelimit.Decision + +/** + * The operations of the shared `CIO` client that backend adapters run directly. `scanTargets` returns every keyspace SCAN must visit, + * and `lockWrite` waits for the replica acknowledgement a lock requires. + */ +private[sage] trait SharedRunner extends ScanTarget { + + // a single keyspace: the cluster runtime overrides this to scan every master + def scanTargets: CIO[Vector[ScanTarget]] = CIO.value(Vector(this)) + + // a standalone server has no replicas to wait for + def lockWrite(command: Command[Boolean], @unused timeout: FiniteDuration, @unused replicaAcknowledgement: Boolean): CIO[Boolean] = + run(command) +} + +private[sage] object SharedRunner { + + // the default of Client.runner, which no Sage client uses + val unavailable: SharedRunner = new SharedRunner { + def run[A](command: Command[A]): CIO[A] = CIO.fail(InvalidArgument("this operation needs a client created by Sage")) + } +} + +private[internal] trait LiveClient(events: Events) extends Client[CIO, String] with SharedRunner { + + final protected inline def tracked[A](command: Command[A])(inline submit: (Try[A] => Unit) => Unit): CIO[A] = + CIO.async[A] { complete => + val t = Events.trackCommand(events, command, complete) + Client.completing(t)(submit(t)) + } + + // receives only non-empty pipelines without blocking commands + protected def submitPipeline[R](p: Pipeline[R]): CIO[Vector[Either[SageException, Any]]] + + // receives only cacheable commands + protected def cachedChecked[A](command: Command[A], ttl: FiniteDuration): CIO[A] + + protected def pubsub: SubscriptionConnection.PubSub + + protected def openTransaction: CIO[LiveTransactionScope] + + // the first connection or topology discovery + protected def establish(): Unit + + protected def shutdown(): Unit + + // a failed start, including an interrupted one, closes everything the client opened + final private[client] def start(): Unit = onThrow(establish())(_ => shutdown()) + + final def close: CIO[Unit] = CIO.blocking(shutdown()) + + // Release the transaction connection after success, failure, or interruption. Return it to the pool only after EXEC or UNWATCH has + // cleared WATCH/MULTI state and no replies remain pending. Discard it when watched keys or commands may still be active. + final def transaction[A](body: TransactionScope[CIO, String] => CIO[A]): CIO[A] = + CIO.acquireReleaseWith(openTransaction)(scope => CIO.blocking(scope.release()))(scope => CIO.unit.flatMap(_ => body(scope))) + + final def cached[A](command: Command[A], ttl: FiniteDuration): CIO[A] = + if (!Client.cacheable(command)) CIO.fail(Client.notCacheable(command)) else cachedChecked(command, ttl) + + final def subscribeChannels[V: ValueCodec](channel: String, rest: String*): CIO[Subscription[CIO, Message[V]]] = + CIO.blocking(Client.channelMessages(pubsub.subscribeChannels(channel +: rest.toVector))) + + final def subscribePatterns[V: ValueCodec](pattern: String, rest: String*): CIO[Subscription[CIO, PatternMessage[V]]] = + CIO.blocking(Client.patternMessages(pubsub.subscribePatterns(pattern +: rest.toVector))) + + final def subscribeShardChannels[V: ValueCodec](channel: String, rest: String*): CIO[Subscription[CIO, Message[V]]] = + CIO.blocking(Client.channelMessages(pubsub.subscribeShard(channel +: rest.toVector))) + + final override private[sage] def runner: SharedRunner = this + + final private[sage] def rateLimitAcquire[RK](executor: RateLimitExecutor[RK], subject: RK, cost: Long, peek: Boolean): CIO[Decision] = + executor.evalSha(this, subject, cost, peek) + + final private[sage] def lockTryWith[LK, A](executor: LockExecutor[LK], key: LK)(body: => CIO[A]): CIO[Option[A]] = + executor.tryWithLock(this, key)(body) + + final private[sage] def lockWith[LK, A](executor: LockExecutor[LK], key: LK, waitTimeout: FiniteDuration)(body: => CIO[A]): CIO[A] = + executor.withLock(this, key, waitTimeout)(body) + + final private[sage] def pipeline[R](p: Pipeline[R]): CIO[R] = checked(p).flatMap(p.finish(_).fold(CIO.fail(_), CIO.value(_))) + + private def checked[R](p: Pipeline[R]): CIO[Vector[Either[SageException, Any]]] = + if (p.commands.isEmpty) CIO.value(Vector.empty) + else if (p.commands.exists(_.isBlocking)) + CIO.fail(InvalidArgument("a Pipeline cannot carry blocking commands; run them individually on the client")) + else submitPipeline(p) +} + +// The cluster and master-replica runtimes route each command to a node according to its DispatchMode, chosen once per call. +abstract private[internal] class RoutedClient( + nodeFactory: Node => MultiplexedConnection.TransportFactory, + scheduler: Scheduler, + replicaRole: MultiplexedConnection.NodeRole, + config: SageConfig, + minRefreshInterval: FiniteDuration, + pollInterval: Option[FiniteDuration], + events: Events +) extends LiveClient(events) { + import RoutedClient.DispatchMode + import RoutedClient.DispatchMode.* + + protected val masterPool = new NodePool(nodeFactory, scheduler, config, MultiplexedConnection.NodeRole.Master, events) + protected val replicaPool = new NodePool(nodeFactory, scheduler, config, replicaRole, events) + protected val reads = new ReadRouting(masterPool, replicaPool, scheduler, config.readFrom, () => triggerRefresh()) + protected val refreshThrottle = new RefreshThrottle(scheduler, minRefreshInterval.toMillis, () => rediscover()) + // set once by close; routing refuses afterwards, so close is terminal like the standalone client's + @volatile private var isClosed = false + @volatile private var polling = Option.empty[Scheduler.Cancelable] + + final protected def closed: Boolean = isClosed + + protected def route[A](command: Command[A], complete: Try[A] => Unit, lease: DedicatedPool.Lease, mode: DispatchMode): Unit + + // the replicas a lock write on this master waits for + protected def replicaCount(master: Node): Int + + // the first topology discovery, from the seeds + protected def discover(): Either[Throwable, Unit] + + protected def rediscover(): Unit + + final protected def triggerRefresh(): Unit = refreshThrottle.request() + + final protected def establish(): Unit = discover().fold(error => throw error, _ => polling = pollInterval.map(scheduler.every(_)(triggerRefresh()))) + + final protected def shutdown(): Unit = { + refreshThrottle.stop() + polling.foreach(_.cancel()) + isClosed = true + pubsub.close() + masterPool.close() + replicaPool.close() + events.close() + } + + final def run[A](command: Command[A]): CIO[A] = + Client.withLeaseIfBlocking(command)(lease => tracked(command)(route(command, _, lease, readMode(command)))) + + final override def lockWrite(command: Command[Boolean], timeout: FiniteDuration, replicaAcknowledgement: Boolean): CIO[Boolean] = + Client.withLockLease(timeout, scheduler) { (lease, deadlineMillis) => + tracked(command)(route(command, _, lease, Confirmed(deadlineMillis, replicaAcknowledgement))) + } + + final protected def cachedChecked[A](command: Command[A], ttl: FiniteDuration): CIO[A] = + if (!config.clientCache.enabled) tracked(command)(route(command, _, null, MasterOnly)) + else + CIO.async[A](complete => + Events.trackCached(events, command, complete)((t, trace) => Client.completing(t)(route(command, t, null, Cached(ttl.toMillis, trace)))) + ) + + final protected def readMode(command: Command[?]): DispatchMode = + if (config.readFrom != ReadFrom.Master && ReadRouting.replicaEligible(command)) ReplicaRead else MasterOnly + + // a pipeline goes to a replica only when every command is eligible + final protected def pipelineMode(commands: Vector[Command[?]]): DispatchMode = + if (config.readFrom != ReadFrom.Master && commands.forall(ReadRouting.replicaEligible)) ReplicaRead else MasterOnly + + // Sends one attempt to node; onReply receives this attempt's result. + final protected def submitOn[A]( + nc: MultiplexedConnection, + node: Node, + command: Command[A], + lease: DedicatedPool.Lease, + mode: DispatchMode, + asking: Boolean, + onReply: Try[A] => Unit + ): Unit = + mode match { + case Cached(ttlMillis, trace) if !asking => nc.cachedSubmit[A](command, ttlMillis, onReply, trace) + // an ASK attempt bypasses the cache + case Cached(_, trace) => nc.submit[A](command, trace.fetching(command, onReply), asking, lease) + case Confirmed(deadlineMillis, replicaAcknowledgement) => + val replication = + new LockReplication(scheduler, replicaCount(node), deadlineMillis, () => triggerRefresh(), replicaAcknowledgement) + nc.pool.useLockWrite(command, asking, onReply, lease, replication) + case ReplicaRead | MasterOnly => nc.submit[A](command, onReply, asking, lease) + } +} + +private[internal] object RoutedClient { + + enum DispatchMode { + case ReplicaRead, MasterOnly + // trace starts the call's span when a read is sent, whichever callback the attempt completes + case Cached(ttlMillis: Long, trace: Events.Fetching) + case Confirmed(deadlineMillis: Long, replicaAcknowledgement: Boolean) + } +} diff --git a/sage-client/shared/src/main/scala/sage/client/internal/LockCommands.scala b/sage-client/shared/src/main/scala/sage/client/internal/LockCommands.scala deleted file mode 100644 index f6cbe811..00000000 --- a/sage-client/shared/src/main/scala/sage/client/internal/LockCommands.scala +++ /dev/null @@ -1,67 +0,0 @@ -package sage.client.internal - -import scala.concurrent.duration.FiniteDuration - -import sage.Bytes -import sage.SageException.DecodeError -import sage.codec.KeyCodec -import sage.commands.* - -final private[client] class LockCommands[K](leaseDuration: FiniteDuration, namespace: String)(using codec: KeyCodec[K]) { - import LockCommands.Operation - - private val prefix = { - val ns = Bytes.utf8(namespace) - Bytes.concat(Vector(Bytes.utf8(s"${ns.length}:"), ns, Bytes.utf8(":"))) - } - - def key(value: K): Bytes = Bytes.concat(Vector(prefix, codec.encode(value))) - - def command(key: Bytes, token: String, operation: Operation, cached: Boolean): Command[Boolean] = - Command( - if (cached) "EVALSHA" else "EVAL", - Vector(2), - Vector( - if (cached) LockCommands.digest else LockCommands.body, - Bytes.utf8("1"), - key, - Bytes.utf8(token), - Bytes.utf8(operation.wireName), - Bytes.utf8(leaseDuration.toMillis.toString) - ), - _.asLong.flatMap { - case 0 => Right(false) - case 1 => Right(true) - case n => Left(DecodeError("lock result 0 or 1", n.toString)) - } - ) -} - -private[client] object LockCommands { - enum Operation(val wireName: String) { - case Acquire extends Operation("acquire") - case Renew extends Operation("renew") - case Release extends Operation("release") - } - - val script: String = - """local token = ARGV[1] - |local operation = ARGV[2] - |if operation == 'acquire' then - | if redis.call('SET', KEYS[1], token, 'NX', 'PX', ARGV[3]) then return 1 end - | if redis.call('GET', KEYS[1]) == token then return redis.call('PEXPIRE', KEYS[1], ARGV[3]) end - | return 0 - |end - |if operation ~= 'renew' and operation ~= 'release' then - | return redis.error_reply('SAGE invalid lock operation') - |end - |if redis.call('GET', KEYS[1]) ~= token then return 0 end - |if operation == 'renew' then return redis.call('PEXPIRE', KEYS[1], ARGV[3]) end - |return redis.call('DEL', KEYS[1]) - |""".stripMargin - - val body: Bytes = Bytes.utf8(script) - val digest: Bytes = Bytes.utf8( - java.security.MessageDigest.getInstance("SHA-1").digest(body.toArray).iterator.map(b => f"${b & 0xff}%02x").mkString - ) -} diff --git a/sage-client/shared/src/main/scala/sage/client/internal/LockExecutor.scala b/sage-client/shared/src/main/scala/sage/client/internal/LockExecutor.scala index e0dc21c6..87f6de56 100644 --- a/sage-client/shared/src/main/scala/sage/client/internal/LockExecutor.scala +++ b/sage-client/shared/src/main/scala/sage/client/internal/LockExecutor.scala @@ -8,18 +8,19 @@ import scala.concurrent.duration.* import kyo.compat.* import sage.Bytes -import sage.SageException.{InvalidArgument, LockLost, ServerError, TimedOut} +import sage.SageException.{DecodeError, InvalidArgument, LockLost, TimedOut} import sage.client.BackoffConfig import sage.codec.KeyCodec +import sage.commands.* final private[client] class LockExecutor[K]( leaseDuration: FiniteDuration, namespace: String, replicaAcknowledgement: Boolean -)(using KeyCodec[K]) { - import LockCommands.Operation +)(using codec: KeyCodec[K]) { + import LockExecutor.Operation - private val commands = new LockCommands[K](leaseDuration, namespace) + private val namespaced = SingleKeyScript.namespaced(namespace) // The usable lease accounts for millisecond rounding, scheduling, clock drift, and time spent waiting for replies. private val leaseNanos = leaseDuration.toMillis * 1000000L private val usableNanos = leaseNanos - leaseNanos / 10L @@ -32,7 +33,7 @@ final private[client] class LockExecutor[K]( def remainingAt(nowNanos: Long): Long = durationNanos - (nowNanos - startedAtNanos) } - final private class Attempt(val key: Bytes) { + final private class Attempt(val runner: SharedRunner, val key: Bytes) { val token: String = UUID.randomUUID().toString val stopped = new AtomicBoolean(false) val mayOwn = new AtomicBoolean(false) @@ -40,16 +41,37 @@ final private[client] class LockExecutor[K]( val releaseDeadline = new AtomicReference[Deadline]() } - def tryWithLock[A](runner: CommandRunner[CIO, String], key: K)(body: => CIO[A]): CIO[Option[A]] = { + private[internal] def key(value: K): Bytes = namespaced(codec.encode(value)) + + private[internal] def command(key: Bytes, token: String, operation: Operation, cached: Boolean): Command[Boolean] = + Command( + LockExecutor.compiled.verb(cached), + SingleKeyScript.KeyIndices, + Vector( + LockExecutor.compiled.reference(cached), + SingleKeyScript.NumKeys, + key, + Bytes.utf8(token), + Bytes.utf8(operation.wireName), + Bytes.utf8(leaseDuration.toMillis.toString) + ), + _.asLong.flatMap { + case 0 => Right(false) + case 1 => Right(true) + case n => Left(DecodeError("lock result 0 or 1", n.toString)) + } + ) + + def tryWithLock[A](runner: SharedRunner, key: K)(body: => CIO[A]): CIO[Option[A]] = { val bodyThunk = () => body - validate(None).flatMap(_ => attempt(runner, commands.key(key), None)(bodyThunk)) + validate(None).flatMap(_ => attempt(runner, this.key(key), None)(bodyThunk)) } - def withLock[A](runner: CommandRunner[CIO, String], key: K, waitTimeout: FiniteDuration)(body: => CIO[A]): CIO[A] = { + def withLock[A](runner: SharedRunner, key: K, waitTimeout: FiniteDuration)(body: => CIO[A]): CIO[A] = { val bodyThunk = () => body validate(Some(waitTimeout)).flatMap { _ => CIO.nowMonotonic.flatMap { started => - val encoded = commands.key(key) + val encoded = this.key(key) val wait = Deadline(started.toNanos, waitTimeout.toNanos) def loop(retry: Int): CIO[A] = attempt(runner, encoded, Some(wait))(bodyThunk).flatMap { case Some(value) => CIO.value(value) @@ -88,32 +110,28 @@ final private[client] class LockExecutor[K]( } private def eval( - runner: CommandRunner[CIO, String], state: Attempt, operation: Operation, confirmationDeadline: Option[Deadline] = None ): CIO[Boolean] = { def run(cached: Boolean): CIO[Boolean] = { - val command = commands.command(state.key, state.token, operation, cached) + val command = this.command(state.key, state.token, operation, cached) confirmationDeadline match { - case None => runner.run(command) + case None => state.runner.run(command) case Some(deadline) => CIO.nowMonotonic.flatMap { now => val remaining = math.min(operationTimeout.toNanos, deadline.remainingAt(now.toNanos)) if (remaining <= 0L) CIO.fail(TimedOut("distributed lock write timed out")) - else runner.lockWrite(command, remaining.nanos, replicaAcknowledgement) + else state.runner.lockWrite(command, remaining.nanos, replicaAcknowledgement) } } } - run(cached = true).recover { - case ServerError("NOSCRIPT", _) => run(cached = false) - case other => CIO.fail(other) - } + Client.withScriptFallback(run) } - private def attempt[A](runner: CommandRunner[CIO, String], key: Bytes, wait: Option[Deadline])(body: () => CIO[A]): CIO[Option[A]] = + private def attempt[A](runner: SharedRunner, key: Bytes, wait: Option[Deadline])(body: () => CIO[A]): CIO[Option[A]] = // Allocate only local state during acquisition so a contended or unresponsive server never masks cancellation of the wait loop. - CIO.acquireReleaseWith(CIO.defer(new Attempt(key)))(cleanup(runner, _)) { state => + CIO.acquireReleaseWith(CIO.defer(new Attempt(runner, key)))(cleanup) { state => val budget = wait.fold(CIO.value(operationTimeout.toNanos))(remainingWait) budget.flatMap { waitRemaining => CIO.nowMonotonic.flatMap { started => @@ -121,7 +139,7 @@ final private[client] class LockExecutor[K]( val deadline = Deadline(started.toNanos, waitRemaining) CIO .timeoutWithError(waitRemaining.nanos)(TimedOut("distributed lock acquisition timed out"))( - acquire(runner, state, deadline, retry = 0) + acquire(state, deadline, retry = 0) ) .flatMap { case false => @@ -129,28 +147,28 @@ final private[client] class LockExecutor[K]( CIO.value(None) case true => val checkWait = wait.fold(CIO.unit)(remainingWait(_).unit) - checkWait.flatMap(_ => runWithAcquiredLock(runner, state)(body)).map(Some(_)) + checkWait.flatMap(_ => runWithAcquiredLock(state)(body)).map(Some(_)) } } } } - private def runWithAcquiredLock[A](runner: CommandRunner[CIO, String], state: Attempt)(body: () => CIO[A]): CIO[A] = + private def runWithAcquiredLock[A](state: Attempt)(body: () => CIO[A]): CIO[A] = remainingLease(state).flatMap { _ => val work = CIO.defer(()).flatMap(_ => body()).flatMap(value => remainingLease(state).map(_ => value)) // Returning errors as values lets the race stop on either body failure or lock loss. - CIO.race(work.liftToTry, renew[A](runner, state).liftToTry).flatMap(CIO.get(_)).flatMap { value => - release(runner, state).map(_ => value) + CIO.race(work.liftToTry, renew[A](state).liftToTry).flatMap(CIO.get(_)).flatMap { value => + release(state).map(_ => value) } } - private def release(runner: CommandRunner[CIO, String], state: Attempt): CIO[Unit] = { + private def release(state: Attempt): CIO[Unit] = { state.stopped.set(true) remainingLease(state).flatMap { remaining => remainingRelease(state).flatMap { releaseRemaining => CIO .timeoutWithError(math.min(remaining, releaseRemaining).nanos)(LockLost("distributed lock release timed out"))( - eval(runner, state, Operation.Release) + eval(state, Operation.Release) ) .flatMap { released => state.mayOwn.set(false) @@ -162,41 +180,40 @@ final private[client] class LockExecutor[K]( } private def acquire( - runner: CommandRunner[CIO, String], state: Attempt, deadline: Deadline, retry: Int ): CIO[Boolean] = CIO.nowMonotonic.flatMap { started => state.renewedAt.set(started.toNanos) - eval(runner, state, Operation.Acquire, Some(deadline)).recover { + eval(state, Operation.Acquire, Some(deadline)).recover { case error if retryableFailure(error) => CIO.nowMonotonic.flatMap { now => val remaining = deadline.remainingAt(now.toNanos) if (remaining <= 0L) CIO.fail(TimedOut("distributed lock acquisition timed out")) - else retryAfter(remaining, failureBackoff, retry)(acquire(runner, state, deadline, _)) + else retryAfter(remaining, failureBackoff, retry)(acquire(state, deadline, _)) } case error => CIO.fail(error) } } - private def renew[A](runner: CommandRunner[CIO, String], state: Attempt): CIO[A] = + private def renew[A](state: Attempt): CIO[A] = remainingLease(state).flatMap { beforeDelay => // Acquisition confirmation can leave less than the usual renewal delay. Wake by the midpoint of the remaining conservative lease. val delay = math.min(renewalDelay.toNanos, math.max(1L, beforeDelay / 2L)).nanos CIO.sleep(delay).flatMap { _ => if (state.stopped.get()) CIO.never else - renewAttempt(runner, state, retry = 0).flatMap(_ => renew[A](runner, state)) + renewAttempt(state, retry = 0).flatMap(_ => renew[A](state)) } } - private def renewAttempt(runner: CommandRunner[CIO, String], state: Attempt, retry: Int): CIO[Unit] = + private def renewAttempt(state: Attempt, retry: Int): CIO[Unit] = remainingLease(state).flatMap { remaining => CIO.nowMonotonic.flatMap { started => CIO .timeoutWithError(remaining.nanos)(LockLost("distributed lock renewal timed out"))( - eval(runner, state, Operation.Renew, Some(Deadline(started.toNanos, remaining))) + eval(state, Operation.Renew, Some(Deadline(started.toNanos, remaining))) ) .flatMap { renewed => if (!renewed) CIO.fail(LockLost("distributed lock ownership was lost during renewal")) @@ -206,14 +223,13 @@ final private[client] class LockExecutor[K]( } } .recover { - case error if retryableFailure(error) => retryRenewal(runner, state, retry, error) + case error if retryableFailure(error) => retryRenewal(state, retry, error) case error => renewalFailed(error) } } } private def retryRenewal( - runner: CommandRunner[CIO, String], state: Attempt, retry: Int, lastError: Throwable @@ -222,7 +238,7 @@ final private[client] class LockExecutor[K]( retryAfter(remaining, failureBackoff, retry) { nextRetry => CIO.defer(()).flatMap { _ => if (state.stopped.get()) CIO.never - else retryRemainingLease(state, lastError).flatMap(_ => renewAttempt(runner, state, nextRetry)) + else retryRemainingLease(state, lastError).flatMap(_ => renewAttempt(state, nextRetry)) } } } @@ -231,7 +247,7 @@ final private[client] class LockExecutor[K]( remainingLease(state).recover { case _ => renewalFailed(lastError) } private def retryAfter[A](remainingNanos: Long, config: BackoffConfig, retry: Int)(next: Int => CIO[A]): CIO[A] = { - val delay = math.min(remainingNanos, Backoff.jitteredMillis(config, retry, Scheduler.real).millis.toNanos) + val delay = math.min(remainingNanos, Scheduler.real.backoffMillis(config, retry).millis.toNanos) CIO.sleep(delay.nanos).flatMap(_ => next(if (retry < Int.MaxValue) retry + 1 else retry)) } @@ -251,14 +267,40 @@ final private[client] class LockExecutor[K]( CIO.fail(lost) } - private def cleanup(runner: CommandRunner[CIO, String], state: Attempt): CIO[Unit] = CIO.defer(()).flatMap { _ => + private def cleanup(state: Attempt): CIO[Unit] = CIO.defer(()).flatMap { _ => // Future/Pekko cannot cancel the losing race branch. This flag also stops their renewal loop after the scope has finished. state.stopped.set(true) if (!state.mayOwn.get()) CIO.unit else remainingRelease(state).flatMap { remaining => if (remaining <= 0) CIO.unit - else CIO.timeout(remaining.nanos)(eval(runner, state, Operation.Release)).unit.recover(_ => CIO.unit) + else CIO.timeout(remaining.nanos)(eval(state, Operation.Release)).unit.recover(_ => CIO.unit) } } } + +private[client] object LockExecutor { + enum Operation(val wireName: String) { + case Acquire extends Operation("acquire") + case Renew extends Operation("renew") + case Release extends Operation("release") + } + + val script: String = + """local token = ARGV[1] + |local operation = ARGV[2] + |if operation == 'acquire' then + | if redis.call('SET', KEYS[1], token, 'NX', 'PX', ARGV[3]) then return 1 end + | if redis.call('GET', KEYS[1]) == token then return redis.call('PEXPIRE', KEYS[1], ARGV[3]) end + | return 0 + |end + |if operation ~= 'renew' and operation ~= 'release' then + | return redis.error_reply('SAGE invalid lock operation') + |end + |if redis.call('GET', KEYS[1]) ~= token then return 0 end + |if operation == 'renew' then return redis.call('PEXPIRE', KEYS[1], ARGV[3]) end + |return redis.call('DEL', KEYS[1]) + |""".stripMargin + + val compiled = SingleKeyScript(script) +} diff --git a/sage-client/shared/src/main/scala/sage/client/internal/LockReplication.scala b/sage-client/shared/src/main/scala/sage/client/internal/LockReplication.scala index cfbcb7c3..dc3052ed 100644 --- a/sage-client/shared/src/main/scala/sage/client/internal/LockReplication.scala +++ b/sage-client/shared/src/main/scala/sage/client/internal/LockReplication.scala @@ -1,12 +1,10 @@ package sage.client.internal -import java.util.concurrent.atomic.AtomicBoolean - import scala.concurrent.duration.* import scala.util.{Failure, Success, Try} import sage.SageException.{ConnectionLost, LockLost, TimedOut} -import sage.commands.{Command, Connection, Reply, Role, Server} +import sage.commands.{Command, Role, Server} /** * Confirms a lock write on the connection that executed it. Each instance handles one write and its replication check. @@ -18,11 +16,7 @@ final private[client] class LockReplication( onConfirmationFailure: () => Unit, replicaAcknowledgement: Boolean ) { - private val confirming = new AtomicBoolean(false) - - def cancelled(): Unit = if (confirming.get()) onConfirmationFailure() - - def submit[A](conn: DedicatedConnection, command: Command[A], asking: Boolean, complete: Try[A] => Unit): Unit = { + def submit[A](conn: DedicatedConnection, command: Command[A], complete: Try[A] => Unit): Unit = { // Lock acquisition retries are safe because they reuse the same ownership token. def confirmationFailed(error: Throwable): Unit = { onConfirmationFailure() @@ -36,8 +30,7 @@ final private[client] class LockReplication( complete(Failure(failure)) } - def confirm(value: A): Unit = { - confirming.set(true) + def confirm(value: A): Unit = conn.submit( Server.role, { @@ -48,15 +41,12 @@ final private[client] class LockReplication( case Failure(error) => confirmationFailed(error) } ) - } val onReply: Try[A] => Unit = { case Success(value) if value == true => confirm(value) case result => complete(result) } - if (asking) - conn.submitRaw(Vector(Connection.asking, command.rawFrame), result => onReply(result.flatMap(frames => Reply.decode(command, frames.last)))) - else conn.submit(command, onReply) + conn.submit(command, onReply) } private def waitForReplicas[A]( diff --git a/sage-client/shared/src/main/scala/sage/client/internal/LoweredClient.scala b/sage-client/shared/src/main/scala/sage/client/internal/LoweredClient.scala index c0ce21d2..5819bb2d 100644 --- a/sage-client/shared/src/main/scala/sage/client/internal/LoweredClient.scala +++ b/sage-client/shared/src/main/scala/sage/client/internal/LoweredClient.scala @@ -31,15 +31,11 @@ abstract class LoweredClient[F[_]](underlying: Client[CIO, String]) extends Clie timeout: FiniteDuration, replicaAcknowledgement: Boolean ): CIO[Boolean] = - underlying.lockWrite(command, timeout, replicaAcknowledgement) - - private val lockRunner: CommandRunner[CIO, String] = new CommandRunner[CIO, String] { - def run[A](command: Command[A]): CIO[A] = lockCommand(command) - override private[sage] def lockWrite( - command: Command[Boolean], - timeout: FiniteDuration, - replicaAcknowledgement: Boolean - ): CIO[Boolean] = + underlying.runner.lockWrite(command, timeout, replicaAcknowledgement) + + private val lockRunner: SharedRunner = new SharedRunner { + def run[A](command: Command[A]): CIO[A] = lockCommand(command) + override def lockWrite(command: Command[Boolean], timeout: FiniteDuration, replicaAcknowledgement: Boolean): CIO[Boolean] = confirmedLockCommand(command, timeout, replicaAcknowledgement) } @@ -47,9 +43,10 @@ abstract class LoweredClient[F[_]](underlying: Client[CIO, String]) extends Clie final def cached[A](command: Command[A], ttl: FiniteDuration): F[A] = lower(underlying.cached(command, ttl)) - final private[sage] def pipeline[Out, R](p: Pipeline[Out, R]): F[Out] = lower(underlying.pipeline(p)) + final private[sage] def pipeline[R](p: Pipeline[R]): F[R] = lower(underlying.pipeline(p)) - final private[sage] def pipelineAttempt[Out, R](p: Pipeline[Out, R]): F[R] = lower(underlying.pipelineAttempt(p)) + // keeps the method signature of 0.4.0 for MiMa; `pipeline` handles both result shapes + final private[sage] def pipelineAttempt[R](p: Pipeline[R]): F[R] = pipeline(p) final def transaction[A](body: TransactionScope[F, String] => F[A]): F[A] = lower(underlying.transaction[A](scope => lift(body(lowerScope(scope))))) @@ -63,9 +60,7 @@ abstract class LoweredClient[F[_]](underlying: Client[CIO, String]) extends Clie final def subscribeShardChannels[V: ValueCodec](channel: String, rest: String*): F[Subscription[F, Message[V]]] = lower(underlying.subscribeShardChannels[V](channel, rest*).map(lowerSub)) - final private[sage] def scanTargets: F[Vector[ScanTarget]] = lower(underlying.scanTargets) - - final private[sage] def runOn[A](target: ScanTarget, command: Command[A]): F[A] = lower(underlying.runOn(target, command)) + final override private[sage] def runner: SharedRunner = underlying.runner final private[sage] def rateLimitAcquire[RK](executor: RateLimitExecutor[RK], subject: RK, cost: Long, peek: Boolean): F[Decision] = lower(underlying.rateLimitAcquire(executor, subject, cost, peek)) @@ -84,11 +79,10 @@ abstract class LoweredClient[F[_]](underlying: Client[CIO, String]) extends Clie private def lowerScope(scope: TransactionScope[CIO, String]): TransactionScope[F, String] = new TransactionScope[F, String] { - def watch[K: KeyCodec](key: K, rest: K*): F[Unit] = lower(scope.watch(key, rest*)) - def run[A](command: Command[A]): F[A] = lower(scope.run(command)) - private[sage] def exec[Out, R](p: Pipeline[Out, R]): F[Option[Out]] = lower(scope.exec(p)) - private[sage] def execAttempt[Out, R](p: Pipeline[Out, R]): F[Option[R]] = lower(scope.execAttempt(p)) - def discard: F[Unit] = lower(scope.discard) + def watch[K: KeyCodec](key: K, rest: K*): F[Unit] = lower(scope.watch(key, rest*)) + def run[A](command: Command[A]): F[A] = lower(scope.run(command)) + private[sage] def exec[R](p: Pipeline[R]): F[Option[R]] = lower(scope.exec(p)) + def discard: F[Unit] = lower(scope.discard) } private def lowerSub[A](sub: Subscription[CIO, A]): Subscription[F, A] = diff --git a/sage-client/shared/src/main/scala/sage/client/internal/MasterReplicaLive.scala b/sage-client/shared/src/main/scala/sage/client/internal/MasterReplicaLive.scala index 457215e0..d089404a 100644 --- a/sage-client/shared/src/main/scala/sage/client/internal/MasterReplicaLive.scala +++ b/sage-client/shared/src/main/scala/sage/client/internal/MasterReplicaLive.scala @@ -1,26 +1,24 @@ package sage.client.internal import java.util.concurrent.atomic.AtomicReference -import java.util.concurrent.locks.ReentrantLock -import scala.concurrent.duration.* +import scala.annotation.tailrec import scala.util.{Failure, Success, Try} import scala.util.control.NonFatal +import RoutedClient.DispatchMode import kyo.compat.* -import sage.{CommandSpan, Message, Outcome, PatternMessage, SageEvent, SageException} -import sage.SageException.{ConnectionFailed, ConnectionLost, InvalidArgument, NotConnected, TimedOut} -import sage.client.{MasterReplicaConfig, ReadFrom, SageConfig} +import sage.{SageEvent, SageException} +import sage.SageException.{ConnectionFailed, ConnectionLost, NotConnected, TimedOut} +import sage.client.{MasterReplicaConfig, SageConfig} import sage.cluster.Node -import sage.codec.ValueCodec -import sage.commands.{Command, Connection, Pipeline, Role, Server} -import sage.ratelimit.Decision +import sage.commands.{Command, Pipeline, Role, Server} /** * The runtime for a non-cluster deployment with one master and its replicas. It discovers their roles by sending `ROLE` to the seed nodes. * Writes, blocking reads, transactions, and `cached` reads go to the master. Other read-only commands use replicas according to the - * [[ReadFrom]] policy, including its fallback behavior. Standalone, master-replica, and cluster deployments use the same `Client` type; the + * [[sage.client.ReadFrom]] policy, including its fallback behavior. Standalone, master-replica, and cluster deployments use the same `Client` type; the * configured topology chooses the runtime. * * The runtime refreshes roles after a command is lost during reconnection, a presumed master returns `READONLY`, a read cannot reach any @@ -31,48 +29,36 @@ import sage.ratelimit.Decision final private[client] class MasterReplicaLive( nodeFactory: Node => MultiplexedConnection.TransportFactory, scheduler: Scheduler, - bootstrap: Vector[Command[?]], config: SageConfig, seeds: Vector[Node], masterReplica: MasterReplicaConfig, events: Events = Events.disabled -) extends Client[CIO, String] { - - private val readFrom = config.readFrom - private val cachingEnabled = config.clientCache.enabled - - // only the master multiplexed connection caches: cached reads run on the master, replicas and dedicated connections never serve them - private val masterPool = pool(caching = true) - private val replicaPool = pool(caching = false) - - private def pool(caching: Boolean): NodePool = { - val cached = caching && cachingEnabled - val poolBootstrap = if (cached) bootstrap :+ Connection.clientTrackingOnOptin else bootstrap - val cacheMaxBytes = if (cached) config.clientCache.maxBytes else 0L - new NodePool( - nodeFactory, - scheduler, - poolBootstrap, - config.reconnect, - config.watchdog, - config.connectTimeout, - config.closeTimeout, - config.dedicatedPool, - cacheMaxBytes, - events, - dedicatedBootstrap = Some(bootstrap) - ) - } - - private val masterNodeRef = new AtomicReference[Node](null) - private val replicasRef = new AtomicReference[Vector[Node]](Vector.empty) - private val reads = new ReadRouting(masterPool, replicaPool, scheduler, readFrom, () => triggerRefresh()) - @volatile private var closed = false +) extends RoutedClient( + nodeFactory, + scheduler, + MultiplexedConnection.NodeRole.Replica, + config, + masterReplica.minRefreshInterval, + masterReplica.topologyRefreshInterval, + events + ) { + + // null until discover installs the first topology + private val topologyRef = new AtomicReference[MasterReplicaLive.ResolvedTopology](null) + + // resolve the master for every connection attempt so subscriptions move to the promoted master after failover. + private val subscribedOn = new SubscriptionConnection.Following(nodeFactory, () => Option(topologyRef.get()).map(_.master)) - private val subLock = new ReentrantLock() - @volatile private var subscriptions: SubscriptionConnection = null - - private val refreshThrottle = new RefreshThrottle(scheduler, masterReplica.minRefreshInterval.toMillis) + // master-replica mode uses one subscription connection for all shard channels, as standalone mode does. + private val subscriptions = new SubscriptionConnection( + subscribedOn.factory, + scheduler, + config, + // subscriptions use a separate socket. Wait for master discovery before opening it; a pooled connection is not required. + () => !closed && topologyRef.get() != null, + // request an immediate refresh. If another refresh is active, wait for it; a later reconnect requests discovery again if the master changed. + SubscriptionConnection.OnLoss.Reconnect(() => refreshThrottle(force = true), events, () => subscribedOn.node) + ) // --- discovery ----------------------------------------------------------------------------------------------------------------------- @@ -80,15 +66,7 @@ final private[client] class MasterReplicaLive( // addresses returned by ROLE. private val pinnedToSeeds = seeds.sizeIs > 1 - private[client] def bootstrapRoles(): Unit = - resolveTopology(seeds) match { - case Right(topology) => - installTopology(topology) - startRefreshPoll() - case Left(error) => - closeAll() - throw error - } + protected def discover(): Either[Throwable, Unit] = resolveTopology(seeds).map(installTopology) private def resolveTopology(discoveredCandidates: => Vector[Node]): Either[Throwable, MasterReplicaLive.ResolvedTopology] = if (pinnedToSeeds) resolvePinned() @@ -96,250 +74,124 @@ final private[client] class MasterReplicaLive( // request ROLE from every supplied endpoint. Omit endpoints that cannot be reached, but report their connection failures through events. private def resolvePinned(): Either[Throwable, MasterReplicaLive.ResolvedTopology] = { - var lastError: Throwable = NotConnected() - val roles = seeds.flatMap { seed => - try probeRole(seed).map(seed -> _) - catch { - case NonFatal(error) => - lastError = error - None - } - } + val probed = seeds.map(seed => seed -> probeRole(seed)) + val roles = probed.collect { case (node, Success(role)) => node -> role } roles.collectFirst { case (node, _: Role.Master) => node } match { case Some(master) => Right(MasterReplicaLive.ResolvedTopology(master, roles.collect { case (node, role) if role.isConnectedReplica => node })) - case None if roles.isEmpty => Left(lastError) + case None if roles.isEmpty => Left(probed.collect { case (_, Failure(error)) => error }.lastOption.getOrElse(NotConnected())) case None => Left(ConnectionFailed("no supplied endpoint reports the master role")) } } // contact candidates until one answers ROLE, then use its advertised master and replica addresses - private def resolveDiscovered(candidates: Vector[Node]): Either[Throwable, MasterReplicaLive.ResolvedTopology] = { - var lastError: Throwable = NotConnected() - val it = candidates.iterator - while (it.hasNext) { - val seed = it.next() - try + @tailrec private def resolveDiscovered( + candidates: Vector[Node], + lastError: Throwable = NotConnected() + ): Either[Throwable, MasterReplicaLive.ResolvedTopology] = + candidates match { + case seed +: rest => resolveFrom(seed) match { - case Some(topology) => return Right(topology) - case None => () + case Success(Some(topology)) => Right(topology) + case Success(None) => resolveDiscovered(rest, lastError) + case Failure(error) => resolveDiscovered(rest, error) } - catch { case NonFatal(error) => lastError = error } + case _ => Left(lastError) } - Left(lastError) - } // probes a node's ROLE; a master answers with its replica list, a replica points at its master (followed once), a sentinel is skipped - private def resolveFrom(node: Node): Option[MasterReplicaLive.ResolvedTopology] = + private def resolveFrom(node: Node): Try[Option[MasterReplicaLive.ResolvedTopology]] = probeRole(node).flatMap { - case Role.Master(_, replicas) => Some(MasterReplicaLive.ResolvedTopology(node, replicas.map(r => Node(r.host, r.port)))) + case Role.Master(_, replicas) => Success(Some(MasterReplicaLive.ResolvedTopology(node, replicas.map(r => Node(r.host, r.port))))) case Role.Replica(host, port, _, _) => val master = Node(host, port) - probeRole(master).collect { case Role.Master(_, replicas) => - MasterReplicaLive.ResolvedTopology(master, replicas.map(r => Node(r.host, r.port))) + probeRole(master).map { + case Role.Master(_, replicas) => Some(MasterReplicaLive.ResolvedTopology(master, replicas.map(r => Node(r.host, r.port)))) + case _ => None } - case _: Role.Sentinel => None + case _: Role.Sentinel => Success(None) } // use an existing live connection for ROLE when possible. Otherwise, open a temporary connection and close it after the probe. - private def probeRole(node: Node): Option[Role] = { - val pooled = pooledFor(node) - if (pooled != null) { - val reply = askRole(pooled) - if (!lostConnection(reply)) return interpretRole(node, reply) + private def probeRole(node: Node): Try[Role] = { + val pooled = pooledFor(node).map(askRole).filterNot(lostConnection) + // a refresh can outlive the start of close, and a closed client must not open a new socket + if (pooled.isEmpty && closed) Failure(NotConnected()) + else { + val reply = pooled.getOrElse( + Try(new MultiplexedConnection(nodeFactory(node), scheduler, config, MultiplexedConnection.NodeRole.Replica, Some(node)).start()) + .flatMap(nc => + try askRole(nc) + finally nc.close() + ) + ) + reply.failed.foreach(reportProbeFailure(node, _)) + reply } - val nc = connectForProbe(node) - try interpretRole(node, askRole(nc)) - finally nc.close() - } - - private def pooledFor(node: Node): NodeClient = { - val master = masterPool.existing(node) - val nc = if (master != null) master else replicaPool.existing(node) - if (nc != null && nc.isLive) nc else null } - private def askRole(nc: NodeClient): Option[Try[Role]] = - Bootstrap.awaitReply[Role](config.connectTimeout.toMillis)(callback => nc.submit(Server.role, asking = false, callback)) - - private def interpretRole(node: Node, reply: Option[Try[Role]]): Option[Role] = - reply match { - case Some(Success(role)) => Some(role) - case Some(Failure(error)) => - reportProbeFailure(node, error) - None - case None => - reportProbeFailure(node, TimedOut(s"ROLE timed out after ${config.connectTimeout.toMillis}ms")) - None - } + private def pooledFor(node: Node): Option[MultiplexedConnection] = + Option(masterPool.existing(node)).orElse(Option(replicaPool.existing(node))).filter(_.isLive) - private def lostConnection(reply: Option[Try[Role]]): Boolean = - reply match { - case Some(Failure(error)) => - Fault.categorize(error) match { - case Fault.Lost(_) => true - case _ => false - } - case _ => false - } + private def askRole(nc: MultiplexedConnection): Try[Role] = + Bootstrap.awaitReply[Role](config.connectTimeout.toMillis, TimedOut(s"ROLE timed out after ${config.connectTimeout.toMillis}ms"))( + nc.submit(Server.role, _) + ) - private def connectForProbe(node: Node): NodeClient = { - // a refresh can outlive the start of close, and a closed client must not open a new socket - if (closed) throw NotConnected() - try - NodeClient.connect( - nodeFactory(node), - scheduler, - bootstrap, - config.reconnect, - config.watchdog, - config.connectTimeout, - config.closeTimeout, - config.dedicatedPool, - node = node, - events = Events.disabled - ) - catch { - case NonFatal(error) => - reportProbeFailure(node, error) - throw error + private def lostConnection(reply: Try[Role]): Boolean = + reply.failed.toOption.map(Fault.categorize).exists { + case Fault.Lost(_) => true + case _ => false } - } private def reportProbeFailure(node: Node, error: Throwable): Unit = events.emit(SageEvent.Connection.ConnectFailed(Some(node), error)) - private val rediscoverWork: () => Unit = () => rediscover() - - private def triggerRefresh(): Unit = refreshThrottle.trigger(rediscoverWork) - - private def startRefreshPoll(): Unit = refreshThrottle.startPolling(masterReplica.topologyRefreshInterval)(triggerRefresh()) - - // request an immediate refresh. If another refresh is active, wait for it; a later reconnect requests discovery again if the master changed. - private def refreshRolesBeforeRehome(): Unit = refreshThrottle(force = true)(rediscover()) - - // a re-discovery queued before close must not probe ROLE on a connection the close cannot reach - private def rediscover(): Unit = - if (!closed) resolveTopology((Option(masterNodeRef.get()).toVector ++ replicasRef.get() ++ seeds).distinct).foreach(installTopology) + protected def rediscover(): Unit = + resolveTopology((Option(topologyRef.get()).toVector.flatMap(t => t.master +: t.replicas) ++ seeds).distinct).foreach(installTopology) private def installTopology(topology: MasterReplicaLive.ResolvedTopology): Unit = { - masterNodeRef.set(topology.master) - replicasRef.set(topology.replicas) + topologyRef.set(topology) replicaPool.retain(topology.replicas.toSet.contains) masterPool.retain(_ == topology.master) reads.retain(_ == topology.master) + // a demoted master still receives PUBLISH through replication, so only a node outside the topology loses the subscription + subscribedOn.retain(node => node == topology.master || topology.replicas.contains(node)) } // --- routing ------------------------------------------------------------------------------------------------------------------------- - def run[A](command: Command[A]): CIO[A] = { - def body(lease: DedicatedPool.Lease): CIO[A] = - CIO.async[A] { complete => - val tracked = Events.trackCommand(events, command, complete) - Client.completing(tracked) { - if (readFrom != ReadFrom.Master && ReadRouting.replicaEligible(command)) sendRead(command, tracked) - else sendMaster(command, tracked, lease) - } - } - Client.withLeaseIfBlocking(command)(body) - } - - override private[sage] def lockWrite( - command: Command[Boolean], - timeout: FiniteDuration, - replicaAcknowledgement: Boolean - ): CIO[Boolean] = - Client.withLockLease(timeout, scheduler) { (lease, deadlineMillis) => - CIO.async { complete => - val tracked = Events.trackCommand(events, command, complete) - Client.completing(tracked) { - onMaster(tracked) { (nc, _, cb) => - val replication = new LockReplication( - scheduler, - replicasRef.get().size, - deadlineMillis, - () => refreshThrottle.request(rediscoverWork), - replicaAcknowledgement - ) - nc.submitLockWrite( - command, - asking = false, - cb, - lease, - replication - ) + // Submit to the master and add its node to the result. Start role discovery if the server is no longer the master. + protected def route[A](command: Command[A], complete: Try[A] => Unit, lease: DedicatedPool.Lease, mode: DispatchMode): Unit = + if (closed) complete(Failure(NotConnected())) + else if (mode == DispatchMode.ReplicaRead) { + val topology = topologyRef.get() + walkRead(command, reads.candidatesFor(topology.master, topology.replicas), topology.master, complete) + } else { + val node = topologyRef.get().master + masterPool.withClient(node) { + triggerRefresh() + complete(Failure(NotConnected())) + } { nc => + submitOn( + nc, + node, + command, + lease, + mode, + asking = false, + result => { + result match { + case Failure(e) if isOwnershipFault(e) => triggerRefresh() + case _ => () + } + Events.completeAt(complete, node)(result) } - } - } - } - - def cached[A](command: Command[A], ttl: FiniteDuration): CIO[A] = - if (!Client.cacheable(command)) CIO.fail(Client.notCacheable(command)) - else if (!cachingEnabled) - CIO.async[A] { complete => - val tracked = Events.trackCommand(events, command, complete) - Client.completing(tracked)(sendMaster(command, tracked)) - } - else - CIO.async[A] { complete => - val deferred = Events.deferSpan(events, command) - Client.completing(complete)(sendMasterCached(command, ttl.toMillis, complete, deferred)) + ) } - - private def sendMaster[A](command: Command[A], complete: Try[A] => Unit, lease: DedicatedPool.Lease = null): Unit = - onMaster(complete)((nc, _, cb) => nc.submit[A](command, asking = false, cb, lease)) - - private def sendMasterCached[A](command: Command[A], ttlMillis: Long, complete: Try[A] => Unit, deferred: () => CommandSpan): Unit = - // if no master is available, cachedSubmit is not called. Complete its deferred span here. - onMaster(complete, onDown = () => Events.settleSpan(Events.startDeferred(deferred), Outcome.Failed(NotConnected()))) { (nc, _, cb) => - nc.cachedSubmit[A](command, ttlMillis, cb, deferred) } - // Submit on the master and add its node to the result. Start role discovery if the server is no longer the master. Call onDown when the - // operation ends before submit is called. - private def onMaster[A](complete: Try[A] => Unit, onDown: () => Unit = () => ())(submit: (NodeClient, Node, Try[A] => Unit) => Unit): Unit = { - if (closed) { - onDown() - complete(Failure(NotConnected())) - return - } - val node = masterNodeRef.get() - val existing = masterPool.existing(node) - if (existing != null) submitMaster(existing, node, complete, submit) - else - scheduler.offload { - val nc = masterPool.getOrEstablishOrNull(node) - if (nc == null) { - triggerRefresh() - onDown() - complete(Failure(NotConnected())) - } else submitMaster(nc, node, complete, submit) - } - } - - private def submitMaster[A](nc: NodeClient, node: Node, complete: Try[A] => Unit, submit: (NodeClient, Node, Try[A] => Unit) => Unit): Unit = - submit( - nc, - node, - { - case s @ Success(_) => - Events.attributeNode(complete, node) - complete(s) - case f @ Failure(e) => - if (isOwnershipFault(e)) triggerRefresh() - Events.attributeNode(complete, node) - complete(f) - } - ) - - private def sendRead[A](command: Command[A], complete: Try[A] => Unit): Unit = { - if (closed) { - complete(Failure(NotConnected())) - return - } - val master = masterNodeRef.get() - walkRead(command, reads.candidatesFor(master, replicasRef.get()), master, complete) - } + protected def replicaCount(master: Node): Int = topologyRef.get().replicas.size private def walkRead[A](command: Command[A], candidates: Vector[Node], master: Node, complete: Try[A] => Unit): Unit = reads.walk(command, candidates, master, complete)((node, error, rest) => onReadFault(ReadRoute(node, master, rest), error, command, complete)) @@ -350,39 +202,27 @@ final private[client] class MasterReplicaLive( command: Command[A], complete: Try[A] => Unit ): Unit = - handleReadFaults(route, Vector(error), RetryExecution.Inline)( + handleReadFaults(route, Vector(error))( remaining => walkRead(command, remaining, route.master, complete), - () => { - Events.attributeNode(complete, route.node) - complete(Failure(error)) - } + () => Events.completeAt(complete, route.node)(Failure(error)) ) final private case class ReadRoute(node: Node, master: Node, remaining: Vector[Node]) - private enum RetryExecution { - case Inline, Offloaded - } - - private def handleReadFaults(route: ReadRoute, errors: Vector[Throwable], retryExecution: RetryExecution)( + private def handleReadFaults(route: ReadRoute, errors: Vector[Throwable])( retry: Vector[Node] => Unit, settle: () => Unit ): Unit = { val ownershipFault = route.node == route.master && errors.exists(isOwnershipFault) if (ownershipFault) triggerRefresh() - if (errors.exists(servesNoRead)) { - def continue(): Unit = - if (route.remaining.nonEmpty) retry(route.remaining) - else { - // an ownership fault already requested the same throttled refresh above - if (!ownershipFault) triggerRefresh() - settle() - } - retryExecution match { - case RetryExecution.Inline => continue() - case RetryExecution.Offloaded => scheduler.offload(continue()) + if (errors.exists(servesNoRead)) + if (route.remaining.nonEmpty) retry(route.remaining) + else { + // an ownership fault already requested the same throttled refresh above + if (!ownershipFault) triggerRefresh() + settle() } - } else settle() + else settle() } private def isOwnershipFault(error: Throwable): Boolean = Fault.categorize(error) match { @@ -398,92 +238,50 @@ final private[client] class MasterReplicaLive( // --- pipelines ----------------------------------------------------------------------------------------------------------------------- - private[sage] def pipeline[Out, R](p: Pipeline[Out, R]): CIO[Out] = submitPipeline(p).flatMap(TxSupport.collapseStrict(_, p.toOut)) - private[sage] def pipelineAttempt[Out, R](p: Pipeline[Out, R]): CIO[R] = submitPipeline(p).map(p.toResults) - - private def submitPipeline[Out, R](p: Pipeline[Out, R]): CIO[Vector[Either[SageException, Any]]] = - if (p.commands.isEmpty) CIO.value(Vector.empty) - else if (p.commands.exists(_.isBlocking)) - CIO.fail(InvalidArgument("a Pipeline cannot carry blocking commands; run them individually on the client")) - else - CIO.async { complete => - val spans = Events.startSpans(events, p.commands) - // route the whole pipeline to a replica only when every command is eligible. Otherwise, route the whole pipeline to the master. - val useReplica = readFrom != ReadFrom.Master && p.commands.forall(ReadRouting.replicaEligible) - val master = masterNodeRef.get() - val refreshOnUnsent = () => triggerRefresh() - if (useReplica) { - val batch = new Client.TrackedBatch(events, p.commands, spans, complete) - def failUnsent(): Unit = - // without a submission, a connection error cannot trigger role discovery. Refresh roles before failing the batch. - batch.failUnsent(refreshOnUnsent) - def submitOn(picked: Option[ReadRouting.Picked]): Unit = - picked match { - case Some(ReadRouting.Picked(node, nc, rest)) => - val route = ReadRoute(node, master, rest) - val attempt = new TxSupport.IndexedCollector[Try[Any]]( - p.commands.length, - results => - handleReadFaults(route, results.collect { case Failure(error) => error }, RetryExecution.Offloaded)( - remaining => reads.pickOne(remaining, master)(submitOn), - () => batch.settleAll(node, results) - ) - ) - val callbacks = Vector.tabulate(p.commands.length)(i => (result: Try[Any]) => attempt.set(i, result)) - // the selected connection died before reserving the batch; retry the whole batch on the remaining candidates - if (!nc.submitAll(p.commands, callbacks)) - if (rest.nonEmpty) reads.pickOne(rest, master)(submitOn) - else failUnsent() - case None => failUnsent() - } - reads.pickOne(reads.candidatesFor(master, replicasRef.get()), master)(submitOn) - } else { - def submitOn(picked: Option[(Node, NodeClient)]): Unit = { - val submit = picked match { - case Some((_, nc)) => nc.submitAll - case None => (_: Vector[Command[?]], _: Vector[Try[Any] => Unit]) => false - } - Client.submitBatchOnOne( - events, - p.commands, - spans, - submit, - complete, - onUnsent = refreshOnUnsent, - node = picked.map(_._1) - ) - } - val existing = masterPool.existing(master) - if (existing != null) submitOn(Some((master, existing))) - else - scheduler.offload { - val nc = masterPool.getOrEstablishOrNull(master) - submitOn(Option(nc).map(master -> _)) - } - } + protected def submitPipeline[R](p: Pipeline[R]): CIO[Vector[Either[SageException, Any]]] = + CIO.async { complete => + val spans = Events.startSpans(events, p.commands) + val topology = topologyRef.get() + val master = topology.master + val batch = new Client.TrackedBatch(events, p.commands, spans, complete) + // without a submission, a connection error cannot trigger role discovery. Refresh roles before failing the batch. + def failUnsent(): Unit = { + triggerRefresh() + batch.failUnsent() } + if (pipelineMode(p.commands) == DispatchMode.ReplicaRead) { + def submitOn(picked: Option[ReadRouting.Picked]): Unit = + picked match { + case Some(ReadRouting.Picked(node, nc, rest)) => + val route = ReadRoute(node, master, rest) + val attempt = new TxSupport.IndexedCollector[Try[Any]]( + p.commands.length, + results => + handleReadFaults(route, results.collect { case Failure(error) => error })( + remaining => scheduler.offload(reads.pickOne(remaining, master)(submitOn)), + () => batch.settleAll(node, results) + ) + ) + val callbacks = Vector.tabulate(p.commands.length)(i => (result: Try[Any]) => attempt.set(i, result)) + // the selected connection died before reserving the batch; retry the whole batch on the remaining candidates + if (!nc.submitAll(p.commands, callbacks)) + if (rest.nonEmpty) reads.pickOne(rest, master)(submitOn) + else failUnsent() + case None => failUnsent() + } + reads.pickOne(reads.candidatesFor(master, topology.replicas), master)(submitOn) + } else + masterPool.withClient(master)(failUnsent())(nc => if (!nc.submitAll(p.commands, batch.callbacks(Some(master)))) failUnsent()) + } // --- transactions (always on the master) --------------------------------------------------------------------------------------------- - def transaction[A](body: TransactionScope[CIO, String] => CIO[A]): CIO[A] = - CIO.acquireReleaseWith(acquireScope)(releaseScope)(lease => CIO.unit.flatMap(_ => body(lease.scope))) - - private def refreshOnTxFault(error: Throwable): Unit = if (isOwnershipFault(error)) triggerRefresh() - - private def acquireScope: CIO[MasterReplicaLive.TxLease] = + protected def openTransaction: CIO[LiveTransactionScope] = CIO.blocking { - val nc = - try masterPool.getOrEstablish(masterNodeRef.get()) - catch { - case e: SageException => - triggerRefresh() - throw e - case NonFatal(_) => - triggerRefresh() - throw ConnectionLost(mayHaveExecuted = false) - } - try new MasterReplicaLive.TxLease(new Client.TxScope(nc.acquireForTransaction(), refreshOnTxFault, events), nc) - catch { + try { + val nc = masterPool.getOrEstablish(topologyRef.get().master) + new Client.TxScope(nc.pool.acquireForTransaction(), nc.pool.releaseTransaction, p => if (p == RefreshPolicy.Forced) triggerRefresh(), events) + } catch { case e: TimedOut => throw e case e: SageException => triggerRefresh() @@ -494,106 +292,12 @@ final private[client] class MasterReplicaLive( } } - private def releaseScope(lease: MasterReplicaLive.TxLease): CIO[Unit] = - CIO.blocking(lease.nc.releaseTransaction(lease.scope.conn, lease.scope.sealAndReusable())) - // --- pub/sub (on the master) --------------------------------------------------------------------------------------------------------- - private def subs(): SubscriptionConnection = { - var s = subscriptions - if (s == null) { - subLock.lock() - try { - if (subscriptions == null) { - // resolve the master for every connection attempt so subscriptions move to the promoted master after failover. - val rehomingFactory: MultiplexedConnection.TransportFactory = - (onFrame, onClosed) => nodeFactory(masterNodeRef.get())(onFrame, onClosed) - subscriptions = new SubscriptionConnection( - rehomingFactory, - bootstrap, - scheduler, - config.reconnect, - config.watchdog, - config.connectTimeout.toMillis, - config.pubsub.bufferSize, - // subscriptions use a separate socket. Wait for master discovery before opening it; a pooled connection is not required. - () => !closed && masterNodeRef.get() != null, - onReconnect = () => refreshRolesBeforeRehome(), - events = events - ) - } - s = subscriptions - } finally subLock.unlock() - } - s - } - - def subscribeChannels[V: ValueCodec](channel: String, rest: String*): CIO[Subscription[CIO, Message[V]]] = - CIO.blocking(Client.channelMessages(subs().subscribeChannels(channel +: rest.toVector))) - - def subscribePatterns[V: ValueCodec](pattern: String, rest: String*): CIO[Subscription[CIO, PatternMessage[V]]] = - CIO.blocking(Client.patternMessages(subs().subscribePatterns(pattern +: rest.toVector))) - - // master-replica mode uses one subscription connection for all shard channels, as standalone mode does. - def subscribeShardChannels[V: ValueCodec](channel: String, rest: String*): CIO[Subscription[CIO, Message[V]]] = - CIO.blocking(Client.channelMessages(subs().subscribeShard(channel +: rest.toVector))) - - // --- scan / lifecycle ---------------------------------------------------------------------------------------------------------------- - - def scanTargets: CIO[Vector[ScanTarget]] = CIO.value(Vector(ScanTarget.any)) - def runOn[A](target: ScanTarget, command: Command[A]): CIO[A] = run(command) - - private[sage] def rateLimitAcquire[RK](executor: RateLimitExecutor[RK], subject: RK, cost: Long, peek: Boolean): CIO[Decision] = - executor.evalSha(this, subject, cost, peek) - - private[sage] def lockTryWith[LK, A](executor: LockExecutor[LK], key: LK)(body: => CIO[A]): CIO[Option[A]] = - executor.tryWithLock(this, key)(body) - - private[sage] def lockWith[LK, A](executor: LockExecutor[LK], key: LK, waitTimeout: FiniteDuration)(body: => CIO[A]): CIO[A] = - executor.withLock(this, key, waitTimeout)(body) - - def close: CIO[Unit] = CIO.blocking(closeAll()) - - private def closeAll(): Unit = { - closed = true - refreshThrottle.stopPolling() - val s = subscriptions - if (s != null) s.close() - masterPool.close() - replicaPool.close() - events.close() - } + protected def pubsub: SubscriptionConnection.PubSub = subscriptions } private[client] object MasterReplicaLive { final private[client] case class ResolvedTopology(master: Node, replicas: Vector[Node]) - - final private[client] class TxLease(val scope: Client.TxScope, val nc: NodeClient) - - def connect( - config: SageConfig, - seeds: Vector[Node], - masterReplica: MasterReplicaConfig, - scheduler: Scheduler, - translate: Throwable => Throwable - ): CIO[Client[CIO, String]] = - CIO.blocking[Client[CIO, String]] { - val bootstrap = Bootstrap.commands(config.auth, config.database, config.clientName) - val factory: Node => MultiplexedConnection.TransportFactory = node => { - val upgrade = Tls.buildUpgrade(config.tls, node.host, node.port) - (onFrame, onClosed) => SocketTransport.connect(node.host, node.port, config.connectTimeout, upgrade, onFrame, onClosed) - } - val events = Events(config.listeners, config.tracer) - val live = - new MasterReplicaLive(factory, scheduler, bootstrap, config, seeds, masterReplica, events) - try { - live.bootstrapRoles() - live - } catch { - case NonFatal(error) => - events.close() - throw translate(error) - } - } } diff --git a/sage-client/shared/src/main/scala/sage/client/internal/MultiplexedConnection.scala b/sage-client/shared/src/main/scala/sage/client/internal/MultiplexedConnection.scala index 941befa5..db8dd520 100644 --- a/sage-client/shared/src/main/scala/sage/client/internal/MultiplexedConnection.scala +++ b/sage-client/shared/src/main/scala/sage/client/internal/MultiplexedConnection.scala @@ -1,17 +1,14 @@ package sage.client.internal -import java.util.concurrent.{ConcurrentLinkedQueue, CountDownLatch, TimeUnit} -import java.util.concurrent.atomic.{AtomicInteger, AtomicReference} +import java.util.concurrent.{CountDownLatch, TimeUnit} import java.util.concurrent.locks.ReentrantLock import scala.annotation.tailrec -import scala.concurrent.duration.* import scala.util.{Failure, Success, Try} -import scala.util.control.NonFatal -import sage.{Bytes, CommandSpan, Outcome, SageEvent, SageException} -import sage.SageException.{ConnectionLost, NotConnected} -import sage.client.{BackoffConfig, WatchdogConfig} +import sage.{Bytes, SageEvent} +import sage.SageException.NotConnected +import sage.client.SageConfig import sage.cluster.Node import sage.commands.{Command, Connection, Invalidation, Reply} import sage.protocol.Frame @@ -22,41 +19,29 @@ import sage.protocol.Frame * [[TransportFactory]] again and resolves the hostname again, allowing DNS changes after failover to select the new master. * * A new connection gets a new `pending` queue, which prevents late frames from a closed connection from affecting the current one. The - * [[Scheduler]] runs reconnect delays outside the reader thread. + * [[Scheduler]] runs reconnect delays outside the reader thread. The connection owns its node's [[DedicatedPool]] and sends blocking + * commands there. */ -final private[client] class MultiplexedConnection private ( +final private[client] class MultiplexedConnection( factory: MultiplexedConnection.TransportFactory, scheduler: Scheduler, - bootstrap: Vector[Command[?]], - backoff: BackoffConfig, - watchdog: WatchdogConfig, - connectTimeout: FiniteDuration, - closeTimeout: FiniteDuration, - cacheMaxBytes: Long, - node: Option[Node], - events: Events + config: SageConfig, + role: MultiplexedConnection.NodeRole, + node: Option[Node] = None, + events: Events = Events.disabled ) { - import MultiplexedConnection.{Generation, State} + import MultiplexedConnection.State // use ReentrantLock because a waiting virtual thread can unmount. A synchronized monitor can pin its carrier thread on JDK versions before 24. - private val lock = new ReentrantLock() - private var state: State = State.Reconnecting - private var current: Conn = null - private var establishing: Conn = null - private var watchdogHandle: Scheduler.Cancelable = null - // increment when a new socket becomes live; the dedicated pool records this value and rejects connections from an earlier generation - private var generation: Generation = Generation.initial - // Store the current live generation, or notLive, in an atomic reference. The dedicated pool can read it while holding the pool lock without - // also acquiring this connection's lifecycle lock. - private val liveEpoch = new AtomicReference[Generation](Generation.notLive) - @volatile private var onLivenessLost: () => Unit = () => () - // keep the reconnect attempt count across short-lived connections, increasing their backoff until a connection remains stable - private var reconnectAttempt: Int = 0 - private var liveSinceMillis: Long = -1L - // whether the live generation accepts CLIENT TRACKING; false makes cached reads run uncached - @volatile private var trackingActive = true - - private[internal] def setOnLivenessLost(hook: () => Unit): Unit = onLivenessLost = hook + private val lock = new ReentrantLock() + // readable by the dedicated pool without this connection's lifecycle lock + @volatile private var state: State = State.Reconnecting + private var current: Conn = null + private var establishing: Conn = null + private val reconnects = new Reconnects(scheduler, config.reconnect, lock) + private val bootstrap = Bootstrap.commands(config) ++ role.setup + private[internal] val pool = + new DedicatedPool(factory, bootstrap, scheduler, () => isLive, config.dedicatedPool, config.connectTimeout.toMillis) private inline def locked[A](inline body: A): A = { lock.lock() @@ -64,16 +49,14 @@ final private[client] class MultiplexedConnection private ( finally lock.unlock() } - // the sole mutator of `state`, keeping `liveEpoch` in step. Must hold `lock`; at a Live edge the caller bumps `generation` first. - private def transition(to: State): Unit = { - state = to - liveEpoch.set(if (to == State.Live) generation else Generation.notLive) - if (to == State.Live) liveSinceMillis = scheduler.nowMillis + private def goLive(conn: Conn): Unit = { + current = conn + state = State.Live + reconnects.live() + conn.watch() + events.emit(SageEvent.Connection.Connected(node)) } - // return the current connection only in the Live state; disconnected or reconnecting callers receive null and fail immediately - private inline def liveConn(): Conn = locked(if (state == State.Live) current else null) - // Add n accepted entries to inFlight while holding the lock that admits the command. This records them before close() reads inFlight, including // commands that have not reached Conn.submit yet (#95). A null result means the connection is not Live. private def reserved(n: Int): Conn = locked(if (state == State.Live) { @@ -81,34 +64,30 @@ final private[client] class MultiplexedConnection private ( current } else null) - def submit[A](command: Command[A], callback: Try[A] => Unit): Unit = { - val conn = reserved(1) - if (conn == null) callback(Failure(NotConnected())) - else conn.submit(command, callback) - } + // A supplied lease lets an interrupted caller release the slot; blocking commands without one use a private lease. ASKING must immediately + // precede its command on the wire (it arms the target node for the next command on the connection). Writing the pair as one batch keeps them + // adjacent and FIFO-matched even though every fiber shares this connection; the ASKING reply is discarded. + def submit[A](command: Command[A], callback: Try[A] => Unit, asking: Boolean = false, lease: DedicatedPool.Lease = null): Unit = + if (command.isBlocking) pool.use(command, callback, if (lease != null) lease else new DedicatedPool.Lease, asking) + else { + val conn = reserved(if (asking) 2 else 1) + if (conn == null) callback(Failure(NotConnected())) + else if (asking) conn.submitAfter(Connection.asking, command, callback) + else conn.write(command, callback) + } - // Use one connection for the cache lookup and server fetch; reconnecting during the fetch reports a connection loss. In master-replica mode, - // deferred contains tracing context captured before offloading. A null value starts tracing on this thread. Local cache hits are not traced. - def cachedSubmit[A](command: Command[A], ttlMillis: Long, callback: Try[A] => Unit, deferred: () => CommandSpan = null): Unit = { + // Use one connection for the cache lookup and server fetch; reconnecting during the fetch reports a connection loss. trace tracks the read + // once it is sent or fails. + def cachedSubmit[A](command: Command[A], ttlMillis: Long, callback: Try[A] => Unit, trace: Events.Fetching): Unit = { // a Fetch sends [CLIENT CACHING YES, read]; a cache hit/wait sends nothing and releases the reservation val conn = reserved(2) - if (conn == null) { - val error = NotConnected() - Events.settleSpan(Events.startOrDefer(events, command, deferred), Outcome.Failed(error)) - callback(Failure(error)) - } else conn.cachedSubmit(command, ttlMillis, callback, deferred) - } - - // ASKING must immediately precede its command on the wire (it arms the target node for the next command on the connection). Writing the - // pair as one batch keeps them adjacent and FIFO-matched even though every fiber shares this connection; the ASKING reply is discarded. - def submitAsking[A](command: Command[A], callback: Try[A] => Unit): Unit = { - val conn = reserved(2) - if (conn == null) callback(Failure(NotConnected())) - else conn.submitAll(Vector(Connection.asking, command), Vector(_ => (), callback.asInstanceOf[Try[Any] => Unit])) + if (conn == null) trace.fetching(command, callback)(Failure(NotConnected())) + else conn.cachedSubmit(command, ttlMillis, callback, trace) } // Enqueues a whole pipeline onto one captured generation. A reconnect cannot split the batch across connections. // Return false when disconnected. The caller handles the unsent batch by rerouting or failing it as appropriate. + // Blocking commands are rejected before this method is called. def submitAll(commands: Vector[Command[?]], callbacks: Vector[Try[Any] => Unit]): Boolean = { val conn = reserved(commands.length) if (conn == null) false @@ -119,61 +98,56 @@ final private[client] class MultiplexedConnection private ( } def close(): Unit = { + pool.close() // `aborting`: a reconnect attempt's in-flight connection, closed so its socket is released before close() returns. val (draining, aborting) = locked { state match { case State.Closed | State.Draining => (null, null) case State.Reconnecting => - transition(State.Closed) - stopWatchdog() + state = State.Closed (null, establishing) case State.Live => - transition(State.Draining) + state = State.Draining (current, null) } } if (aborting != null) aborting.close() if (draining != null) { - draining.beginDrain().await(closeTimeout.toMillis, TimeUnit.MILLISECONDS) - locked(stopWatchdog()) + try draining.beginDrain().await(config.closeTimeout.toMillis, TimeUnit.MILLISECONDS): Unit + catch { case _: InterruptedException => Thread.currentThread().interrupt() } draining.close() } } - private[internal] def currentState: State = locked(state) + private[internal] def currentState: State = state - private[internal] def isLive: Boolean = liveEpoch.get() != Generation.notLive + private[internal] def isLive: Boolean = state == State.Live - // return the current generation for a new DedicatedConnection while this connection is Live; None reports that it is unavailable - private[internal] def liveGeneration(): Option[Generation] = { - val g = liveEpoch.get() - if (g == Generation.notLive) None else Some(g) + /** + * Connects to the node and runs bootstrap synchronously. Like a standalone client, this method throws without retrying when the first + * handshake fails, and closes the connection and its pool. The caller moves this blocking connection attempt off its thread and treats a + * failure as an unreachable node. + */ + def start(): MultiplexedConnection = { + // the first connect propagates a handshake failure; only reconnects retry + onThrow(install(establish()))(_ => close()) + this } - // Compare a DedicatedConnection's generation with the current live generation. The `notLive` value rejects every generation while - // reconnecting, including the previous number before the next connection increments it. - private[internal] def isCurrent(g: Generation): Boolean = liveEpoch.get() == g - - private def connectInitial(): Unit = { - val conn = establish() // the first connect propagates a handshake failure; only reconnects retry - // Emit while holding the lock to keep the event ordered with the state transition. If the socket drops immediately afterward, - // Disconnected is enqueued after Connected because the same lock serializes both events. - val teardown = locked { - // if close runs during establishment, close the new connection instead of changing the state to Live - if (state == State.Closed) conn - else if (conn.isTerminated) { - scheduleReconnect(0) - null - } else { - current = conn - generation = generation.next - transition(State.Live) - startWatchdog() - events.emit(SageEvent.Connection.Connected(node)) - null - } + // Go live unless close ran during establishment or the connection died in setup; a dead one is replaced by a reconnect. Emit while holding + // the lock to keep the event ordered with the state transition. If the socket drops immediately afterward, Disconnected is enqueued after + // Connected because the same lock serializes both events. + private def install(conn: Conn): Unit = { + val live = locked { + if (state == State.Reconnecting && !conn.isDead) { + goLive(conn) + true + } else false + } + if (!live) { + conn.close() + locked(if (state == State.Reconnecting) reconnect()) } - if (teardown != null) teardown.close() } private def establish(): Conn = { @@ -183,21 +157,13 @@ final private[client] class MultiplexedConnection private ( if (state == State.Closed) conn.close() } try { - conn.start() - var tracking = true - // Call reserve(1) for each command because submit does not increment the count. The reply later calls retire() to balance it. A - // half-open peer can accept the socket without answering HELLO, so the runner limits each wait and lets the reconnect loop continue. - Bootstrap.run( - bootstrap, - connectTimeout.toMillis, - (c, cb) => { - conn.reserve(1) - conn.submit(c, cb) - }, - () => conn.close(), - onTolerated = c => if (Connection.isClientTracking(c)) tracking = false - ) - trackingActive = tracking + // A half-open peer can accept the socket without answering HELLO, so the handshake limits each wait and lets the reconnect loop continue. + conn.handshake(bootstrap, config.connectTimeout.toMillis) + // a server that denies tracking (an ACL restriction, a proxy) still connects and serves cached reads uncached (ADR-0045) + conn.cache = Option + .when(role.caches && config.clientCache.enabled)(config.clientCache.maxBytes) + .filter(_ => conn.step(Connection.clientTrackingOnOptin, config.connectTimeout.toMillis).isEmpty) + .map(new ClientCache(_)) conn } finally locked { if (establishing eq conn) establishing = null } } @@ -207,317 +173,111 @@ final private[client] class MultiplexedConnection private ( if (c != null) c.flushCache() } - private def scheduleReconnect(attempt: Int): Unit = { - reconnectAttempt = attempt - scheduler.after(Backoff.jitteredMillis(backoff, attempt, scheduler).millis)(attemptReconnect(attempt)) - } - - // reset the attempt count after a connection remains live for the maximum backoff period; shorter connections keep increasing the delay - private def nextReconnectAttempt(): Int = { - val stable = liveSinceMillis >= 0L && scheduler.nowMillis - liveSinceMillis >= backoff.maxDelay.toMillis - reconnectAttempt = if (stable) 0 else reconnectAttempt + 1 - reconnectAttempt - } - - private def attemptReconnect(attempt: Int): Unit = { - val proceed = locked(state == State.Reconnecting) - if (proceed) - try { - val conn = establish() - val live = locked { - if (state == State.Reconnecting && !conn.isTerminated) { - current = conn - generation = generation.next - transition(State.Live) - startWatchdog() - events.emit(SageEvent.Connection.Connected(node)) - true - } else false - } - if (!live) { - conn.close() - locked(if (state == State.Reconnecting) scheduleReconnect(attempt + 1)) - } - } catch { - case NonFatal(error) => - locked(if (state == State.Reconnecting) { - events.emit(SageEvent.Connection.ReconnectFailed(node, error)) - scheduleReconnect(attempt + 1) - }) - } - } + // must hold lock + private def reconnect(): Unit = + reconnects.schedule(state == State.Reconnecting, error => events.emit(SageEvent.Connection.ReconnectFailed(node, error)))(install(establish())) // ignore connections that fail before becoming `current`; the establishment caller handles those failures - private def onConnTerminated(conn: Conn): Unit = { - // emit Disconnected only when the current Live connection ends. Holding the lock orders it after that connection's Connected event. - val lostLiveness = locked { + private def onConnTerminated(conn: Conn): Unit = + // Emit Disconnected only when the current Live connection ends. Holding the lock orders it after that connection's Connected event, and + // the pool learns of the loss before any reconnect can go live. + locked { if (conn eq current) state match { case State.Live => - transition(State.Reconnecting) - scheduleReconnect(nextReconnectAttempt()) + state = State.Reconnecting + pool.onLivenessLost() + reconnect() events.emit(SageEvent.Connection.Disconnected(node)) - true case State.Draining => - transition(State.Closed) - stopWatchdog() - false - case State.Reconnecting | State.Closed => false + state = State.Closed + case State.Reconnecting | State.Closed => () } - else false } - // notify the pool after releasing the lifecycle lock - if (lostLiveness) onLivenessLost() - } - - private def startWatchdog(): Unit = - if (watchdog.enabled && watchdogHandle == null) - watchdogHandle = scheduler.every(watchdog.pingInterval)(watchdogTick()) - - private def stopWatchdog(): Unit = - if (watchdogHandle != null) { - watchdogHandle.cancel() - watchdogHandle = null - } - - private def watchdogTick(): Unit = { - val conn = liveConn() - if (conn != null) conn.checkLiveness(scheduler.nowMillis, watchdog.pingInterval.toMillis, watchdog.pingTimeout.toMillis) - } - - private def decodeFrame[A](command: Command[A], frame: Frame): Try[A] = Reply.decode(command, frame) - final private class Conn { + final private class Conn extends WatchedPipe(factory, scheduler, config.watchdog) { - private val pending = new ConcurrentLinkedQueue[Entry[?]]() - // count entries when they are accepted by submit, including commands still queued for writing, so close waits for all accepted work (#95) - private val inFlight = new AtomicInteger(0) - private val transportRef = new AtomicReference[Transport]() - // each connection has its own cache. Reconnecting creates a new Conn and discards the previous cached values. - private val cache = new ClientCache(cacheMaxBytes) - @volatile private var lastReplyAtMillis: Long = scheduler.nowMillis + // each connection has its own cache. Reconnecting creates a new Conn and discards the previous cached values. None runs cached reads + // uncached, either because caching is off or because the server refused CLIENT TRACKING. + @volatile var cache: Option[ClientCache] = None @volatile private var drainLatch: CountDownLatch = null - @volatile private var aborted: Boolean = false - @volatile private var terminated: Boolean = false - - def isTerminated: Boolean = terminated - - // Publish transportRef before the blocking connect starts so close can abort it. aborted records a close that happens before transportRef - // is published. - def start(): Unit = { - val transport = factory(onFrame, onClosed) - transportRef.set(transport) - if (aborted) transport.close() - else transport.start() - } - - // inFlight is reserved by the caller (reserve below) before the send, so neither submit nor submitAll touches the counter here. - def submit[A](command: Command[A], callback: Try[A] => Unit): Unit = - transportRef.get().send(new Entry(command, callback)) // Concatenate the pipeline into one Transport.Item. The writer processes an item atomically, which keeps the pipeline in one socket write // and prevents other sends from being inserted between its commands. - def submitAll(commands: Vector[Command[?]], callbacks: Vector[Try[Any] => Unit]): Unit = { - val entries = Vector.tabulate(commands.length)(i => new Entry(commands(i), callbacks(i))) - transportRef.get().send(new Batch(entries)) - } + def submitAll(commands: Vector[Command[?]], callbacks: Vector[Try[Any] => Unit]): Unit = + sendAll(Vector.tabulate(commands.length)(i => new Entry(commands(i), callbacks(i)))) - // increment inFlight for accepted entries. retire() decrements it after completion, and release() decrements it for unsent entries. - def reserve(n: Int): Unit = - inFlight.addAndGet(n): Unit + // writes `prefix`, whose reply is discarded, and `command` as one item so no other command is written between them + def submitAfter[A](prefix: Command[Unit], command: Command[A], callback: Try[A] => Unit): Unit = + sendAll(Vector(new Entry(prefix, _ => ()), new Entry(command, callback))) // OPTIN tracking applies CLIENT CACHING YES only to the next command. Submit it together with the cached read to keep them adjacent. An // identity decoder passes the raw reply Frame to the cache, and each waiter then uses its own command decoder. - def cachedSubmit[A](command: Command[A], ttlMillis: Long, callback: Try[A] => Unit, deferred: () => CommandSpan): Unit = { - // tracking off: run uncached, releasing one of the two slots reserved for the (now unsent) caching prefix - if (!trackingActive) { - release(1) - uncached(command, callback, deferred) - return - } - val commandBytes = command.encode - val keys = command.keys - def deliver(frame: Frame): Unit = callback(decodeFrame(command, frame)) - val waiter: Try[Frame] => Unit = { - case Success(frame) => deliver(frame) - case Failure(error) => callback(Failure(error)) - } - @tailrec def attempt(): Unit = - cache.acquire(commandBytes, keys, scheduler.nowMillis, waiter) match { - // Hit returns a local value. Wait joins an in-flight fetch. Neither case sends another server request, so report a cache hit and - // release the two slots reserved for the unsent fetch. - case ClientCache.Acquire.Hit(frame, epoch) => - if (cache.isCurrent(epoch)) { + def cachedSubmit[A](command: Command[A], ttlMillis: Long, callback: Try[A] => Unit, trace: Events.Fetching): Unit = + cache match { + // tracking off: run uncached, releasing one of the two slots reserved for the (now unsent) caching prefix + case None => + release(1) + write(command, trace.fetching(command, callback)) + case Some(cache) => + val commandBytes = command.encode + val keys = command.keys + val reply = new MultiplexedConnection.CachedReply(command, callback) + @tailrec def lookup(): Unit = cache.acquire(commandBytes, keys, scheduler.nowMillis, reply) match { + // A flush since the lookup (MOVED, a dropped connection, FLUSHALL) retires the hit. Looking it up again fetches from the server, + // which redirects the read, or fails it with ConnectionLost when the connection is gone. + case hit: ClientCache.Acquire.Hit if !cache.isCurrent(hit) => lookup() + // Hit returns a local value. Wait joins an in-flight fetch. Neither case sends another server request, so report a cache hit and + // release the two slots reserved for the unsent fetch. + case ClientCache.Acquire.Hit(frame, _) => release(2) if (events.emitsEvents) events.emit(SageEvent.Cache.Hit(command.name)) - deliver(frame) - } - // reroute a hit retired by a topology change or a dead connection; refetch here one retired by a server flush (ownership unchanged) - else if (terminated || cache.rerouteRetired(epoch)) { + reply.deliver(frame) + case ClientCache.Acquire.Wait => release(2) - callback(Failure(ConnectionLost(mayHaveExecuted = false))) - } else attempt() - case ClientCache.Acquire.Wait => - release(2) - if (events.emitsEvents) events.emit(SageEvent.Cache.Hit(command.name)) - case ClientCache.Acquire.Fetch => - if (events.emitsEvents) events.emit(SageEvent.Cache.Miss(command.name)) - traced(command, deferred) { settle => - val raw = Command[Frame](command.name, command.keyIndices, command.args, frame => Right(frame)) - val onReply: Try[Frame] => Unit = { result => - result match { - case Success(frame) => cache.store(commandBytes, keys, frame, scheduler.nowMillis, ttlMillis) - case Failure(error) => cache.fail(commandBytes, error) - } - // the outcome reflects the decoded reply, not the raw identity-decoded frame - settle(Outcome.of(result.flatMap(decodeFrame(command, _)))) - } - try submitAll(Vector(Connection.clientCachingYes, raw), Vector(_ => (), onReply.asInstanceOf[Try[Any] => Unit])) - catch { - // nothing was sent: release the slots the entries won't retire, and fail the fetch as a reply failure would - case NonFatal(error) => - release(2) - onReply(Failure(error)) + if (events.emitsEvents) events.emit(SageEvent.Cache.Hit(command.name)) + case ClientCache.Acquire.Fetch(fetching) => + if (events.emitsEvents) events.emit(SageEvent.Cache.Miss(command.name)) + reply.callback = trace.fetching(command, reply.callback) + val onReply: Try[Frame] => Unit = { + case Success(frame) => cache.store(fetching, frame, scheduler.nowMillis, ttlMillis) + case Failure(error) => cache.fail(fetching, error) } - } - } - attempt() - } - - private def uncached[A](command: Command[A], callback: Try[A] => Unit, deferred: () => CommandSpan): Unit = - traced(command, deferred)(settle => - submit( - command, - (result: Try[A]) => { - settle(Outcome.of(result)) - callback(result) + submitAfter(Connection.clientCachingYes, command.rawFrame, onReply) } - ) - ) - - // record tracing and CommandCompleted events around a command sent to the server. send calls the supplied function with the outcome. - private def traced(command: Command[?], deferred: () => CommandSpan)(send: ((=> Outcome) => Unit) => Unit): Unit = { - val span = Events.startOrDefer(events, command, deferred) - node.foreach(Events.routeSpan(span, _)) - val started = System.nanoTime() - send { outcome => - if (events.enabled) { - val settled = outcome - Events.settleSpan(span, settled) - if (events.emitsEvents) - events.emit(SageEvent.CommandCompleted(command.name, node, FiniteDuration(System.nanoTime() - started, NANOSECONDS), settled)) - } + trace.lookingUp() + lookup() } - } - def flushCache(): Unit = cache.flushForReroute() - - def close(): Unit = { - aborted = true - val transport = transportRef.get() - if (transport != null) transport.close() - } + def flushCache(): Unit = cache.foreach(_.flush()) def beginDrain(): CountDownLatch = { + unwatch() val latch = new CountDownLatch(1) drainLatch = latch - if (inFlight.get() == 0) latch.countDown() + if (isIdle) latch.countDown() latch } - private def retire(): Unit = release(1) - - private def release(n: Int): Unit = - if (inFlight.addAndGet(-n) == 0) { - val latch = drainLatch - if (latch != null) latch.countDown() - } - - def checkLiveness(now: Long, intervalMillis: Long, timeoutMillis: Long): Unit = { - val head = pending.peek() - if (head != null) { - // offload: close() blocks joining I/O threads, and the watchdog tick runs on the shared timer thread, which must not block - if (now - head.sentAtMillis >= timeoutMillis) scheduler.after(Duration.Zero)(close()) - } else if (now - lastReplyAtMillis >= intervalMillis) { - reserve(1) - submit(Connection.ping(None), _ => ()) - } + override protected def onDrained(): Unit = { + val latch = drainLatch + if (latch != null) latch.countDown() } - // Out-of-band frames do not consume a pending entry. A READONLY reply fails its command and closes the connection because an in-place - // failover can leave the old master connected but unable to accept writes - private def onFrame(frame: Frame): Unit = - frame match { - case Frame.Push(elements) => - // a push confirms only that reads are working. Leave lastReplyAtMillis unchanged so push-only traffic still receives idle PING checks. - Invalidation.decode(elements) match { - case Some(Invalidation.Evict(keys)) => keys.foreach(cache.invalidate) - case Some(Invalidation.FlushAll) => cache.flush() - case None => () - } - case reply => - lastReplyAtMillis = scheduler.nowMillis - val entry = pending.poll() - if (entry == null) close() - else entry.complete(reply) - if (Poison.isReadonly(reply)) close() + protected def onPush(elements: Vector[Frame]): Unit = + Invalidation.decode(elements) match { + case Some(Invalidation.Evict(keys)) => cache.foreach(c => keys.foreach(c.invalidate)) + case Some(Invalidation.FlushAll) => cache.foreach(_.flush()) + case None => () } - private def onClosed(): Unit = { - terminated = true + override protected def onClosed(): Unit = { // a dropped connection loses all further invalidations, so its cache can no longer be trusted for a hit - cache.flush() - var entry = pending.poll() - while (entry != null) { - entry.fail(ConnectionLost(mayHaveExecuted = true)) - entry = pending.poll() - } + cache.foreach(_.flush()) + super.onClosed() if (drainLatch != null) drainLatch.countDown() onConnTerminated(this) } - - final private class Entry[A](command: Command[A], callback: Try[A] => Unit) extends Transport.Item { - - @volatile var sentAtMillis: Long = 0L - - var payload: Bytes = command.encode - - override def clearPayload(): Unit = payload = Bytes.empty - - def writeAttempted(): Unit = { - sentAtMillis = scheduler.nowMillis - pending.add(this): Unit - } - - def dropped(): Unit = { - callback(Failure(ConnectionLost(mayHaveExecuted = false))) - retire() - } - - // decodeFrame guards against throwing user decoders: an escaped exception would otherwise lose the callback and hang the awaiting fiber - def complete(frame: Frame): Unit = { - callback(decodeFrame(command, frame)) - retire() - } - - def fail(error: SageException): Unit = { - callback(Failure(error)) - retire() - } - } - - // Concatenate a pipeline into one transport write and notify each entry individually when the write succeeds or fails; the transport - // writes or drops the complete batch without splitting it across socket writes - final private class Batch(entries: Vector[Entry[Any]]) extends Transport.Item { - - val payload: Bytes = Bytes.concatBy(entries)(_.payload) - - override def clearPayload(): Unit = entries.foreach(_.clearPayload()) - - def writeAttempted(): Unit = entries.foreach(_.writeAttempted()) - - def dropped(): Unit = entries.foreach(_.dropped()) - } } } @@ -525,40 +285,27 @@ private[client] object MultiplexedConnection { type TransportFactory = (Frame => Unit, () => Unit) => Transport - enum State { - case Live, Reconnecting, Draining, Closed + // Completes a cached read from the raw reply frame. A read sent to the server switches to its tracked callback before the send, which + // orders the switch before any reply. + final private class CachedReply[A](command: Command[A], var callback: Try[A] => Unit) extends (Try[Frame] => Unit) { + def deliver(frame: Frame): Unit = callback(Reply.decode(command, frame)) + + def apply(result: Try[Frame]): Unit = + result match { + case Success(frame) => deliver(frame) + case Failure(error) => callback(Failure(error)) + } } - // a monotonic generation for the current socket; the pool records it on each DedicatedConnection and rejects the connection after it changes - opaque type Generation = Long - object Generation { - val initial: Generation = 0L - // distinct from every real generation, making all recorded generations invalid while disconnected - val notLive: Generation = -1L - extension (g: Generation) def next: Generation = g + 1L + enum State { + case Live, Reconnecting, Draining, Closed } - /** - * Connects and runs the bootstrap synchronously; throws (no retry) if the first handshake fails. - */ - def connect( - factory: TransportFactory, - scheduler: Scheduler, - bootstrap: Vector[Command[?]], - backoff: BackoffConfig, - watchdog: WatchdogConfig, - connectTimeout: FiniteDuration, - closeTimeout: FiniteDuration, - cacheMaxBytes: Long = 0L, - node: Option[Node] = None, - events: Events = Events.disabled, - onConstructed: MultiplexedConnection => Unit = _ => () - ): MultiplexedConnection = { - val connection = - new MultiplexedConnection(factory, scheduler, bootstrap, backoff, watchdog, connectTimeout, closeTimeout, cacheMaxBytes, node, events) - // expose the instance before the blocking connect so an owner can abort it from close() - onConstructed(connection) - connection.connectInitial() - connection + // Only master connections cache, because cached reads run on the master. Cluster replicas send READONLY during setup to serve reads for + // their master's slots. + enum NodeRole(val setup: Vector[Command[?]], val caches: Boolean) { + case Master extends NodeRole(Vector.empty, caches = true) + case Replica extends NodeRole(Vector.empty, caches = false) + case ClusterReplica extends NodeRole(Vector(Connection.readonly), caches = false) } } diff --git a/sage-client/shared/src/main/scala/sage/client/internal/NodeClient.scala b/sage-client/shared/src/main/scala/sage/client/internal/NodeClient.scala deleted file mode 100644 index 1b61c47d..00000000 --- a/sage-client/shared/src/main/scala/sage/client/internal/NodeClient.scala +++ /dev/null @@ -1,95 +0,0 @@ -package sage.client.internal - -import scala.concurrent.duration.FiniteDuration -import scala.util.Try - -import sage.CommandSpan -import sage.client.{BackoffConfig, DedicatedPoolConfig, WatchdogConfig} -import sage.cluster.Node -import sage.commands.Command - -/** - * One node's connection bundle: a Multiplexed Connection + Dedicated Pool pinned to that node. A standalone client holds one; the cluster - * runtime holds one per master it routes to. `asking` prefixes the command with `ASKING` to follow an `ASK` redirect. - */ -final private[client] class NodeClient(connection: MultiplexedConnection, pool: DedicatedPool) { - - // A supplied lease lets an interrupted caller release the slot; blocking commands without one use a private lease. - def submit[A](command: Command[A], asking: Boolean, callback: Try[A] => Unit, lease: DedicatedPool.Lease = null): Unit = - if (command.isBlocking) { - val l = if (lease != null) lease else new DedicatedPool.Lease - if (asking) pool.useAsking(command, callback, l) else pool.use(command, callback, l) - } else if (asking) connection.submitAsking(command, callback) - else connection.submit(command, callback) - - def submitLockWrite[A]( - command: Command[A], - asking: Boolean, - callback: Try[A] => Unit, - lease: DedicatedPool.Lease, - replication: LockReplication - ): Unit = pool.useLockWrite(command, asking, callback, lease, replication) - - def cachedSubmit[A](command: Command[A], ttlMillis: Long, callback: Try[A] => Unit, deferred: () => CommandSpan = null): Unit = - connection.cachedSubmit(command, ttlMillis, callback, deferred) - - // Submit this node's part of a pipeline in one round trip. A false result means the connection accepted no commands, and the caller should - // route the complete batch again. Blocking commands are rejected before this method is called. - def submitAll(commands: Vector[Command[?]], callbacks: Vector[Try[Any] => Unit]): Boolean = - connection.submitAll(commands, callbacks) - - def flushCache(): Unit = connection.flushCache() - - def acquireForTransaction(): DedicatedConnection = pool.acquireForTransaction() - - def releaseTransaction(connection: DedicatedConnection, reusable: Boolean): Unit = pool.releaseTransaction(connection, reusable) - - def isLive: Boolean = connection.isLive - - def close(): Unit = { - pool.close() - connection.close() - } -} - -private[client] object NodeClient { - - /** - * Connects to the node and runs bootstrap synchronously. Like a standalone client, this method throws without retrying when the first - * handshake fails. The caller moves this blocking connection attempt off its thread and treats a failure as an unreachable node. - */ - def connect( - factory: MultiplexedConnection.TransportFactory, - scheduler: Scheduler, - bootstrap: Vector[Command[?]], - reconnect: BackoffConfig, - watchdog: WatchdogConfig, - connectTimeout: FiniteDuration, - closeTimeout: FiniteDuration, - dedicatedPool: DedicatedPoolConfig, - cacheMaxBytes: Long = 0L, - node: Node, - events: Events = Events.disabled, - dedicatedBootstrap: Option[Vector[Command[?]]] = None, - onConstructed: MultiplexedConnection => Unit = _ => () - ): NodeClient = { - val connection = - MultiplexedConnection - .connect( - factory, - scheduler, - bootstrap, - reconnect, - watchdog, - connectTimeout, - closeTimeout, - cacheMaxBytes, - Some(node), - events, - onConstructed - ) - val pool = - DedicatedPool.forConnection(factory, dedicatedBootstrap.getOrElse(bootstrap), scheduler, connection, dedicatedPool, connectTimeout.toMillis) - new NodeClient(connection, pool) - } -} diff --git a/sage-client/shared/src/main/scala/sage/client/internal/NodePool.scala b/sage-client/shared/src/main/scala/sage/client/internal/NodePool.scala index 52ff0e19..c3f7e349 100644 --- a/sage-client/shared/src/main/scala/sage/client/internal/NodePool.scala +++ b/sage-client/shared/src/main/scala/sage/client/internal/NodePool.scala @@ -1,6 +1,6 @@ package sage.client.internal -import java.util.concurrent.CountDownLatch +import java.util.concurrent.{CompletableFuture, ExecutionException} import java.util.concurrent.locks.ReentrantLock import scala.collection.mutable @@ -10,35 +10,27 @@ import scala.util.control.NonFatal import sage.SageEvent import sage.SageException.NotConnected -import sage.client.{BackoffConfig, DedicatedPoolConfig, WatchdogConfig} +import sage.client.SageConfig import sage.cluster.Node -import sage.commands.Command /** - * Stores one [[NodeClient]] for each [[Node]]. Concurrent callers for the same node share one connection attempt and receive the same result. + * Stores one [[MultiplexedConnection]] for each [[Node]]. Concurrent callers for the same node share one connection attempt and receive the same result. * If an attempt finishes after [[close]], its connection is closed. Master-replica clients use these pools for both roles. Cluster clients * use one for replicas and a separate pool for masters because master failures affect redirects and topology refresh. - * - * The `bootstrap` is fixed per pool, so a replica pool can append `READONLY` while a master pool stays read-write. */ final private[client] class NodePool( nodeFactory: Node => MultiplexedConnection.TransportFactory, scheduler: Scheduler, - bootstrap: Vector[Command[?]], - reconnect: BackoffConfig, - watchdog: WatchdogConfig, - connectTimeout: FiniteDuration, - closeTimeout: FiniteDuration, - dedicatedPool: DedicatedPoolConfig, - cacheMaxBytes: Long = 0L, - events: Events = Events.disabled, - dedicatedBootstrap: Option[Vector[Command[?]]] = None + config: SageConfig, + role: MultiplexedConnection.NodeRole, + events: Events = Events.disabled ) { private val lock = new ReentrantLock() // lock-free reads; every mutation stays under `lock` - private val established = new java.util.concurrent.ConcurrentHashMap[Node, NodeClient]() - private val pendingEstablish = mutable.HashMap.empty[Node, NodePool.Establish] + private val established = new java.util.concurrent.ConcurrentHashMap[Node, MultiplexedConnection]() + // one attempt shared by concurrent callers for a node; the first result is final, so a late establishment after close is ignored + private val pendingEstablish = mutable.HashMap.empty[Node, CompletableFuture[MultiplexedConnection]] // connections whose socket is still being opened, so close() can abort one still connecting private val establishing = mutable.Set.empty[MultiplexedConnection] @volatile private var closed = false @@ -52,14 +44,14 @@ final private[client] class NodePool( /** * Returns the established client for the node, or `null`. This method never blocks. */ - def existing(node: Node): NodeClient = established.get(node) + def existing(node: Node): MultiplexedConnection = established.get(node) def firstLiveNode: Option[Node] = established.asScala.collectFirst { case (node, nc) if nc.isLive => node } - def foreachEstablished(f: NodeClient => Unit): Unit = established.values.forEach(nc => f(nc)) + def foreachEstablished(f: MultiplexedConnection => Unit): Unit = established.values.forEach(nc => f(nc)) private[internal] def pendingWaiterCount(node: Node): Int = - locked(pendingEstablish.get(node).fold(0)(_.waiterCount)) + locked(pendingEstablish.get(node).fold(0)(_.getNumberOfDependents)) // live nodes first, so a refresh prefers a known-good node def candidatesByLiveness: Vector[Node] = { @@ -67,91 +59,82 @@ final private[client] class NodePool( live.map(_._1) ++ others.map(_._1) } - def getOrEstablish(node: Node): NodeClient = { - val fast = established.get(node) + def getOrEstablish(node: Node): MultiplexedConnection = { + val fast = established.get(node) if (fast != null) return fast - var existing: NodeClient = null - var waitOn: NodePool.Establish = null - var mine: NodePool.Establish = null - locked { + // Left joins an attempt in flight, Right owns a new one + val attempt = locked { if (closed) throw NotConnected() - existing = established.get(node) - if (existing == null) - pendingEstablish.get(node) match { - case Some(p) => waitOn = p - case None => - mine = new NodePool.Establish - pendingEstablish.put(node, mine): Unit - } + val existing = established.get(node) + if (existing != null) return existing + pendingEstablish.get(node).toLeft { + val mine = new CompletableFuture[MultiplexedConnection] + pendingEstablish.put(node, mine) + mine + } } - if (existing != null) existing - else if (waitOn != null) waitOn.get() - else { - val connRef = new java.util.concurrent.atomic.AtomicReference[MultiplexedConnection]() - val nc = - try - NodeClient.connect( - nodeFactory(node), - scheduler, - bootstrap, - reconnect, - watchdog, - connectTimeout, - closeTimeout, - dedicatedPool, - cacheMaxBytes, - node, - events, - dedicatedBootstrap, - onConstructed = conn => { - connRef.set(conn) - val poolClosed = locked { - establishing += conn - closed - } - if (poolClosed) conn.close() - } - ) - catch { - case error: Throwable => - locked { - val conn = connRef.get() - if (conn != null) establishing -= conn - if (pendingEstablish.get(node).exists(_ eq mine)) { pendingEstablish.remove(node): Unit } - } - mine.fail(error) - if (!closed) events.emit(SageEvent.Connection.ConnectFailed(Some(node), error)) - throw error + attempt match { + case Left(waitOn) => + try waitOn.get() + catch { case e: ExecutionException => throw e.getCause } + case Right(mine) => + var nc: MultiplexedConnection = null + onThrow { + nc = new MultiplexedConnection(nodeFactory(node), scheduler, config, role, Some(node), events) + // register before the blocking connect so that close() can abort it + if (locked { establishing += nc; closed }) nc.close() + nc.start(): Unit + } { error => + locked { + establishing -= nc + if (pendingEstablish.get(node).exists(_ eq mine)) { pendingEstablish.remove(node): Unit } + } + // a joiner gets NotConnected for this thread's interrupt, as it does for an abandoned attempt + mine.completeExceptionally(error match { case NonFatal(e) => e; case _ => NotConnected() }) + if (!closed) events.emit(SageEvent.Connection.ConnectFailed(Some(node), error)) + } + // Publish the client only while this attempt is current. A retain, close, or newer attempt supersedes it, in which case it is closed below. + val publish = locked { + establishing -= nc + val current = pendingEstablish.get(node).exists(_ eq mine) + if (current) { pendingEstablish.remove(node): Unit } + if (current && !closed) { + established.put(node, nc) + true + } else false + } + if (publish) { + mine.complete(nc) + nc + } else { + mine.completeExceptionally(NotConnected()) + nc.close() + throw NotConnected() } - // Publish the client only while this attempt is current. A retain, close, or newer attempt supersedes it, in which case it is closed below. - val publish = locked { - val conn = connRef.get() - if (conn != null) establishing -= conn - val current = pendingEstablish.get(node).exists(_ eq mine) - if (current) { pendingEstablish.remove(node): Unit } - if (current && !closed) { - established.put(node, nc) - true - } else false - } - if (publish) { - mine.succeed(nc) - nc - } else { - nc.close() - mine.fail(NotConnected()) - throw NotConnected() - } } } /** * As [[getOrEstablish]], blocking to connect if need be, but `null` rather than throwing when the connect fails. */ - def getOrEstablishOrNull(node: Node): NodeClient = + def getOrEstablishOrNull(node: Node): MultiplexedConnection = try getOrEstablish(node) catch { case NonFatal(_) => null } + /** + * Runs `use` on the caller's thread when the node already has a client. Otherwise connects on the scheduler and runs `use` with the new + * client, or `unreachable` when the connect fails. Inlining keeps the established path free of a closure allocation. + */ + inline def withClient(node: Node)(inline unreachable: => Unit)(inline use: MultiplexedConnection => Unit): Unit = { + val nc = existing(node) + if (nc != null) use(nc) + else + scheduler.offload { + val connected = getOrEstablishOrNull(node) + if (connected != null) use(connected) else unreachable + } + } + // remove and close clients for nodes rejected by keep. Also fail connection attempts for those nodes. Schedule closes outside the pool lock. def retain(keep: Node => Boolean): Unit = { val (gone, rejected) = locked { @@ -160,7 +143,7 @@ final private[client] class NodePool( (absent, rejectedPending) } gone.foreach(nc => scheduler.after(Duration.Zero)(nc.close())) - rejected.foreach(_.fail(NotConnected())) + rejected.foreach(_.completeExceptionally(NotConnected())) } def close(): Unit = { @@ -175,42 +158,8 @@ final private[client] class NodePool( } // Fail callers waiting for a connection immediately instead of making them wait for the connection timeout; an opening connection // observes `closed` when it finishes and closes the node - waiters.foreach(_.fail(NotConnected())) + waiters.foreach(_.completeExceptionally(NotConnected())) opening.foreach(_.close()) all.foreach(_.close()) } } - -private[client] object NodePool { - - // Share one connection attempt among concurrent callers. One caller opens the connection while the others wait for the same result. - // The first result is final; once close has failed the waiters, a late establishment is ignored. - final private class Establish { - private val latch = new CountDownLatch(1) - private val settled = new java.util.concurrent.atomic.AtomicBoolean(false) - private val waiters = new java.util.concurrent.atomic.AtomicInteger(0) - @volatile private var result: Either[Throwable, NodeClient] = null - - def waiterCount: Int = waiters.get() - - def succeed(nc: NodeClient): Unit = settle(Right(nc)) - def fail(error: Throwable): Unit = settle(Left(error)) - - private def settle(outcome: Either[Throwable, NodeClient]): Unit = - if (settled.compareAndSet(false, true)) { - result = outcome - latch.countDown() - } - - def get(): NodeClient = { - waiters.incrementAndGet() - try { - latch.await() - result match { - case Right(nc) => nc - case Left(error) => throw error - } - } finally waiters.decrementAndGet(): Unit - } - } -} diff --git a/sage-client/shared/src/main/scala/sage/client/internal/Paged.scala b/sage-client/shared/src/main/scala/sage/client/internal/Paged.scala index baf05771..180a5bb7 100644 --- a/sage-client/shared/src/main/scala/sage/client/internal/Paged.scala +++ b/sage-client/shared/src/main/scala/sage/client/internal/Paged.scala @@ -7,7 +7,8 @@ import scala.concurrent.duration.FiniteDuration import kyo.compat.* import sage.BlockTimeout -import sage.commands.{ScanCursor, ScanPage, StreamEntry, StreamId, StreamRangeId, XAutoClaimResult} +import sage.codec.{KeyCodec, ValueCodec} +import sage.commands.* /** * Shared paging logic for streaming helpers such as `scanAll`, `xRangeAll`, and `xConsume`. Each builder accepts a `CIO` function that @@ -26,51 +27,108 @@ private[sage] object Paged { // a finite poll lets xTail and xConsume check for cancellation between blocking reads. val defaultPoll: BlockTimeout = BlockTimeout.After(FiniteDuration(5, TimeUnit.SECONDS)) + final case class Pages[S, A](init: S, step: Step[S, A]) + + // scan each target in sequence with its own node-local cursor. A cluster scan visits every master that owns slots. + def scanAll[K: KeyCodec](r: SharedRunner, pattern: Option[String], count: Option[Long], ofType: Option[RedisType]): Pages[ScanStep, K] = + Pages(ScanStep.Begin, acrossTargets(r.scanTargets)(target => cursor => target.run(Keys.scan[K](cursor, pattern, count, ofType)))) + + // HSCAN, SSCAN, or ZSCAN of one key + def scanKey[A](r: SharedRunner)(page: ScanCursor => Command[ScanPage[A]]): Pages[Option[ScanCursor], A] = + Pages(Some(ScanCursor.start), byCursor(cursor => r.run(page(cursor)))) + + def xRangeAll[K: KeyCodec, F: KeyCodec, V: ValueCodec]( + r: SharedRunner, + key: K, + start: StreamRangeId, + end: StreamRangeId, + batch: Long + ): Pages[Option[StreamRangeId], StreamEntry[F, V]] = + Pages(Some(start), byRange(batch)(from => r.run(Streams.xRange[K, F, V](key, from, end, Some(batch))))) + + def xAutoClaimAll[K: KeyCodec, F: KeyCodec, V: ValueCodec]( + r: SharedRunner, + key: K, + group: String, + consumer: String, + minIdle: FiniteDuration, + start: StreamId, + count: Option[Long] + ): Pages[Option[StreamId], StreamEntry[F, V]] = + Pages(Some(start), byAutoClaim(from => r.run(Streams.xAutoClaim[K, F, V](key, group, consumer, minIdle, from, count)))) + + def xTail[K: KeyCodec, F: KeyCodec, V: ValueCodec]( + r: SharedRunner, + key: K, + from: StreamId, + count: Option[Long], + block: BlockTimeout + ): Pages[StreamId, StreamEntry[F, V]] = + Pages(from, tail(last => r.run(Streams.xRead[K, F, V]((key, ReadId.After(last)))(count = count, block = Some(block))).map(_.flatMap(_._2)))) + + def xConsume[K: KeyCodec, F: KeyCodec, V: ValueCodec]( + r: SharedRunner, + group: String, + consumer: String, + key: K, + count: Option[Long], + block: BlockTimeout + ): Pages[Either[StreamId, Unit], StreamEntry[F, V]] = { + def read(id: GroupReadId, block: Option[BlockTimeout]): CIO[Vector[StreamEntry[F, V]]] = + r.run(Streams.xReadGroup[K, F, V](group, consumer)((key, id))(count = count, block = block)).map(_.flatMap(_._2)) + Pages(Left(StreamId.Zero), consume(after => read(GroupReadId.After(after), None), read(GroupReadId.New, Some(block)))) + } + /** * Pages through HSCAN, SSCAN, ZSCAN, or one SCAN target until the server returns a zero cursor. A filtered scan can return an empty page * with a non-zero cursor, so an empty page does not end iteration. */ - def byCursor[A](fetch: ScanCursor => CIO[ScanPage[A]]): Step[Option[ScanCursor], A] = { - case None => CIO.value(None) - case Some(cursor) => fetch(cursor).map(page => Some((page.items, page.next))) - } + def byCursor[A](fetch: ScanCursor => CIO[ScanPage[A]]): Step[Option[ScanCursor], A] = + resumable(cursor => fetch(cursor).map(page => (page.items, page.next))) /** * Scans every cluster target in turn, completing one node-local cursor before moving to the next. `Begin` discovers the targets, and an * empty target list ends the stream immediately. */ def acrossTargets[A](scanTargets: CIO[Vector[ScanTarget]])(fetch: ScanTarget => ScanCursor => CIO[ScanPage[A]]): Step[ScanStep, A] = { - case ScanStep.Begin => - scanTargets.map(targets => if (targets.isEmpty) None else Some((Vector.empty[A], ScanStep.Visit(ScanCursor.start, targets)))) - case ScanStep.Visit(cursor, remaining) => - fetch(remaining.head)(cursor).map { page => + case ScanStep.Begin => + scanTargets.map { + case target +: rest => Some((Vector.empty[A], ScanStep.Visit(ScanCursor.start, target, rest))) + case _ => None + } + case ScanStep.Visit(cursor, target, rest) => + fetch(target)(cursor).map { page => page.next match { - case Some(next) => Some((page.items, ScanStep.Visit(next, remaining))) - case None => Some((page.items, if (remaining.tail.isEmpty) ScanStep.End else ScanStep.Visit(ScanCursor.start, remaining.tail))) + case Some(next) => Some((page.items, ScanStep.Visit(next, target, rest))) + case None => + rest match { + case next +: others => Some((page.items, ScanStep.Visit(ScanCursor.start, next, others))) + case _ => Some((page.items, ScanStep.End)) + } } } - case ScanStep.End => CIO.value(None) + case ScanStep.End => CIO.value(None) } /** * XRANGE paging: advance past the last id each page; a short page (fewer than `batch`) or an empty page ends the stream. */ - def byRange[F, V](batch: Long)(fetch: StreamRangeId => CIO[Vector[StreamEntry[F, V]]]): Step[Option[StreamRangeId], StreamEntry[F, V]] = { - case None => CIO.value(None) - case Some(from) => - fetch(from).map { entries => - if (entries.isEmpty) None - else Some((entries, if (entries.length < batch) None else Some(StreamRangeId.Exclusive(entries.last.id)))) - } - } + def byRange[F, V](batch: Long)(fetch: StreamRangeId => CIO[Vector[StreamEntry[F, V]]]): Step[Option[StreamRangeId], StreamEntry[F, V]] = + resumable(from => + fetch(from).map(entries => (entries, if (entries.isEmpty || entries.length < batch) None else Some(StreamRangeId.Exclusive(entries.last.id)))) + ) /** * Pages through XAUTOCLAIM until its cursor returns to `StreamId.Zero`. Entries whose data has already been deleted are omitted. */ - def byAutoClaim[F, V](fetch: StreamId => CIO[XAutoClaimResult[F, V]]): Step[Option[StreamId], StreamEntry[F, V]] = { - case None => CIO.value(None) - case Some(from) => - fetch(from).map(result => Some((result.entries.filter(_.fields.nonEmpty), if (result.cursor == StreamId.Zero) None else Some(result.cursor)))) + def byAutoClaim[F, V](fetch: StreamId => CIO[XAutoClaimResult[F, V]]): Step[Option[StreamId], StreamEntry[F, V]] = + resumable(from => + fetch(from).map(result => (result.entries.filter(_.fields.nonEmpty), Option.when(result.cursor != StreamId.Zero)(result.cursor))) + ) + + private def resumable[C, A](fetch: C => CIO[(Vector[A], Option[C])]): Step[Option[C], A] = { + case None => CIO.value(None) + case Some(c) => fetch(c).map(Some(_)) } /** diff --git a/sage-client/shared/src/main/scala/sage/client/internal/Placement.scala b/sage-client/shared/src/main/scala/sage/client/internal/Placement.scala deleted file mode 100644 index 940602ab..00000000 --- a/sage-client/shared/src/main/scala/sage/client/internal/Placement.scala +++ /dev/null @@ -1,97 +0,0 @@ -package sage.client.internal - -import java.util.concurrent.locks.ReentrantLock - -import scala.collection.mutable -import scala.util.control.NonFatal - -import SubscriptionConnection.{Kind, Sink} - -import sage.cluster.Node - -/** - * Tracks which node handles each shard channel for one subscription. A node that owns several slot ranges needs a separate `SSUBSCRIBE` - * for each slot because a cross-slot subscription returns `CROSSSLOT`. Placement plans therefore group channels by node and slot. - * - * Plan updates are atomic with respect to unsubscribe and reassignment: - * - [[place]] performs the initial subscription. It records successful groups and leaves failed groups for the caller to retry. - * - [[reconcile]] removes obsolete groups, attaches new ones, records successful changes, and reports whether any attachment failed. - */ -final private[internal] class Placement(sink: Sink, requested: Vector[String]) { - - private val lock = new ReentrantLock() - private var placedAt: Map[Node, Set[String]] = Map.empty - - private inline def locked[A](inline body: A): A = { - lock.lock() - try body - finally lock.unlock() - } - - def place(plan: Placement.Plan, conns: Placement.Conns): Unit = - locked { - plan.foreach { case (node, groups) => - conns.ensure(node).foreach { conn => - groups.foreach { group => - // a concurrent eviction can close `conn` before attach. Leave the channel pending for a retry instead of failing reconciliation. - try { - conn.attach(sink, group, Kind.Shard) - placedAt = placedAt.updatedWith(node)(prev => Some(prev.getOrElse(Set.empty) ++ group)) - } catch { case NonFatal(_) => () } - } - } - } - } - - // Return true when any connection or attachment failed. The caller omits unowned slots from the plan and retries until every requested - // channel is attached. - def reconcile(plan: Placement.Plan, conns: Placement.Conns): Boolean = - locked { - val desired = plan.view.mapValues(_.flatten.toSet).toMap - var incomplete = false - placedAt.foreach { case (node, had) => - val gone = (had -- desired.getOrElse(node, Set.empty)).toVector - if (gone.nonEmpty) conns.get(node).foreach(_.detach(sink, gone, Kind.Shard)) - } - val actual = mutable.HashMap.empty[Node, Set[String]] - plan.foreach { case (node, groups) => - conns.ensure(node) match { - case None => incomplete = true - case Some(conn) => - groups.foreach { group => - try { - conn.attach(sink, group, Kind.Shard) - actual.update(node, actual.getOrElse(node, Set.empty) ++ group) - } catch { case NonFatal(_) => incomplete = true } - } - } - } - placedAt = actual.toMap - incomplete - } - - // count distinct attached channels across all nodes. Counting each node separately could hide a missing channel when another is recorded twice. - def fullyPlaced: Boolean = locked(placedAt.valuesIterator.flatten.toSet.size) >= requested.distinct.size -} - -private[internal] object Placement { - - // a Node's channels grouped so each inner Vector is one SSUBSCRIBE (one Slot), never spanning Slots - type Plan = Map[Node, Vector[Vector[String]]] - - /** - * The Sharded Subscription Connections a placement attaches to, looked up under the manager's lock. `ensure` is empty once the manager is closed. - */ - trait Conns { - def ensure(node: Node): Option[ShardConn] - def get(node: Node): Option[ShardConn] - } - - /** - * What a placement needs of a connection: register/unregister a sink under names. [[SubscriptionConnection]] is the production adapter. - */ - trait ShardConn { - def attach(sink: Sink, names: Vector[String], kind: Kind): Unit - def detach(sink: Sink, names: Vector[String], kind: Kind): Boolean - } -} diff --git a/sage-client/shared/src/main/scala/sage/client/internal/Poison.scala b/sage-client/shared/src/main/scala/sage/client/internal/Poison.scala deleted file mode 100644 index 310a80fe..00000000 --- a/sage-client/shared/src/main/scala/sage/client/internal/Poison.scala +++ /dev/null @@ -1,23 +0,0 @@ -package sage.client.internal - -import sage.protocol.Frame - -/** - * A `-READONLY` reply means that the server became a replica without dropping the socket. The connection still looks healthy but rejects - * writes, so the client must replace it. Only `READONLY` poisons a connection here. `LOADING` resolves without a reconnect, and cluster - * reply codes have separate handling. - */ -private[internal] object Poison { - - def isReadonly(frame: Frame): Boolean = - frame match { - case Frame.SimpleError(message) => errorCode(message) == "READONLY" - case Frame.BulkError(message) => errorCode(message.asUtf8String) == "READONLY" - case _ => false - } - - private def errorCode(message: String): String = { - val space = message.indexOf(' ') - if (space < 0) message else message.substring(0, space) - } -} diff --git a/sage-client/shared/src/main/scala/sage/client/internal/RateLimitExecutor.scala b/sage-client/shared/src/main/scala/sage/client/internal/RateLimitExecutor.scala index 00faf720..53c6f54b 100644 --- a/sage-client/shared/src/main/scala/sage/client/internal/RateLimitExecutor.scala +++ b/sage-client/shared/src/main/scala/sage/client/internal/RateLimitExecutor.scala @@ -2,7 +2,7 @@ package sage.client.internal import kyo.compat.* -import sage.SageException.{InvalidArgument, ServerError} +import sage.SageException.InvalidArgument import sage.commands.Command import sage.ratelimit.{Decision, RateLimiter} @@ -20,9 +20,6 @@ final private[client] class RateLimitExecutor[K](definition: RateLimiter[K]) { definition.validate(cost) match { case Some(problem) => CIO.fail(InvalidArgument(problem)) case None => - runner.run(definition.evalSha(subject, cost, peek)).recover { - case ServerError(code, _) if code == "NOSCRIPT" => runner.run(definition.evalScript(subject, cost, peek)) - case other => CIO.fail(other) - } + Client.withScriptFallback(cached => runner.run(definition.eval(cached, subject, cost, peek))) } } diff --git a/sage-client/shared/src/main/scala/sage/client/internal/ReadRouting.scala b/sage-client/shared/src/main/scala/sage/client/internal/ReadRouting.scala index 505a4ad9..38f78e05 100644 --- a/sage-client/shared/src/main/scala/sage/client/internal/ReadRouting.scala +++ b/sage-client/shared/src/main/scala/sage/client/internal/ReadRouting.scala @@ -16,7 +16,7 @@ import sage.commands.Command */ private[client] object ReadRouting { - final case class Picked(node: Node, client: NodeClient, remaining: Vector[Node]) + final case class Picked(node: Node, client: MultiplexedConnection, remaining: Vector[Node]) // Allow ordinary non-blocking reads. Exclude writes, blocking reads, and cursor-bound scans because their cursors are node-local. Cached // reads apply their own routing rule before calling this method. @@ -58,18 +58,6 @@ final private[client] class ReadRouting( private val cursors = new ConcurrentHashMap[Node, AtomicInteger]() - private enum CandidateState { - case Unknown - case Unavailable - case Connected(client: NodeClient) - } - - private enum Selection { - case Found(picked: ReadRouting.Picked) - case NeedsEstablish - case Exhausted - } - /** * The candidates for `master`'s replicas under the policy, advancing that master's cursor once. */ @@ -101,17 +89,10 @@ final private[client] class ReadRouting( ): Unit = candidates match { case node +: rest => - val pool = poolFor(node, master) - val existing = pool.existing(node) - if (existing != null) - if (existing.isLive) submit(existing, node, command, rest, complete)(onFault) - else walk(command, rest, master, complete)(onFault) - else - scheduler.offload { - val nc = pool.getOrEstablishOrNull(node) - if (nc == null || !nc.isLive) walk(command, rest, master, complete)(onFault) - else submit(nc, node, command, rest, complete)(onFault) - } + poolFor(node, master).withClient(node)(walk(command, rest, master, complete)(onFault)) { nc => + if (!nc.isLive) walk(command, rest, master, complete)(onFault) + else submit(nc, node, command, rest, complete)(onFault) + } case _ => triggerRefresh() complete(Failure(NotConnected())) @@ -123,58 +104,24 @@ final private[client] class ReadRouting( * establishment is offloaded. */ def pickOne(candidates: Vector[Node], master: Node)(onPick: Option[ReadRouting.Picked] => Unit): Unit = - select(candidates, master, existingCandidate) match { - case Selection.Found(picked) => onPick(Some(picked)) - case Selection.Exhausted => onPick(None) - case Selection.NeedsEstablish => - scheduler.offload { - // retry the full order, allowing a previously disconnected candidate to reconnect before checking the first unknown one - select(candidates, master, establishCandidate) match { - case Selection.Found(picked) => onPick(Some(picked)) - case _ => onPick(None) - } + candidates match { + case node +: rest => + poolFor(node, master).withClient(node)(pickOne(rest, master)(onPick)) { nc => + if (nc.isLive) onPick(Some(ReadRouting.Picked(node, nc, rest))) else pickOne(rest, master)(onPick) } + case _ => onPick(None) } - private def submit[A](nc: NodeClient, node: Node, command: Command[A], rest: Vector[Node], complete: Try[A] => Unit)( + private def submit[A](nc: MultiplexedConnection, node: Node, command: Command[A], rest: Vector[Node], complete: Try[A] => Unit)( onFault: (Node, Throwable, Vector[Node]) => Unit ): Unit = nc.submit[A]( command, - asking = false, { - case success @ Success(_) => - Events.attributeNode(complete, node) - complete(success) + case success @ Success(_) => Events.completeAt(complete, node)(success) case Failure(error) => scheduler.offload(onFault(node, error, rest)) } ) - private def select(candidates: Vector[Node], master: Node, lookup: (NodePool, Node) => CandidateState): Selection = { - var remaining = candidates - while (remaining.nonEmpty) { - val node = remaining.head - val pool = poolFor(node, master) - lookup(pool, node) match { - case CandidateState.Unknown => return Selection.NeedsEstablish - case CandidateState.Unavailable => () - case CandidateState.Connected(nc) => - if (nc.isLive) return Selection.Found(ReadRouting.Picked(node, nc, remaining.tail)) - } - remaining = remaining.tail - } - Selection.Exhausted - } - - private def existingCandidate(pool: NodePool, node: Node): CandidateState = { - val nc = pool.existing(node) - if (nc == null) CandidateState.Unknown else CandidateState.Connected(nc) - } - - private def establishCandidate(pool: NodePool, node: Node): CandidateState = { - val nc = pool.getOrEstablishOrNull(node) - if (nc == null) CandidateState.Unavailable else CandidateState.Connected(nc) - } - private def poolFor(node: Node, master: Node): NodePool = if (node == master) masterPool else replicaPool } diff --git a/sage-client/shared/src/main/scala/sage/client/internal/RefreshThrottle.scala b/sage-client/shared/src/main/scala/sage/client/internal/RefreshThrottle.scala index 55ffe6ed..fdc382a2 100644 --- a/sage-client/shared/src/main/scala/sage/client/internal/RefreshThrottle.scala +++ b/sage-client/shared/src/main/scala/sage/client/internal/RefreshThrottle.scala @@ -5,146 +5,105 @@ import java.util.concurrent.locks.ReentrantLock import scala.concurrent.duration.* /** - * Coordinates discovery refreshes for the cluster and master-replica runtimes. Only one refresh runs at a time. [[apply]] waits for a current - * refresh to finish and then returns. Non-forced calls to `apply` and `trigger` within `minRefreshMs` of the previous refresh are skipped. - * `request` retains work until it can run. A forced call ignores the minimum interval. The first refresh can run immediately. - * - * It also starts and stops the optional background polling task. + * Coordinates discovery refreshes (`work`) for the cluster and master-replica runtimes. Only one refresh runs at a time. [[apply]] waits for a current + * refresh to finish and then returns. A non-forced call to `apply` within `minRefreshMs` of the previous refresh is skipped. + * `request` retains a refresh until it can run. A forced call ignores the minimum interval. The first refresh can run immediately. + * After [[stop]], no refresh starts, including one that was already queued. */ -final private[client] class RefreshThrottle(scheduler: Scheduler, minRefreshMs: Long) { +final private[client] class RefreshThrottle(scheduler: Scheduler, minRefreshMs: Long, work: () => Unit) { + import RefreshThrottle.Phase - private val lock = new ReentrantLock() - private val done = lock.newCondition() - // volatile so `throttled` can answer without the lock; every mutation still happens under it - @volatile private var refreshing = false - @volatile private var lastRefreshMs = scheduler.nowMillis - minRefreshMs - private var ticker = null: Scheduler.Cancelable - private var stopped = false - private var requested: () => Unit = null - private var requestScheduled = false + private val lock = new ReentrantLock() + private val done = lock.newCondition() + private var lastRefreshMs = scheduler.nowMillis - minRefreshMs + // volatile so a request already pending returns without the lock; every change still happens under it + @volatile private var phase: Phase = Phase.Idle - def apply(force: Boolean)(work: => Unit): Unit = - if (claim(force, wait = true)) run(work) - - /** - * Schedules a non-forced refresh when no refresh is active and the minimum interval has passed. It returns immediately while another - * refresh is active, a retained request is already scheduled, or the interval has not passed. Routing may call this for every read when no - * replica is available. Taking a pre-created callback avoids allocating a new closure for each call that returns without scheduling work. - */ - def trigger(work: () => Unit): Unit = - if (claim(force = false, wait = false)) - try scheduler.after(Duration.Zero)(run(work())) - catch { - case error: Throwable => - finish() - throw error - } - - /** - * Retains a refresh request until it can run. Requests during a refresh or its minimum interval coalesce into one later refresh. - */ - def request(work: () => Unit): Unit = { + private inline def locked[A](inline body: A): A = { lock.lock() - try - if (!stopped) { - requested = work - scheduleRequest() - } + try body finally lock.unlock() } - // Called under the mutex so concurrent confirmation failures share one scheduled refresh. - private def scheduleRequest(): Unit = - if (!refreshing && !requestScheduled && requested != null && !stopped) { - requestScheduled = true - val delayMillis = math.max(0L, minRefreshMs - (scheduler.nowMillis - lastRefreshMs)) - try scheduler.after(delayMillis.millis)(runRequested()) - catch { - case error: Throwable => - requestScheduled = false - throw error - } - } - - private def runRequested(): Unit = { - lock.lock() - val work = try { - requestScheduled = false - if (stopped || refreshing) null - else if (scheduler.nowMillis - lastRefreshMs < minRefreshMs) { - scheduleRequest() - null - } else { - val next = requested - requested = null - if (next != null) refreshing = true - next - } - } finally lock.unlock() - if (work != null) run(work()) - } + // wait for any current refresh to finish before callers read `topologyRef` + def apply(force: Boolean): Unit = + if (claim(force)) run() /** - * Starts the background poll when an interval is configured; `tick` runs on the timer thread, so it must only queue work. + * Retains a refresh request until it can run. Requests during a refresh or its minimum interval coalesce into one later refresh. Routing + * may call this for every read when no replica is available. */ - def startPolling(interval: Option[FiniteDuration])(tick: => Unit): Unit = - interval.foreach { period => - val handle = scheduler.every(period)(tick) - lock.lock() - val keep = - try - if (stopped) false - else { - ticker = handle - true - } - finally lock.unlock() - if (!keep) handle.cancel() + def request(): Unit = + phase match { + case Phase.Idle | Phase.Running(false) => + locked(phase match { + case Phase.Idle => schedule() + case Phase.Running(false) => phase = Phase.Running(again = true) + case _ => () + }) + case _ => () } - def stopPolling(): Unit = { - lock.lock() - val handle = - try { - stopped = true - requested = null - val current = ticker - ticker = null - current - } finally lock.unlock() - if (handle != null) handle.cancel() + // Must hold lock. A timer whose Scheduled phase was replaced does nothing when it fires. + private def schedule(): Unit = { + val mine = Phase.Scheduled() + phase = mine + val delayMillis = math.max(0L, minRefreshMs - (scheduler.nowMillis - lastRefreshMs)) + onThrow(scheduler.after(delayMillis.millis)(runScheduled(mine)))(_ => phase = Phase.Idle) } - private def claim(force: Boolean, wait: Boolean): Boolean = { - if (!force && !wait && throttled) return false - lock.lock() - try - if (refreshing && wait) { - while (refreshing) done.awaitUninterruptibly() - false - } else if (refreshing || (!wait && requestScheduled) || (!force && scheduler.nowMillis - lastRefreshMs < minRefreshMs)) false - else { - refreshing = true - true - } - finally lock.unlock() - } + private def runScheduled(mine: Phase): Unit = + if ( + locked((phase eq mine) && { + val due = scheduler.nowMillis - lastRefreshMs >= minRefreshMs + if (due) phase = Phase.Running(again = false) else schedule() + due + }) + ) run() + + def stop(): Unit = locked { phase = Phase.Stopped } - private def throttled: Boolean = refreshing || scheduler.nowMillis - lastRefreshMs < minRefreshMs + // a refresh claimed while one is scheduled keeps the scheduled one as a follow-up + private def claim(force: Boolean): Boolean = + locked(phase match { + case Phase.Stopped => false + case Phase.Running(_) => + while (phase.isRunning) done.awaitUninterruptibly() + false + case waiting => + val due = force || scheduler.nowMillis - lastRefreshMs >= minRefreshMs + if (due) phase = Phase.Running(again = waiting != Phase.Idle) + due + }) - private def run(work: => Unit): Unit = - try work + // a refresh queued before stop must not open new connections + private def run(): Unit = + try if (phase != Phase.Stopped) work() finally finish() - private def finish(): Unit = { - lock.lock() - try { - // Record the completion time before publishing refreshing = false. Volatile write ordering ensures a lock-free throttled check that sees - // false also sees the updated completion time. + private def finish(): Unit = + locked { lastRefreshMs = scheduler.nowMillis - refreshing = false done.signalAll() - scheduleRequest() - } finally lock.unlock() + phase match { + case Phase.Running(true) => schedule() + case Phase.Running(false) => phase = Phase.Idle + case _ => () + } + } +} + +private object RefreshThrottle { + + // Scheduled has an identity per timer; Running(again) records a request that arrived during the refresh + enum Phase { + case Idle, Stopped + case Scheduled() + case Running(again: Boolean) + + def isRunning: Boolean = this match { + case Running(_) => true + case _ => false + } } } diff --git a/sage-client/shared/src/main/scala/sage/client/internal/Scheduler.scala b/sage-client/shared/src/main/scala/sage/client/internal/Scheduler.scala index 8b3593fa..894ff0fb 100644 --- a/sage-client/shared/src/main/scala/sage/client/internal/Scheduler.scala +++ b/sage-client/shared/src/main/scala/sage/client/internal/Scheduler.scala @@ -1,10 +1,13 @@ package sage.client.internal import java.util.concurrent.{Executors, ScheduledExecutorService, ScheduledFuture, ThreadLocalRandom, TimeUnit} +import java.util.concurrent.locks.ReentrantLock -import scala.concurrent.duration.{Duration, FiniteDuration} +import scala.concurrent.duration.* import scala.util.control.NonFatal +import sage.client.BackoffConfig + /** * The clock and timer abstraction under the reconnect loop and the watchdog, injected so tests drive virtual time. `nowMillis` is monotonic. A * one-shot `after` task may block (connect, bootstrap); a periodic `every` tick must not. @@ -23,6 +26,51 @@ private[client] trait Scheduler { def every(interval: FiniteDuration)(task: => Unit): Scheduler.Cancelable final def offload(task: => Unit): Unit = after(Duration.Zero)(task) + + // Reconnects and cluster-redirect retries share this formula: exponential backoff capped at maxDelay, then full jitter in [0, base]. + final def backoffMillis(config: BackoffConfig, attempt: Int): Long = { + val capped = config.maxDelay.toMillis + val raw = config.initialDelay.toMillis.toDouble * math.pow(config.multiplier, attempt.toDouble) + val base = if (raw.isInfinite || raw >= capped.toDouble) capped else math.max(0L, raw.toLong) + jitterMillis(base + 1) + } + + final def afterBackoff(config: BackoffConfig, attempt: Int)(task: => Unit): Unit = after(backoffMillis(config, attempt).millis)(task) +} + +/** + * The reconnect loop of one connection or subscription manager, guarded by its owner's `lock`. Attempts keep counting across connections + * that drop before staying live for `maxDelay`, so a connection the server keeps closing backs off instead of reconnecting at the initial + * delay. At most one retry waits at a time; a loss while one waits joins it. + */ +final private[internal] class Reconnects(scheduler: Scheduler, config: BackoffConfig, lock: ReentrantLock) { + private var attempt = 0 + private var liveSince = -1L + private var waiting = false + + def live(): Unit = liveSince = scheduler.nowMillis + + // Must hold lock. After the backoff, runs `retry` if `wanted` still holds; a failed retry is reported and scheduled again while it holds. + // With `immediately`, the first attempt runs without waiting. + def schedule(wanted: => Boolean, onFailure: Throwable => Unit, immediately: Boolean = false)(retry: => Unit): Unit = + if (!waiting) { + if (liveSince >= 0L && scheduler.nowMillis - liveSince >= config.maxDelay.toMillis) attempt = 0 + val delay = if (immediately && attempt == 0) 0L else scheduler.backoffMillis(config, attempt) + attempt += 1 + liveSince = -1L + waiting = true + scheduler.after(delay.millis) { + if (holding { waiting = false; wanted }) + try retry + catch { case NonFatal(error) => holding(if (wanted) { onFailure(error); schedule(wanted, onFailure)(retry) }) } + } + } + + private def holding[A](body: => A): A = { + lock.lock() + try body + finally lock.unlock() + } } private[client] object Scheduler { diff --git a/sage-client/shared/src/main/scala/sage/client/internal/SocketTransport.scala b/sage-client/shared/src/main/scala/sage/client/internal/SocketTransport.scala index 28ddff09..68094596 100644 --- a/sage-client/shared/src/main/scala/sage/client/internal/SocketTransport.scala +++ b/sage-client/shared/src/main/scala/sage/client/internal/SocketTransport.scala @@ -46,17 +46,13 @@ final private[client] class SocketTransport private ( private[internal] val writer: Thread = Thread.ofVirtual().name(s"sage-writer-$id").unstarted(() => writeLoop()) def start(): Unit = { - try { + onThrow { socket.setTcpNoDelay(true) socket.connect(new InetSocketAddress(host, port), connectTimeoutMillis) // resolves the hostname here, fresh on every attempt socket.setSoTimeout(connectTimeoutMillis) // bound a TLS handshake's reads like the connect ioSocket = upgrade(socket) socket.setSoTimeout(0) // steady state: blocking reads with no timeout - } catch { - case NonFatal(error) => - terminate() - throw error - } + }(_ => terminate()) val launched = locked { if (closed.get()) false else { @@ -70,15 +66,20 @@ final private[client] class SocketTransport private ( } def send(item: Transport.Item): Unit = { - queue.put(item) + // offer, unlike put, does not throw for an interrupted caller + queue.offer(item): Unit // terminate may empty the queue just before this item is added. Check closed after the add so this item is failed as well. if (closed.get()) drainQueue() } + // An I/O thread closing from a callback returns without waiting for the other I/O thread, which may own teardown and be joining it. def close(): Unit = { terminate() - if (Thread.currentThread() ne reader) reader.join() - if (Thread.currentThread() ne writer) writer.join() + val current = Thread.currentThread() + if ((current ne reader) && (current ne writer)) { + join(reader) + join(writer) + } } private def readLoop(): Unit = { @@ -177,21 +178,27 @@ final private[client] class SocketTransport private ( case _: IOException => () } if (locked(threadsStarted)) { - if (Thread.currentThread() ne writer) { - writer.interrupt() - writer.join() - } + if (Thread.currentThread() ne writer) writer.interrupt() + join(writer) // Stop the reader before onClosed; otherwise, an in-flight reply can race cleanup of the consumer's pending commands (#94). // Interrupting releases a reader waiting for backpressure before the join. - if (Thread.currentThread() ne reader) { - reader.interrupt() - reader.join() - } + if (Thread.currentThread() ne reader) reader.interrupt() + join(reader) } drainQueue() onClosed() } + // Waits even when the caller is interrupted, so the cleanup that follows always runs, then restores the interrupt. + private def join(thread: Thread): Unit = + if (Thread.currentThread() ne thread) { + var interrupted = false + while (thread.isAlive) + try thread.join() + catch { case _: InterruptedException => interrupted = true } + if (interrupted) Thread.currentThread().interrupt() + } + private def drainQueue(): Unit = { var item = queue.poll() while (item != null) { diff --git a/sage-client/shared/src/main/scala/sage/client/internal/Subscription.scala b/sage-client/shared/src/main/scala/sage/client/internal/Subscription.scala index b7254289..cc665ccd 100644 --- a/sage-client/shared/src/main/scala/sage/client/internal/Subscription.scala +++ b/sage-client/shared/src/main/scala/sage/client/internal/Subscription.scala @@ -1,8 +1,9 @@ package sage.client.internal /** - * The shared subscription API used by backend-specific streams. `next` returns the next message or `None` when the subscription ends, and - * `close` unsubscribes. Each backend calls `close` from its stream finalizer when the stream's scope closes. + * The shared subscription API used by backend-specific streams. `next` returns the next message or `None` when the subscription ends. It + * fails with the server's error when the server refuses a subscribed name, including when a reconnect subscribes it again. `close` + * unsubscribes. Each backend calls `close` from its stream finalizer when the stream's scope closes. */ trait Subscription[F[_], A] { diff --git a/sage-client/shared/src/main/scala/sage/client/internal/SubscriptionConnection.scala b/sage-client/shared/src/main/scala/sage/client/internal/SubscriptionConnection.scala index b582ee20..75ba550e 100644 --- a/sage-client/shared/src/main/scala/sage/client/internal/SubscriptionConnection.scala +++ b/sage-client/shared/src/main/scala/sage/client/internal/SubscriptionConnection.scala @@ -1,19 +1,17 @@ package sage.client.internal import java.util.concurrent.TimeUnit -import java.util.concurrent.atomic.AtomicReference import java.util.concurrent.locks.ReentrantLock import scala.collection.mutable -import scala.concurrent.duration.* import scala.util.{Failure, Success, Try} import scala.util.control.NonFatal -import sage.{Bytes, SageEvent, SageException} -import sage.SageException.{ConnectionFailed, ConnectionLost, NotConnected} -import sage.client.{BackoffConfig, WatchdogConfig} +import sage.SageEvent +import sage.SageException.{NotConnected, ServerError} +import sage.client.SageConfig import sage.cluster.Node -import sage.commands.{Command, Connection, Pubsub, Reply} +import sage.commands.Pubsub import sage.protocol.Frame /** @@ -22,49 +20,33 @@ import sage.protocol.Frame * connection also wait, but command connections are unaffected. The watchdog does not close the connection while its reader is waiting * on this backpressure. * - * In standalone and master-replica mode, the connection owns its subscribers and restores them after reconnecting. `onReconnect` lets a - * master-replica client discover the current master before each attempt. In cluster mode, [[ClusterSubscriptions]] owns the subscribers - * and assigns them to nodes. When a cluster connection closes, the manager assigns its subscribers again using the latest topology. + * With [[OnLoss.Reconnect]] (standalone, master-replica, and cluster classic subscriptions), the connection owns its subscribers and restores + * them after reconnecting; `beforeAttempt` lets the client discover the current master before each attempt. For shard channels in a + * cluster, [[ClusterSubscriptions]] owns the subscribers and assigns them to nodes. When such a connection closes, the manager assigns its + * subscribers again using the latest topology. */ final private[client] class SubscriptionConnection( factory: MultiplexedConnection.TransportFactory, - bootstrap: Vector[Command[?]], scheduler: Scheduler, - backoff: BackoffConfig, - watchdog: WatchdogConfig, - connectTimeoutMillis: Long, - bufferSize: Int, + config: SageConfig, isLive: () => Boolean, - cluster: Boolean = false, - onTerminated: () => Unit = () => (), - onReconnect: () => Unit = () => (), - events: Events = Events.disabled, - node: Option[Node] = None -) extends Placement.ShardConn { + onLoss: SubscriptionConnection.OnLoss = SubscriptionConnection.OnLoss.Reconnect(() => (), Events.disabled) +) extends ClusterSubscriptions.ShardConn + with SubscriptionConnection.PubSub { import SubscriptionConnection.* - private val lock = new ReentrantLock() - private val established = lock.newCondition() - private val confirmed = lock.newCondition() - private var state: State = State.Idle - private var current: Conn = null + private enum State { + case Idle, Establishing, Reconnecting, Closed + case Live(conn: Conn) + } + + private val lock = new ReentrantLock() + private val changed = lock.newCondition() + private var state: State = State.Idle // connections still being opened; a set, since a reconnect and a fresh attach can be establishing at once - private val establishing = mutable.Set.empty[Conn] - private val sinksByKind = Array.fill(Kind.values.length)(mutable.HashMap.empty[String, mutable.LinkedHashSet[Sink]]) - - // The server confirms each subscribed name with one push frame in send order. Standalone subscriptions return after confirmation. Cluster - // attachment waits up to the connection timeout and may return before confirmation, which can arrive later. These counters are guarded by `lock`. - private var subscribeSent: Long = 0L - private var subscribeConfirmed: Long = 0L - // Set this generation's resubscribe acknowledgement count in goLive before waking waiters. Later subscriptions do not change what an - // existing waiter expects. - private var liveTarget: Long = -1L - - private var watchdogHandle: Scheduler.Cancelable = null - @volatile private var readerBlocked: Boolean = false - @volatile private var lastReplyAtMillis: Long = scheduler.nowMillis - @volatile private var lastBackpressureMillis: Long = 0L - @volatile private var pingSentAtMillis: Long = 0L + private val establishing = mutable.Set.empty[Conn] + private val sinksByKind = Array.fill(Kind.values.length)(mutable.HashMap.empty[String, Name]) + private val reconnects = new Reconnects(scheduler, config.reconnect, lock) private inline def locked[A](inline body: A): A = { lock.lock() @@ -72,25 +54,24 @@ final private[client] class SubscriptionConnection( finally lock.unlock() } - private def sinksFor(kind: Kind): mutable.HashMap[String, mutable.LinkedHashSet[Sink]] = sinksByKind(kind.ordinal) + private def sinksFor(kind: Kind): mutable.HashMap[String, Name] = sinksByKind(kind.ordinal) + + // must hold lock. A refused, dropped or emptied name is replaced by a new Name when subscribed again. + private def registered(name: Name): Boolean = sinksFor(name.kind).get(name.name).exists(_ eq name) // --- standalone conveniences: the connection owns the sink ------------------------------------------------------------------------------- - def subscribeChannels(channels: Vector[String]): RawSubscription = ownedSubscription(channels, Kind.Channel) + def subscribeChannels(channels: Vector[String]): RawSubscription = owned(channels, Kind.Channel, failIfUnconfirmed = true) - def subscribePatterns(patterns: Vector[String]): RawSubscription = ownedSubscription(patterns, Kind.Pattern) + def subscribePatterns(patterns: Vector[String]): RawSubscription = owned(patterns, Kind.Pattern, failIfUnconfirmed = true) - def subscribeShard(channels: Vector[String]): RawSubscription = ownedSubscription(channels, Kind.Shard) + def subscribeShard(channels: Vector[String]): RawSubscription = owned(channels, Kind.Shard, failIfUnconfirmed = true) - private def ownedSubscription(names: Vector[String], kind: Kind): RawSubscription = { - val sink = new Sink(names, kind, bufferSize) + // cluster classic subscriptions pass failIfUnconfirmed = false, since confirmation in a cluster is best-effort + def owned(names: Vector[String], kind: Kind, failIfUnconfirmed: Boolean): RawSubscription = { + val sink = new Sink(names, kind, config.pubsub.bufferSize) // closeOwned removes a sink that attachInternal registered before awaitActive failed, preventing it from being restored after reconnecting - try attachInternal(sink, names, kind, failIfUnconfirmed = true) - catch { - case e: Throwable => - closeOwned(sink) - throw e - } + onThrow { attachInternal(sink, names, failIfUnconfirmed); sink.failure.foreach(throw _) }(_ => closeOwned(sink)) new RawSubscription(sink, () => closeOwned(sink)) } @@ -98,36 +79,31 @@ final private[client] class SubscriptionConnection( /** * Registers `sink` under `names` and subscribes names that are not already active. It waits up to the connection timeout for confirmation; - * if the timeout expires, the method returns and the connection can confirm later. For shard subscriptions, the caller must pass - * names from one slot so a single `SSUBSCRIBE` does not cross slots. + * if the timeout expires, the method returns and the connection can confirm later. A name the server rejects leaves the registry. */ - def attach(sink: Sink, names: Vector[String], kind: Kind): Unit = attachInternal(sink, names, kind, failIfUnconfirmed = false) + def attach(sink: Sink, names: Vector[String]): Unit = attachInternal(sink, names, failIfUnconfirmed = false) - private def attachInternal(sink: Sink, names: Vector[String], kind: Kind, failIfUnconfirmed: Boolean): Unit = { - var doEstablish = false - // the acknowledgement count this attachment waits for; -1 means goLive will send the subscription and set the target - var confirmTarget = -1L + private def attachInternal(sink: Sink, names: Vector[String], failIfUnconfirmed: Boolean): Unit = { + var doEstablish = false + var waitFor = Vector.empty[Name] lock.lock() try { var settled = false while (!settled) state match { case State.Closed => throw NotConnected() - case State.Establishing => established.await() - case State.Live => - val fresh = register(sink, names, kind) - // when every name is already subscribed, wait only for acknowledgements already confirmed and ignore other pending subscriptions - if (fresh.nonEmpty) { - sendSubscribe(current, kind, fresh) - confirmTarget = subscribeSent - } else confirmTarget = subscribeConfirmed + case State.Establishing => changed.await() + case State.Live(conn) => + val (all, created) = register(sink, names) + conn.subscribe(created) + waitFor = all settled = true case State.Reconnecting => - register(sink, names, kind) // the next successful reconnect resubscribes everything currently registered + waitFor = register(sink, names)._1 // the next successful reconnect resubscribes everything currently registered settled = true case State.Idle => if (!isLive()) throw NotConnected() - register(sink, names, kind) + waitFor = register(sink, names)._1 state = State.Establishing doEstablish = true settled = true @@ -135,130 +111,116 @@ final private[client] class SubscriptionConnection( } finally lock.unlock() if (doEstablish) - try goLive(establish()) - catch { - case e: Throwable => - // If establishment fails after registering the sink, remove it while the connection remains in the Establishing state; a concurrent - // close or goLive call may have already changed the state and completed the cleanup - locked(if (state == State.Establishing) { - deregister(sink, names, kind) - state = State.Idle - established.signalAll() - }) - throw e + onThrow(goLive(establish())) { _ => + // If establishment fails after registering the sink, remove it while the connection remains in the Establishing state; a concurrent + // close or goLive call may have already changed the state and completed the cleanup + locked(if (state == State.Establishing) { + deregister(sink, names) + state = State.Idle + changed.signalAll() + }) } - awaitActive(failIfUnconfirmed, confirmTarget) + awaitActive(waitFor, failIfUnconfirmed) } /** - * Deregisters `sink` from `names` and unsubscribes the names left with no subscriber. Returns true when the connection now holds no sinks - * at all, so the manager can evict and close it. + * Deregisters `sink` from `names` and unsubscribes the names left with no subscriber. */ - def detach(sink: Sink, names: Vector[String], kind: Kind): Boolean = + def detach(sink: Sink, names: Vector[String]): Unit = locked { - val emptied = deregister(sink, names, kind) - // best-effort: swallow a failed (or interrupted) unsubscribe write so the caller still learns emptiness and terminates the sink - if (emptied.nonEmpty && state == State.Live) - try current.send(kind.unsubscribeWire(emptied)) - catch { case NonFatal(_) | _: InterruptedException => () } - isEmptyUnlocked + val emptied = deregister(sink, names) + liveConn.foreach(_.unsubscribe(sink.kind, emptied)) } + def namesOf(sink: Sink): Vector[String] = + locked(sink.names.distinct.filter(name => sinksFor(sink.kind).get(name).exists(_.sinks.contains(sink)))) + def isEmpty: Boolean = locked(isEmptyUnlocked) private def isEmptyUnlocked: Boolean = sinksByKind.forall(_.isEmpty) + // must hold lock + private def liveConn: Option[Conn] = + state match { + case State.Live(conn) => Some(conn) + case _ => None + } + // --- shared establish/dispatch machinery ------------------------------------------------------------------------------------------------- - // Mark a new socket Live and subscribe it to every registered name. Reset confirmation counters for the new connection. In cluster mode, - // each connection has shard channels for at most one slot, which keeps its SSUBSCRIBE within that slot. - private def goLive(conn: Conn): Unit = { - var reconnect = false - var notify = false + // Mark a new socket Live and subscribe it to every registered name. + private def goLive(conn: Conn): Unit = // conn.close waits for the reader, and onConnClosed needs lock. Close conn after releasing lock. - var teardown: Conn = null - var failure: Throwable = null - locked { - if (state != State.Establishing && state != State.Reconnecting) teardown = conn - else if (conn.isTerminated) { - if (cluster) { - stopWatchdog() - current = null - state = State.Closed - notify = true - } else { - state = State.Reconnecting - reconnect = true - } - established.signalAll() - confirmed.signalAll() - } else { - current = conn - subscribeSent = 0L - subscribeConfirmed = 0L - pingSentAtMillis = 0L - lastReplyAtMillis = scheduler.nowMillis - val pending = Kind.values.map(kind => kind -> sinksFor(kind).keys.toVector) - try - pending.foreach { case (kind, names) => - if (names.nonEmpty) sendSubscribe(conn, kind, names) - } - catch { - // If writing the subscriptions fails, clear current and close this connection before reporting the failure; this prevents an older - // connection from dispatching after its replacement becomes active - case e: Throwable => - current = null - teardown = conn - failure = e - } - if (failure == null) - if (pending.forall(_._2.isEmpty)) { + locked(state match { + case State.Establishing | State.Reconnecting => + if (conn.isDead) lost() + else { + val pending = sinksByKind.toVector.flatMap(_.values) + conn.subscribe(pending) + changed.signalAll() + if (pending.isEmpty) { // if all subscribers close during establishment, close the new connection and return to Idle - teardown = conn - current = null state = State.Idle + () => conn.close() } else { - state = State.Live - startWatchdog() - liveTarget = subscribeSent + state = State.Live(conn) + reconnects.live() + conn.watch() + () => () } - established.signalAll() - confirmed.signalAll() - } - } - if (teardown != null) teardown.close() - if (reconnect) scheduleReconnect(0) - if (notify) onTerminated() - if (failure != null) throw failure - } + } + case _ => () => conn.close() + })() - private def sendSubscribe(conn: Conn, kind: Kind, names: Vector[String]): Unit = { - conn.send(kind.subscribeWire(names)) - subscribeSent += names.size - } + // Runs once per subscribed name with its confirmation, its error reply, or ConnectionLost when the connection ends first. A refusal (NOPERM, + // ERR) ends the name's subscriptions with the server's error. After any other error, such as BUSY, LOADING or MOVED, a cluster shard + // channel is placed again, and any other connection closes so that its reconnect subscribes the name again. + private def confirm(conn: Conn, name: Name)(result: Try[Unit]): Unit = + locked { + changed.signalAll() + result match { + case Success(_) => + name.confirmedOn = Some(conn) + () => () + case Failure(error: ServerError) if registered(name) => + onLoss match { + case _ if refuses(error) => + sinksFor(name.kind) -= name.name + val ends = name.sinks.toVector.map(_.end(Some(error))) + () => ends.foreach(_()) + case OnLoss.Report(_, _, onMoved) => + sinksFor(name.kind) -= name.name + onMoved + case _: OnLoss.Reconnect => () => scheduler.offload(conn.close()) + } + case _ => () => () // a lost reply is sent again by the next connection + } + }() - // Wait up to the connection timeout for subscribeConfirmed to reach the target. A target of -1 is replaced with liveTarget after the - // connection becomes live, covering subscriptions sent by goLive after reconnecting. Owned subscriptions fail with NotConnected when - // confirmation does not arrive before the deadline. - private def awaitActive(failIfUnconfirmed: Boolean, target0: Long): Unit = { + // Wait up to the connection timeout until the live connection confirmed every name that is still registered. Owned subscriptions fail + // with NotConnected when confirmation does not arrive before the deadline. + private def awaitActive(names: Vector[Name], failIfUnconfirmed: Boolean): Unit = { var active = false lock.lock() try { - val deadline = scheduler.nowMillis + connectTimeoutMillis - var target = target0 - var settled = false + val deadline = scheduler.nowMillis + config.connectTimeout.toMillis + // names before `done` are settled on `checkedOn`; a new live connection must confirm them again + var checkedOn = Option.empty[Conn] + var done = 0 + var settled = false while (!settled) - state match { - case State.Closed => settled = true - case State.Live => - if (target < 0) target = liveTarget - if (subscribeConfirmed >= target) { - active = true - settled = true - } else if (awaitOrTimeout(deadline)) settled = true - case _ => - target = -1L // reconnecting resets the counters; use the next liveTarget when the connection becomes live - if (awaitOrTimeout(deadline)) settled = true + if (state == State.Closed) settled = true + else { + val live = liveConn + if (live != checkedOn) { + checkedOn = live + done = 0 + } + while (done < names.size && (!registered(names(done)) || live.exists(conn => names(done).confirmedOn.exists(_ eq conn)))) done += 1 + if (done == names.size) { + active = true + settled = true + } else if (awaitOrTimeout(deadline)) settled = true } } finally lock.unlock() if (failIfUnconfirmed && !active) throw NotConnected() @@ -269,7 +231,7 @@ final private[client] class SubscriptionConnection( val remaining = deadline - scheduler.nowMillis if (remaining <= 0) true else { - confirmed.await(remaining, TimeUnit.MILLISECONDS) + changed.await(remaining, TimeUnit.MILLISECONDS) false } } @@ -281,211 +243,114 @@ final private[client] class SubscriptionConnection( if (state == State.Closed) conn.close() } try { - try conn.start() - catch { - case e: SageException => throw e - case NonFatal(e) => - val failed = ConnectionFailed(s"could not open the subscription connection: $e") - failed.initCause(e) - throw failed - } - runBootstrap(conn) + try conn.handshake(Bootstrap.commands(config), config.connectTimeout.toMillis) + catch { case NonFatal(e) => throw Client.translateHandshake(e) } conn } finally locked(establishing -= conn): Unit } - // keep bootstrap completion on its connection because two connections may bootstrap concurrently; clear it afterward to ignore later PONG replies - private def runBootstrap(conn: Conn): Unit = - try - Bootstrap.run( - bootstrap, - connectTimeoutMillis, - (command, cb) => { - conn.armBootstrap(result => cb(result.flatMap(frame => Reply.decode(command, frame)))) - if (conn.isTerminated) { conn.completeBootstrap(Failure(ConnectionLost(mayHaveExecuted = false))): Unit } - else conn.send(command.encode) - }, - () => conn.close() - ) - finally conn.clearBootstrap() - - private def scheduleReconnect(attempt: Int): Unit = - scheduler.after(Backoff.jitteredMillis(backoff, attempt, scheduler).millis)(attemptReconnect(attempt)) - - private def attemptReconnect(attempt: Int): Unit = { - val proceed = locked(state == State.Reconnecting) - if (proceed) { - onReconnect() - try goLive(establish()) - catch { - case NonFatal(error) => - locked(if (state == State.Reconnecting) { - events.emit(SageEvent.Connection.ReconnectFailed(node, error)) - scheduleReconnect(attempt + 1) - }) - } - } - } - private def onConnClosed(conn: Conn): Unit = - if (cluster) { - // Cluster connections do not reconnect themselves. The manager uses the latest topology to reassign their subscribers. During slot - // migration, the server sends `sunsubscribe` and disconnects, making closure the reliable signal to do this. - val notify = locked { - if (conn ne current) false - else - state match { - case State.Live | State.Establishing => - stopWatchdog() - current = null - state = State.Closed - established.signalAll() - confirmed.signalAll() - true - case _ => false - } - } - if (notify) onTerminated() - } else { - val reconnect = locked { - if (conn ne current) false - else - state match { - case State.Live | State.Reconnecting => - state = State.Reconnecting - confirmed.signalAll() - true - case _ => false - } - } - if (reconnect) scheduleReconnect(0) - } - - private def onFrame(conn: Conn, frame: Frame): Unit = - frame match { - case Frame.Push(elements) => - // a push confirms only that reads are working. Leave lastReplyAtMillis unchanged so push-only traffic still receives idle PING checks. - Pubsub.decode(elements) match { - case Some(Pubsub.Event.Message(channel, payload)) => - dispatch(sinksFor(Kind.Channel), channel, Delivery.Channel(channel, payload)) - case Some(Pubsub.Event.ShardMessage(channel, payload)) => - dispatch(sinksFor(Kind.Shard), channel, Delivery.Channel(channel, payload)) - case Some(Pubsub.Event.PatternMessage(pattern, ch, payload)) => - dispatch(sinksFor(Kind.Pattern), pattern, Delivery.Pattern(pattern, ch, payload)) - case Some(_: Pubsub.Event.Subscribed) => - // conn eq current: a late ack from a superseded generation must not advance this generation's count - locked(if (conn eq current) { - subscribeConfirmed += 1 - confirmed.signalAll() - }) - case _ => () // an Unsubscribed ack is informational; re-homing is disconnect-driven + locked(state match { + case State.Live(c) if c eq conn => lost() + case _ => () => () + })() + + // Must hold lock; returns the action to run after releasing it. A cluster shard connection does not reconnect itself: the manager reassigns its + // subscribers using the latest topology. + private def lost(): () => Unit = { + changed.signalAll() + onLoss match { + case OnLoss.Report(onTerminated, _, _) => + state = State.Closed + () => onTerminated(this) + case mode: OnLoss.Reconnect => + state = State.Reconnecting + reconnects.schedule( + state == State.Reconnecting, + error => mode.events.emit(SageEvent.Connection.ReconnectFailed(mode.node(), error)), + mode.immediately + ) { + mode.beforeAttempt() + goLive(establish()) } - case reply => // non-push reply: bootstrap HELLO, watchdog PONG, or an unexpected error - lastReplyAtMillis = scheduler.nowMillis - if (!conn.completeBootstrap(Success(reply))) - reply match { - // an error such as MOVED is not a PONG. Close the connection so subscription placement is recalculated. - case _: Frame.SimpleError | _: Frame.BulkError => scheduler.after(Duration.Zero)(conn.close()) // off the reader thread: close() joins it - case _ => pingSentAtMillis = 0L - } + () => () } + } // snapshot the sinks under the lock, then deliver outside it: a blocking put (backpressure) must never hold the registry lock - private def dispatch(map: mutable.HashMap[String, mutable.LinkedHashSet[Sink]], key: String, delivery: Delivery): Unit = { - val targets = locked(map.get(key).map(_.toVector).getOrElse(Vector.empty)) + private def dispatch(conn: Conn, map: mutable.HashMap[String, Name], key: String, delivery: Delivery): Unit = { + val targets = locked(map.get(key).map(_.sinks.toVector).getOrElse(Vector.empty)) if (targets.nonEmpty) { - readerBlocked = true + conn.readerBlocked = true try { var blocked = false targets.foreach(sink => if (sink.offer(delivery)) blocked = true) - if (blocked) lastBackpressureMillis = scheduler.nowMillis - } finally readerBlocked = false + if (blocked) conn.lastBackpressureMillis = scheduler.nowMillis + } finally conn.readerBlocked = false } } - // in standalone mode, close the sink after unsubscribing it. Close the socket when the last sink is removed. + // Close the socket without unsubscribing when the last sink is removed. Terminate the sink first: closing the socket waits for the reader, + // which may be blocked offering to this sink. private def closeOwned(sink: Sink): Unit = { - var teardown: Conn = null - var failure: Throwable = null + sink.terminate() locked { - val emptied = deregister(sink, sink.names, sink.kind) - try if (emptied.nonEmpty && state == State.Live) current.send(sink.kind.unsubscribeWire(emptied)) - catch { case e: Throwable => failure = e } - if (isEmptyUnlocked && (state == State.Live || state == State.Reconnecting)) { - stopWatchdog() - teardown = current - current = null + val emptied = deregister(sink, sink.names) + if (isEmptyUnlocked && (liveConn.nonEmpty || state == State.Reconnecting)) { + val teardown = liveConn state = State.Idle + teardown + } else { + liveConn.foreach(_.unsubscribe(sink.kind, emptied)) + None } - } - sink.terminate() - if (teardown != null) teardown.close() - if (failure != null) throw failure + }.foreach(_.close()) } // must hold lock. Change the state to Closed and return the current and establishing connections for the caller to close. private def markClosed(): Vector[Conn] = { - val conns = (Option(current) ++ establishing).toVector + val conns = (liveConn ++ establishing).toVector establishing.clear() state = State.Closed - stopWatchdog() - current = null - established.signalAll() - confirmed.signalAll() + changed.signalAll() conns } // check for subscribers and set Closed under one lock, preventing attach from registering a subscriber between those operations def closeIfEmpty(): Boolean = { - var toClose: Vector[Conn] = Vector.empty - val closing = locked { - if (!isEmptyUnlocked) false - else { - toClose = markClosed() - true - } - } - toClose.foreach(_.close()) - closing + val toClose = locked(Option.when(isEmptyUnlocked)(markClosed())) + toClose.foreach(_.foreach(_.close())) + toClose.nonEmpty } - // close the socket and watchdog but keep subscribers available for reassignment. - def shutdown(): Unit = tearDown(terminateSinks = false) - - def close(): Unit = tearDown(terminateSinks = true) - - private def tearDown(terminateSinks: Boolean): Unit = { - var sinks: Set[Sink] = Set.empty - val toClose = locked { - if (terminateSinks) sinks = sinksByKind.iterator.flatMap(_.values.flatten).toSet + def close(): Unit = { + val (sinks, toClose) = locked { + val sinks = sinksByKind.iterator.flatMap(_.values.flatMap(_.sinks)).toSet val conns = markClosed() sinksByKind.foreach(_.clear()) - conns + (sinks, conns) } // Terminate sinks before closing connections. Connection close waits for the reader, and the reader may be waiting in Sink.offer until - // its sink is closed. Closing the connection first would deadlock. In cluster mode, the manager has already terminated the sinks. + // its sink is closed. Closing the connection first would deadlock. For shard connections, the cluster manager has already terminated the sinks. sinks.foreach(_.terminate()) toClose.foreach(_.close()) } - private def register(sink: Sink, names: Vector[String], kind: Kind): Vector[String] = { - val map = sinksFor(kind) - val fresh = Vector.newBuilder[String] - names.foreach { name => - val set = map.getOrElseUpdate(name, mutable.LinkedHashSet.empty) - if (set.isEmpty) fresh += name - set += sink - } - fresh.result() + // returns the Names of `names` and the ones this call created, which still need a subscribe + private def register(sink: Sink, names: Vector[String]): (Vector[Name], Vector[Name]) = { + val created = Vector.newBuilder[Name] + val all = names.map(name => sinksFor(sink.kind).getOrElseUpdate(name, { val n = new Name(sink.kind, name); created += n; n })) + all.foreach(_.sinks += sink) + (all, created.result()) } - private def deregister(sink: Sink, names: Vector[String], kind: Kind): Vector[String] = { - val map = sinksFor(kind) + private def deregister(sink: Sink, names: Vector[String]): Vector[String] = { + val map = sinksFor(sink.kind) val emptied = Vector.newBuilder[String] names.foreach { name => - map.get(name).foreach { set => - set -= sink - if (set.isEmpty) { + map.get(name).foreach { n => + n.sinks -= sink + if (n.sinks.isEmpty) { map -= name emptied += name } @@ -494,119 +359,133 @@ final private[client] class SubscriptionConnection( emptied.result() } - private def startWatchdog(): Unit = - if (watchdog.enabled && watchdogHandle == null) - watchdogHandle = scheduler.every(watchdog.pingInterval)(watchdogTick()) + // One registered name and its sinks, guarded by `lock`. `confirmedOn` is the last connection that confirmed it, so after a reconnect the + // name counts as active only once the new connection confirms it. + final private class Name(val kind: Kind, val name: String) { + val sinks = mutable.LinkedHashSet.empty[Sink] + var confirmedOn: Option[Conn] = None + } - private def stopWatchdog(): Unit = - if (watchdogHandle != null) { - watchdogHandle.cancel() - watchdogHandle = null + // Replies (bootstrap, subscribed and unsubscribed names, and watchdog PING) match their entries in write order; a reply with nothing pending + // closes the connection. + final private class Conn extends WatchedPipe(factory, scheduler, config.watchdog) { + + // Each name is its own entry, answered in write order by its confirmation push or by an error reply such as NOPERM or MOVED. + def subscribe(names: Vector[Name]): Unit = sendEach(names.map(name => new Entry(name.kind.subscribe(name.name), confirm(this, name)))) + + // An UNSUBSCRIBE or PUNSUBSCRIBE is answered by its own push. The push of an SUNSUBSCRIBE is ambiguous: a cluster node sends the same + // push when it drops a channel whose slot moved, and then answers the SUNSUBSCRIBE with another push (Redis) or MOVED (Valkey). One HELLO + // written after the SUNSUBSCRIBEs ends their replies: an error before the HELLO's reply is one of theirs, and their pushes go to onPush. + // Each name has its own SUNSUBSCRIBE because Valkey answers CROSSSLOT to one naming channels in several slots. + def unsubscribe(kind: Kind, names: Vector[String]): Unit = + if (kind != Kind.Shard) sendEach(names.map(name => new Entry(kind.unsubscribe(name), _ => ()))) + else if (names.nonEmpty) sendEach(names.map(name => new Entry(kind.unsubscribe(name), untilHello)) :+ new Entry(Pubsub.helloInfo, afterHello)) + + // used on the reader thread only, while the HELLO's reply is passed on to its own entry + private var passing = false + private var reached = false + + // The first SUNSUBSCRIBE still pending receives the HELLO's reply and passes it on in a loop, not recursively, so a long batch cannot + // overflow the stack. + private val untilHello: Try[Frame] => Unit = { + case Success(reply) if !passing => + passing = true + reached = false + try while (!reached && !isDead) answer(reply) + finally passing = false + case _ => () } - private def watchdogTick(): Unit = { - if (readerBlocked) return // deliberate backpressure on a slow consumer; the connection is alive, not stuck - val conn = locked(if (state == State.Live) current else null) - if (conn != null) { - val now = scheduler.nowMillis - if (pingSentAtMillis != 0L) { - // Recent backpressure may have kept the reader from reaching the queued PONG. When the sink has room, an unanswered PING still closes - // the connection after the timeout. - val backpressured = now - lastBackpressureMillis < watchdog.pingTimeout.toMillis - if (!backpressured && now - pingSentAtMillis >= watchdog.pingTimeout.toMillis) scheduler.after(Duration.Zero)(conn.close()) - } else if (now - lastReplyAtMillis >= watchdog.pingInterval.toMillis) { - pingSentAtMillis = now - conn.send(Connection.ping(None).encode) + // ACL does not refuse an argument-less HELLO, so an error here means that replies no longer match their entries + private val afterHello: Try[Frame] => Unit = { result => + reached = true + result match { + case Failure(_: ServerError) => this.close() + case _ => () } } - } - - final private class Conn { - private val transportRef = new AtomicReference[Transport]() - @volatile private var terminated = false - @volatile private var aborted = false - private val bootstrapWaiter = new AtomicReference[Try[Frame] => Unit]() - - def isTerminated: Boolean = terminated + private def sendEach(entries: Vector[Entry[?]]): Unit = + if (entries.nonEmpty) { + reserve(entries.size) + sendAll(entries) + } - def start(): Unit = { - val transport = factory(frame => onFrame(this, frame), () => onTerminated()) - transportRef.set(transport) - if (aborted) transport.close() - else transport.start() - } + @volatile var readerBlocked: Boolean = false + @volatile var lastBackpressureMillis: Long = 0L + + // Recent backpressure may have kept the reader from reaching the probe's reply. When the sink has room, an unanswered probe still closes + // the connection after the timeout. + override protected def tick(): Unit = + if (!readerBlocked) // deliberate backpressure on a slow consumer; the connection is alive, not stuck + checkLiveness(lastBackpressureMillis + config.watchdog.pingTimeout.toMillis) + + protected def onPush(elements: Vector[Frame]): Unit = + Pubsub.decode(elements).foreach { + case Pubsub.Event.Confirmed => answer(Frame.Push(elements)) + case Pubsub.Event.Delivered(kind, subscription, delivery) => dispatch(this, sinksFor(kind), subscription, delivery) + // The server dropped the channel only if it is still registered and confirmed here. An unsubscribed channel has left the registry. + case Pubsub.Event.ShardUnsubscribed(channel) => + onLoss match { + case OnLoss.Report(_, onDropped, _) + if locked( + sinksFor(Kind.Shard).get(channel).exists(_.confirmedOn.exists(_ eq this)) && sinksFor(Kind.Shard).remove(channel).isDefined + ) => + onDropped() + case _ => () + } + } - private def onTerminated(): Unit = { - terminated = true - completeBootstrap(Failure(ConnectionLost(mayHaveExecuted = false))) + override protected def onClosed(): Unit = { + super.onClosed() onConnClosed(this) } - - def armBootstrap(waiter: Try[Frame] => Unit): Unit = bootstrapWaiter.set(waiter) - def clearBootstrap(): Unit = bootstrapWaiter.set(null) - - def completeBootstrap(result: Try[Frame]): Boolean = { - val waiter = bootstrapWaiter.getAndSet(null) - if (waiter != null) { - waiter(result) - true - } else false - } - - def send(payload: Bytes): Unit = { - val transport = transportRef.get() - if (transport != null) transport.send(new RawItem(payload)) - } - - def close(): Unit = { - aborted = true - val transport = transportRef.get() - if (transport != null) transport.close() - } } } private[client] object SubscriptionConnection { - private enum State { - case Idle, Establishing, Live, Reconnecting, Closed + // What a lost connection does: reconnect and restore its subscribers, or report the loss (cluster shard channels). + enum OnLoss { + // `node` names the node of the last attempt in reported failures. With `immediately`, the first attempt after a stable period does not wait. + case Reconnect(beforeAttempt: () => Unit, events: Events, node: () => Option[Node] = () => None, immediately: Boolean = false) + // onDropped runs when the server drops a shard channel without closing the connection, as it does when the channel's slot moves, and + // onMoved when the server refuses a shard channel with a redirect or another retryable error + case Report(onTerminated: SubscriptionConnection => Unit, onDropped: () => Unit, onMoved: () => Unit) } - /** - * The three subscription kinds, each with its wire encoders: classic channels (`SUBSCRIBE`), glob patterns (`PSUBSCRIBE`), and shard - * channels (`SSUBSCRIBE`). - */ - private[internal] enum Kind { - case Channel, Pattern, Shard - - def subscribeWire(names: Vector[String]): Bytes = - this match { - case Channel => Pubsub.subscribe(names) - case Pattern => Pubsub.psubscribe(names) - case Shard => Pubsub.ssubscribe(names) - } + export Pubsub.{Delivery, Kind} + + // The errors a subscribe gets every time it is sent: an ACL denial, or a command or argument the server does not support. + private def refuses(error: ServerError): Boolean = error.code == "NOPERM" || error.code == "ERR" - def unsubscribeWire(names: Vector[String]): Bytes = - this match { - case Channel => Pubsub.unsubscribe(names) - case Pattern => Pubsub.punsubscribe(names) - case Shard => Pubsub.sunsubscribe(names) + // Connects each attempt to the node `pick` names at that time. `retain` closes that connection once its node leaves the deployment, and + // the connection then reconnects to a current node. + final class Following(nodeFactory: Node => MultiplexedConnection.TransportFactory, pick: () => Option[Node]) { + + @volatile private var on: Option[(Node, Transport)] = None + // the node of the latest attempt, recorded before connecting so that a failed attempt reports it + @volatile var node: Option[Node] = None + + val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => { + node = pick() + node match { + case Some(target) => + val transport = nodeFactory(target)(onFrame, onClosed) + on = Some(target -> transport) + transport + case None => throw NotConnected() } - } + } - /** - * A raw delivery sent to a subscription buffer. Shard channel messages use [[Channel]] because they contain the same channel and payload. - */ - enum Delivery { - case Channel(channel: String, payload: Bytes) - case Pattern(pattern: String, channel: String, payload: Bytes) + def retain(listed: Node => Boolean): Unit = on.foreach { case (node, transport) => if (!listed(node)) transport.close() } } - // pub/sub writes (SUBSCRIBE/UNSUBSCRIBE) are confirmed by push frames, not a per-write reply, so the write hooks are no-ops - final private class RawItem(val payload: Bytes) extends Transport.Item { - def writeAttempted(): Unit = () - def dropped(): Unit = () + trait PubSub { + def subscribeChannels(channels: Vector[String]): RawSubscription + def subscribePatterns(patterns: Vector[String]): RawSubscription + def subscribeShard(channels: Vector[String]): RawSubscription + def close(): Unit } /** @@ -617,12 +496,13 @@ private[client] object SubscriptionConnection { */ final private[internal] class Sink(val names: Vector[String], val kind: Kind, capacity: Int) { - private val cap = math.max(1, capacity) private val lock = new ReentrantLock() private val notFull = lock.newCondition() - private val backlog = new java.util.ArrayDeque[Delivery](cap) + private val backlog = new java.util.ArrayDeque[Delivery](capacity) private var waiter: Option[Delivery] => Unit = null - private var closed = false + // `failure` holds the server's error when it ended the subscription by refusing a name; both are written under `lock`, failure first + @volatile var failure: Option[Throwable] = None + @volatile private var ended = false def next(callback: Option[Delivery] => Unit): Unit = { var ready: Option[Delivery] = null // null means the callback was stored; a non-null value is delivered immediately @@ -634,7 +514,7 @@ private[client] object SubscriptionConnection { if (head != null) { notFull.signal() ready = Some(head) - } else if (closed) ready = None + } else if (ended) ready = None else waiter = callback } finally lock.unlock() if (ready != null) callback(ready) @@ -653,12 +533,12 @@ private[client] object SubscriptionConnection { try { var settled = false while (!settled) - if (closed) settled = true + if (ended) settled = true else if (waiter != null) { hungry = waiter waiter = null settled = true - } else if (backlog.size < cap) { + } else if (backlog.size < capacity) { backlog.add(delivery) settled = true } else { @@ -671,26 +551,37 @@ private[client] object SubscriptionConnection { blocked } - def terminate(): Unit = { + def terminate(): Unit = end(None)() + + // Ends the subscription, with the server's error if it refused a name, and returns the call that wakes a waiting consumer, to run after + // the caller's locks are released. Only the first end counts. + def end(error: Option[Throwable]): () => Unit = { var pending: Option[Delivery] => Unit = null lock.lock() try { - closed = true + if (!ended) { + failure = error + ended = true + } backlog.clear() pending = waiter waiter = null notFull.signalAll() // release a reader blocked on backpressure } finally lock.unlock() - if (pending != null) pending(None) + () => if (pending != null) pending(None) } } final class RawSubscription private[internal] (sink: Sink, onClose: () => Unit) { + // None once the subscription has ended def next(callback: Option[Delivery] => Unit): Unit = sink.next(callback) def cancelNext(callback: Option[Delivery] => Unit): Unit = sink.cancelNext(callback) + // the server's error when it ended the subscription by refusing a name + def failure: Option[Throwable] = sink.failure + def close(): Unit = onClose() } } diff --git a/sage-client/shared/src/main/scala/sage/client/internal/Tls.scala b/sage-client/shared/src/main/scala/sage/client/internal/Tls.scala index a0e2bda2..a22fe32f 100644 --- a/sage-client/shared/src/main/scala/sage/client/internal/Tls.scala +++ b/sage-client/shared/src/main/scala/sage/client/internal/Tls.scala @@ -18,9 +18,9 @@ private[client] object Tls { private val plaintext: Socket => Socket = socket => socket /** - * Builds the `SSLContext` once for a client and reports invalid trust configuration as a [[TlsError]]. The returned function wraps a - * connected socket in an `SSLSocket` and performs the TLS handshake. The existing connection timeout and close behavior still apply. - * `host` must match the socket destination because it is used for SNI and hostname verification. + * Builds the `SSLContext` once, reading any trust file now, and reports invalid trust configuration as a [[TlsError]]. The returned + * function wraps a connected socket in an `SSLSocket` and performs the TLS handshake. The existing connection timeout and close behavior + * still apply. `host` must match the socket destination because it is used for SNI and hostname verification. */ def buildUpgrade(tls: Option[TlsConfig], host: String, port: Int): Socket => Socket = tls match { @@ -96,7 +96,7 @@ private[client] object Tls { } private def keyStoreType(path: Path): String = { - val name = path.getFileName.toString.toLowerCase + val name = path.getFileName.toString.toLowerCase(java.util.Locale.ROOT) if (name.endsWith(".p12") || name.endsWith(".pfx")) "PKCS12" else if (name.endsWith(".jks")) "JKS" else KeyStore.getDefaultType diff --git a/sage-client/shared/src/main/scala/sage/client/internal/TransactionScope.scala b/sage-client/shared/src/main/scala/sage/client/internal/TransactionScope.scala index 918692b1..3a3a047e 100644 --- a/sage-client/shared/src/main/scala/sage/client/internal/TransactionScope.scala +++ b/sage-client/shared/src/main/scala/sage/client/internal/TransactionScope.scala @@ -18,9 +18,7 @@ trait TransactionScope[F[_], K] extends CommandRunner[F, K] { */ def watch[K: KeyCodec](key: K, rest: K*): F[Unit] - private[sage] def exec[Out, R](pipeline: Pipeline[Out, R]): F[Option[Out]] - - private[sage] def execAttempt[Out, R](pipeline: Pipeline[Out, R]): F[Option[R]] + private[sage] def exec[R](pipeline: Pipeline[R]): F[Option[R]] /** * Executes a fixed-arity batch of commands atomically (`MULTI`/`EXEC`), yielding a result tuple that mirrors the argument tuple @@ -40,13 +38,13 @@ trait TransactionScope[F[_], K] extends CommandRunner[F, K] { * Like the tuple [[exec]], but yields the per-position results (each slot a `Right`/`Left`) on commit. `None` still means aborted. */ def execAttempt[T <: NonEmptyTuple](commands: T)(using Tuple.IsMappedBy[Command][T]): F[Option[Tuple.Map[Tuple.InverseMap[T, Command], Attempt]]] = - execAttempt(Pipeline.fromTuple(commands)) + exec(Pipeline.fromTupleAttempt(commands)) /** * Like the `Seq` [[exec]], but yields the per-position results (each slot a `Right`/`Left`) on commit. `None` still means aborted. */ def execAttempt[A](commands: Seq[Command[A]]): F[Option[Vector[Attempt[A]]]] = - execAttempt(Pipeline.sequence(commands)) + exec(Pipeline.sequenceAttempt(commands)) /** * Abandons the scope without committing, clearing any watched keys so the connection can be recycled (issues `UNWATCH`). @@ -60,11 +58,10 @@ trait TransactionScope[F[_], K] extends CommandRunner[F, K] { override def as[K2](using KeyCodec[K2]): TransactionScope[F, K2] = { val self = this new TransactionScope[F, K2] { - def run[A](command: Command[A]): F[A] = self.run(command) - def watch[K0: KeyCodec](key: K0, rest: K0*): F[Unit] = self.watch(key, rest*) - private[sage] def exec[Out, R](pipeline: Pipeline[Out, R]): F[Option[Out]] = self.exec(pipeline) - private[sage] def execAttempt[Out, R](pipeline: Pipeline[Out, R]): F[Option[R]] = self.execAttempt(pipeline) - def discard: F[Unit] = self.discard + def run[A](command: Command[A]): F[A] = self.run(command) + def watch[K0: KeyCodec](key: K0, rest: K0*): F[Unit] = self.watch(key, rest*) + private[sage] def exec[R](pipeline: Pipeline[R]): F[Option[R]] = self.exec(pipeline) + def discard: F[Unit] = self.discard } } } diff --git a/sage-client/shared/src/main/scala/sage/client/internal/Transport.scala b/sage-client/shared/src/main/scala/sage/client/internal/Transport.scala index 17e456a2..6b622527 100644 --- a/sage-client/shared/src/main/scala/sage/client/internal/Transport.scala +++ b/sage-client/shared/src/main/scala/sage/client/internal/Transport.scala @@ -1,6 +1,27 @@ package sage.client.internal -import sage.Bytes +import java.util.concurrent.ConcurrentLinkedQueue +import java.util.concurrent.atomic.{AtomicInteger, AtomicReference} + +import scala.concurrent.duration.Duration +import scala.util.{Failure, Success, Try} + +import sage.{Bytes, SageException} +import sage.SageException.{ConnectionLost, ServerError} +import sage.client.WatchdogConfig +import sage.commands.{Command, Connection, Reply} +import sage.protocol.Frame + +/** + * Runs `undo` when `body` throws anything, including an `InterruptedException` from a cancelled wait, then rethrows. + */ +private[internal] inline def onThrow[A](inline body: A)(inline undo: Throwable => Unit): A = + try body + catch { + case error: Throwable => + undo(error) + throw error + } /** * Defines the interface between connection logic and socket I/O. The connection submits [[Transport.Item]] values and receives parsed frames @@ -15,13 +36,14 @@ private[client] trait Transport { def start(): Unit /** - * Adds an item to the write queue and returns immediately. The transport calls exactly one of `writeAttempted` or `dropped`. On the write - * path, it calls `clearPayload` after capturing the bytes to write. + * Adds an item to the write queue and returns immediately without throwing, even on an interrupted thread. The transport calls exactly + * one of `writeAttempted` or `dropped`. On the write path, it calls `clearPayload` after capturing the bytes to write. */ def send(item: Transport.Item): Unit /** - * Idempotent. Blocks until the I/O threads have terminated and `onClosed` has run. + * Idempotent. Blocks until the I/O threads have terminated and `onClosed` has run. It never throws, even on an interrupted thread, and + * keeps the caller's interrupt flag set. */ def close(): Unit } @@ -48,3 +70,221 @@ private[client] object Transport { def dropped(): Unit } } + +/** + * Owns one transport from `start` until it terminates and matches replies to commands in write order. `close` may run before `start`, + * which then closes the transport at once. `isDead` becomes true when the connection is closed, drops a write, terminates, or is retired + * by its pool. A command counts as in flight from `reserve` until its reply, failure or drop, including while the transport still queues + * it, so a caller that checks `isQuiescent` or waits for `onDrained` sees queued work. + */ +abstract private[internal] class Pipe(factory: MultiplexedConnection.TransportFactory, scheduler: Scheduler) { + + private val transportRef = new AtomicReference[Transport]() + @volatile private var dead = false + + final def isDead: Boolean = dead + + final private[internal] def markDead(): Unit = dead = true + + // Publish transportRef before the blocking connect starts so close can abort it. + final def start(): Unit = { + val transport = factory(onFrame, () => { dead = true; onClosed() }) + transportRef.set(transport) + if (dead) transport.close() + else transport.start() + } + + private def send(item: Transport.Item): Unit = transportRef.get().send(item) + + // the caller has reserved the entries + final protected def sendAll(entries: Vector[Entry[?]]): Unit = send(new Batch(entries)) + + def close(): Unit = { + dead = true + val transport = transportRef.get() + if (transport != null) transport.close() + } + + private val pending = new ConcurrentLinkedQueue[Entry[?]]() + private val inFlight = new AtomicInteger(0) + + final def reserve(n: Int): Unit = inFlight.addAndGet(n): Unit + + final protected def release(n: Int): Unit = if (inFlight.addAndGet(-n) == 0) onDrained() + + final protected def isIdle: Boolean = inFlight.get() == 0 + + final def isQuiescent: Boolean = !isDead && isIdle + + protected def onDrained(): Unit = () + + final def submit[A](command: Command[A], callback: Try[A] => Unit): Unit = { + reserve(1) + write(command, callback) + } + + /** + * Starts the transport and runs the setup commands in turn, blocking up to `timeoutMillis` for each reply. Closes the connection and + * throws on a timeout or a failed reply, except that a [[Bootstrap.bestEffort]] command's `ServerError` is tolerated. + */ + final def handshake(commands: Vector[Command[?]], timeoutMillis: Long): Unit = { + start() + commands.foreach { command => + step(command, timeoutMillis).filterNot(_ => Bootstrap.bestEffort(command)).foreach { error => + close() + throw error + } + } + } + + /** + * Runs one setup command and returns the server's error reply, or `None` on success. Closes the connection and throws on a timeout or any + * other failure. + */ + final def step(command: Command[?], timeoutMillis: Long): Option[ServerError] = + onThrow { + Bootstrap.awaitReply[Any](timeoutMillis, ConnectionLost(mayHaveExecuted = false))(submit(command, _)) match { + case Failure(error: ServerError) => Some(error) + // a lost setup reply means the caller's own command was never sent + case Failure(_: ConnectionLost) => throw ConnectionLost(mayHaveExecuted = false) + case Failure(error) => throw error + case Success(_) => None + } + }(_ => close()) + + // the caller has reserved the command + final def write[A](command: Command[A], callback: Try[A] => Unit): Unit = send(new Entry(command, callback)) + + final protected def oldestWriteMillis: Option[Long] = Option(pending.peek()).map(_.sentAtMillis) + + protected def onFrame(frame: Frame): Unit = + frame match { + case _: Frame.Push => () + case reply => answer(reply) + } + + // completes the oldest pending entry; a reply with nothing pending means the stream desynced + final protected def answer(reply: Frame): Unit = { + val waiter = pending.poll() + if (waiter == null) close() + else waiter.complete(reply) + } + + protected def onReadOnly(): Unit = close() + + protected def onClosed(): Unit = { + var waiter = pending.poll() + while (waiter != null) { + waiter.fail(ConnectionLost(mayHaveExecuted = true)) + waiter = pending.poll() + } + } + + final protected class Entry[A](command: Command[A], callback: Try[A] => Unit) extends Transport.Item { + + @volatile var sentAtMillis: Long = 0L + + var payload: Bytes = command.encode + + override def clearPayload(): Unit = payload = Bytes.empty + + def writeAttempted(): Unit = { + sentAtMillis = scheduler.nowMillis + pending.add(this): Unit + } + + // the transport calls dropped during teardown before onClosed; mark the connection dead first so a pool cannot reuse it meanwhile + def dropped(): Unit = { + markDead() + settle(Failure(ConnectionLost(mayHaveExecuted = false))) + } + + // Reply.decode guards against throwing user decoders: an escaped exception would otherwise lose the callback and hang the awaiting fiber + def complete(frame: Frame): Unit = + Reply.decode(command, frame) match { + // A READONLY reply fails its command and retires the connection, because an in-place failover can leave the old master connected + // but unable to accept writes. Mark it dead before delivering the reply so that a pool discards it on release. + case failure @ Failure(ServerError("READONLY", _)) => + markDead() + settle(failure) + onReadOnly() + case result => settle(result) + } + + def fail(error: SageException): Unit = settle(Failure(error)) + + private def settle(result: Try[A]): Unit = { + release(1) + callback(result) + } + } + + // Concatenate a pipeline into one transport write and notify each entry individually when the write succeeds or fails; the transport + // writes or drops the complete batch without splitting it across socket writes + final private class Batch(entries: Vector[Entry[?]]) extends Transport.Item { + + val payload: Bytes = Bytes.concatBy(entries)(_.payload) + + override def clearPayload(): Unit = entries.foreach(_.clearPayload()) + + def writeAttempted(): Unit = entries.foreach(_.writeAttempted()) + + def dropped(): Unit = entries.foreach(_.dropped()) + } +} + +/** + * A [[Pipe]] checked by a watchdog. A push frame does not count as a reply, so push-only traffic still receives idle PING checks. + */ +abstract private[internal] class WatchedPipe(factory: MultiplexedConnection.TransportFactory, scheduler: Scheduler, watchdog: WatchdogConfig) + extends Pipe(factory, scheduler) { + + @volatile private var lastReplyAtMillis: Long = scheduler.nowMillis + // WatchedPipe.Stopped once unwatched, so a watch that loses the race with close cancels its own timer + private val ticker = new AtomicReference[Scheduler.Cancelable]() + + protected def onPush(elements: Vector[Frame]): Unit + + final override protected def onFrame(frame: Frame): Unit = + frame match { + case Frame.Push(elements) => onPush(elements) + case reply => + lastReplyAtMillis = scheduler.nowMillis + super.onFrame(reply) + } + + // Runs tick every PING interval from go-live until unwatch or close. + final def watch(): Unit = + if (watchdog.enabled) { + val handle = scheduler.every(watchdog.pingInterval)(tick()) + if (!ticker.compareAndSet(null, handle)) handle.cancel() + } + + final def unwatch(): Unit = { + val handle = ticker.getAndSet(WatchedPipe.Stopped) + if (handle != null) handle.cancel() + } + + protected def tick(): Unit = checkLiveness(0L) + + override protected def onClosed(): Unit = { + unwatch() + super.onClosed() + } + + // Close when the oldest unanswered write is older than the timeout and graceUntilMillis has passed; otherwise PING when idle. + final protected def checkLiveness(graceUntilMillis: Long): Unit = { + val now = scheduler.nowMillis + oldestWriteMillis match { + // offload: close() blocks joining I/O threads, and the watchdog tick runs on the shared timer thread, which must not block + case Some(sentAtMillis) => + if (now >= graceUntilMillis && now - sentAtMillis >= watchdog.pingTimeout.toMillis) scheduler.after(Duration.Zero)(close()) + // any reply to the PING, even an error such as NOPERM or LOADING, proves the connection is alive + case None => if (now - lastReplyAtMillis >= watchdog.pingInterval.toMillis) submit(Connection.ping(None), _ => ()) + } + } +} + +private object WatchedPipe { + val Stopped: Scheduler.Cancelable = () => () +} diff --git a/sage-client/shared/src/main/scala/sage/client/internal/TxSupport.scala b/sage-client/shared/src/main/scala/sage/client/internal/TxSupport.scala index 51cdc6da..25a7ba72 100644 --- a/sage-client/shared/src/main/scala/sage/client/internal/TxSupport.scala +++ b/sage-client/shared/src/main/scala/sage/client/internal/TxSupport.scala @@ -1,14 +1,16 @@ package sage.client.internal -import java.util.concurrent.atomic.{AtomicInteger, AtomicReferenceArray} +import java.util.concurrent.atomic.{AtomicBoolean, AtomicInteger, AtomicReferenceArray} +import java.util.concurrent.locks.ReentrantLock import scala.util.{Failure, Success, Try} import kyo.compat.* import sage.SageException -import sage.SageException.{DecodeError, ProtocolError, ServerError, TransactionDiscarded} -import sage.commands.{Command, Reply} +import sage.SageException.{DecodeError, InvalidArgument, ProtocolError, ServerError, TransactionDiscarded} +import sage.codec.KeyCodec +import sage.commands.{Command, Connection, Pipeline, Reply} import sage.protocol.Frame /** @@ -17,18 +19,6 @@ import sage.protocol.Frame */ private[internal] object TxSupport { - def collapseStrict[Out](results: Vector[Either[SageException, Any]], toOut: Vector[Any] => Out): CIO[Out] = { - val values = Vector.newBuilder[Any] - values.sizeHint(results.length) - val it = results.iterator - while (it.hasNext) - it.next() match { - case Right(value) => values += value - case Left(error) => return CIO.fail(error) - } - CIO.value(toOut(values.result())) - } - // decoders should return Either; another exception indicates a decoder bug and becomes DecodeError to preserve per-command results def toEither(result: Try[Any]): Either[SageException, Any] = result match { @@ -37,45 +27,31 @@ private[internal] object TxSupport { case Failure(other) => Left(DecodeError.fromThrowable(other)) } + // the replies to MULTI and to each queued command, then the reply to EXEC; a Failure is the server's error reply + final case class ExecReplies(queued: Vector[Try[Frame]], exec: Try[Frame]) + // Return None when EXEC reports that a watched key changed. Otherwise, return each command's decoded result. A queueing error fails the - // effect before the transaction executes. - def interpretExec(commands: Vector[Command[?]], frames: Vector[Frame]): CIO[Option[Vector[Either[SageException, Any]]]] = { - val n = commands.length - // frames: MULTI reply, then one queue reply per command, then the EXEC reply - val queueError = (0 to n).iterator.map(i => errorOf(frames(i))).collectFirst { case Some(message) => message } - queueError match { - case Some(message) => CIO.fail(TransactionDiscarded(message)) - case None => - frames(n + 1) match { - case Frame.Null => CIO.value(None) - case Frame.Array(elems) if elems.length == n => - CIO.value(Some(Vector.tabulate(n)(i => toEither(Reply.decode(commands(i), elems(i)))))) - case Frame.Array(elems) => - CIO.fail(ProtocolError(s"EXEC returned ${elems.length} results for $n queued commands")) - case other => - errorOf(other) match { - case Some(message) => CIO.fail(TransactionDiscarded(message)) - case None => CIO.fail(ProtocolError(s"unexpected EXEC reply: ${Frame.describe(other)}")) - } + // transaction before it executes. + def interpretExec(commands: Vector[Command[?]], replies: ExecReplies): Either[SageException, Option[Vector[Either[SageException, Any]]]] = + replies.queued.collectFirst { case Failure(error) => error } match { + case Some(error) => Left(TransactionDiscarded(error.getMessage)) + case None => + val n = commands.length + replies.exec match { + case Failure(error) => Left(TransactionDiscarded(error.getMessage)) + case Success(Frame.Null) => Right(None) + case Success(Frame.Array(elems)) if elems.length == n => + Right(Some(Vector.tabulate(n)(i => toEither(Reply.decode(commands(i), elems(i)))))) + case Success(Frame.Array(elems)) => + Left(ProtocolError(s"EXEC returned ${elems.length} results for $n queued commands")) + case Success(other) => Left(ProtocolError(s"unexpected EXEC reply: ${Frame.describe(other)}")) } } - } - def errorOf(frame: Frame): Option[String] = - frame match { - case Frame.SimpleError(message) => Some(message) - case Frame.BulkError(message) => Some(message.asUtf8String) - case _ => None - } - - // Return error frames from either the top-level transaction replies or the array returned by EXEC. - def execErrors(frames: Vector[Frame]): Iterator[ServerError] = { - val nested = frames.lastOption match { - case Some(Frame.Array(elems)) => elems.iterator - case _ => Iterator.empty[Frame] - } - (frames.iterator ++ nested).flatMap(errorOf).map(ServerError.of) - } + // Return error replies from either the top-level transaction replies or the commands' results inside EXEC. + def execErrors(replies: ExecReplies, interpreted: Either[SageException, Option[Vector[Either[SageException, Any]]]]): Iterator[ServerError] = + (replies.queued.iterator ++ Iterator.single(replies.exec)).collect { case Failure(error: ServerError) => error } ++ + interpreted.toOption.flatten.iterator.flatten.collect { case Left(error: ServerError) => error } // using a transaction scope after its block ends is an invalid state and returns IllegalStateException instead of a SageException def scopeReleasedError: IllegalStateException = @@ -96,3 +72,122 @@ private[internal] object TxSupport { } } } + +private[internal] trait LiveTransactionScope(events: Events, refresh: RefreshPolicy => Unit) extends TransactionScope[CIO, String] { + + protected val lock = new ReentrantLock() + protected var released = false + + // true after WATCH is attempted and false after EXEC or UNWATCH; prevents reuse while the server may still track watched keys + protected val armed = new AtomicBoolean(false) + + // what selects the transaction connection; the cluster derives a slot from the commands, standalone needs nothing + protected type Target + protected def targetOf(command: Command[?]): Target + protected def targetOf(commands: Vector[Command[?]]): Target + + // Run `use` on the transaction connection while holding `lock`, or complete with the reason no connection is available. A command + // accepted before release is recorded as in flight before [[release]] checks the connection, and commands submitted after release are + // rejected. + protected def withConn[A](target: Target, complete: Try[A] => Unit)(use: DedicatedConnection => Unit): Unit + + final protected def onFault(error: Throwable): Unit = refresh(Fault.categorize(error).refreshPolicy) + + final protected def faulting[A](complete: Try[A] => Unit): Try[A] => Unit = { + case failure @ Failure(error) => + onFault(error) + complete(failure) + case success => complete(success) + } + + final def watch[K: KeyCodec](key: K, rest: K*): CIO[Unit] = { + val command = Connection.watch(key, rest*) + CIO.async[Unit] { complete => + val tracked = Events.trackSpan(events, command, complete) + withConn(targetOf(command), tracked) { conn => + armed.set(true) + conn.submit(command, faulting(tracked)) + } + } + } + + def run[A](command: Command[A]): CIO[A] = + if (command.isBlocking) + CIO.fail(InvalidArgument("a Transaction cannot run blocking commands; run them individually on the client")) + else + CIO.async[A] { complete => + val tracked = Events.trackSpan(events, command, complete) + withConn(targetOf(command), tracked)(_.submit(command, faulting(tracked))) + } + + final def discard: CIO[Unit] = + CIO.async[Unit] { complete => + lock.lock() + try + if (released) complete(Failure(TxSupport.scopeReleasedError)) + else { + val conn = leasedConn + if (conn == null) complete(Success(())) // the transaction has not leased a connection or sent WATCH + else + Client.completing(complete) { + armed.set(false) + conn.submit(Connection.unwatch, faulting(complete)) + } + } + finally lock.unlock() + } + + // receives only pipelines without blocking commands, from a scope that is not released + protected def sendMultiExec[R](p: Pipeline[R]): CIO[TxSupport.ExecReplies] = + CIO.async[TxSupport.ExecReplies] { complete => + val tracked = Events.trackSpan(events, Connection.multi, complete) + withConn(targetOf(p.commands), tracked)(_.submitExec(p.commands, faulting(tracked))) + } + + // the leased connection, or null while none is leased; read under `lock` + protected def leasedConn: DedicatedConnection + + protected def giveBack(conn: DedicatedConnection, reusable: Boolean): Unit + + // Reject further operations, then release the transaction connection. Reuse it only when healthy, with no pending commands or watched + // keys. A transaction that did not submit any commands has no connection to release. + final private[internal] def release(): Unit = { + lock.lock() + val (conn, reusable) = + try { + released = true + val c = leasedConn + (c, c != null && c.isQuiescent && !armed.get) + } finally lock.unlock() + if (conn != null) giveBack(conn, reusable) + } + + protected def isReleased: Boolean = { + lock.lock() + try released + finally lock.unlock() + } + + final private[sage] def exec[R](p: Pipeline[R]): CIO[Option[R]] = + runExec(p).flatMap { + case None => CIO.value(None) + case Some(results) => p.finish(results).map(Some(_)).fold(CIO.fail(_), CIO.value(_)) + } + + // return None when EXEC reports a WATCH abort and Some with one decoded result per command; a queueing error fails the effect before execution + private def runExec[R](p: Pipeline[R]): CIO[Option[Vector[Either[SageException, Any]]]] = + if (isReleased) + CIO.fail(TxSupport.scopeReleasedError) + // skip MULTI/EXEC for an empty pipeline only when WATCH is inactive; watched keys still require EXEC to detect concurrent changes + else if (p.commands.isEmpty && !armed.get) + CIO.value(Some(Vector.empty)) + else if (p.commands.exists(_.isBlocking)) + CIO.fail(InvalidArgument("a Transaction cannot carry blocking commands; run them individually on the client")) + else + sendMultiExec(p).flatMap { replies => + armed.set(false) // EXEC clears WATCH/MULTI state server-side whether it committed or aborted + val interpreted = TxSupport.interpretExec(p.commands, replies) + TxSupport.execErrors(replies, interpreted).map(Fault.categorize(_).refreshPolicy).maxByOption(_.ordinal).foreach(refresh) + interpreted.fold(CIO.fail(_), CIO.value(_)) + } +} diff --git a/sage-client/shared/src/test/scala/sage/client/SageConfigSpec.scala b/sage-client/shared/src/test/scala/sage/client/SageConfigSpec.scala index af744bbe..e6a7a15b 100644 --- a/sage-client/shared/src/test/scala/sage/client/SageConfigSpec.scala +++ b/sage-client/shared/src/test/scala/sage/client/SageConfigSpec.scala @@ -17,6 +17,15 @@ class SageConfigSpec extends munit.FunSuite { assertEquals(parsed("redis://h").tls, None) } + test("the scheme is case-insensitive under any default locale, so REDIS and REDISS parse with a Turkish locale") { + val previous = java.util.Locale.getDefault + java.util.Locale.setDefault(java.util.Locale.forLanguageTag("tr-TR")) + try { + assertEquals(parsed("REDIS://h").tls, None) + assertEquals(parsed("REDISS://h").tls, Some(TlsConfig(TrustSource.System))) + } finally java.util.Locale.setDefault(previous) + } + test("userinfo becomes auth, with the default user when only a password is given") { assertEquals(parsed("redis://alice:secret@h").auth, Some(AuthConfig("secret", "alice"))) assertEquals(parsed("redis://:secret@h").auth, Some(AuthConfig("secret", "default"))) diff --git a/sage-client/shared/src/test/scala/sage/client/internal/BootstrapSpec.scala b/sage-client/shared/src/test/scala/sage/client/internal/BootstrapSpec.scala index 4b219ccd..fcb2b655 100644 --- a/sage-client/shared/src/test/scala/sage/client/internal/BootstrapSpec.scala +++ b/sage-client/shared/src/test/scala/sage/client/internal/BootstrapSpec.scala @@ -1,15 +1,16 @@ package sage.client.internal -import scala.util.{Failure, Success} - import sage.SageException.{ConnectionLost, ServerError} -import sage.client.AuthConfig +import sage.client.{AuthConfig, SageConfig} import sage.commands.Connection +import sage.protocol.Frame class BootstrapSpec extends munit.FunSuite { private def lines(auth: Option[AuthConfig], database: Int, clientName: Option[String]): Vector[String] = - Bootstrap.commands(auth, database, clientName).map(c => (c.name +: c.args.map(_.asUtf8String)).mkString(" ")) + Bootstrap + .commands(SageConfig(auth = auth, database = database, clientName = clientName)) + .map(c => (c.name +: c.args.map(_.asUtf8String)).mkString(" ")) test("the default bootstrap is HELLO then library identification, no SELECT") { val cmds = lines(None, 0, None) @@ -30,63 +31,52 @@ class BootstrapSpec extends munit.FunSuite { } test("HELLO carries AUTH when credentials are configured") { - val first = Bootstrap.commands(Some(AuthConfig("pw", "alice")), 0, None).head + val first = Bootstrap.commands(SageConfig(auth = Some(AuthConfig("pw", "alice")))).head assertEquals(first.name, "HELLO") assert(first.args.map(_.asUtf8String).containsSlice(Vector("AUTH", "alice", "pw"))) } + // a dedicated connection whose transport answers each written command through `reply`; the second value reports whether it was closed + private def connection(reply: (String, FakeTransport) => Seq[Frame]): (DedicatedConnection, () => Boolean) = { + var transport: FakeTransport = null + val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => { + transport = new FakeTransport(onFrame, onClosed, payload => reply(payload.asUtf8String, transport)) + transport + } + (new DedicatedConnection(factory, new ManualScheduler), () => transport.closeCount > 0) + } + + private def setupReply(setInfo: FakeTransport => Seq[Frame]): (String, FakeTransport) => Seq[Frame] = (payload, transport) => + if (payload.contains("HELLO")) Seq(Replies.hello) else if (payload.contains("SETINFO")) setInfo(transport) else Seq(Replies.ok) + test("a CLIENT SETINFO error does not abort setup, so a pre-7.2 server still connects") { - val unknown = Failure(ServerError("ERR", "Unknown subcommand or wrong number of arguments for 'SETINFO'.")) - var closed = false - Bootstrap.run( - Bootstrap.commands(None, 0, None), - connectTimeoutMillis = 1000, - submit = (c, cb) => cb(if (c.name == "CLIENT" && c.args.head.asUtf8String == "SETINFO") unknown else Success(())), - close = () => closed = true - ) - assert(!closed, "connection must stay open when only library identification fails") + val (conn, closed) = connection(setupReply(_ => Seq(Frame.SimpleError("ERR Unknown subcommand or wrong number of arguments for 'SETINFO'.")))) + conn.handshake(Bootstrap.commands(SageConfig()), 1000) + assert(!closed(), "connection must stay open when only library identification fails") } - test("a CLIENT TRACKING error is tolerated and reported, so a server that denies tracking still connects (ADR-0045)") { - val denied = Failure(ServerError("ERR", "This instance has cluster support disabled")) - var closed = false - var tolerated = Vector.empty[String] - Bootstrap.run( - Bootstrap.commands(None, 0, None) :+ Connection.clientTrackingOnOptin, - connectTimeoutMillis = 1000, - submit = (c, cb) => cb(if (c.name == "CLIENT" && c.args.head.asUtf8String == "TRACKING") denied else Success(())), - close = () => closed = true, - onTolerated = c => tolerated = tolerated :+ s"${c.name} ${c.args.head.asUtf8String}" - ) - assert(!closed, "connection must stay open when only tracking is denied") - assertEquals(tolerated, Vector("CLIENT TRACKING")) + test("a step returns a server error reply without closing, so a server that denies tracking still connects (ADR-0045)") { + val (conn, closed) = connection((_, _) => Seq(Frame.SimpleError("ERR This instance has cluster support disabled"))) + conn.start() + val result = conn.step(Connection.clientTrackingOnOptin, 1000) + assert(!closed(), "connection must stay open when only tracking is denied") + assertEquals(result, Some(ServerError("ERR", "This instance has cluster support disabled"))) } test("a CLIENT SETINFO connection loss is NOT tolerated: it closes and throws") { - var closed = false - val thrown = intercept[ConnectionLost] { - Bootstrap.run( - Bootstrap.commands(None, 0, None), - connectTimeoutMillis = 1000, - submit = (c, cb) => - cb( - if (c.name == "CLIENT" && c.args.head.asUtf8String == "SETINFO") Failure(ConnectionLost(mayHaveExecuted = false)) - else Success(()) - ), - close = () => closed = true - ) - } + val (conn, closed) = connection(setupReply { transport => + transport.close() + Nil + }) + val thrown = intercept[ConnectionLost](conn.handshake(Bootstrap.commands(SageConfig()), 1000)) assertEquals(thrown.mayHaveExecuted, false) - assert(closed, "a broken connection must be discarded even on a best-effort command") + assert(closed(), "a broken connection must be discarded even on a best-effort command") } test("a load-bearing command error aborts setup and closes the connection") { - val error = ServerError("NOAUTH", "Authentication required.") - var closed = false - val thrown = intercept[ServerError] { - Bootstrap.run(Vector(Connection.hello(None)), 1000, (_, cb) => cb(Failure(error)), () => closed = true) - } - assertEquals(thrown, error) - assert(closed, "connection must be closed when a load-bearing command fails") + val (conn, closed) = connection((_, _) => Seq(Frame.SimpleError("NOAUTH Authentication required."))) + val thrown = intercept[ServerError](conn.handshake(Vector(Connection.hello(None)), 1000)) + assertEquals(thrown, ServerError("NOAUTH", "Authentication required.")) + assert(closed(), "connection must be closed when a load-bearing command fails") } } diff --git a/sage-client/shared/src/test/scala/sage/client/internal/ClientCacheSpec.scala b/sage-client/shared/src/test/scala/sage/client/internal/ClientCacheSpec.scala index 37c70797..d570326f 100644 --- a/sage-client/shared/src/test/scala/sage/client/internal/ClientCacheSpec.scala +++ b/sage-client/shared/src/test/scala/sage/client/internal/ClientCacheSpec.scala @@ -3,29 +3,50 @@ package sage.client.internal import scala.util.{Failure, Success, Try} import sage.Bytes +import sage.commands.{Command, Execution} import sage.protocol.Frame class ClientCacheSpec extends munit.FunSuite { - private def key(s: String): Bytes = Bytes.utf8(s) - private def frame(s: String): Frame = Frame.BulkString(Bytes.utf8(s)) - private def collector(): (Try[Frame] => Unit, () => Option[Try[Frame]]) = { + private def key(s: String): Bytes = Bytes.utf8(s) + private def frame(s: String): Frame = Frame.BulkString(Bytes.utf8(s)) + private def collector(): (Try[Frame] => Unit, () => Option[Try[Frame]]) = { var slot: Option[Try[Frame]] = None ((r: Try[Frame]) => slot = Some(r), () => slot) } - private def hitFrame(acquired: ClientCache.Acquire): Frame = acquired match { + private def hitFrame(acquired: ClientCache.Acquire): Frame = acquired match { case ClientCache.Acquire.Hit(frame, _) => frame case other => fail(s"expected a hit, got $other") } + private def fetching(acquired: ClientCache.Acquire): ClientCache.Fetching = acquired match { + case ClientCache.Acquire.Fetch(ticket) => ticket + case other => fail(s"expected a fetch, got $other") + } + private def assertFetch(acquired: ClientCache.Acquire): Unit = fetching(acquired): Unit + // seed an entry through the acquire/store path a cache miss takes + private def seed(cache: ClientCache, cmd: Bytes, tracked: Bytes, value: Frame, now: Long, ttlMillis: Long): Unit = + cache.store(fetching(cache.acquire(cmd, Vector(tracked), now, _ => ())), value, now, ttlMillis) + + test("a command whose declared key positions fall outside its arguments is not cacheable") { + val decode: Frame => Either[sage.SageException.DecodeError, Frame] = Right(_) + assert(Client.cacheable(Command("GET", Vector(0), Vector(key("k")), decode, isReadOnly = true, cacheable = true))) + assert(!Client.cacheable(Command("GET", Vector(5), Vector(key("k")), decode, isReadOnly = true, cacheable = true))) + } + + test("a blocking command is not cacheable, since a cached read runs on the shared multiplexed connection") { + val decode: Frame => Either[sage.SageException.DecodeError, Frame] = Right(_) + val blocking = Command("BLPOP", Vector(0), Vector(key("k"), key("0")), decode, Execution.Blocking, isReadOnly = true, cacheable = true) + assert(!Client.cacheable(blocking)) + } test("first miss fetches, a concurrent miss waits, and the stored reply reaches both") { val cache = new ClientCache(1024) val cmd = key("GET foo") val (w1, get1) = collector() val (w2, get2) = collector() - assertEquals(cache.acquire(cmd, Vector(key("foo")), 0L, w1), ClientCache.Acquire.Fetch) + val ticket = fetching(cache.acquire(cmd, Vector(key("foo")), 0L, w1)) assertEquals(cache.acquire(cmd, Vector(key("foo")), 0L, w2), ClientCache.Acquire.Wait) - cache.store(cmd, Vector(key("foo")), frame("bar"), 0L, 1000L) + cache.store(ticket, frame("bar"), 0L, 1000L) assertEquals(get1(), Some(Success(frame("bar")))) assertEquals(get2(), Some(Success(frame("bar")))) assertEquals(hitFrame(cache.acquire(cmd, Vector(key("foo")), 0L, _ => ())), frame("bar")) @@ -34,65 +55,70 @@ class ClientCacheSpec extends munit.FunSuite { test("absolute TTL: a hit before expiry, a refetch at expiry") { val cache = new ClientCache(1024) val cmd = key("GET foo") - cache.store(cmd, Vector(key("foo")), frame("bar"), 0L, 1000L) + seed(cache, cmd, key("foo"), frame("bar"), 0L, 1000L) assertEquals(hitFrame(cache.acquire(cmd, Vector(key("foo")), 999L, _ => ())), frame("bar")) - assertEquals(cache.acquire(cmd, Vector(key("foo")), 1000L, _ => ()), ClientCache.Acquire.Fetch) + assertFetch(cache.acquire(cmd, Vector(key("foo")), 1000L, _ => ())) } test("an invalidation for a tracked key evicts every entry that touched it") { val cache = new ClientCache(1024) val get = key("GET foo") val range = key("GETRANGE foo 0 1") - cache.store(get, Vector(key("foo")), frame("bar"), 0L, 10000L) - cache.store(range, Vector(key("foo")), frame("ba"), 0L, 10000L) + seed(cache, get, key("foo"), frame("bar"), 0L, 10000L) + seed(cache, range, key("foo"), frame("ba"), 0L, 10000L) cache.invalidate(key("foo")) - assertEquals(cache.acquire(get, Vector(key("foo")), 0L, _ => ()), ClientCache.Acquire.Fetch) - assertEquals(cache.acquire(range, Vector(key("foo")), 0L, _ => ()), ClientCache.Acquire.Fetch) + assertFetch(cache.acquire(get, Vector(key("foo")), 0L, _ => ())) + assertFetch(cache.acquire(range, Vector(key("foo")), 0L, _ => ())) } test("flush drops the whole cache") { val cache = new ClientCache(1024) val cmd = key("GET foo") - cache.store(cmd, Vector(key("foo")), frame("bar"), 0L, 10000L) + seed(cache, cmd, key("foo"), frame("bar"), 0L, 10000L) cache.flush() - assertEquals(cache.acquire(cmd, Vector(key("foo")), 0L, _ => ()), ClientCache.Acquire.Fetch) + assertFetch(cache.acquire(cmd, Vector(key("foo")), 0L, _ => ())) } - test("a hit carries the epoch it was acquired under, and a flush retires it") { - val cache = new ClientCache(1024) - val cmd = key("GET foo") - cache.store(cmd, Vector(key("foo")), frame("bar"), 0L, 10000L) - val ClientCache.Acquire.Hit(_, epoch) = cache.acquire(cmd, Vector(key("foo")), 0L, _ => ()): @unchecked - assert(cache.isCurrent(epoch)) + test("a hit taken before a flush is retired, and looking it up again fetches from the server") { + val cache = new ClientCache(1024) + val cmd = key("GET foo") + seed(cache, cmd, key("foo"), frame("bar"), 0L, 10000L) + val hit = cache.acquire(cmd, Vector(key("foo")), 0L, _ => ()).asInstanceOf[ClientCache.Acquire.Hit] + assert(cache.isCurrent(hit)) cache.flush() - assert(!cache.isCurrent(epoch)) + assert(!cache.isCurrent(hit), "a hit retired by a flush must not be delivered") + assertFetch(cache.acquire(cmd, Vector(key("foo")), 0L, _ => ())) } - test("a server flush retires a hit for refetch, a topology flush retires it for reroute") { - val cache = new ClientCache(1024) - val cmd = key("GET foo") - cache.store(cmd, Vector(key("foo")), frame("bar"), 0L, 10000L) - val ClientCache.Acquire.Hit(_, served) = cache.acquire(cmd, Vector(key("foo")), 0L, _ => ()): @unchecked + test("a hit taken after a flush is current") { + val cache = new ClientCache(1024) + val cmd = key("GET foo") cache.flush() - assert(!cache.isCurrent(served)) - assert(!cache.rerouteRetired(served), "a server flush leaves ownership unchanged: refetch, not reroute") - - cache.store(cmd, Vector(key("foo")), frame("bar"), 0L, 10000L) - val ClientCache.Acquire.Hit(_, moved) = cache.acquire(cmd, Vector(key("foo")), 0L, _ => ()): @unchecked - cache.flushForReroute() - assert(!cache.isCurrent(moved)) - assert(cache.rerouteRetired(moved), "a topology flush requires rerouting to the new owner") + seed(cache, cmd, key("foo"), frame("bar"), 0L, 10000L) + val hit = cache.acquire(cmd, Vector(key("foo")), 0L, _ => ()).asInstanceOf[ClientCache.Acquire.Hit] + assert(cache.isCurrent(hit)) } test("an invalidation mid-flight delivers the reply but does not cache it") { val cache = new ClientCache(1024) val cmd = key("GET foo") val (w1, get1) = collector() - assertEquals(cache.acquire(cmd, Vector(key("foo")), 0L, w1), ClientCache.Acquire.Fetch) - cache.invalidate(key("foo")) // arrives before the fetch completes - cache.store(cmd, Vector(key("foo")), frame("bar"), 0L, 10000L) - assertEquals(get1(), Some(Success(frame("bar")))) // waiter still gets the value - assertEquals(cache.acquire(cmd, Vector(key("foo")), 0L, _ => ()), ClientCache.Acquire.Fetch) // but it was not stored + val ticket = fetching(cache.acquire(cmd, Vector(key("foo")), 0L, w1)) + cache.invalidate(key("foo")) // arrives before the fetch completes + cache.store(ticket, frame("bar"), 0L, 10000L) + assertEquals(get1(), Some(Success(frame("bar")))) // waiter still gets the value + assertFetch(cache.acquire(cmd, Vector(key("foo")), 0L, _ => ())) // but it was not stored + } + + test("a flush mid-flight (MOVED, topology change, dropped connection) delivers the reply but does not cache it") { + val cache = new ClientCache(1024) + val cmd = key("GET foo") + val (w1, get1) = collector() + val ticket = fetching(cache.acquire(cmd, Vector(key("foo")), 0L, w1)) + cache.flush() + cache.store(ticket, frame("bar"), 0L, 10000L) + assertEquals(get1(), Some(Success(frame("bar")))) + assertFetch(cache.acquire(cmd, Vector(key("foo")), 0L, _ => ())) } test("a failed fetch reaches every waiter and stores nothing") { @@ -100,18 +126,18 @@ class ClientCacheSpec extends munit.FunSuite { val cmd = key("GET foo") val (w1, get1) = collector() val boom = new RuntimeException("boom") - assertEquals(cache.acquire(cmd, Vector(key("foo")), 0L, w1), ClientCache.Acquire.Fetch) - cache.fail(cmd, boom) + val ticket = fetching(cache.acquire(cmd, Vector(key("foo")), 0L, w1)) + cache.fail(ticket, boom) assertEquals(get1(), Some(Failure(boom))) - assertEquals(cache.acquire(cmd, Vector(key("foo")), 0L, _ => ()), ClientCache.Acquire.Fetch) + assertFetch(cache.acquire(cmd, Vector(key("foo")), 0L, _ => ())) } test("the bytes cap evicts the least-recently-used entry") { val cache = new ClientCache(100) // each 50-byte value is ~66 bytes stored, so two together exceed the cap val big = frame("x" * 50) - cache.store(key("GET a"), Vector(key("a")), big, 0L, 10000L) - cache.store(key("GET b"), Vector(key("b")), big, 0L, 10000L) - assertEquals(cache.acquire(key("GET a"), Vector(key("a")), 0L, _ => ()), ClientCache.Acquire.Fetch) // evicted + seed(cache, key("GET a"), key("a"), big, 0L, 10000L) + seed(cache, key("GET b"), key("b"), big, 0L, 10000L) + assertFetch(cache.acquire(key("GET a"), Vector(key("a")), 0L, _ => ())) // evicted assertEquals(hitFrame(cache.acquire(key("GET b"), Vector(key("b")), 0L, _ => ())), big) } } diff --git a/sage-client/shared/src/test/scala/sage/client/internal/DedicatedPoolSpec.scala b/sage-client/shared/src/test/scala/sage/client/internal/DedicatedPoolSpec.scala index 6ff3a848..d09feb29 100644 --- a/sage-client/shared/src/test/scala/sage/client/internal/DedicatedPoolSpec.scala +++ b/sage-client/shared/src/test/scala/sage/client/internal/DedicatedPoolSpec.scala @@ -8,7 +8,7 @@ import Replies.bulk import sage.Bytes import sage.SageException.{ConnectionLost, NotConnected, TimedOut} -import sage.client.{BackoffConfig, DedicatedPoolConfig, WatchdogConfig} +import sage.client.{BackoffConfig, CacheConfig, DedicatedPoolConfig, SageConfig, WatchdogConfig} import sage.cluster.Node import sage.commands.{BlockTimeout, Connection, Lists, Server} import sage.protocol.Frame @@ -17,32 +17,28 @@ class DedicatedPoolSpec extends munit.FunSuite { private val popReply: Frame = Frame.Array(Vector(bulk("k"), bulk("v"))) - // HELLO always answers so the bootstrap succeeds; the blocking command's reply is the test's to script - private def replyWith(blocking: Seq[Frame]): Bytes => Seq[Frame] = - payload => if (payload.asUtf8String.contains("HELLO")) Seq(Replies.hello) else blocking + // the setup always succeeds; the blocking command's reply is the test's to script + private def replyWith(blocking: Seq[Frame]): Bytes => Seq[Frame] = Replies.withSetup(_ => blocking) private def make( respond: Bytes => Seq[Frame], isLive: () => Boolean = () => true, - liveGeneration: () => Option[MultiplexedConnection.Generation] = () => Some(MultiplexedConnection.Generation.initial), config: DedicatedPoolConfig = DedicatedPoolConfig() ): (DedicatedPool, ManualScheduler, mutable.ArrayBuffer[FakeTransport]) = { - val scheduler = new ManualScheduler - val transports = mutable.ArrayBuffer.empty[FakeTransport] - val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => { + val scheduler = new ManualScheduler + val transports = mutable.ArrayBuffer.empty[FakeTransport] + val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => { val transport = new FakeTransport(onFrame, onClosed, respond) transports += transport transport } - // mirrors the real MultiplexedConnection: a generation is current when the connection is live and the recorded generation matches - val isCurrent: MultiplexedConnection.Generation => Boolean = g => liveGeneration().contains(g) - val pool = - new DedicatedPool(factory, Vector(Connection.hello()), scheduler, isLive, liveGeneration, isCurrent, config, 1000L) + val pool = new DedicatedPool(factory, Vector(Connection.hello()), scheduler, isLive, config, 1000L) (pool, scheduler, transports) } private val lockWrite = - new LockCommands[String](3.seconds, "lock").command(Bytes.utf8("key"), "owner", LockCommands.Operation.Acquire, cached = true) + new LockExecutor[String](3.seconds, "lock", replicaAcknowledgement = true) + .command(Bytes.utf8("key"), "owner", LockExecutor.Operation.Acquire, cached = true) private def replication( scheduler: Scheduler, @@ -148,7 +144,8 @@ class DedicatedPoolSpec extends munit.FunSuite { var result: Option[Try[Boolean]] = None pool.useLockWrite(lockWrite, true, r => result = Some(r), new DedicatedPool.Lease, replication(scheduler, 1, 100L)) scheduler.advance(Duration.Zero) - assert(transports.head.written.last.sameBytes(Bytes.concat(Vector(Connection.asking.encode, lockWrite.encode)))) + val wire = Bytes.concat(transports.head.written).asUtf8String + assert(wire.endsWith(Bytes.concat(Vector(Connection.asking.encode, lockWrite.encode)).asUtf8String), wire) transports.head.emit(Replies.ok) transports.head.emit(Frame.Integer(1)) transports.head.emit(Replies.masterRole(Node("replica", 6380))) @@ -177,21 +174,36 @@ class DedicatedPoolSpec extends munit.FunSuite { } test("a stalled confirmation leaves ordinary commands free and cancellation releases the pool slot") { - val (pool, scheduler, transports) = make(replyWith(Nil), config = DedicatedPoolConfig(maxConnections = 1)) - val shared = MultiplexedConnection.connect( - (onFrame, onClosed) => new FakeTransport(onFrame, onClosed, _ => Seq(Frame.SimpleString("PONG"))), + val scheduler = new ManualScheduler + // the first connection is the multiplexed one; the rest are the pool's dedicated connections + val transports = mutable.ArrayBuffer.empty[FakeTransport] + var created = 0 + val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => { + created += 1 + if (created == 1) new FakeTransport(onFrame, onClosed, Replies.withSetup(_ => Seq(Frame.SimpleString("PONG")))) + else { + val transport = new FakeTransport(onFrame, onClosed, replyWith(Nil)) + transports += transport + transport + } + } + val node = new MultiplexedConnection( + factory, scheduler, - Vector.empty, - BackoffConfig(), - WatchdogConfig(enabled = false), - 1.second, - Duration.Zero - ) - val node = new NodeClient(shared, pool) - val lease = new DedicatedPool.Lease - var result: Option[Try[Boolean]] = None - var refreshed = false - node.submitLockWrite( + SageConfig( + reconnect = BackoffConfig(), + watchdog = WatchdogConfig(enabled = false), + connectTimeout = 1.second, + closeTimeout = Duration.Zero, + clientCache = CacheConfig(enabled = false), + dedicatedPool = DedicatedPoolConfig(maxConnections = 1) + ), + MultiplexedConnection.NodeRole.Master + ).start() + val lease = new DedicatedPool.Lease + var result: Option[Try[Boolean]] = None + var refreshed = false + node.pool.useLockWrite( lockWrite, false, r => result = Some(r), @@ -201,15 +213,15 @@ class DedicatedPoolSpec extends munit.FunSuite { scheduler.advance(Duration.Zero) transports.head.emit(Frame.Integer(1)) transports.head.emit(Replies.masterRole(Node("replica", 6380))) - var ping: Option[Try[String]] = None - node.submit(Connection.ping(), false, r => ping = Some(r)) + var ping: Option[Try[String]] = None + node.submit(Connection.ping(), r => ping = Some(r)) assertEquals(ping, Some(Success("PONG"))) assertEquals(result, None) lease.cancel() scheduler.advance(Duration.Zero) assertEquals(result, Some(Failure(ConnectionLost(mayHaveExecuted = true)))) assert(refreshed) - node.submitLockWrite(lockWrite, false, _ => (), new DedicatedPool.Lease, replication(scheduler, 0, 100L)) + node.pool.useLockWrite(lockWrite, false, _ => (), new DedicatedPool.Lease, replication(scheduler, 0, 100L)) scheduler.advance(Duration.Zero) assertEquals(transports.size, 2) assertEquals(transports.head.closeCount, 1) @@ -227,6 +239,19 @@ class DedicatedPoolSpec extends munit.FunSuite { assertEquals(transports.size, 1) } + test("a socket that closes before the setup reply fails the blocking command as never sent") { + val scheduler = new ManualScheduler + val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => { + lazy val transport: FakeTransport = new FakeTransport(onFrame, onClosed, _ => { transport.close(); Nil }) + transport + } + val pool = new DedicatedPool(factory, Vector(Connection.hello()), scheduler, () => true, DedicatedPoolConfig(), 1000L) + var result: Option[Try[Option[(String, String)]]] = None + pool.use(blPop, r => result = Some(r)) + scheduler.advance(Duration.Zero) + assertEquals(result, Some(Failure(ConnectionLost(mayHaveExecuted = false)))) + } + test("a released connection is reused rather than reopened") { val (pool, scheduler, transports) = make(replyWith(Seq(popReply))) pool.use(blPop, _ => ()) @@ -385,12 +410,12 @@ class DedicatedPoolSpec extends munit.FunSuite { transports += t t } - val conn = DedicatedConnection.create(factory, 1000L) - conn.establish(Vector(Connection.hello())) + val conn = new DedicatedConnection(factory, new ManualScheduler) + conn.handshake(Vector(Connection.hello()), 1000L) transports.head.autoWrite = false // keep the command in the queue so transport teardown reports it as unsent var healthyWhenFailed: Option[Boolean] = None - conn.submit(blPop, _ => healthyWhenFailed = Some(conn.isHealthy)) + conn.submit(blPop, _ => healthyWhenFailed = Some(!conn.isDead)) transports.head.close() assertEquals(healthyWhenFailed, Some(false)) } @@ -420,51 +445,48 @@ class DedicatedPoolSpec extends munit.FunSuite { assertEquals(result, Some(Failure(ConnectionLost(mayHaveExecuted = true)))) } - test("an idle connection from a previous generation is discarded rather than reused") { - var live = Option(MultiplexedConnection.Generation.initial) - val (pool, scheduler, transports) = make(replyWith(Seq(popReply)), liveGeneration = () => live) + test("an idle connection opened before a liveness loss is discarded rather than reused") { + val (pool, scheduler, transports) = make(replyWith(Seq(popReply))) pool.use(blPop, _ => ()) scheduler.advance(Duration.Zero) assertEquals(transports.size, 1) - live = Some(MultiplexedConnection.Generation.initial.next) // the multiplexed connection reconnected (e.g. failover) under the pool + pool.onLivenessLost() // the multiplexed connection reconnected (e.g. failover) under the pool pool.use(blPop, _ => ()) scheduler.advance(Duration.Zero) assertEquals(transports.size, 2) } - test("a connection built across a reconnect is admitted under the new generation, not discarded") { - var live = Option(MultiplexedConnection.Generation.initial) - var bumped = false - // the multiplexed connection reconnects (generation bumps) while the first dedicated connection is running its HELLO bootstrap + test("a connection built across a reconnect is admitted, not discarded") { + var lose: () => Unit = () => () + // the multiplexed connection reconnects while the first dedicated connection is running its HELLO bootstrap val respond: Bytes => Seq[Frame] = payload => if (payload.asUtf8String.contains("HELLO")) { - if (!bumped) { - bumped = true - live = Some(MultiplexedConnection.Generation.initial.next) - } + lose() + lose = () => () Seq(Replies.hello) } else Seq(popReply) - val (pool, scheduler, transports) = make(respond, liveGeneration = () => live) + val (pool, scheduler, transports) = make(respond) + lose = () => pool.onLivenessLost() var result: Option[Try[Option[(String, String)]]] = None pool.use(blPop, r => result = Some(r)) scheduler.advance(Duration.Zero) assertEquals(result, Some(Success(Some(("k", "v"))))) - // recording the generation after establishment makes the connection current. The pool keeps it and does not retry. + // the liveness loss preceded admission, so the pool keeps the connection and reuses it + pool.use(blPop, _ => ()) + scheduler.advance(Duration.Zero) assertEquals(transports.size, 1) } test("an idle connection is not reused when the connection leaves Live between lease and acquire") { var live = true - val gen = MultiplexedConnection.Generation.initial - val (pool, scheduler, transports) = - make(replyWith(Seq(popReply)), isLive = () => live, liveGeneration = () => if (live) Some(gen) else None) + val (pool, scheduler, transports) = make(replyWith(Seq(popReply)), isLive = () => live) pool.use(blPop, _ => ()) scheduler.advance(Duration.Zero) - assertEquals(transports.size, 1) // established and returned to idle at generation `gen` + assertEquals(transports.size, 1) // The lease check observes Live, but the multiplexed connection starts reconnecting before the offloaded acquire runs. The idle connection - // has the same generation but is no longer live, so the pool refuses it. + // was not retired but is no longer live, so the pool refuses it. var result: Option[Try[Option[(String, String)]]] = None pool.use(blPop, r => result = Some(r)) live = false @@ -475,10 +497,8 @@ class DedicatedPoolSpec extends munit.FunSuite { test("an exhausted pool fails fast NotConnected, not TimedOut, when the connection is not live") { var live = true - val gen = MultiplexedConnection.Generation.initial val config = DedicatedPoolConfig(maxConnections = 1, acquireTimeout = 50.millis, idleTimeout = Duration.Inf) - val (pool, scheduler, transports) = - make(replyWith(Nil), isLive = () => live, liveGeneration = () => if (live) Some(gen) else None, config = config) + val (pool, scheduler, transports) = make(replyWith(Nil), isLive = () => live, config = config) pool.use(blPop, _ => ()) // the only slot is held by a BLPOP that is still waiting for a reply scheduler.advance(Duration.Zero) assertEquals(transports.size, 1) @@ -555,17 +575,21 @@ class DedicatedPoolSpec extends munit.FunSuite { transports += t t } - val connection = MultiplexedConnection.connect( + val config = DedicatedPoolConfig(maxConnections = 1, acquireTimeout = 10.seconds, idleTimeout = Duration.Inf) + val connection = new MultiplexedConnection( factory, scheduler, - Vector(Connection.hello()), - BackoffConfig(1.milli, 1.milli, 1.0), - WatchdogConfig(enabled = false), - 1.second, - Duration.Zero - ) - val config = DedicatedPoolConfig(maxConnections = 1, acquireTimeout = 10.seconds, idleTimeout = Duration.Inf) - val pool = DedicatedPool.forConnection(factory, Vector(Connection.hello()), scheduler, connection, config, 1000L) + SageConfig( + reconnect = BackoffConfig(1.milli, 1.milli, 1.0), + watchdog = WatchdogConfig(enabled = false), + connectTimeout = 1.second, + closeTimeout = Duration.Zero, + clientCache = CacheConfig(enabled = false), + dedicatedPool = config + ), + MultiplexedConnection.NodeRole.Master + ).start() + val pool = connection.pool val held = pool.acquireForTransaction() @@ -591,9 +615,8 @@ class DedicatedPoolSpec extends munit.FunSuite { connecting.set(transport) transport } - val gen = MultiplexedConnection.Generation.initial val pool = - new DedicatedPool(factory, Vector(Connection.hello()), scheduler, () => true, () => Some(gen), _ == gen, DedicatedPoolConfig(), 1000L) + new DedicatedPool(factory, Vector(Connection.hello()), scheduler, () => true, DedicatedPoolConfig(), 1000L) pool.use(blPop, _ => ()) val establishing = new Thread(() => scheduler.advance(Duration.Zero)) // blocks inside ConnectingTransport.start() diff --git a/sage-client/shared/src/test/scala/sage/client/internal/EventsSpec.scala b/sage-client/shared/src/test/scala/sage/client/internal/EventsSpec.scala index df02cfb1..0994a10e 100644 --- a/sage-client/shared/src/test/scala/sage/client/internal/EventsSpec.scala +++ b/sage-client/shared/src/test/scala/sage/client/internal/EventsSpec.scala @@ -9,7 +9,7 @@ import scala.jdk.CollectionConverters.* import scala.util.{Failure, Success} import sage.{Bytes, CommandTracer, SageEvent, SageListener} -import sage.client.{BackoffConfig, WatchdogConfig} +import sage.client.{BackoffConfig, CacheConfig, SageConfig, WatchdogConfig} import sage.cluster.Node import sage.commands.{Connection, Strings} import sage.protocol.Frame @@ -40,17 +40,30 @@ class EventsSpec extends munit.FunSuite { private def connect( node: Option[Node], events: Events, - respond: Bytes => Seq[Frame] = _ => Nil + respond: Bytes => Seq[Frame] = _ => Nil, + clientCache: CacheConfig = CacheConfig(enabled = false) ): (MultiplexedConnection, mutable.ArrayBuffer[FakeTransport]) = { val transports = mutable.ArrayBuffer.empty[FakeTransport] val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => { - val transport = new FakeTransport(onFrame, onClosed, respond) + val transport = new FakeTransport(onFrame, onClosed, Replies.withSetup(respond)) transports += transport transport } val connection = - MultiplexedConnection - .connect(factory, new ManualScheduler, Vector.empty, fixedBackoff, noWatchdog, 1.second, Duration.Zero, 1L << 20, node, events) + new MultiplexedConnection( + factory, + new ManualScheduler, + SageConfig( + reconnect = fixedBackoff, + watchdog = noWatchdog, + connectTimeout = 1.second, + closeTimeout = Duration.Zero, + clientCache = clientCache + ), + MultiplexedConnection.NodeRole.Master, + node, + events + ).start() (connection, transports) } @@ -198,6 +211,19 @@ class EventsSpec extends munit.FunSuite { // --- command completion ---------------------------------------------------------------------------------------------------------------- + test("a submit interrupted on the caller's thread still completes the tracked command") { + val rec = new Recording + var settled = Option.empty[scala.util.Try[String]] + val tracked = Events.trackCommand[String](rec, Connection.ping(None), r => settled = Some(r)) + val interrupted = new InterruptedException + Client.completing(tracked)(throw interrupted) + assertEquals(settled, Some(Failure(interrupted))) + rec.events match { + case Vector(SageEvent.CommandCompleted("PING", None, _, sage.Outcome.Failed(`interrupted`))) => () + case other => fail(s"unexpected: $other") + } + } + test("trackCommand emits a completion with name and outcome, and is transparent when disabled") { val rec = new Recording var settled = Option.empty[String] @@ -224,46 +250,17 @@ class EventsSpec extends munit.FunSuite { } } - test("submitBatchOnOne attributes each completion to the selected node, and routes its span there") { + test("a tracked batch attributes each completion to the selected node, and routes its span there") { val tracer = new RecordingTracer val rec = new Recording(Some(tracer)) val node = Node("replica", 7001) val commands = Vector(Connection.ping(None)) - Client.submitBatchOnOne( - rec, - commands, - Events.startSpans(rec, commands), - (_, cbs) => { - cbs.foreach(_(Success("PONG"))) - true - }, - _ => (), - onUnsent = () => (), - node = Some(node) - ) + val batch = new Client.TrackedBatch(rec, commands, Events.startSpans(rec, commands), _ => ()) + batch.callbacks(Some(node)).foreach(_(Success("PONG"))) assertEquals(rec.events.collect { case c: SageEvent.CommandCompleted => c.node }, Vector(Some(node))) assert(tracer.log.contains(s"routed:${node.host}:${node.port}"), s"expected routedTo the selected node, got ${tracer.log.toVector}") } - test("submitBatchOnOne attributes no node when the batch is not submitted, even if a target was selected") { - val tracer = new RecordingTracer - val rec = new Recording(Some(tracer)) - val commands = Vector(Connection.ping(None), Connection.ping(None)) - var completed: scala.util.Try[Vector[Either[sage.SageException, Any]]] = null - Client.submitBatchOnOne( - rec, - commands, - Events.startSpans(rec, commands), - (_, _) => false, - r => completed = r, - onUnsent = () => (), - node = Some(Node("replica", 7001)) - ) - assert(completed != null && completed.isFailure, s"an unsubmitted batch must fail the effect, got $completed") - assertEquals(rec.events.collect { case c: SageEvent.CommandCompleted => c.node }, Vector(None, None)) - assert(!tracer.log.exists(_.startsWith("routed:")), s"an unsent batch must route no span, got ${tracer.log.toVector}") - } - // --- tracing --------------------------------------------------------------------------------------------------------------------------- test("a tracer-only bus is enabled but emits no events, and drives the span lifecycle through trackCommand") { @@ -360,16 +357,29 @@ class EventsSpec extends munit.FunSuite { val node = Some(Node("shard-a", 6379)) val rec = new Recording val scheduler = new ManualScheduler - var healthy = true // first establish's PING succeeds; the reconnect's is rejected, as a rotated password would be + var healthy = true // first establish's HELLO succeeds; the reconnect's is rejected, as a rotated password would be val transports = mutable.ArrayBuffer.empty[FakeTransport] val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => { - val respond: Bytes => Seq[Frame] = _ => Seq(if (healthy) Frame.SimpleString("PONG") else Frame.SimpleError("WRONGPASS invalid password")) + val respond: Bytes => Seq[Frame] = + payload => if (healthy) Replies.withSetup(_ => Nil)(payload) else Seq(Frame.SimpleError("WRONGPASS invalid password")) val t = new FakeTransport(onFrame, onClosed, respond) transports += t t } - MultiplexedConnection - .connect(factory, scheduler, Vector(Connection.ping()), fixedBackoff, noWatchdog, 1.second, Duration.Zero, 1L << 20, node, rec): Unit + new MultiplexedConnection( + factory, + scheduler, + SageConfig( + clientCache = CacheConfig(enabled = false), + reconnect = fixedBackoff, + watchdog = noWatchdog, + connectTimeout = 1.second, + closeTimeout = Duration.Zero + ), + MultiplexedConnection.NodeRole.Master, + node, + rec + ).start(): Unit healthy = false transports.head.close() @@ -393,7 +403,7 @@ class EventsSpec extends munit.FunSuite { val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => { attempt += 1 if (attempt == 1) { - val t = new FakeTransport(onFrame, onClosed, _ => Seq(Frame.SimpleString("PONG"))) + val t = new FakeTransport(onFrame, onClosed, Replies.withSetup(_ => Nil)) first.set(t) t } else @@ -406,8 +416,20 @@ class EventsSpec extends munit.FunSuite { def close(): Unit = () } } - connection = MultiplexedConnection - .connect(factory, scheduler, Vector(Connection.ping()), fixedBackoff, noWatchdog, 1.second, Duration.Zero, 1L << 20, node, rec) + connection = new MultiplexedConnection( + factory, + scheduler, + SageConfig( + clientCache = CacheConfig(enabled = false), + reconnect = fixedBackoff, + watchdog = noWatchdog, + connectTimeout = 1.second, + closeTimeout = Duration.Zero + ), + MultiplexedConnection.NodeRole.Master, + node, + rec + ).start() first.get().close() scheduler.advance(1.milli) @@ -428,15 +450,18 @@ class EventsSpec extends munit.FunSuite { val (connection, _) = connect( None, rec, - respond = _ => { - writes += 1 - // the first write is the [CLIENT CACHING YES, GET] batch: reply OK to the marker, the value to the read - if (writes == 1) Seq(Frame.SimpleString("OK"), Frame.BulkString(Bytes.utf8("v"))) else Nil - } + respond = payload => + if (payload.asUtf8String.contains("TRACKING")) Seq(Frame.SimpleString("OK")) + else { + writes += 1 + // the first write is the [CLIENT CACHING YES, GET] batch: reply OK to the marker, the value to the read + if (writes == 1) Seq(Frame.SimpleString("OK"), Frame.BulkString(Bytes.utf8("v"))) else Nil + }, + clientCache = CacheConfig(enabled = true, maxBytes = 1L << 20) ) val get = Strings.get[String, String]("k") - connection.cachedSubmit(get, 60000L, (_: scala.util.Try[Option[String]]) => ()) - connection.cachedSubmit(get, 60000L, (_: scala.util.Try[Option[String]]) => ()) + for (_ <- 1 to 2) + connection.cachedSubmit(get, 60000L, (_: scala.util.Try[Option[String]]) => (), Events.fetchTracking(rec)) val evs = rec.events assertEquals(evs.collect { case c: SageEvent.Cache => c }, Vector(SageEvent.Cache.Miss("GET"), SageEvent.Cache.Hit("GET"))) // the miss touched the server, so it also produces one CommandCompleted; the hit produces none @@ -445,4 +470,31 @@ class EventsSpec extends munit.FunSuite { Vector(("GET", None, sage.Outcome.Succeeded)) ) } + + test("the parts of a cached read that start fetching at once start one span") { + val entered = new CountDownLatch(1) + val release = new CountDownLatch(1) + val starts = new java.util.concurrent.atomic.AtomicInteger + val tracer = new CommandTracer { + def onCommand(command: sage.commands.Command[?]): sage.CommandSpan = { + starts.incrementAndGet() + entered.countDown() + release.await() + sage.CommandSpan.noop + } + } + val get = Strings.get[String, String]("k") + Events.trackCached[Option[String]](new Recording(Some(tracer)), get, _ => ()) { (_, trace) => + def fetch(): Thread = Thread.ofPlatform().start(() => trace.fetching(get, (_: scala.util.Try[Option[String]]) => ()): Unit) + val first = fetch() + entered.await() + val second = fetch() + // the second part either waits for the first to finish starting the span or starts its own and parks in the tracer + while (!Set(Thread.State.BLOCKED, Thread.State.WAITING, Thread.State.TERMINATED).contains(second.getState)) Thread.onSpinWait() + release.countDown() + first.join() + second.join() + } + assertEquals(starts.get(), 1) + } } diff --git a/sage-client/shared/src/test/scala/sage/client/internal/FakeTransport.scala b/sage-client/shared/src/test/scala/sage/client/internal/FakeTransport.scala index 59d488e8..7095899b 100644 --- a/sage-client/shared/src/test/scala/sage/client/internal/FakeTransport.scala +++ b/sage-client/shared/src/test/scala/sage/client/internal/FakeTransport.scala @@ -22,6 +22,9 @@ final class FakeTransport( def written: Vector[Bytes] = synchronized(writes.toVector) + // the writes after the connection setup + def sent: Vector[Bytes] = written.filterNot(Replies.isSetup) + var closeCount: Int = 0 def start(): Unit = () diff --git a/sage-client/shared/src/test/scala/sage/client/internal/FaultSpec.scala b/sage-client/shared/src/test/scala/sage/client/internal/FaultSpec.scala index ffaf809f..e4cac98c 100644 --- a/sage-client/shared/src/test/scala/sage/client/internal/FaultSpec.scala +++ b/sage-client/shared/src/test/scala/sage/client/internal/FaultSpec.scala @@ -1,21 +1,21 @@ package sage.client.internal import sage.SageException.{ConnectionLost, NotConnected, ServerError} -import sage.cluster.{Node, Redirect, RedirectKind, Slot} +import sage.cluster.{Node, Redirect, RedirectKind} class FaultSpec extends munit.FunSuite { test("a MOVED reply categorizes as Redirected carrying the parsed redirect") { assertEquals( Fault.categorize(ServerError("MOVED", "3999 127.0.0.1:6379")), - Fault.Redirected(Redirect(RedirectKind.Moved, Slot.at(3999).get, Node("127.0.0.1", 6379))) + Fault.Redirected(Redirect(RedirectKind.Moved, Node("127.0.0.1", 6379))) ) } test("an ASK reply categorizes as Redirected") { assertEquals( Fault.categorize(ServerError("ASK", "42 127.0.0.1:7000")), - Fault.Redirected(Redirect(RedirectKind.Ask, Slot.at(42).get, Node("127.0.0.1", 7000))) + Fault.Redirected(Redirect(RedirectKind.Ask, Node("127.0.0.1", 7000))) ) } @@ -64,7 +64,7 @@ class FaultSpec extends munit.FunSuite { } test("an ownership or connection change forces a refresh past the throttle window") { - val moved = Fault.Redirected(Redirect(RedirectKind.Moved, Slot.at(1).get, Node("a", 6379))) + val moved = Fault.Redirected(Redirect(RedirectKind.Moved, Node("a", 6379))) assertEquals(moved.refreshPolicy, RefreshPolicy.Forced) assertEquals(Fault.Demoted.refreshPolicy, RefreshPolicy.Forced) assertEquals(Fault.Lost(mayHaveExecuted = false).refreshPolicy, RefreshPolicy.Forced) @@ -72,7 +72,7 @@ class FaultSpec extends munit.FunSuite { } test("an ASK refreshes throttled: it leaves slot ownership unchanged") { - val ask = Fault.Redirected(Redirect(RedirectKind.Ask, Slot.at(1).get, Node("a", 6379))) + val ask = Fault.Redirected(Redirect(RedirectKind.Ask, Node("a", 6379))) assertEquals(ask.refreshPolicy, RefreshPolicy.Throttled) } diff --git a/sage-client/shared/src/test/scala/sage/client/internal/LockCancellationSpec.scala b/sage-client/shared/src/test/scala/sage/client/internal/LockCancellationSpec.scala index 1ac9645b..8e382e01 100644 --- a/sage-client/shared/src/test/scala/sage/client/internal/LockCancellationSpec.scala +++ b/sage-client/shared/src/test/scala/sage/client/internal/LockCancellationSpec.scala @@ -18,10 +18,10 @@ abstract class LockCancellationSpec extends munit.FunSuite { override val munitTimeout = 10.seconds private given ExecutionContext = munitExecutionContext - protected def tryWithLock[A](commands: CommandRunner[CIO, String], lease: FiniteDuration)(body: CIO[A]): CIO[Option[A]] = + protected def tryWithLock[A](commands: SharedRunner, lease: FiniteDuration)(body: CIO[A]): CIO[Option[A]] = new LockExecutor[String](lease, "cancel", replicaAcknowledgement = true).tryWithLock(commands, "key")(body) - protected def withLock[A](commands: CommandRunner[CIO, String], lease: FiniteDuration, wait: FiniteDuration)(body: CIO[A]): CIO[A] = + protected def withLock[A](commands: SharedRunner, lease: FiniteDuration, wait: FiniteDuration)(body: CIO[A]): CIO[A] = new LockExecutor[String](lease, "cancel", replicaAcknowledgement = true).withLock(commands, "key", wait)(body) protected def runner( @@ -29,8 +29,8 @@ abstract class LockCancellationSpec extends munit.FunSuite { renewals: AtomicInteger, stallRelease: Boolean, releaseCompleted: () => Unit = () => () - ): CommandRunner[CIO, String] = - new CommandRunner[CIO, String] { + ): SharedRunner = + new SharedRunner { def run[A](command: Command[A]): CIO[A] = CIO.defer(()).flatMap { _ => command.args(4).asUtf8String match { case "renew" => renewals.incrementAndGet() @@ -80,7 +80,7 @@ abstract class LockCancellationSpec extends munit.FunSuite { val bodyStarted = new AtomicBoolean(false) val bodyStopped = new AtomicBoolean(false) val cleanupCompleted = Promise[Unit]() - val commands = new CommandRunner[CIO, String] { + val commands = new SharedRunner { def run[A](command: Command[A]): CIO[A] = command .decode(Frame.Integer(if (command.args(4).asUtf8String == "renew" && bodyStarted.get()) 0 else 1)) @@ -108,30 +108,33 @@ abstract class LockCancellationSpec extends munit.FunSuite { } test("a contended acquisition can be cancelled before its wait timeout") { - val busy = new CommandRunner[CIO, String] { + val busy = new SharedRunner { def run[A](command: Command[A]): CIO[A] = command.decode(Frame.Integer(0)).fold(CIO.fail(_), CIO.value(_)) } CIO.timeout(100.millis)(withLock(busy, 3.seconds, 30.seconds)(CIO.unit)).unsafeRun.map(result => assertEquals(result, None)) } + + test("a key-type view keeps the runner of the client it wraps") { + val client = new LockTestClient(SharedRunner.unavailable) + assert(client.as[Array[Byte]].runner eq client) + } } -final class LockTestClient(runner: CommandRunner[CIO, String]) extends Client[CIO, String] { +final class LockTestClient(commands: SharedRunner) extends Client[CIO, String] with SharedRunner { private def unsupported[A]: CIO[A] = CIO.fail(new AssertionError("unexpected client operation in a lock test")) - def run[A](command: Command[A]): CIO[A] = runner.run(command) + def run[A](command: Command[A]): CIO[A] = commands.run(command) def cached[A](command: Command[A], ttl: FiniteDuration): CIO[A] = unsupported - private[sage] def pipeline[Out, R](p: Pipeline[Out, R]): CIO[Out] = unsupported - private[sage] def pipelineAttempt[Out, R](p: Pipeline[Out, R]): CIO[R] = unsupported + private[sage] def pipeline[R](p: Pipeline[R]): CIO[R] = unsupported def transaction[A](body: TransactionScope[CIO, String] => CIO[A]): CIO[A] = unsupported def subscribeChannels[V: ValueCodec](channel: String, rest: String*): CIO[Subscription[CIO, Message[V]]] = unsupported def subscribePatterns[V: ValueCodec](pattern: String, rest: String*): CIO[Subscription[CIO, PatternMessage[V]]] = unsupported def subscribeShardChannels[V: ValueCodec](channel: String, rest: String*): CIO[Subscription[CIO, Message[V]]] = unsupported - private[sage] def scanTargets: CIO[Vector[ScanTarget]] = unsupported - private[sage] def runOn[A](target: ScanTarget, command: Command[A]): CIO[A] = unsupported private[sage] def rateLimitAcquire[RK](executor: RateLimitExecutor[RK], subject: RK, cost: Long, peek: Boolean): CIO[Decision] = unsupported private[sage] def lockTryWith[LK, A](executor: LockExecutor[LK], key: LK)(body: => CIO[A]): CIO[Option[A]] = executor.tryWithLock(this, key)(body) private[sage] def lockWith[LK, A](executor: LockExecutor[LK], key: LK, waitTimeout: FiniteDuration)(body: => CIO[A]): CIO[A] = executor.withLock(this, key, waitTimeout)(body) + override private[sage] def runner: SharedRunner = this def close: CIO[Unit] = CIO.unit } diff --git a/sage-client/shared/src/test/scala/sage/client/internal/LockExecutorSpec.scala b/sage-client/shared/src/test/scala/sage/client/internal/LockExecutorSpec.scala index 378d0f34..14e662eb 100644 --- a/sage-client/shared/src/test/scala/sage/client/internal/LockExecutorSpec.scala +++ b/sage-client/shared/src/test/scala/sage/client/internal/LockExecutorSpec.scala @@ -49,7 +49,7 @@ class LockExecutorSpec extends munit.FunSuite { .unsafeRun } - private class Store extends CommandRunner[CIO, String] { + private class Store extends SharedRunner { private var entries = Map.empty[String, (String, Long)] private var loaded = false val operations = new ConcurrentLinkedQueue[String]() @@ -148,17 +148,17 @@ class LockExecutorSpec extends munit.FunSuite { new LockExecutor[String](lease, namespace, replicaAcknowledgement = true) test("namespace framing distinguishes ambiguous prefixes and preserves binary key bytes") { - val a = new LockCommands[String](1.second, "a").key("b:c") - val b = new LockCommands[String](1.second, "a:b").key("c") + val a = new LockExecutor[String](1.second, "a", replicaAcknowledgement = true).key("b:c") + val b = new LockExecutor[String](1.second, "a:b", replicaAcknowledgement = true).key("c") assert(!a.sameBytes(b)) val raw = Array[Byte](0, -1, 42) - val binary = new LockCommands[Array[Byte]](1.second, "é").key(raw) + val binary = new LockExecutor[Array[Byte]](1.second, "é", replicaAcknowledgement = true).key(raw) assert(binary.sameBytes(Bytes.concat(Vector(Bytes.utf8("2:é:"), Bytes.fromArray(raw))))) } test("all scripts route to one key on its master and decode only valid outcomes") { - val commands = new LockCommands[String](1.second, "lock") - for (operation <- LockCommands.Operation.values) { + val commands = executor(1.second) + for (operation <- LockExecutor.Operation.values) { val command = commands.command(commands.key("subject"), "owner", operation, cached = true) assertEquals(command.name, "EVALSHA") assertEquals(command.keys.map(_.asUtf8String), Vector("4:lock:subject")) @@ -191,7 +191,7 @@ class LockExecutorSpec extends munit.FunSuite { test("a successful renewal keeps the scope running and releases it afterwards") { val store = new Store - executor(300.millis) + executor(900.millis) .tryWithLock(store, "key") { store.awaitRenewal.map(_ => assert(store.held)) } @@ -222,7 +222,7 @@ class LockExecutorSpec extends munit.FunSuite { test(s"renewal retries a transient $failureName while the lease remains valid") { val store = new Store store.renewalErrors.add(failure) - executor(300.millis) + executor(900.millis) .tryWithLock(store, "key") { store.awaitRenewal.map(_ => assert(store.held, "the lock expired while renewal was retrying")) } @@ -364,7 +364,7 @@ class LockExecutorSpec extends munit.FunSuite { test("sustained contention reduces acquisition requests while respecting the wait timeout") { val attempts = new AtomicInteger(0) - val commands = new CommandRunner[CIO, String] { + val commands = new SharedRunner { def run[A](command: Command[A]): CIO[A] = CIO.defer(()).flatMap { _ => attempts.incrementAndGet() command.decode(Frame.Integer(0)).fold(CIO.fail(_), CIO.value(_)) diff --git a/sage-client/shared/src/test/scala/sage/client/internal/MultiplexedConnectionSpec.scala b/sage-client/shared/src/test/scala/sage/client/internal/MultiplexedConnectionSpec.scala index 4a52220c..e583fa25 100644 --- a/sage-client/shared/src/test/scala/sage/client/internal/MultiplexedConnectionSpec.scala +++ b/sage-client/shared/src/test/scala/sage/client/internal/MultiplexedConnectionSpec.scala @@ -8,7 +8,7 @@ import Replies.bulk import sage.Bytes import sage.SageException.{ConnectionLost, DecodeError, NotConnected, ServerError} -import sage.client.{BackoffConfig, WatchdogConfig} +import sage.client.{BackoffConfig, CacheConfig, SageConfig, WatchdogConfig} import sage.commands.{Command, Connection, Strings} import sage.protocol.Frame @@ -17,25 +17,71 @@ class MultiplexedConnectionSpec extends munit.FunSuite { private val fixedBackoff = BackoffConfig(initialDelay = 1.milli, maxDelay = 1.milli, multiplier = 1.0) private val noWatchdog = WatchdogConfig(enabled = false) + private val noCache = CacheConfig(enabled = false) + + private val cachingConfig = SageConfig( + reconnect = fixedBackoff, + watchdog = noWatchdog, + connectTimeout = 1.second, + closeTimeout = Duration.Zero, + clientCache = CacheConfig(enabled = true, maxBytes = 1L << 20) + ) + private def make( autoWrite: Boolean = true, respond: Bytes => Seq[Frame] = _ => Nil, watchdog: WatchdogConfig = noWatchdog, - closeTimeout: FiniteDuration = Duration.Zero, - bootstrap: Vector[Command[?]] = Vector.empty + closeTimeout: FiniteDuration = Duration.Zero ): (MultiplexedConnection, ManualScheduler, mutable.ArrayBuffer[FakeTransport]) = { val scheduler = new ManualScheduler val transports = mutable.ArrayBuffer.empty[FakeTransport] val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => { - val transport = new FakeTransport(onFrame, onClosed, respond, autoWrite) + val transport = new FakeTransport(onFrame, onClosed, Replies.withSetup(respond)) transports += transport transport } val connection = - MultiplexedConnection.connect(factory, scheduler, bootstrap, fixedBackoff, watchdog, 1.second, closeTimeout) + new MultiplexedConnection( + factory, + scheduler, + SageConfig(clientCache = noCache, reconnect = fixedBackoff, watchdog = watchdog, connectTimeout = 1.second, closeTimeout = closeTimeout), + MultiplexedConnection.NodeRole.Master + ).start() + // the setup is written before the test takes control of writes + transports.foreach(_.autoWrite = autoWrite) (connection, scheduler, transports) } + test("an interrupted handshake closes the connection it opened") { + val transports = mutable.ArrayBuffer.empty[FakeTransport] + // the reply to HELLO never comes, and the waiting thread is interrupted instead + val respond: Bytes => Seq[Frame] = payload => { + if (payload.asUtf8String.contains("HELLO")) Thread.currentThread().interrupt() + Nil + } + val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => { + val transport = new FakeTransport(onFrame, onClosed, respond) + transports += transport + transport + } + val connection = new MultiplexedConnection( + factory, + new ManualScheduler, + SageConfig(clientCache = noCache, reconnect = fixedBackoff, watchdog = noWatchdog, connectTimeout = 1.second, closeTimeout = Duration.Zero), + MultiplexedConnection.NodeRole.Master + ) + @volatile var failure: Throwable = null + Thread + .ofVirtual() + .start { () => + try connection.start(): Unit + catch { case e: Throwable => failure = e } + } + .join() + assert(failure.isInstanceOf[InterruptedException], String.valueOf(failure)) + assertEquals(transports.map(_.closeCount).toVector, Vector(1)) + } + test("matches replies to commands in FIFO order") { val (connection, _, transports) = make() var first: Option[Try[String]] = None @@ -167,7 +213,7 @@ class MultiplexedConnectionSpec extends munit.FunSuite { val callbacks = Vector.tabulate(3)(i => (r: Try[Any]) => results(i) = r) val submitted = connection.submitAll(commands, callbacks) assert(submitted) - assertEquals(transports.head.written.length, 1) // one round-trip: the whole pipeline is one write, not three + assertEquals(transports.head.sent.length, 1) // one round-trip: the whole pipeline is one write, not three transports.head.emit(Frame.SimpleString("a")) transports.head.emit(Frame.SimpleString("b")) transports.head.emit(Frame.SimpleString("c")) @@ -180,8 +226,8 @@ class MultiplexedConnectionSpec extends munit.FunSuite { connection.submitAll(commands, Vector.fill(2)((_: Try[Any]) => ())) transports.head.writeNext() val expected = Bytes.concat(commands.map(_.encode)) - assertEquals(transports.head.written.length, 1) - assert(transports.head.written.head.sameBytes(expected)) + assertEquals(transports.head.sent.length, 1) + assert(transports.head.sent.head.sameBytes(expected)) } test("a batch returns false when not connected, submitting nothing") { @@ -268,12 +314,17 @@ class MultiplexedConnectionSpec extends munit.FunSuite { val scheduler = new ManualScheduler val transports = mutable.ArrayBuffer.empty[FakeTransport] val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => { - val transport = new FakeTransport(onFrame, onClosed, _ => Nil, autoWrite = true) + val transport = new FakeTransport(onFrame, onClosed, Replies.withSetup(_ => Nil)) transports += transport transport } val connection = - MultiplexedConnection.connect(factory, scheduler, Vector.empty[Command[?]], backoff, noWatchdog, 1.second, Duration.Zero) + new MultiplexedConnection( + factory, + scheduler, + SageConfig(clientCache = noCache, reconnect = backoff, watchdog = noWatchdog, connectTimeout = 1.second, closeTimeout = Duration.Zero), + MultiplexedConnection.NodeRole.Master + ).start() import MultiplexedConnection.State scheduler.advance(11.seconds) @@ -289,26 +340,31 @@ class MultiplexedConnectionSpec extends munit.FunSuite { val scheduler = new ManualScheduler val transports = mutable.ArrayBuffer.empty[FakeTransport] val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => { - val transport = new FakeTransport(onFrame, onClosed, _ => Nil, autoWrite = true) + val transport = new FakeTransport(onFrame, onClosed, Replies.withSetup(_ => Nil)) transports += transport transport } val connection = - MultiplexedConnection.connect(factory, scheduler, Vector.empty[Command[?]], backoff, noWatchdog, 1.second, Duration.Zero) + new MultiplexedConnection( + factory, + scheduler, + SageConfig(clientCache = noCache, reconnect = backoff, watchdog = noWatchdog, connectTimeout = 1.second, closeTimeout = Duration.Zero), + MultiplexedConnection.NodeRole.Master + ).start() import MultiplexedConnection.State // when a Live connection drops before the stable interval, keep the attempt count and increase the next backoff delay assertEquals(connection.currentState, State.Live) - transports.last.emit(Frame.SimpleString("stray")) // flap 1 -> attempt 1 -> 20ms + transports.last.emit(Frame.SimpleString("stray")) // flap 1 -> attempt 0 -> 10ms assertEquals(connection.currentState, State.Reconnecting) - scheduler.advance(19.millis) - assertEquals(connection.currentState, State.Reconnecting, "still backing off before the 20ms attempt-1 delay") + scheduler.advance(9.millis) + assertEquals(connection.currentState, State.Reconnecting, "still backing off before the 10ms attempt-0 delay") scheduler.advance(1.milli) assertEquals(connection.currentState, State.Live) - transports.last.emit(Frame.SimpleString("stray")) // the second short-lived connection uses attempt 2 and a 40 ms delay - scheduler.advance(39.millis) - assertEquals(connection.currentState, State.Reconnecting, "the backoff doubled to 40ms rather than resetting to 10ms") + transports.last.emit(Frame.SimpleString("stray")) // the second short-lived connection uses attempt 1 and a 20 ms delay + scheduler.advance(19.millis) + assertEquals(connection.currentState, State.Reconnecting, "the backoff doubled to 20ms rather than resetting to 10ms") scheduler.advance(1.milli) assertEquals(connection.currentState, State.Live) @@ -325,10 +381,11 @@ class MultiplexedConnectionSpec extends munit.FunSuite { test("a socket dying in the establish->Live window is not published Live, and the reconnect loop recovers") { final class DyingAfterBootstrap(onFrame: Frame => Unit, onClosed: () => Unit) extends Transport { def start(): Unit = () + // answers the setup, then drops the connection right after the last setup command def send(item: Transport.Item): Unit = { item.writeAttempted() - onFrame(Frame.SimpleString("PONG")) - onClosed() + Replies.withSetup(_ => Nil)(item.payload).foreach(onFrame) + if (item.payload.asUtf8String.contains("LIB-VER")) onClosed() } def close(): Unit = () } @@ -340,12 +397,17 @@ class MultiplexedConnectionSpec extends munit.FunSuite { first = false new DyingAfterBootstrap(onFrame, onClosed) } else { - val t = new FakeTransport(onFrame, onClosed, _ => Seq(Frame.SimpleString("PONG"))) + val t = new FakeTransport(onFrame, onClosed, Replies.withSetup(_ => Seq(Frame.SimpleString("PONG")))) healthy += t t } val connection = - MultiplexedConnection.connect(factory, scheduler, Vector(Connection.ping()), fixedBackoff, noWatchdog, 1.second, Duration.Zero) + new MultiplexedConnection( + factory, + scheduler, + SageConfig(clientCache = noCache, reconnect = fixedBackoff, watchdog = noWatchdog, connectTimeout = 1.second, closeTimeout = Duration.Zero), + MultiplexedConnection.NodeRole.Master + ).start() assertEquals(connection.currentState, MultiplexedConnection.State.Reconnecting) @@ -359,20 +421,21 @@ class MultiplexedConnectionSpec extends munit.FunSuite { assertEquals(afterRecovery, Some(Success("PONG"))) } - test("isCurrent gates on liveness: a stamp is current only while Live, never during a reconnect window at the same generation") { + test("losing liveness retires the pool's dedicated connections, so one opened before a reconnect is not reused") { val (connection, scheduler, transports) = make() - val g = connection.liveGeneration().getOrElse(fail("expected a live generation")) - assert(connection.isCurrent(g)) + val dedicated = connection.pool.acquireForTransaction() + connection.pool.releaseTransaction(dedicated, reusable = true) + assertEquals(transports.size, 2) - transports.head.emit(Frame.SimpleString("PONG")) // stray frame -> discarded -> Reconnecting, generation not yet bumped - assertEquals(connection.currentState, MultiplexedConnection.State.Reconnecting) - assertEquals(connection.liveGeneration(), None) - assert(!connection.isCurrent(g), "the stamp must not read as current during a reconnect window, even at the same generation") + transports.head.emit(Frame.SimpleString("PONG")) // stray frame -> discarded -> Reconnecting + assert(!connection.isLive) + scheduler.advance(1.milli) // reconnects -> Live + assert(connection.isLive) - scheduler.advance(1.milli) // reconnects -> Live, generation bumps - assertEquals(connection.currentState, MultiplexedConnection.State.Live) - assert(!connection.isCurrent(g), "the old stamp stays stale once the generation bumps") - assert(connection.isCurrent(connection.liveGeneration().getOrElse(fail("expected a live generation")))) + val fresh = connection.pool.acquireForTransaction() + assert(fresh ne dedicated) + assertEquals(transports.size, 4) + connection.pool.releaseTransaction(fresh, reusable = true) } test("every reconnect attempt re-resolves the endpoint, honoring a repoint between attempts") { @@ -383,14 +446,19 @@ class MultiplexedConnectionSpec extends munit.FunSuite { val scheduler = new ManualScheduler val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => { seen += endpoint - val reply: Bytes => Seq[Frame] = _ => if (healthy) Seq(Frame.SimpleString("PONG")) else Seq(Frame.SimpleError("LOADING")) + val reply: Bytes => Seq[Frame] = payload => if (healthy) Replies.withSetup(_ => Nil)(payload) else Seq(Frame.SimpleError("LOADING")) val transport = new FakeTransport(onFrame, onClosed, reply) transports += transport transport } - // use PING as the bootstrap command. It fails while the current endpoint is unhealthy, which makes every retry resolve the endpoint again. + // HELLO fails while the current endpoint is unhealthy, which makes every retry resolve the endpoint again. val connection = - MultiplexedConnection.connect(factory, scheduler, Vector(Connection.ping()), fixedBackoff, noWatchdog, 1.second, Duration.Zero) + new MultiplexedConnection( + factory, + scheduler, + SageConfig(clientCache = noCache, reconnect = fixedBackoff, watchdog = noWatchdog, connectTimeout = 1.second, closeTimeout = Duration.Zero), + MultiplexedConnection.NodeRole.Master + ).start() assertEquals(seen.toList, List("master-1")) assertEquals(connection.currentState, MultiplexedConnection.State.Live) @@ -428,15 +496,20 @@ class MultiplexedConnectionSpec extends munit.FunSuite { val transports = mutable.ArrayBuffer.empty[FakeTransport] val scheduler = new ManualScheduler val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => { - // generation 0 answers the bootstrap PING; later generations never reply, so establish() blocks until close() aborts it + // generation 0 answers the setup; later generations never reply, so establish() blocks until close() aborts it val first = transports.isEmpty - val reply: Bytes => Seq[Frame] = _ => if (first) Seq(Frame.SimpleString("PONG")) else Nil + val reply: Bytes => Seq[Frame] = if (first) Replies.withSetup(_ => Nil) else _ => Nil val transport = new FakeTransport(onFrame, onClosed, reply) transports += transport transport } val connection = - MultiplexedConnection.connect(factory, scheduler, Vector(Connection.ping()), fixedBackoff, noWatchdog, 5.seconds, Duration.Zero) + new MultiplexedConnection( + factory, + scheduler, + SageConfig(clientCache = noCache, reconnect = fixedBackoff, watchdog = noWatchdog, connectTimeout = 5.seconds, closeTimeout = Duration.Zero), + MultiplexedConnection.NodeRole.Master + ).start() transports.head.emit(Frame.SimpleString("stray")) val reconnect = new Thread(() => scheduler.advance(1.milli)) // advance blocks inside establish() awaiting the bootstrap reply @@ -459,7 +532,7 @@ class MultiplexedConnectionSpec extends munit.FunSuite { val scheduler = new ManualScheduler val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => if (head == null) { - head = new FakeTransport(onFrame, onClosed, _ => Seq(Frame.SimpleString("PONG"))) + head = new FakeTransport(onFrame, onClosed, Replies.withSetup(_ => Seq(Frame.SimpleString("PONG")))) head } else { val transport = new ConnectingTransport(onClosed) @@ -467,7 +540,12 @@ class MultiplexedConnectionSpec extends munit.FunSuite { transport } val connection = - MultiplexedConnection.connect(factory, scheduler, Vector(Connection.ping()), fixedBackoff, noWatchdog, 5.seconds, Duration.Zero) + new MultiplexedConnection( + factory, + scheduler, + SageConfig(clientCache = noCache, reconnect = fixedBackoff, watchdog = noWatchdog, connectTimeout = 5.seconds, closeTimeout = Duration.Zero), + MultiplexedConnection.NodeRole.Master + ).start() head.emit(Frame.SimpleString("stray")) val reconnect = new Thread(() => scheduler.advance(1.milli)) // blocks inside ConnectingTransport.start() @@ -515,18 +593,28 @@ class MultiplexedConnectionSpec extends munit.FunSuite { val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => if (first) { first = false + // answers the setup, then drops the connection right after the last setup command new Transport { - def start(): Unit = onClosed() - def send(item: Transport.Item): Unit = () + def start(): Unit = () + def send(item: Transport.Item): Unit = { + item.writeAttempted() + Replies.withSetup(_ => Nil)(item.payload).foreach(onFrame) + if (item.payload.asUtf8String.contains("LIB-VER")) onClosed() + } def close(): Unit = () } } else { - val t = new FakeTransport(onFrame, onClosed, _ => Nil, autoWrite = true) + val t = new FakeTransport(onFrame, onClosed, Replies.withSetup(_ => Nil)) live += t t } val connection = - MultiplexedConnection.connect(factory, scheduler, Vector.empty[Command[?]], fixedBackoff, watchdog, 1.second, Duration.Zero) + new MultiplexedConnection( + factory, + scheduler, + SageConfig(clientCache = noCache, reconnect = fixedBackoff, watchdog = watchdog, connectTimeout = 1.second, closeTimeout = Duration.Zero), + MultiplexedConnection.NodeRole.Master + ).start() scheduler.advance(1.milli) assertEquals(connection.currentState, MultiplexedConnection.State.Live) @@ -536,6 +624,22 @@ class MultiplexedConnectionSpec extends munit.FunSuite { assertEquals(live.head.written.count(_.asUtf8String.contains("PING")), 1) } + test("a close interrupted while draining still closes the socket, returns normally and keeps the interrupt") { + val (connection, _, transports) = make(autoWrite = false, closeTimeout = 2.seconds) + connection.submit(Connection.ping(), _ => ()) + @volatile var interruptedAfter = false + Thread + .ofVirtual() + .start { () => + Thread.currentThread().interrupt() + connection.close() + interruptedAfter = Thread.currentThread().isInterrupted + } + .join() + assertEquals(transports.head.closeCount, 1) + assert(interruptedAfter) + } + test("graceful drain lets an in-flight reply complete before close finishes") { val (connection, _, transports) = make(autoWrite = false, closeTimeout = 2.seconds) var result: Option[Try[String]] = None @@ -590,13 +694,20 @@ class MultiplexedConnectionSpec extends munit.FunSuite { val scheduler = new ManualScheduler val transports = mutable.ArrayBuffer.empty[FakeTransport] val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => { - val transport = new FakeTransport(onFrame, onClosed) + // answer the CLIENT TRACKING setup write, so every cached connection's first write is that command + val transport = + new FakeTransport(onFrame, onClosed, Replies.withSetup(p => if (p.asUtf8String.contains("TRACKING")) Seq(Frame.SimpleString("OK")) else Nil)) transports += transport transport } val connection = - MultiplexedConnection - .connect(factory, scheduler, Vector.empty, fixedBackoff, noWatchdog, 1.second, Duration.Zero, cacheMaxBytes = 1L << 20, events = events) + new MultiplexedConnection( + factory, + scheduler, + cachingConfig, + MultiplexedConnection.NodeRole.Master, + events = events + ).start() (connection, scheduler, transports) } @@ -608,21 +719,21 @@ class MultiplexedConnectionSpec extends munit.FunSuite { val get = Strings.get[String, String]("foo") var first: Option[Try[Option[String]]] = None - connection.cachedSubmit(get, 60000L, r => first = Some(r)) - assertEquals(transports.head.written.length, 1) // one batch: [CLIENT CACHING YES, GET foo] - transports.head.emit(Frame.SimpleString("OK")) // CLIENT CACHING YES reply, discarded + connection.cachedSubmit(get, 60000L, r => first = Some(r), Events.untraced) + assertEquals(transports.head.sent.length, 2) // CLIENT TRACKING, then one batch: [CLIENT CACHING YES, GET foo] + transports.head.emit(Frame.SimpleString("OK")) // CLIENT CACHING YES reply, discarded transports.head.emit(bulk("bar")) assertEquals(first, Some(Success(Some("bar")))) var second: Option[Try[Option[String]]] = None - connection.cachedSubmit(get, 60000L, r => second = Some(r)) + connection.cachedSubmit(get, 60000L, r => second = Some(r), Events.untraced) assertEquals(second, Some(Success(Some("bar")))) // served locally - assertEquals(transports.head.written.length, 1) // no new round-trip + assertEquals(transports.head.sent.length, 2) // no new round-trip transports.head.emit(invalidationOf("foo")) var third: Option[Try[Option[String]]] = None - connection.cachedSubmit(get, 60000L, r => third = Some(r)) - assertEquals(transports.head.written.length, 2) // evicted -> refetch + connection.cachedSubmit(get, 60000L, r => third = Some(r), Events.untraced) + assertEquals(transports.head.sent.length, 3) // evicted -> refetch transports.head.emit(Frame.SimpleString("OK")) transports.head.emit(bulk("baz")) assertEquals(third, Some(Success(Some("baz")))) @@ -632,16 +743,16 @@ class MultiplexedConnectionSpec extends munit.FunSuite { val (connection, _, transports) = cachedConnection() val get = Strings.get[String, String]("foo") - connection.cachedSubmit(get, 60000L, _ => ()) + connection.cachedSubmit(get, 60000L, _ => (), Events.untraced) transports.head.emit(Frame.SimpleString("OK")) transports.head.emit(bulk("bar")) - assertEquals(transports.head.written.length, 1) + assertEquals(transports.head.sent.length, 2) transports.head.emit(Frame.Push(Vector(bulk("invalidate"), Frame.Null))) // FLUSHALL/tracking-drop form var afterFlush: Option[Try[Option[String]]] = None - connection.cachedSubmit(get, 60000L, r => afterFlush = Some(r)) - assertEquals(transports.head.written.length, 2) // flushed -> refetch + connection.cachedSubmit(get, 60000L, r => afterFlush = Some(r), Events.untraced) + assertEquals(transports.head.sent.length, 3) // flushed -> refetch transports.head.emit(Frame.SimpleString("OK")) transports.head.emit(bulk("baz")) assertEquals(afterFlush, Some(Success(Some("baz")))) @@ -651,21 +762,21 @@ class MultiplexedConnectionSpec extends munit.FunSuite { val (connection, scheduler, transports) = cachedConnection() val get = Strings.get[String, String]("foo") - connection.cachedSubmit(get, 1000L, _ => ()) + connection.cachedSubmit(get, 1000L, _ => (), Events.untraced) transports.head.emit(Frame.SimpleString("OK")) transports.head.emit(bulk("bar")) - assertEquals(transports.head.written.length, 1) + assertEquals(transports.head.sent.length, 2) - scheduler.advance(1001.millis) // past the TTL - connection.cachedSubmit(get, 1000L, _ => ()) - assertEquals(transports.head.written.length, 2) // expired -> refetch + scheduler.advance(1001.millis) // past the TTL + connection.cachedSubmit(get, 1000L, _ => (), Events.untraced) + assertEquals(transports.head.sent.length, 3) // expired -> refetch } test("a reconnect flushes the cache: tracking state is connection-bound") { val (connection, scheduler, transports) = cachedConnection() val get = Strings.get[String, String]("foo") - connection.cachedSubmit(get, 60000L, _ => ()) + connection.cachedSubmit(get, 60000L, _ => (), Events.untraced) transports.head.emit(Frame.SimpleString("OK")) transports.head.emit(bulk("bar")) @@ -675,8 +786,8 @@ class MultiplexedConnectionSpec extends munit.FunSuite { assertEquals(transports.size, 2) var afterReconnect: Option[Try[Option[String]]] = None - connection.cachedSubmit(get, 60000L, r => afterReconnect = Some(r)) - assertEquals(transports(1).written.length, 1) // fresh generation, empty cache -> refetch + connection.cachedSubmit(get, 60000L, r => afterReconnect = Some(r), Events.untraced) + assertEquals(transports(1).sent.length, 2) // fresh generation, empty cache -> refetch transports(1).emit(Frame.SimpleString("OK")) transports(1).emit(bulk("baz")) assertEquals(afterReconnect, Some(Success(Some("baz")))) @@ -688,21 +799,22 @@ class MultiplexedConnectionSpec extends munit.FunSuite { val (connection, _, transports) = cachedConnection(events) val get = Strings.get[String, String]("foo") - connection.cachedSubmit(get, 60000L, _ => ()) // miss -> fetch + tracedRead(connection, events, get)(_ => ()) // miss -> fetch transports.head.emit(Frame.SimpleString("OK")) transports.head.emit(bulk("bar")) assertEquals(tracer.log.toVector, Vector("start:GET", "settled:Succeeded")) - connection.cachedSubmit(get, 60000L, _ => ()) // served locally + tracedRead(connection, events, get)(_ => ()) // served locally assertEquals(tracer.log.toVector, Vector("start:GET", "settled:Succeeded")) // unchanged: no span for a hit } test("a cached read that fails fast (not connected) settles a Failed span, like an ordinary command") { val tracer = new RecordingTracer - val (connection, _, _) = cachedConnection(Events(Vector.empty, Some(tracer))) + val events = Events(Vector.empty, Some(tracer)) + val (connection, _, _) = cachedConnection(events) connection.close() var result: Option[Try[Option[String]]] = None - connection.cachedSubmit(Strings.get[String, String]("foo"), 60000L, r => result = Some(r)) + tracedRead(connection, events, Strings.get[String, String]("foo"))(r => result = Some(r)) assert(result.exists(_.isFailure)) assertEquals(tracer.log.head, "start:GET") assert(tracer.log.last.startsWith("settled:Failed")) @@ -718,26 +830,22 @@ class MultiplexedConnectionSpec extends munit.FunSuite { else Nil } val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => { - val transport = new FakeTransport(onFrame, onClosed, respond) + val transport = new FakeTransport(onFrame, onClosed, Replies.withSetup(respond)) transports += transport transport } - val connection = MultiplexedConnection.connect( + val connection = new MultiplexedConnection( factory, scheduler, - Vector(Connection.clientTrackingOnOptin), - fixedBackoff, - noWatchdog, - 1.second, - Duration.Zero, - cacheMaxBytes = 1L << 20 - ) + cachingConfig, + MultiplexedConnection.NodeRole.Master + ).start() val get = Strings.get[String, String]("foo") var first: Option[Try[Option[String]]] = None - connection.cachedSubmit(get, 60000L, r => first = Some(r)) + connection.cachedSubmit(get, 60000L, r => first = Some(r), Events.untraced) var second: Option[Try[Option[String]]] = None - connection.cachedSubmit(get, 60000L, r => second = Some(r)) + connection.cachedSubmit(get, 60000L, r => second = Some(r), Events.untraced) assertEquals(first, Some(Success(Some("bar")))) assertEquals(second, Some(Success(Some("bar")))) @@ -745,4 +853,8 @@ class MultiplexedConnectionSpec extends munit.FunSuite { assert(!writes.exists(_.contains("CACHING")), "no CLIENT CACHING YES when tracking is unavailable") assertEquals(writes.count(_.contains("GET")), 2, "each cached read re-contacts the server: nothing is cached without tracking") } + + // traces the read once per call, as the client runtimes do + private def tracedRead[A](connection: MultiplexedConnection, events: Events, command: Command[A])(callback: Try[A] => Unit): Unit = + connection.cachedSubmit(command, 60000L, callback, Events.fetchTracking(events)) } diff --git a/sage-client/shared/src/test/scala/sage/client/internal/NodePoolSpec.scala b/sage-client/shared/src/test/scala/sage/client/internal/NodePoolSpec.scala index 6be91a6d..3ee5a739 100644 --- a/sage-client/shared/src/test/scala/sage/client/internal/NodePoolSpec.scala +++ b/sage-client/shared/src/test/scala/sage/client/internal/NodePoolSpec.scala @@ -5,31 +5,31 @@ import java.util.concurrent.{CountDownLatch, TimeUnit} import java.util.concurrent.atomic.{AtomicInteger, AtomicReference, AtomicReferenceArray} import scala.concurrent.duration.* +import scala.jdk.CollectionConverters.* import scala.util.Try import sage.Bytes import sage.SageException.NotConnected -import sage.client.{BackoffConfig, DedicatedPoolConfig, WatchdogConfig} +import sage.client.{CacheConfig, SageConfig, WatchdogConfig} import sage.cluster.Node -import sage.commands.Connection import sage.protocol.Frame class NodePoolSpec extends munit.FunSuite { - private def respond(payload: Bytes): Seq[Frame] = - if (payload.asUtf8String.contains("HELLO")) Seq(Replies.hello) else Nil + private def respond(payload: Bytes): Seq[Frame] = Replies.withSetup(_ => Nil)(payload) private def newPool(factory: Node => MultiplexedConnection.TransportFactory, events: Events = Events.disabled): NodePool = new NodePool( factory, Scheduler.real, - Vector(Connection.hello(None)), - BackoffConfig(), - WatchdogConfig(enabled = false), - 1.second, - Duration.Zero, - DedicatedPoolConfig(), - events = events + SageConfig( + watchdog = WatchdogConfig(enabled = false), + connectTimeout = 1.second, + closeTimeout = Duration.Zero, + clientCache = CacheConfig(enabled = false) + ), + MultiplexedConnection.NodeRole.Master, + events ) // for the nth connection attempt, signal `reached(n)` and wait for `release(n)` before opening the transport stored at `transport(n)` @@ -52,7 +52,7 @@ class NodePoolSpec extends munit.FunSuite { def transport(i: Int): FakeTransport = opened.get(i) } - private def existingLive(pool: NodePool, node: Node): Option[NodeClient] = Option(pool.existing(node)).filter(_.isLive) + private def existingLive(pool: NodePool, node: Node): Option[MultiplexedConnection] = Option(pool.existing(node)).filter(_.isLive) private def awaitTrue(cond: => Boolean, clue: String): Unit = { val deadline = System.nanoTime() + 2.seconds.toNanos @@ -87,13 +87,13 @@ class NodePoolSpec extends munit.FunSuite { val gated = new GatedFactory(1) val pool = newPool(gated.factory) - val establisher = new AtomicReference[Try[NodeClient]]() + val establisher = new AtomicReference[Try[MultiplexedConnection]]() val establishing = new Thread(() => establisher.set(Try(pool.getOrEstablish(node))), "establisher") establishing.start() gated.awaitReached(0) val waiters = (1 to 3).map { i => - val result = new AtomicReference[Try[NodeClient]]() + val result = new AtomicReference[Try[MultiplexedConnection]]() val thread = new Thread(() => result.set(Try(pool.getOrEstablish(node))), s"waiter-$i") thread.start() (thread, result) @@ -115,7 +115,7 @@ class NodePoolSpec extends munit.FunSuite { s"establisher: ${establisher.get()}" ) - awaitTrue(gated.transport(0).closeCount > 0, "the discarded NodeClient was not closed") + awaitTrue(gated.transport(0).closeCount > 0, "the discarded MultiplexedConnection was not closed") assert(existingLive(pool, node).isEmpty, "the rejected node leaked into the pool") assert(!pool.candidatesByLiveness.contains(node), "the rejected node leaked into refresh candidates") pool.close() @@ -126,7 +126,7 @@ class NodePoolSpec extends munit.FunSuite { val gated = new GatedFactory(1) val pool = newPool(gated.factory) - val result = new AtomicReference[Try[NodeClient]]() + val result = new AtomicReference[Try[MultiplexedConnection]]() val establishing = new Thread(() => result.set(Try(pool.getOrEstablish(node))), "establisher") establishing.start() gated.awaitReached(0) @@ -145,14 +145,14 @@ class NodePoolSpec extends munit.FunSuite { val gated = new GatedFactory(2) val pool = newPool(gated.factory) - val first = new AtomicReference[Try[NodeClient]]() + val first = new AtomicReference[Try[MultiplexedConnection]]() val firstThread = new Thread(() => first.set(Try(pool.getOrEstablish(node))), "attempt-1") firstThread.start() gated.awaitReached(0) pool.retain(_ => false) // clears attempt 1 from pending while its connect is still in flight - val second = new AtomicReference[Try[NodeClient]]() + val second = new AtomicReference[Try[MultiplexedConnection]]() val secondThread = new Thread(() => second.set(Try(pool.getOrEstablish(node))), "attempt-2") secondThread.start() gated.awaitReached(1) // a fresh attempt, since attempt 1 is no longer pending @@ -184,7 +184,7 @@ class NodePoolSpec extends munit.FunSuite { val pool = newPool(factory) val node = Node("slow", 6379) - val result = new AtomicReference[Try[NodeClient]]() + val result = new AtomicReference[Try[MultiplexedConnection]]() val establishing = new Thread(() => result.set(Try(pool.getOrEstablish(node))), "establisher") establishing.start() awaitTrue(connecting.get() != null, "the establish never started") @@ -198,6 +198,76 @@ class NodePoolSpec extends munit.FunSuite { assert(result.get() != null && result.get().isFailure, s"the aborted establisher should fail: ${result.get()}") } + test("an interrupted owner fails a joined attempt with NotConnected instead of handing over its interrupt") { + val node = Node("gated", 6379) + val gated = new GatedFactory(1) + val pool = newPool(gated.factory) + val owner = new Thread(() => Try(pool.getOrEstablish(node)): Unit, "owner") + owner.start() + gated.awaitReached(0) + val outcome = new AtomicReference[String]("not run") + val joiner = new Thread( + () => + outcome.set( + try String.valueOf(pool.getOrEstablishOrNull(node)) + catch { case e: Throwable => e.toString } + ), + "joiner" + ) + joiner.start() + awaitTrue(pool.pendingWaiterCount(node) == 1, "the joiner never blocked on the owner's attempt") + owner.interrupt() // the owner is blocked opening its transport + joiner.join() + owner.join() + assertEquals(outcome.get(), "null") + } + + test("a close on an interrupted thread closes every node and keeps the interrupt") { + val transports = new java.util.concurrent.ConcurrentLinkedQueue[FakeTransport]() + val pool = newPool { _ => (onFrame, onClosed) => + val t = new FakeTransport(onFrame, onClosed, respond) + transports.add(t) + t + } + pool.getOrEstablish(Node("a", 6379)): Unit + pool.getOrEstablish(Node("b", 6379)): Unit + @volatile var failure: Throwable = null + @volatile var interruptedAfter = false + Thread + .ofVirtual() + .start { () => + Thread.currentThread().interrupt() + try pool.close() + catch { case e: Throwable => failure = e } + interruptedAfter = Thread.currentThread().isInterrupted + } + .join() + assertEquals(transports.asScala.toList.map(_.closeCount), List(1, 1)) + assertEquals(failure, null) + assert(interruptedAfter, "close should keep the caller's interrupt") + } + + test("a connection that fails to construct ends its attempt, so the next caller connects instead of waiting forever") { + val node = Node("n", 6379) + val broken = new java.util.concurrent.atomic.AtomicBoolean(true) + val factory: Node => MultiplexedConnection.TransportFactory = _ => + if (broken.getAndSet(false)) throw new IllegalStateException("no transport") + else (onFrame, onClosed) => new FakeTransport(onFrame, onClosed, respond) + val pool = newPool(factory) + intercept[IllegalStateException](pool.getOrEstablish(node)) + val second = new AtomicReference[MultiplexedConnection]() + val caller = new Thread(() => second.set(pool.getOrEstablish(node))) + caller.start() + caller.join(2000) + if (caller.isAlive) { + pool.close() + caller.join() + fail("the next caller waited on the attempt that failed to construct its connection") + } + assert(second.get() != null && second.get().isLive) + pool.close() + } + test("a failed establishment reports the node and cause through ConnectFailed") { val node = Node("unreachable", 6379) val recorder = new ConnectFailureRecorder diff --git a/sage-client/shared/src/test/scala/sage/client/internal/PagedSpec.scala b/sage-client/shared/src/test/scala/sage/client/internal/PagedSpec.scala index 814570b0..2af645b3 100644 --- a/sage-client/shared/src/test/scala/sage/client/internal/PagedSpec.scala +++ b/sage-client/shared/src/test/scala/sage/client/internal/PagedSpec.scala @@ -35,7 +35,7 @@ class PagedSpec extends munit.FunSuite { test("acrossTargets walks every target to its own zero cursor, then ends; empty targets end immediately") { val fetch: ScanTarget => ScanCursor => CIO[ScanPage[String]] = _ => _ => CIO.value(ScanPage(Vector("x"), None)) - val step = Paged.acrossTargets(CIO.value(Vector(ScanTarget.any, ScanTarget.any)))(fetch) + val step = Paged.acrossTargets(CIO.value(Vector(SharedRunner.unavailable, SharedRunner.unavailable)))(fetch) val none = Paged.acrossTargets[String](CIO.value(Vector.empty[ScanTarget]))(fetch) for { begin <- step(ScanStep.Begin).unsafeRun @@ -46,13 +46,13 @@ class PagedSpec extends munit.FunSuite { } yield { assertEquals(begin.get._1, Vector.empty[String]) begin.get._2 match { - case ScanStep.Visit(_, remaining) => assertEquals(remaining.length, 2) - case other => fail(s"expected Visit over two targets, got $other") + case ScanStep.Visit(_, _, rest) => assertEquals(rest.length, 1) + case other => fail(s"expected Visit over two targets, got $other") } assertEquals(visit1.get._1, Vector("x")) visit1.get._2 match { - case ScanStep.Visit(_, remaining) => assertEquals(remaining.length, 1) - case other => fail(s"expected Visit over the last target, got $other") + case ScanStep.Visit(_, _, rest) => assertEquals(rest.length, 0) + case other => fail(s"expected Visit over the last target, got $other") } assertEquals(visit2.get._1, Vector("x")) assertEquals(visit2.get._2, ScanStep.End) @@ -77,7 +77,7 @@ class PagedSpec extends munit.FunSuite { assertEquals(page1, Some((full, Some(StreamRangeId.Exclusive(StreamId(3L, 0L)))))) assertEquals(page2, Some((Vector(entry(4)), None))) assertEquals(end, None) - assertEquals(empty, None) + assertEquals(empty, Some((Vector.empty, None))) } } diff --git a/sage-client/shared/src/test/scala/sage/client/internal/PlacementSpec.scala b/sage-client/shared/src/test/scala/sage/client/internal/PlacementSpec.scala index 918e50b1..cb1703cd 100644 --- a/sage-client/shared/src/test/scala/sage/client/internal/PlacementSpec.scala +++ b/sage-client/shared/src/test/scala/sage/client/internal/PlacementSpec.scala @@ -11,74 +11,95 @@ class PlacementSpec extends munit.FunSuite { private val n1 = Node("h1", 1) private val n2 = Node("h2", 2) - // a fake owner that records which names are currently subscribed; attach throws for any name in `failOn` - final private class FakeConn extends Placement.ShardConn { + // a fake owner that records which names are currently subscribed; like the server, it rejects each name in `failOn` and keeps the others + final private class FakeConn extends ClusterSubscriptions.ShardConn { val subscribed = mutable.LinkedHashSet.empty[String] var failOn: Set[String] = Set.empty + var closed = false - def attach(sink: Sink, names: Vector[String], kind: Kind): Unit = { - if (names.exists(failOn)) throw new RuntimeException("attach refused") - subscribed ++= names - } - def detach(sink: Sink, names: Vector[String], kind: Kind): Boolean = { - subscribed --= names - subscribed.isEmpty - } + def attach(sink: Sink, names: Vector[String]): Unit = + if (closed) throw sage.SageException.NotConnected() else subscribed ++= names.filterNot(failOn) + def detach(sink: Sink, names: Vector[String]): Unit = subscribed --= names + def namesOf(sink: Sink): Vector[String] = subscribed.toVector } - final private class FakePool(unavailable: Set[Node] = Set.empty) extends Placement.Conns { - val byNode = mutable.HashMap.empty[Node, FakeConn] - def ensure(node: Node): Option[Placement.ShardConn] = if (unavailable(node)) None else Some(byNode.getOrElseUpdate(node, new FakeConn)) - def get(node: Node): Option[Placement.ShardConn] = byNode.get(node) + final private class FakePool(unavailable: Set[Node] = Set.empty) { + val byNode = mutable.HashMap.empty[Node, FakeConn] + def ensure(node: Node): Option[ClusterSubscriptions.ShardConn] = + if (unavailable(node)) None else Some(byNode.getOrElseUpdate(node, new FakeConn)) } - private def sink(names: String*): Sink = new Sink(names.toVector, Kind.Shard, 16) + // the placement of one sink, as ClusterSubscriptions drives it + final private class Placement(sink: Sink) { + def reconcile(plan: ClusterSubscriptions.Plan, pool: FakePool): Boolean = + ClusterSubscriptions.reconcile(sink, plan, pool.byNode.toVector, pool.ensure) + } - private def groups(g: Vector[String]*): Vector[Vector[String]] = g.toVector + private def sink(names: String*): Sink = new Sink(names.toVector, Kind.Shard, 16) - test("place attaches each group on its owner and reports full coverage") { + test("reconcile attaches each channel on its owner and reports full coverage") { val pool = new FakePool - val placement = new Placement(sink("a", "b"), Vector("a", "b")) + val placement = new Placement(sink("a", "b")) - placement.place(Map(n1 -> groups(Vector("a"), Vector("b"))), pool) + assert(placement.reconcile(Map(n1 -> Vector("a", "b")), pool), "every channel landed, so the placement is full") assertEquals(pool.byNode(n1).subscribed.toSet, Set("a", "b")) - assert(placement.fullyPlaced, "every channel landed, so the placement is full") } - test("place leaves an unowned channel unplaced — coverage is not full") { + test("reconcile leaves a channel unplaced when its owner is unavailable — coverage is not full") { val pool = new FakePool(unavailable = Set(n2)) - val placement = new Placement(sink("a", "b"), Vector("a", "b")) + val placement = new Placement(sink("a", "b")) - placement.place(Map(n1 -> groups(Vector("a")), n2 -> groups(Vector("b"))), pool) + assert( + !placement.reconcile(Map(n1 -> Vector("a"), n2 -> Vector("b")), pool), + "b's owner was unavailable, so coverage is incomplete" + ) assertEquals(pool.byNode(n1).subscribed.toSet, Set("a")) - assert(!placement.fullyPlaced, "b's owner was unavailable, so coverage is incomplete") } - test("place leaves a failed attach unplaced instead of propagating, recording what landed for roll-back") { + test("reconcile counts distinct channels, so one recorded on two owners cannot mask an unplaced channel") { + val placement = new Placement(sink("a", "b")) + + assert( + !placement.reconcile(Map(n1 -> Vector("a"), n2 -> Vector("a")), new FakePool), + "b never landed; a recorded on n1 and n2 must not count as full coverage" + ) + } + + test("reconcile leaves a failed attach unplaced instead of propagating, recording what landed for roll-back") { val pool = new FakePool - val placement = new Placement(sink("a", "b"), Vector("a", "b")) + val placement = new Placement(sink("a", "b")) pool.byNode.getOrElseUpdate(n1, new FakeConn).failOn = Set("b") - placement.place(Map(n1 -> groups(Vector("a"), Vector("b"))), pool) + assert(!placement.reconcile(Map(n1 -> Vector("a", "b")), pool), "b is unplaced, so coverage is incomplete and the caller retries") assertEquals(pool.byNode(n1).subscribed.toSet, Set("a"), "a landed; b's attach failed but did not propagate") - assert(!placement.fullyPlaced, "b is unplaced, so coverage is incomplete and the caller retries") // roll-back: reconcile to the empty plan detaches exactly what was placed placement.reconcile(Map.empty, pool) assert(pool.byNode(n1).subscribed.isEmpty, "the empty plan detaches the landed channel") } + test("reconcile leaves a channel pending when its connection throws on attach, and places the other channels") { + val pool = new FakePool + val placement = new Placement(sink("a", "b")) + pool.byNode.getOrElseUpdate(n2, new FakeConn).closed = true + + assert(!placement.reconcile(Map(n1 -> Vector("a"), n2 -> Vector("b")), pool), "b's connection threw, so coverage is incomplete") + + assertEquals(pool.byNode(n1).subscribed.toSet, Set("a")) + assert(pool.byNode(n2).subscribed.isEmpty) + } + test("reconcile re-homes a channel to its new owner, leaving nothing on the old one") { val pool = new FakePool - val placement = new Placement(sink("a", "b"), Vector("a", "b")) + val placement = new Placement(sink("a", "b")) - placement.reconcile(Map(n1 -> groups(Vector("a"), Vector("b"))), pool) + placement.reconcile(Map(n1 -> Vector("a", "b")), pool) assertEquals(pool.byNode(n1).subscribed.toSet, Set("a", "b")) // b migrates to n2 - placement.reconcile(Map(n1 -> groups(Vector("a")), n2 -> groups(Vector("b"))), pool) + placement.reconcile(Map(n1 -> Vector("a"), n2 -> Vector("b")), pool) assertEquals(pool.byNode(n1).subscribed.toSet, Set("a"), "b is detached from its old owner — no stale subscription") assertEquals(pool.byNode(n2).subscribed.toSet, Set("b")) @@ -86,43 +107,30 @@ class PlacementSpec extends munit.FunSuite { test("reconcile reports incomplete when an attach fails, then converges once the owner accepts") { val pool = new FakePool - val placement = new Placement(sink("a", "b"), Vector("a", "b")) + val placement = new Placement(sink("a", "b")) val owner = pool.byNode.getOrElseUpdate(n1, new FakeConn) owner.failOn = Set("b") - assert(placement.reconcile(Map(n1 -> groups(Vector("a"), Vector("b"))), pool), "b's attach failed, so the pass is incomplete") + assert(!placement.reconcile(Map(n1 -> Vector("a", "b")), pool), "b's attach failed, so the pass is incomplete") assertEquals(owner.subscribed.toSet, Set("a")) - assert(!placement.fullyPlaced) owner.failOn = Set.empty - assert(!placement.reconcile(Map(n1 -> groups(Vector("a"), Vector("b"))), pool), "retry now lands b — complete") + assert(placement.reconcile(Map(n1 -> Vector("a", "b")), pool), "retry now lands b — complete") assertEquals(owner.subscribed.toSet, Set("a", "b")) - assert(placement.fullyPlaced) - } - - test("fullyPlaced counts distinct channels, so one double-recorded across owners cannot mask an unplaced channel") { - val pool = new FakePool - val placement = new Placement(sink("a", "b"), Vector("a", "b")) - - placement.reconcile(Map(n2 -> groups(Vector("a"))), pool) // fresh topology records a on n2 - placement.place(Map(n1 -> groups(Vector("a"))), pool) // a stale concurrent place re-records a on n1, never touching b - - assert(!placement.fullyPlaced, "b never landed; the duplicate a on n1 and n2 must not count as full coverage") } test("a partial attach followed by a topology shift never duplicates a channel across owners") { val pool = new FakePool - val placement = new Placement(sink("a", "b"), Vector("a", "b")) + val placement = new Placement(sink("a", "b")) pool.byNode.getOrElseUpdate(n1, new FakeConn).failOn = Set("b") // first pass: a is assigned to n1, while b is refused - assert(placement.reconcile(Map(n1 -> groups(Vector("a"), Vector("b"))), pool)) + assert(!placement.reconcile(Map(n1 -> Vector("a", "b")), pool)) // before the retry, b migrates to n2; a stays on n1 - placement.reconcile(Map(n1 -> groups(Vector("a")), n2 -> groups(Vector("b"))), pool) + assert(placement.reconcile(Map(n1 -> Vector("a"), n2 -> Vector("b")), pool)) assertEquals(pool.byNode(n1).subscribed.toSet, Set("a"), "a is undisturbed on n1") assertEquals(pool.byNode(n2).subscribed.toSet, Set("b"), "b lands on its new owner") - assert(placement.fullyPlaced) } } diff --git a/sage-client/shared/src/test/scala/sage/client/internal/ReadRoutingSpec.scala b/sage-client/shared/src/test/scala/sage/client/internal/ReadRoutingSpec.scala index df73f115..efc35583 100644 --- a/sage-client/shared/src/test/scala/sage/client/internal/ReadRoutingSpec.scala +++ b/sage-client/shared/src/test/scala/sage/client/internal/ReadRoutingSpec.scala @@ -10,7 +10,7 @@ import scala.util.{Success, Try} import sage.Bytes import sage.SageException.NotConnected -import sage.client.{BackoffConfig, DedicatedPoolConfig, ReadFrom, WatchdogConfig} +import sage.client.{CacheConfig, ReadFrom, SageConfig, WatchdogConfig} import sage.cluster.Node import sage.commands.{Command, Connection, Execution} import sage.protocol.Frame @@ -74,8 +74,7 @@ class ReadRoutingSpec extends munit.FunSuite { val refreshes = new AtomicInteger() private val transports = new ConcurrentHashMap[Node, FakeTransport]() private def respond(node: Node)(payload: Bytes): Seq[Frame] = - if (payload.asUtf8String.contains("HELLO")) Seq(Replies.hello) - else if (refusing(node)) Seq(Frame.SimpleError("LOADING the dataset is loading")) + if (refusing(node)) Seq(Frame.SimpleError("LOADING the dataset is loading")) else { val text = payload.asUtf8String val pings = text.sliding("PING".length).count(_ == "PING") @@ -85,7 +84,7 @@ class ReadRoutingSpec extends munit.FunSuite { (onFrame, onClosed) => if (unreachable(node)) throw new IOException(s"unreachable $node") else { - val transport = new FakeTransport(onFrame, onClosed, respond(node)) + val transport = new FakeTransport(onFrame, onClosed, Replies.withSetup(respond(node))) transports.put(node, transport) transport } @@ -93,12 +92,13 @@ class ReadRoutingSpec extends munit.FunSuite { new NodePool( factory, scheduler, - Vector(Connection.hello(None)), - BackoffConfig(), - WatchdogConfig(enabled = false), - 1.second, - Duration.Zero, - DedicatedPoolConfig() + SageConfig( + watchdog = WatchdogConfig(enabled = false), + connectTimeout = 1.second, + closeTimeout = Duration.Zero, + clientCache = CacheConfig(enabled = false) + ), + MultiplexedConnection.NodeRole.Master ) val masterPool = newPool() val replicaPool = newPool() diff --git a/sage-client/shared/src/test/scala/sage/client/internal/RecordingTracer.scala b/sage-client/shared/src/test/scala/sage/client/internal/RecordingTracer.scala index 15afb40f..52b27eec 100644 --- a/sage-client/shared/src/test/scala/sage/client/internal/RecordingTracer.scala +++ b/sage-client/shared/src/test/scala/sage/client/internal/RecordingTracer.scala @@ -9,11 +9,13 @@ import sage.commands.Command // record each span's start, routed node, and outcome so tests can compare the lifecycle without OpenTelemetry final class RecordingTracer extends CommandTracer { val log = mutable.ArrayBuffer.empty[String] + // spans of one command can start and settle on different threads + private def record(entry: String): Unit = log.synchronized(log += entry): Unit def onCommand(command: Command[?]): CommandSpan = { - log += s"start:${command.name}" + record(s"start:${command.name}") new CommandSpan { - def routedTo(node: Node): Unit = log += s"routed:${node.host}:${node.port}" - def settled(outcome: Outcome): Unit = log += s"settled:$outcome" + def routedTo(node: Node): Unit = record(s"routed:${node.host}:${node.port}") + def settled(outcome: Outcome): Unit = record(s"settled:$outcome") } } } diff --git a/sage-client/shared/src/test/scala/sage/client/internal/RefreshThrottleSpec.scala b/sage-client/shared/src/test/scala/sage/client/internal/RefreshThrottleSpec.scala index 80562f64..ab78a3be 100644 --- a/sage-client/shared/src/test/scala/sage/client/internal/RefreshThrottleSpec.scala +++ b/sage-client/shared/src/test/scala/sage/client/internal/RefreshThrottleSpec.scala @@ -9,16 +9,15 @@ class RefreshThrottleSpec extends munit.FunSuite { test("retained requests coalesce and run after the minimum interval without another failure") { val scheduler = new ManualScheduler - val throttle = new RefreshThrottle(scheduler, minRefreshMs = 1000L) var runs = 0 - val work = () => runs += 1 - throttle.request(work) - throttle.trigger(() => fail("an ordinary trigger duplicated an immediate retained refresh")) + val throttle = new RefreshThrottle(scheduler, minRefreshMs = 1000L, () => runs += 1) + throttle.request() + throttle.request() scheduler.advance(Duration.Zero) - (1 to 100).foreach(_ => throttle.request(work)) - throttle.trigger(() => fail("an ordinary trigger duplicated a retained refresh")) + assertEquals(runs, 1, "a second request duplicated an immediate refresh") + (1 to 100).foreach(_ => throttle.request()) scheduler.advance(999.millis) - assertEquals(runs, 1) + assertEquals(runs, 1, "a request inside the window ran before it closed") scheduler.advance(1.milli) assertEquals(runs, 2) scheduler.advance(2.seconds) @@ -26,45 +25,61 @@ class RefreshThrottleSpec extends munit.FunSuite { } test("a retained request during refresh runs after completion and respects a later forced refresh") { - val scheduler = new ManualScheduler - val throttle = new RefreshThrottle(scheduler, minRefreshMs = 1000L) - var runs = 0 - throttle.trigger(() => throttle.request(() => runs += 1)) + val scheduler = new ManualScheduler + var runs = 0 + var throttle: RefreshThrottle = null + throttle = new RefreshThrottle( + scheduler, + minRefreshMs = 1000L, + () => { + runs += 1 + if (runs == 1) throttle.request() + } + ) + throttle.request() scheduler.advance(Duration.Zero) scheduler.advance(500.millis) - throttle(force = true)(()) + throttle(force = true) scheduler.advance(999.millis) - assertEquals(runs, 0) + assertEquals(runs, 2) scheduler.advance(1.milli) - assertEquals(runs, 1) + assertEquals(runs, 3) } test("closing discards a retained refresh") { val scheduler = new ManualScheduler - val throttle = new RefreshThrottle(scheduler, minRefreshMs = 1000L) - throttle(force = true)(()) - throttle.request(() => fail("refresh ran after close")) - throttle.stopPolling() + var runs = 0 + val throttle = new RefreshThrottle(scheduler, minRefreshMs = 1000L, () => runs += 1) + throttle(force = true) + throttle.request() + throttle.stop() scheduler.advance(2.seconds) + assertEquals(runs, 1, "refresh ran after close") } - test("a non-blocking trigger schedules once while a refresh is in flight without changing blocking callers") { + test("requests during an in-flight refresh coalesce into one follow-up refresh, and a blocking caller waits") { val scheduler = new CountingScheduler - val throttle = new RefreshThrottle(scheduler, minRefreshMs = 0L) val started = new CountDownLatch(1) val release = new CountDownLatch(1) val finished = new CountDownLatch(1) - - throttle.trigger { () => - started.countDown() - release.await() - finished.countDown() - } - assert(started.await(2, TimeUnit.SECONDS), "the triggered refresh did not start") + val runs = new AtomicInteger(0) + val throttle = new RefreshThrottle( + scheduler, + minRefreshMs = 0L, + () => + if (runs.incrementAndGet() == 1) { + started.countDown() + release.await() + finished.countDown() + } + ) + + throttle.request() + assert(started.await(2, TimeUnit.SECONDS), "the requested refresh did not start") var i = 0 while (i < 1000) { - throttle.trigger(() => fail("an in-flight trigger must not run later work")) + throttle.request() i += 1 } assertEquals(scheduler.zeroDelays.get(), 1, "an in-flight refresh should own the only scheduler offload") @@ -73,52 +88,53 @@ class RefreshThrottleSpec extends munit.FunSuite { val blockingDone = new CountDownLatch(1) val waiter = Thread.ofVirtual().start { () => blockingStarted.countDown() - throttle(force = false)(()) + throttle(force = false) blockingDone.countDown() } assert(blockingStarted.await(2, TimeUnit.SECONDS), "the blocking caller did not start") assert(!blockingDone.await(100, TimeUnit.MILLISECONDS), "the existing blocking entry point must still wait for the in-flight refresh") release.countDown() - assert(finished.await(2, TimeUnit.SECONDS), "the triggered refresh did not finish") + assert(finished.await(2, TimeUnit.SECONDS), "the requested refresh did not finish") assert(blockingDone.await(2, TimeUnit.SECONDS), "the blocking caller did not resume") waiter.join() + while (runs.get() < 2) Thread.onSpinWait() + throttle.stop() + assertEquals(runs.get(), 2, "the requests made during the refresh did not run exactly one follow-up") } - test("the window closes as soon as a refresh finishes and reopens once it elapses") { + test("a request inside the refresh window runs once the window closes") { val scheduler = new ManualScheduler - val throttle = new RefreshThrottle(scheduler, minRefreshMs = 1000L) val runs = new AtomicInteger(0) - val work = () => runs.incrementAndGet(): Unit + val throttle = new RefreshThrottle(scheduler, minRefreshMs = 1000L, () => runs.incrementAndGet(): Unit) - throttle.trigger(work) + throttle.request() scheduler.advance(Duration.Zero) - assertEquals(runs.get(), 1, "the first trigger did not run") + assertEquals(runs.get(), 1, "the first request did not run") - throttle.trigger(work) + throttle.request() scheduler.advance(500.millis) - assertEquals(runs.get(), 1, "a trigger inside the refresh window ran its work") + assertEquals(runs.get(), 1, "a request inside the refresh window ran before it closed") - scheduler.advance(600.millis) - throttle.trigger(work) - scheduler.advance(Duration.Zero) - assertEquals(runs.get(), 2, "the window did not reopen once minRefreshInterval elapsed") + scheduler.advance(500.millis) + assertEquals(runs.get(), 2, "the retained request did not run when the window closed") } - test("a trigger inside the refresh window offloads nothing") { + test("requests inside the refresh window schedule no immediate offload") { val scheduler = new CountingScheduler - val throttle = new RefreshThrottle(scheduler, minRefreshMs = 60000L) val ran = new CountDownLatch(1) + val throttle = new RefreshThrottle(scheduler, minRefreshMs = 60000L, () => ran.countDown()) - throttle.trigger(() => ran.countDown()) - assert(ran.await(2, TimeUnit.SECONDS), "the first trigger did not run") + throttle.request() + assert(ran.await(2, TimeUnit.SECONDS), "the first request did not run") val offloads = scheduler.zeroDelays.get() var i = 0 while (i < 1000) { - throttle.trigger(() => fail("a throttled trigger must not run work")) + throttle.request() i += 1 } - assertEquals(scheduler.zeroDelays.get(), offloads, "a throttled trigger offloaded") + assertEquals(scheduler.zeroDelays.get(), offloads, "a throttled request offloaded") + throttle.stop() } } diff --git a/sage-client/shared/src/test/scala/sage/client/internal/Replies.scala b/sage-client/shared/src/test/scala/sage/client/internal/Replies.scala index f48ef223..61554625 100644 --- a/sage-client/shared/src/test/scala/sage/client/internal/Replies.scala +++ b/sage-client/shared/src/test/scala/sage/client/internal/Replies.scala @@ -15,6 +15,20 @@ object Replies { val pong: Frame = Frame.SimpleString("PONG") val queued: Frame = Frame.SimpleString("QUEUED") + // the setup HELLO selects RESP3; a HELLO without arguments, written after an SUNSUBSCRIBE, is left to `respond` + private def isSetupHello(payload: Bytes): Boolean = payload.asUtf8String.contains("\r\nHELLO\r\n$1\r\n3\r\n") + + // the setup commands every connection writes before its first command + def isSetup(payload: Bytes): Boolean = isSetupHello(payload) || payload.asUtf8String.contains("\r\nSETINFO\r\n") + + /** + * Answers the connection setup (`HELLO 3`, `CLIENT SETINFO`) and passes every other command to `respond`. + */ + def withSetup(respond: Bytes => Seq[Frame]): Bytes => Seq[Frame] = payload => + if (!isSetup(payload)) respond(payload) + else if (isSetupHello(payload)) Seq(hello) + else Seq(ok) + val hello: Frame = Frame.Map( Vector( diff --git a/sage-client/shared/src/test/scala/sage/client/internal/SocketTransportSpec.scala b/sage-client/shared/src/test/scala/sage/client/internal/SocketTransportSpec.scala index e98e6a41..c2afd256 100644 --- a/sage-client/shared/src/test/scala/sage/client/internal/SocketTransportSpec.scala +++ b/sage-client/shared/src/test/scala/sage/client/internal/SocketTransportSpec.scala @@ -1,9 +1,9 @@ package sage.client.internal -import java.io.InputStream -import java.net.{ServerSocket, Socket} +import java.io.{InputStream, IOException} +import java.net.{InetAddress, ServerSocket, Socket} import java.nio.charset.StandardCharsets -import java.util.concurrent.{CountDownLatch, TimeUnit} +import java.util.concurrent.{CountDownLatch, Semaphore, TimeUnit} import java.util.concurrent.atomic.AtomicBoolean import scala.concurrent.duration.* @@ -31,9 +31,10 @@ class SocketTransportSpec extends munit.FunSuite { onFrame: Frame => Unit = _ => (), beforeStart: SocketTransport => Unit = _ => () )(body: (SocketTransport, Socket) => Unit): Unit = { - val server = new ServerSocket(0) + // A wildcard bind can share its port with another process's loopback listener, which then takes the connection. + val server = new ServerSocket(0, 50, InetAddress.getLoopbackAddress) try { - val transport = SocketTransport.connect("127.0.0.1", server.getLocalPort, 5.seconds, identity, onFrame, onClosed) + val transport = SocketTransport.connect(server.getInetAddress.getHostAddress, server.getLocalPort, 5.seconds, identity, onFrame, onClosed) beforeStart(transport) transport.start() val peer = server.accept() @@ -82,6 +83,60 @@ class SocketTransportSpec extends munit.FunSuite { } } + test("send on an interrupted thread queues the item instead of throwing") { + withTransport(onClosed = () => ()) { (transport, peer) => + val item = new RecordingItem("PING\r\n") + @volatile var failure: Throwable = null + val sender = Thread.ofVirtual().start { () => + Thread.currentThread().interrupt() + try transport.send(item) + catch { case e: Throwable => failure = e } + } + sender.join() + assertEquals(failure, null) + assertEquals(readExactly(peer.getInputStream, 6), "PING\r\n") + assertEquals(item.writeAttempts, 1) + transport.close() + } + } + + test("close on an interrupted thread still waits for the I/O threads, drops queued items and runs onClosed") { + @volatile var closedCount = 0 + val inWrite = new CountDownLatch(1) + val release = new Semaphore(0) + withTransport(onClosed = () => closedCount += 1) { (transport, _) => + transport.send(new Transport.Item { + val payload: Bytes = Bytes.utf8("PING\r\n") + def writeAttempted(): Unit = { + inWrite.countDown() + release.acquireUninterruptibly() + } + def dropped(): Unit = () + }) + assert(inWrite.await(5, TimeUnit.SECONDS), "the writer should reach the first write") + val queued = new RecordingItem("PING\r\n") + transport.send(queued) + @volatile var failure: Throwable = null + @volatile var interruptedAfter = false + val closer = Thread.ofVirtual().start { () => + Thread.currentThread().interrupt() + try transport.close() + catch { case e: Throwable => failure = e } + interruptedAfter = Thread.currentThread().isInterrupted + } + awaitUntil(closer.getState == Thread.State.WAITING || !closer.isAlive, "close to wait for the writer or return") + release.release() + closer.join() + awaitUntil(!transport.writer.isAlive, "the writer to exit") + assertEquals(closedCount, 1) + assertEquals(queued.drops, 1) + assertEquals(failure, null) + assert(interruptedAfter, "close should restore the caller's interrupt") + assert(!transport.reader.isAlive) + assert(!transport.writer.isAlive) + } + } + test("items queued together are batched into a single socket write") { val items = (1 to 10).map(i => new RecordingItem(s"PING $i\r\n")) withTransport(onClosed = () => (), beforeStart = transport => items.foreach(transport.send)) { (transport, peer) => @@ -166,6 +221,35 @@ class SocketTransportSpec extends munit.FunSuite { } } + test("a reader that closes the transport while the writer tears it down does not deadlock") { + @volatile var closedCount = 0 + @volatile var transportRef: SocketTransport = null + val inOnFrame = new CountDownLatch(1) + withTransport( + onClosed = () => closedCount += 1, + onFrame = _ => { + inOnFrame.countDown() + // the writer's teardown interrupts the reader before joining it; a READONLY reply then closes the transport from the reader + val deadline = System.nanoTime() + 5.seconds.toNanos + while (!Thread.currentThread().isInterrupted && System.nanoTime() < deadline) Thread.onSpinWait() + transportRef.close() + }, + beforeStart = transportRef = _ + ) { (transport, peer) => + peer.getOutputStream.write("+OK\r\n".getBytes(StandardCharsets.UTF_8)) + peer.getOutputStream.flush() + assert(inOnFrame.await(5, TimeUnit.SECONDS), "reader should reach frame delivery") + transport.send(new Transport.Item { + val payload: Bytes = Bytes.utf8("PING\r\n") + def writeAttempted(): Unit = throw new IOException("write failed") + def dropped(): Unit = () + }) + awaitUntil(closedCount == 1, "onClosed after both I/O threads stop") + assert(!transport.reader.isAlive) + assert(!transport.writer.isAlive) + } + } + test("a malformed frame poisons the connection") { @volatile var closedCount = 0 withTransport(onClosed = () => closedCount += 1) { (transport, peer) => diff --git a/sage-client/shared/src/test/scala/sage/client/internal/SubscriptionConnectionSpec.scala b/sage-client/shared/src/test/scala/sage/client/internal/SubscriptionConnectionSpec.scala index ad8c92b9..2da1a488 100644 --- a/sage-client/shared/src/test/scala/sage/client/internal/SubscriptionConnectionSpec.scala +++ b/sage-client/shared/src/test/scala/sage/client/internal/SubscriptionConnectionSpec.scala @@ -8,10 +8,10 @@ import scala.concurrent.duration.* import Replies.bulk -import sage.Bytes -import sage.SageException.NotConnected -import sage.client.{BackoffConfig, WatchdogConfig} -import sage.commands.{Command, Connection} +import sage.{Bytes, Message, PatternMessage} +import sage.SageException.{NotConnected, ServerError, TlsError} +import sage.client.{BackoffConfig, PubSubConfig, SageConfig, WatchdogConfig} +import sage.cluster.{ClusterTopology, Node, Slot, SlotRange} import sage.protocol.Frame class SubscriptionConnectionSpec extends munit.FunSuite { @@ -19,19 +19,33 @@ class SubscriptionConnectionSpec extends munit.FunSuite { private val fixedBackoff = BackoffConfig(initialDelay = 1.milli, maxDelay = 1.milli, multiplier = 1.0) private val noWatchdog = WatchdogConfig(enabled = false) - // models the server's replies: PONG to the bootstrap/watchdog PING, and one subscribe/psubscribe confirmation push per subscribed name - private val serverResponder: Bytes => Seq[Frame] = { payload => - val s = payload.asUtf8String - if (s.contains("\r\nPING\r\n")) Seq(Frame.SimpleString("PONG")) - else if (s.contains("\r\nSUBSCRIBE\r\n")) confirmations("subscribe", s) - else if (s.contains("\r\nPSUBSCRIBE\r\n")) confirmations("psubscribe", s) - else if (s.contains("\r\nSSUBSCRIBE\r\n")) confirmations("ssubscribe", s) - else Nil + // models the server's replies: the setup, PONG to a PING, server information to the HELLO written after an SUNSUBSCRIBE, and one + // confirmation push per subscribed or unsubscribed name + private val serverResponder: Bytes => Seq[Frame] = Replies.withSetup(payload => commands(payload).flatMap(reply)) + + private def reply(command: List[String]): Seq[Frame] = + command match { + case List("PING") => Seq(Replies.pong) + case List("HELLO") => Seq(Replies.hello) + case List(verb, name) if verb.endsWith("SUBSCRIBE") => Seq(Frame.Push(Vector(bulk(verb.toLowerCase), bulk(name), Frame.Integer(1)))) + case _ => Nil + } + + // the commands in a write, each as its arguments + private def commands(payload: Bytes): List[List[String]] = { + def parse(tokens: List[String]): List[List[String]] = + tokens match { + case head :: rest if head.startsWith("*") => + val n = head.tail.toInt + rest.take(2 * n).grouped(2).map(_(1)).toList :: parse(rest.drop(2 * n)) + case _ => Nil + } + parse(payload.asUtf8String.split("\r\n").toList) } - // a SUBSCRIBE/PSUBSCRIBE is a RESP array `*K`; K-1 of those elements are channel names, each acknowledged by one confirmation push + // each subscribed name is its own command in the write, acknowledged by one confirmation push private def confirmations(kind: String, payload: String): Seq[Frame] = { - val names = payload.drop(1).takeWhile(_ != '\r').toInt - 1 + val names = payload.split("\r\n").count(_ == kind.toUpperCase) (1 to names).map(i => Frame.Push(Vector(bulk(kind), bulk("?"), Frame.Integer(i.toLong)))) } @@ -40,16 +54,25 @@ class SubscriptionConnectionSpec extends munit.FunSuite { watchdog: WatchdogConfig = noWatchdog, bufferSize: Int = 16, respond: Bytes => Seq[Frame] = serverResponder, - bootstrap: Vector[Command[?]] = Vector(Connection.ping()) + reconnect: BackoffConfig = fixedBackoff, + refuseHello: () => Boolean = () => false ): (SubscriptionConnection, ManualScheduler, mutable.ArrayBuffer[FakeTransport]) = { val scheduler = new ManualScheduler val transports = mutable.ArrayBuffer.empty[FakeTransport] + val refused = Seq(Frame.SimpleError("ERR server is down")) val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => { - val transport = new FakeTransport(onFrame, onClosed, respond, autoWrite = true) + val reply: Bytes => Seq[Frame] = payload => + if (refuseHello() && payload.asUtf8String.contains("\r\nHELLO\r\n")) refused else Replies.withSetup(respond)(payload) + val transport = new FakeTransport(onFrame, onClosed, reply) transports += transport transport } - val connection = new SubscriptionConnection(factory, bootstrap, scheduler, fixedBackoff, watchdog, 1000L, bufferSize, isLive) + val connection = new SubscriptionConnection( + factory, + scheduler, + SageConfig(reconnect = reconnect, watchdog = watchdog, connectTimeout = 1000.millis, pubsub = PubSubConfig(bufferSize = bufferSize)), + isLive + ) (connection, scheduler, transports) } @@ -85,35 +108,24 @@ class SubscriptionConnectionSpec extends munit.FunSuite { box.get() } - private def nextBlocking(sink: SubscriptionConnection.Sink): Option[SubscriptionConnection.Delivery] = { - val box = new AtomicReference[Option[SubscriptionConnection.Delivery]]() - val latch = new CountDownLatch(1) - sink.next { delivery => - box.set(delivery) - latch.countDown() - } - latch.await() - box.get() - } - // Bytes uses reference equality for `==`. Destructure the delivery and compare its payload as text. private def assertChannel(delivery: Option[SubscriptionConnection.Delivery], channel: String, payload: String)(using munit.Location): Unit = delivery match { - case Some(SubscriptionConnection.Delivery.Channel(ch, p)) => + case Some(Message(ch, p)) => assertEquals(ch, channel) assertEquals(p.asUtf8String, payload) - case other => fail(s"expected a channel delivery, got $other") + case other => fail(s"expected a channel delivery, got $other") } private def assertPattern(delivery: Option[SubscriptionConnection.Delivery], pattern: String, channel: String, payload: String)( using munit.Location ): Unit = delivery match { - case Some(SubscriptionConnection.Delivery.Pattern(pat, ch, p)) => + case Some(PatternMessage(pat, ch, p)) => assertEquals(pat, pattern) assertEquals(ch, channel) assertEquals(p.asUtf8String, payload) - case other => fail(s"expected a pattern delivery, got $other") + case other => fail(s"expected a pattern delivery, got $other") } test("first subscribe establishes the connection and sends SUBSCRIBE, then delivers a message") { @@ -144,17 +156,38 @@ class SubscriptionConnectionSpec extends munit.FunSuite { assertChannel(nextBlocking(b), "news", "x") } - test("UNSUBSCRIBE is sent only when the last subscriber of a channel closes, and the connection is torn down") { + test("a second subscriber to a channel whose SUBSCRIBE is still unconfirmed waits for that confirmation") { + val held = (payload: Bytes) => if (payload.asUtf8String.contains("news")) Nil else serverResponder(payload) + val (connection, _, transports) = make(respond = held) + connection.subscribeChannels(Vector("other")) + def subscribeInThread(): Thread = { + val thread = new Thread(() => connection.subscribeChannels(Vector("news")): Unit) + thread.start() + while (thread.getState != Thread.State.TIMED_WAITING && thread.getState != Thread.State.TERMINATED) Thread.onSpinWait() + thread + } + val first = subscribeInThread() + val second = subscribeInThread() + assertEquals(second.getState, Thread.State.TIMED_WAITING, "the second subscriber returned before the server confirmed the channel") + transports.head.emit(Frame.Push(Vector(bulk("subscribe"), bulk("news"), Frame.Integer(2L)))) + first.join() + second.join() + } + + test("UNSUBSCRIBE is sent only when the last subscriber of a channel closes, and closing the last subscription closes the connection") { val (connection, _, transports) = make() val a = connection.subscribeChannels(Vector("news")) val b = connection.subscribeChannels(Vector("news")) + val other = connection.subscribeChannels(Vector("other")) a.close() assert(!wrote(transports.head, "UNSUBSCRIBE"), "another subscriber still wants the channel") + b.close() + assert(wrote(transports.head, "UNSUBSCRIBE"), "the last subscriber of the channel leaving unsubscribes") assertEquals(transports.head.closeCount, 0) - b.close() - assert(wrote(transports.head, "UNSUBSCRIBE"), "the last subscriber leaving unsubscribes") + other.close() + assertEquals(transports.head.written.count(_.asUtf8String.contains("\r\nUNSUBSCRIBE\r\n")), 1, "the closing connection needs no UNSUBSCRIBE") assertEquals(transports.head.closeCount, 1) } @@ -190,6 +223,215 @@ class SubscriptionConnectionSpec extends munit.FunSuite { assertChannel(nextBlocking(sub), "news", "after-reconnect") } + test("a reconnect loop from an earlier loss stops once a later loss starts its own") { + @volatile var refusing = false + val (connection, scheduler, transports) = make(refuseHello = () => refusing) + val first = connection.subscribeChannels(Vector("news")) + refusing = true + transports(0).close() // its loop waits out the backoff + first.close() + refusing = false + connection.subscribeChannels(Vector("news")) + refusing = true + transports(1).close() + scheduler.advance(1.milli) + assertEquals(transports.size, 3, "only the latest loss's loop may attempt a reconnect") + } + + test("a connection that keeps dropping right after going live backs off further on each loss") { + val (connection, scheduler, transports) = make(reconnect = BackoffConfig(initialDelay = 10.millis, maxDelay = 1.second, multiplier = 2.0)) + connection.subscribeChannels(Vector("news")) + transports(0).close() + scheduler.advance(10.millis) + assertEquals(transports.size, 2, "the first loss retries after the initial delay") + transports(1).close() + scheduler.advance(10.millis) + assertEquals(transports.size, 2, "the second loss must wait longer than the first") + scheduler.advance(10.millis) + assertEquals(transports.size, 3) + } + + test("a subscribe the server rejects fails with its error and keeps the connection for other subscribers") { + val transports = mutable.ArrayBuffer.empty[FakeTransport] + val respond: Bytes => Seq[Frame] = payload => + if (payload.asUtf8String.contains("secret")) Seq(Frame.SimpleError("NOPERM no permissions to access the 'secret' channel")) + else serverResponder(payload) + val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => { + val t = new FakeTransport(onFrame, onClosed, respond) + transports += t + t + } + // a real clock, so a subscribe that never confirms gives up after connectTimeout + val connection = new SubscriptionConnection( + factory, + Scheduler.real, + SageConfig(reconnect = fixedBackoff, watchdog = noWatchdog, connectTimeout = 200.millis, pubsub = PubSubConfig(bufferSize = 16)), + () => true + ) + val news = connection.subscribeChannels(Vector("news")) + val error = intercept[ServerError](connection.subscribeChannels(Vector("secret"))) + assertEquals(error.code, "NOPERM") + assertEquals((transports.size, transports.head.closeCount), (1, 0)) + transports.head.emit(message("news", "still here")) + assertChannel(nextBlocking(news), "news", "still here") + } + + test("a name the server rejected can be subscribed again once the server accepts it") { + @volatile var denied = true + val transports = mutable.ArrayBuffer.empty[FakeTransport] + val respond: Bytes => Seq[Frame] = payload => + if (denied && payload.asUtf8String.contains("secret")) Seq(Frame.SimpleError("NOPERM no permissions to access the 'secret' channel")) + else serverResponder(payload) + val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => { + val t = new FakeTransport(onFrame, onClosed, respond) + transports += t + t + } + val connection = new SubscriptionConnection( + factory, + Scheduler.real, + SageConfig(reconnect = fixedBackoff, watchdog = noWatchdog, connectTimeout = 1.second, pubsub = PubSubConfig(bufferSize = 16)), + () => true + ) + connection.subscribeChannels(Vector("news")) + intercept[ServerError](connection.subscribeChannels(Vector("secret"))) + denied = false + val secret = connection.subscribeChannels(Vector("secret")) + transports.head.emit(message("secret", "granted")) + assertChannel(nextBlocking(secret), "secret", "granted") + } + + test("a subscribe whose SUBSCRIBE is dropped unwritten is not reported as confirmed") { + final class DroppingSubscribe(onFrame: Frame => Unit) extends Transport { + var closeCount = 0 + def start(): Unit = () + // a terminating socket drops a queued write without writing it; onClosed comes later + def send(item: Transport.Item): Unit = + if (item.payload.asUtf8String.contains("\r\nSUBSCRIBE\r\n")) item.dropped() + else { + item.writeAttempted() + serverResponder(item.payload).foreach(onFrame) + } + def close(): Unit = closeCount += 1 + } + val transports = mutable.ArrayBuffer.empty[DroppingSubscribe] + val connection = new SubscriptionConnection( + (onFrame, _) => { + val t = new DroppingSubscribe(onFrame) + transports += t + t + }, + Scheduler.real, + SageConfig(reconnect = fixedBackoff, watchdog = noWatchdog, connectTimeout = 200.millis, pubsub = PubSubConfig(bufferSize = 16)), + () => true + ) + intercept[NotConnected](connection.subscribeChannels(Vector("news"))) + assertEquals(transports.map(_.closeCount), mutable.ArrayBuffer(1)) + assert(connection.isEmpty) + } + + test("a name the server rejects when a reconnect restores it ends its subscription with the server's error") { + @volatile var denied = false + val respond: Bytes => Seq[Frame] = payload => + if (denied && payload.asUtf8String.contains("secret")) Seq(Frame.SimpleError("NOPERM no permissions to access the 'secret' channel")) + else serverResponder(payload) + val (connection, scheduler, transports) = make(respond = respond) + val secret = connection.subscribeChannels(Vector("secret")) + denied = true + transports(0).close() + scheduler.advance(1.milli) + val outcome = new AtomicReference[Option[SubscriptionConnection.Delivery]]() + secret.next(outcome.set) + (outcome.get(), secret.failure) match { + case (None, Some(error: ServerError)) => assertEquals(error.code, "NOPERM") + case other => fail(s"the stream must end with the server's error instead of staying open and silent, got $other") + } + } + + test("a name a busy server refuses while a reconnect restores it is restored by a later reconnect") { + @volatile var busy = false + val respond: Bytes => Seq[Frame] = payload => + if (busy && payload.asUtf8String.contains("\r\nSUBSCRIBE\r\n")) Seq(Frame.SimpleError("BUSY Redis is busy running a script")) + else serverResponder(payload) + val (connection, scheduler, transports) = make(respond = respond) + val news = connection.subscribeChannels(Vector("news")) + busy = true + transports(0).close() + scheduler.advance(1.milli) + assertEquals((transports.size, transports(1).closeCount), (2, 1)) + busy = false + scheduler.advance(1.milli) + transports(2).emit(message("news", "after-busy")) + val outcome = new AtomicReference[Option[SubscriptionConnection.Delivery]]() + news.next(outcome.set) + assertChannel(outcome.get(), "news", "after-busy") + } + + test("a shard channel a busy owner refuses is placed again instead of ending its subscription") { + val respond: Bytes => Seq[Frame] = Replies.withSetup(_ => Seq(Frame.SimpleError("BUSY Redis is busy running a script"))) + val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => new FakeTransport(onFrame, onClosed, respond) + var moved = 0 + val connection = new SubscriptionConnection( + factory, + new ManualScheduler, + SageConfig(reconnect = fixedBackoff, watchdog = noWatchdog, connectTimeout = 1000.millis, pubsub = PubSubConfig(bufferSize = 16)), + () => true, + SubscriptionConnection.OnLoss.Report(_ => (), () => (), () => moved += 1) + ) + val sink = new SubscriptionConnection.Sink(Vector("x"), SubscriptionConnection.Kind.Shard, 16) + connection.attach(sink, sink.names) + assertEquals((moved, sink.failure, connection.namesOf(sink)), (1, None, Vector.empty[String])) + } + + test("closing a subscription wakes its parked consumer with None and tears the socket down, even when its UNSUBSCRIBE is dropped") { + final class DroppingUnsubscribe(onFrame: Frame => Unit, onClosed: () => Unit) extends Transport { + var closeCount = 0 + var droppedCount = 0 + def start(): Unit = () + def send(item: Transport.Item): Unit = + if (item.payload.asUtf8String.contains("\r\nUNSUBSCRIBE\r\n")) { + droppedCount += 1 + item.dropped() + } else { + item.writeAttempted() + Replies.withSetup(serverResponder)(item.payload).foreach(onFrame) + } + def close(): Unit = { + closeCount += 1 + if (closeCount == 1) onClosed() + } + } + val transports = mutable.ArrayBuffer.empty[DroppingUnsubscribe] + val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => { + val t = new DroppingUnsubscribe(onFrame, onClosed) + transports += t + t + } + val connection = new SubscriptionConnection( + factory, + new ManualScheduler, + SageConfig(reconnect = fixedBackoff, watchdog = noWatchdog, connectTimeout = 1000.millis, pubsub = PubSubConfig(bufferSize = 16)), + () => true + ) + val sub = connection.subscribeChannels(Vector("news")) + val other = connection.subscribeChannels(Vector("other")) + val woken = new AtomicReference[Option[SubscriptionConnection.Delivery]]() + sub.next(woken.set) + sub.close() + assertEquals(woken.get(), None) + assertEquals(transports.head.droppedCount, 1) + other.close() + assertEquals(transports.head.closeCount, 1) + assert(connection.isEmpty) + } + + test("a failed first establish deregisters the subscription and closes its socket") { + val (connection, _, transports) = make(refuseHello = () => true) + intercept[sage.SageException](connection.subscribeChannels(Vector("news"))) + assertEquals(transports.map(_.closeCount), mutable.ArrayBuffer(1)) + assert(connection.isEmpty, "the failed subscription must not be restored by a later connection") + } + test("closing the client terminates active subscription streams") { val (connection, _, transports) = make() val sub = connection.subscribeChannels(Vector("news")) @@ -201,12 +443,11 @@ class SubscriptionConnectionSpec extends munit.FunSuite { test("a watchdog kill does not poison the next generation (no stale pingSentAtMillis kill-loop)") { val silentToPing: Bytes => Seq[Frame] = { payload => val s = payload.asUtf8String - if (s.contains("\r\nSELECT\r\n")) Seq(Frame.SimpleString("OK")) - else if (s.contains("\r\nSUBSCRIBE\r\n")) confirmations("subscribe", s) + if (s.contains("\r\nSUBSCRIBE\r\n")) confirmations("subscribe", s) else Nil } val watchdog = WatchdogConfig(pingInterval = 100.millis, pingTimeout = 50.millis, enabled = true) - val (connection, scheduler, transports) = make(watchdog = watchdog, respond = silentToPing, bootstrap = Vector(Connection.select(0))) + val (connection, scheduler, transports) = make(watchdog = watchdog, respond = silentToPing) val sub = connection.subscribeChannels(Vector("news")) assertEquals(transports.size, 1) @@ -214,7 +455,9 @@ class SubscriptionConnectionSpec extends munit.FunSuite { assertEquals(transports.size, 2, "the killed connection reconnects to a fresh transport") assertEquals(transports.head.closeCount, 1, "the unresponsive original connection was killed by the watchdog") - scheduler.advance(200.millis) + // one watchdog tick on the fresh generation, which only sends its own probe + scheduler.advance(140.millis) + assert(wrote(transports(1), "PING"), "the fresh generation's watchdog ticked") assertEquals(transports(1).closeCount, 0, "the fresh generation must not be killed by a ping left outstanding on the dead one") assertEquals(transports.size, 2, "no further reconnect churn") @@ -225,12 +468,11 @@ class SubscriptionConnectionSpec extends munit.FunSuite { test("a steady stream of message pushes does not suppress the keepalive PING (push proves only the read path)") { val silentToPing: Bytes => Seq[Frame] = { payload => val s = payload.asUtf8String - if (s.contains("\r\nSELECT\r\n")) Seq(Frame.SimpleString("OK")) - else if (s.contains("\r\nSUBSCRIBE\r\n")) confirmations("subscribe", s) - else Nil // do not answer PING; the watchdog eventually closes the connection after the keepalive times out + if (s.contains("\r\nSUBSCRIBE\r\n")) confirmations("subscribe", s) + else Nil // do not answer the probe; the watchdog eventually closes the connection after the keepalive times out } val watchdog = WatchdogConfig(pingInterval = 100.millis, pingTimeout = 50.millis, enabled = true) - val (connection, scheduler, transports) = make(watchdog = watchdog, respond = silentToPing, bootstrap = Vector(Connection.select(0))) + val (connection, scheduler, transports) = make(watchdog = watchdog, respond = silentToPing) connection.subscribeChannels(Vector("news")) assertEquals(transports.size, 1) @@ -245,6 +487,20 @@ class SubscriptionConnectionSpec extends munit.FunSuite { assertEquals(transports.size, 2, "the killed connection reconnected") } + test("the watchdog of a pub/sub-only user, whom the server refuses PING, keeps the connection") { + val respond: Bytes => Seq[Frame] = payload => + commands(payload).flatMap { + case List("PING") => Seq(Frame.SimpleError("NOPERM User u has no permissions to run the 'ping' command")) + case List("SUBSCRIBE", name) => Seq(Frame.Push(Vector(bulk("subscribe"), bulk(name), Frame.Integer(1)))) + case _ => Nil + } + val watchdog = WatchdogConfig(pingInterval = 100.millis, pingTimeout = 50.millis, enabled = true) + val (connection, scheduler, transports) = make(watchdog = watchdog, respond = respond) + connection.subscribeChannels(Vector("news")) + scheduler.advance(350.millis) + assertEquals((transports.size, transports.head.closeCount), (1, 0)) + } + test("subscribing to multiple channels in one call delivers messages from any of them") { val (connection, _, transports) = make() val sub = connection.subscribeChannels(Vector("a", "b")) @@ -257,34 +513,20 @@ class SubscriptionConnectionSpec extends munit.FunSuite { test("closeIfEmpty keeps a connection holding a sink, closes it once empty, and then rejects a racing attach") { val (connection, _, transports) = make() val sink = new SubscriptionConnection.Sink(Vector("orders"), SubscriptionConnection.Kind.Shard, 16) - connection.attach(sink, Vector("orders"), SubscriptionConnection.Kind.Shard) + connection.attach(sink, Vector("orders")) assert(!connection.closeIfEmpty(), "a connection carrying a live sink must not be evicted") assertEquals(transports.head.closeCount, 0) - connection.detach(sink, Vector("orders"), SubscriptionConnection.Kind.Shard) + connection.detach(sink, Vector("orders")) assert(connection.closeIfEmpty(), "with its last sink gone the connection is empty and closes") assertEquals(transports.head.closeCount, 1) val late = new SubscriptionConnection.Sink(Vector("late"), SubscriptionConnection.Kind.Shard, 16) - intercept[NotConnected](connection.attach(late, Vector("late"), SubscriptionConnection.Kind.Shard)) - } - - test("shutdown drops the socket but leaves the sink usable, so a re-home can re-attach it to a fresh connection") { - val (conn1, _, transports1) = make() - val sink = new SubscriptionConnection.Sink(Vector("news"), SubscriptionConnection.Kind.Channel, 16) - conn1.attach(sink, Vector("news"), SubscriptionConnection.Kind.Channel) - - conn1.shutdown() - assertEquals(transports1.head.closeCount, 1, "the abandoned socket is closed") - - val (conn2, _, transports2) = make() - conn2.attach(sink, Vector("news"), SubscriptionConnection.Kind.Channel) - transports2.head.emit(message("news", "after-rehome")) - assertChannel(nextBlocking(sink), "news", "after-rehome") + intercept[NotConnected](connection.attach(late, Vector("late"))) } - test("an unexpected error reply (e.g. MOVED on a subscribe) drops the connection instead of being swallowed as a PONG") { + test("an error reply with nothing pending drops the connection instead of being swallowed as a PONG") { var terminated = false val scheduler = new ManualScheduler val transports = mutable.ArrayBuffer.empty[FakeTransport] @@ -295,37 +537,278 @@ class SubscriptionConnectionSpec extends munit.FunSuite { } val connection = new SubscriptionConnection( factory, - Vector(Connection.ping()), scheduler, - fixedBackoff, - noWatchdog, - 1000L, - 16, + SageConfig(reconnect = fixedBackoff, watchdog = noWatchdog, connectTimeout = 1000.millis, pubsub = PubSubConfig(bufferSize = 16)), () => true, - cluster = true, - onTerminated = () => terminated = true + SubscriptionConnection.OnLoss.Report(_ => terminated = true, () => (), () => ()) ) val sink = new SubscriptionConnection.Sink(Vector("orders"), SubscriptionConnection.Kind.Shard, 16) - connection.attach(sink, Vector("orders"), SubscriptionConnection.Kind.Shard) + connection.attach(sink, Vector("orders")) assertEquals(transports.size, 1) transports.head.emit(Frame.SimpleError("MOVED 1234 10.0.0.2:6379")) - scheduler.advance(Duration.Zero) // the close is scheduled off the reader thread assert(terminated, "a MOVED error on the subscribe connection must drop it so the manager re-homes") assertEquals(transports.head.closeCount, 1) } + test("the confirmation of our own SUNSUBSCRIBE leaves a channel subscribed again meanwhile in place") { + val scheduler = new ManualScheduler + val transports = mutable.ArrayBuffer.empty[FakeTransport] + // the reader sees the SUNSUBSCRIBE reply only once the next SSUBSCRIBE is written, and in write order, before its confirmation + var unanswered = Seq.empty[Frame] + val respond: Bytes => Seq[Frame] = payload => + if (payload.asUtf8String.contains("\r\nSUNSUBSCRIBE\r\n")) { + unanswered = Seq(Frame.Push(Vector(bulk("sunsubscribe"), bulk("x"), Frame.Integer(1))), Replies.hello); Nil + } else { val replies = unanswered ++ serverResponder(payload); unanswered = Nil; replies } + val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => { + val t = new FakeTransport(onFrame, onClosed, respond) + transports += t + t + } + var dropped = 0 + val connection = new SubscriptionConnection( + factory, + scheduler, + SageConfig(reconnect = fixedBackoff, watchdog = noWatchdog, connectTimeout = 1000.millis, pubsub = PubSubConfig(bufferSize = 16)), + () => true, + SubscriptionConnection.OnLoss.Report(_ => (), () => dropped += 1, () => ()) + ) + def shardSink(name: String) = new SubscriptionConnection.Sink(Vector(name), SubscriptionConnection.Kind.Shard, 16) + val (first, other, second) = (shardSink("x"), shardSink("y"), shardSink("x")) + connection.attach(first, Vector("x")) + connection.attach(other, Vector("y")) + connection.detach(first, Vector("x")) + connection.attach(second, Vector("x")) + + assertEquals(dropped, 0) + transports.head.emit(shardMessage("x", "hello")) + assertChannel(nextBlocking(new SubscriptionConnection.RawSubscription(second, () => ())), "x", "hello") + } + + test("a shard channel dropped while another channel's SSUBSCRIBE is unanswered is placed again, and the SSUBSCRIBE keeps its own reply") { + val scheduler = new ManualScheduler + val transports = mutable.ArrayBuffer.empty[FakeTransport] + // x's slot moves away just before the server reads SSUBSCRIBE y + val respond: Bytes => Seq[Frame] = payload => + if (payload.asUtf8String.contains("\r\ny\r\n")) + Seq( + Frame.Push(Vector(bulk("sunsubscribe"), bulk("x"), Frame.Integer(0))), + Frame.Push(Vector(bulk("ssubscribe"), bulk("y"), Frame.Integer(1))) + ) + else serverResponder(payload) + val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => { + val t = new FakeTransport(onFrame, onClosed, respond) + transports += t + t + } + var (terminated, dropped) = (false, 0) + val connection = new SubscriptionConnection( + factory, + scheduler, + SageConfig(reconnect = fixedBackoff, watchdog = noWatchdog, connectTimeout = 1000.millis, pubsub = PubSubConfig(bufferSize = 16)), + () => true, + SubscriptionConnection.OnLoss.Report(_ => terminated = true, () => dropped += 1, () => ()) + ) + def shardSink(name: String) = new SubscriptionConnection.Sink(Vector(name), SubscriptionConnection.Kind.Shard, 16) + val (x, y) = (shardSink("x"), shardSink("y")) + connection.attach(x, x.names) + connection.attach(y, y.names) + + assertEquals((terminated, dropped, transports.head.closeCount, connection.namesOf(x)), (false, 1, 0, Vector.empty[String])) + transports.head.emit(shardMessage("y", "hello")) + assertChannel(nextBlocking(new SubscriptionConnection.RawSubscription(y, () => ())), "y", "hello") + } + + test("SUNSUBSCRIBE errors like Valkey's leave the shard connection and its other subscribers in place") { + // like Valkey 9.1: a SUNSUBSCRIBE naming channels in two slots gets CROSSSLOT, and one for a slot the node no longer serves gets MOVED, + // after the push that drops the channel when its slot moved while the SUNSUBSCRIBE was on its way + val redirect = Frame.SimpleError("MOVED 1234 10.0.0.2:6379") + val respond: Bytes => Seq[Frame] = Replies.withSetup { payload => + commands(payload).flatMap { + case List("SSUBSCRIBE", name) => Seq(Frame.Push(Vector(bulk("ssubscribe"), bulk(name), Frame.Integer(1)))) + case List("SUNSUBSCRIBE", "moved") => Seq(redirect) + case List("SUNSUBSCRIBE", "migrated") => Seq(Frame.Push(Vector(bulk("sunsubscribe"), bulk("migrated"), Frame.Integer(0))), redirect) + case List("SUNSUBSCRIBE", name) => Seq(Frame.Push(Vector(bulk("sunsubscribe"), bulk(name), Frame.Integer(0)))) + case "SUNSUBSCRIBE" :: _ => Seq(Frame.SimpleError("CROSSSLOT Keys in request don't hash to the same slot")) + case List("HELLO") => Seq(Replies.hello) + case _ => Nil + } + } + val scheduler = new ManualScheduler + val transports = mutable.ArrayBuffer.empty[FakeTransport] + val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => { + val t = new FakeTransport(onFrame, onClosed, respond) + transports += t + t + } + var (terminated, moved) = (false, 0) + val connection = new SubscriptionConnection( + factory, + scheduler, + SageConfig(reconnect = fixedBackoff, watchdog = noWatchdog, connectTimeout = 1000.millis, pubsub = PubSubConfig(bufferSize = 16)), + () => true, + SubscriptionConnection.OnLoss.Report(_ => terminated = true, () => (), () => moved += 1) + ) + def shardSink(names: String*) = new SubscriptionConnection.Sink(names.toVector, SubscriptionConnection.Kind.Shard, 16) + val (multi, unowned, migrated, other) = (shardSink("x", "y"), shardSink("moved"), shardSink("migrated"), shardSink("z")) + Seq(multi, unowned, migrated, other).foreach(sink => connection.attach(sink, sink.names)) + Seq(multi, unowned, migrated).foreach(sink => connection.detach(sink, sink.names)) + + assert(!terminated, "an error reply to an unsubscribe must not end the shared shard connection") + assertEquals((transports.size, transports.head.closeCount, moved), (1, 0, 0)) + transports.head.emit(shardMessage("z", "still here")) + assertChannel(nextBlocking(new SubscriptionConnection.RawSubscription(other, () => ())), "z", "still here") + val late = shardSink("w") + connection.attach(late, late.names) + assertEquals(connection.namesOf(late), Vector("w")) + } + + test("detaching many shard channels writes one HELLO, and later replies still match their commands") { + val respond: Bytes => Seq[Frame] = Replies.withSetup { payload => + commands(payload).flatMap { + case List("SUNSUBSCRIBE", "c5") => Seq(Frame.SimpleError("MOVED 1234 10.0.0.2:6379")) + case command => reply(command) + } + } + val transports = mutable.ArrayBuffer.empty[FakeTransport] + val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => { + val t = new FakeTransport(onFrame, onClosed, respond) + transports += t + t + } + var terminated = false + val connection = new SubscriptionConnection( + factory, + new ManualScheduler, + SageConfig(reconnect = fixedBackoff, watchdog = noWatchdog, connectTimeout = 1000.millis, pubsub = PubSubConfig(bufferSize = 16)), + () => true, + SubscriptionConnection.OnLoss.Report(_ => terminated = true, () => (), () => ()) + ) + val names = Vector.tabulate(1000)(i => s"c$i") + val (many, kept, late) = + ( + new SubscriptionConnection.Sink(names, SubscriptionConnection.Kind.Shard, 16), + new SubscriptionConnection.Sink(Vector("kept"), SubscriptionConnection.Kind.Shard, 16), + new SubscriptionConnection.Sink(Vector("late"), SubscriptionConnection.Kind.Shard, 16) + ) + Seq(many, kept).foreach(sink => connection.attach(sink, sink.names)) + connection.detach(many, names) + + val hellos = transports.head.sent.map(_.asUtf8String.split("\r\n").count(_ == "HELLO")).sum + assertEquals((hellos, terminated, transports.head.closeCount), (1, false, 0)) + connection.attach(late, late.names) + assertEquals(connection.namesOf(late), Vector("late")) + transports.head.emit(shardMessage("kept", "still here")) + assertChannel(nextBlocking(new SubscriptionConnection.RawSubscription(kept, () => ())), "kept", "still here") + } + + test("an unsubscribe settles only its own replies for a pub/sub-only user, whom the server refuses PING") { + // like a user granted only +@pubsub: PING gets NOPERM, while (UN)SUBSCRIBE succeed + val respond: Bytes => Seq[Frame] = payload => + commands(payload).flatMap { + case List("PING") => Seq(Frame.SimpleError("NOPERM User u has no permissions to run the 'ping' command")) + case List(verb @ ("SUBSCRIBE" | "UNSUBSCRIBE"), name) => Seq(Frame.Push(Vector(bulk(verb.toLowerCase), bulk(name), Frame.Integer(1)))) + case _ => Nil + } + val (connection, _, transports) = make(respond = respond) + val kept = connection.subscribeChannels(Vector("kept")) + connection.subscribeChannels(Vector("gone")).close() + val late = new AtomicReference[SubscriptionConnection.RawSubscription]() + val subscribing = new Thread(() => + try late.set(connection.subscribeChannels(Vector("late"))) + catch { case _: Throwable => () } + ) + subscribing.start() + subscribing.join(2000) + if (subscribing.isAlive) { + connection.close() + subscribing.join() + fail("the late subscribe's confirmation answered another command") + } + + assert(late.get() != null, "the late subscribe failed") + assertEquals(transports.head.closeCount, 0) + transports.head.emit(message("late", "hello")) + assertChannel(nextBlocking(late.get()), "late", "hello") + kept.close() + } + + // Like the ACL user `+subscribe +unsubscribe +psubscribe +punsubscribe +ssubscribe +sunsubscribe +ping +hello +client allchannels`: every + // other command gets NOPERM, and `denials` counts them. + final private class LeastPrivilege { + @volatile var denials = 0 + val respond: Bytes => Seq[Frame] = payload => + commands(payload).flatMap { + case command @ (verb :: _) if verb.endsWith("SUBSCRIBE") || verb == "PING" || verb == "HELLO" => reply(command) + case verb :: _ => + denials += 1 + Seq(Frame.SimpleError(s"NOPERM User u has no permissions to run the '${verb.toLowerCase}' command")) + case Nil => Nil + } + } + + test("unsubscribing keeps the connection of a user granted only pub/sub, PING, HELLO and CLIENT, and later replies match") { + val user = new LeastPrivilege + val (connection, _, transports) = make(respond = user.respond) + val kept = connection.subscribeChannels(Vector("kept")) + val gone = + Vector(connection.subscribeChannels(Vector("gone")), connection.subscribePatterns(Vector("gone.*")), connection.subscribeShard(Vector("gone"))) + gone.foreach(_.close()) + assertEquals((transports.size, transports.head.closeCount, user.denials), (1, 0, 0)) + + val late = connection.subscribeChannels(Vector("late")) + transports.head.emit(message("late", "hello")) + assertChannel(nextBlocking(late), "late", "hello") + transports.head.emit(message("kept", "still here")) + assertChannel(nextBlocking(kept), "kept", "still here") + } + + test("an SUNSUBSCRIBE keeps the cluster shard connection of a user granted only pub/sub, PING, HELLO and CLIENT") { + val user = new LeastPrivilege + val transports = mutable.ArrayBuffer.empty[FakeTransport] + val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => { + val t = new FakeTransport(onFrame, onClosed, Replies.withSetup(user.respond)) + transports += t + t + } + var terminated = false + val connection = new SubscriptionConnection( + factory, + new ManualScheduler, + SageConfig(reconnect = fixedBackoff, watchdog = noWatchdog, connectTimeout = 1000.millis, pubsub = PubSubConfig(bufferSize = 16)), + () => true, + SubscriptionConnection.OnLoss.Report(_ => terminated = true, () => (), () => ()) + ) + def shardSink(name: String) = new SubscriptionConnection.Sink(Vector(name), SubscriptionConnection.Kind.Shard, 16) + val (gone, kept, late) = (shardSink("gone"), shardSink("kept"), shardSink("late")) + Seq(gone, kept).foreach(sink => connection.attach(sink, sink.names)) + connection.detach(gone, gone.names) + assertEquals((terminated, transports.size, transports.head.closeCount, user.denials), (false, 1, 0, 0)) + + connection.attach(late, late.names) + assertEquals(connection.namesOf(late), Vector("late")) + transports.head.emit(shardMessage("kept", "still here")) + assertChannel(nextBlocking(new SubscriptionConnection.RawSubscription(kept, () => ())), "kept", "still here") + } + + test("the watchdog of a user granted only pub/sub, PING, HELLO and CLIENT probes with PING, which the ACL allows") { + val user = new LeastPrivilege + val watchdog = WatchdogConfig(pingInterval = 100.millis, pingTimeout = 50.millis, enabled = true) + val (connection, scheduler, transports) = make(watchdog = watchdog, respond = user.respond) + connection.subscribeChannels(Vector("news")) + scheduler.advance(350.millis) + assert(wrote(transports.head, "PING"), "the watchdog never probed") + assertEquals((transports.size, transports.head.closeCount, user.denials), (1, 0, 0)) + } + test("a socket dying in the establish->live window fires onTerminated in cluster mode instead of going silently Live") { final class DyingAfterBootstrap(onFrame: Frame => Unit, onClosed: () => Unit) extends Transport { - private var bootstrapped = false def start(): Unit = () + // answers the setup, then drops the connection right after the last setup command def send(item: Transport.Item): Unit = { item.writeAttempted() serverResponder(item.payload).foreach(onFrame) - if (!bootstrapped) { - bootstrapped = true - onClosed() - } + if (item.payload.asUtf8String.contains("LIB-VER")) onClosed() } def close(): Unit = () } @@ -334,19 +817,14 @@ class SubscriptionConnectionSpec extends munit.FunSuite { val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => new DyingAfterBootstrap(onFrame, onClosed) val connection = new SubscriptionConnection( factory, - Vector(Connection.ping()), scheduler, - fixedBackoff, - noWatchdog, - 1000L, - 16, + SageConfig(reconnect = fixedBackoff, watchdog = noWatchdog, connectTimeout = 1000.millis, pubsub = PubSubConfig(bufferSize = 16)), () => true, - cluster = true, - onTerminated = () => terminatedCalled = true + SubscriptionConnection.OnLoss.Report(_ => terminatedCalled = true, () => (), () => ()) ) val sink = new SubscriptionConnection.Sink(Vector("news"), SubscriptionConnection.Kind.Channel, 16) - connection.attach(sink, Vector("news"), SubscriptionConnection.Kind.Channel) + connection.attach(sink, Vector("news")) assert(terminatedCalled, "a connection that died in the establish->live window must notify the manager, not go silently Live") } @@ -370,7 +848,7 @@ class SubscriptionConnectionSpec extends munit.FunSuite { box.set(delivery) latch.countDown() } - sink.offer(SubscriptionConnection.Delivery.Channel("news", Bytes.utf8("hello"))) + sink.offer(Message("news", Bytes.utf8("hello"))) latch.await() assertChannel(box.get(), "news", "hello") } @@ -383,71 +861,6 @@ class SubscriptionConnectionSpec extends munit.FunSuite { intercept[IllegalStateException](sink.next(_ => ())) } - test("a resubscribe write that fails during goLive tears the socket down instead of orphaning a subscribed connection") { - final class FailingSubscribe(onFrame: Frame => Unit, onClosed: () => Unit) extends Transport { - var closeCount = 0 - def start(): Unit = () - def send(item: Transport.Item): Unit = { - if (item.payload.asUtf8String.contains("\r\nSUBSCRIBE\r\n")) throw new RuntimeException("write failed") - item.writeAttempted() - serverResponder(item.payload).foreach(onFrame) - } - def close(): Unit = { - closeCount += 1 - onClosed() - } - } - val transports = mutable.ArrayBuffer.empty[FailingSubscribe] - val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => { - val t = new FailingSubscribe(onFrame, onClosed) - transports += t - t - } - val connection = - new SubscriptionConnection(factory, Vector(Connection.ping()), new ManualScheduler, fixedBackoff, noWatchdog, 1000L, 16, () => true) - - intercept[RuntimeException](connection.subscribeChannels(Vector("news"))) - assertEquals(transports.head.closeCount, 1, "the socket whose resubscribe failed must be closed, not left dispatching") - assert(connection.isEmpty, "the phantom sink is deregistered so the connection holds nothing") - } - - test("a failed unsubscribe write during close still wakes the parked consumer and tears the socket down") { - final class FailingUnsubscribe(onFrame: Frame => Unit, onClosed: () => Unit) extends Transport { - var closeCount = 0 - def start(): Unit = () - def send(item: Transport.Item): Unit = { - if (item.payload.asUtf8String.contains("\r\nUNSUBSCRIBE\r\n")) throw new RuntimeException("write failed") - item.writeAttempted() - serverResponder(item.payload).foreach(onFrame) - } - def close(): Unit = { - closeCount += 1 - onClosed() - } - } - val transports = mutable.ArrayBuffer.empty[FailingUnsubscribe] - val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => { - val t = new FailingUnsubscribe(onFrame, onClosed) - transports += t - t - } - val connection = - new SubscriptionConnection(factory, Vector(Connection.ping()), new ManualScheduler, fixedBackoff, noWatchdog, 1000L, 16, () => true) - - val sub = connection.subscribeChannels(Vector("news")) - val box = new AtomicReference[Option[SubscriptionConnection.Delivery]]() - val latch = new CountDownLatch(1) - sub.next { delivery => - box.set(delivery) - latch.countDown() - } - - intercept[RuntimeException](sub.close()) - assert(latch.await(1, java.util.concurrent.TimeUnit.SECONDS), "the parked consumer must be woken, not left hanging") - assertEquals(box.get(), None) - assertEquals(transports.head.closeCount, 1) - } - test("a connect failure on the subscribe path surfaces as a modeled ConnectionFailed, not a raw defect") { val factory: MultiplexedConnection.TransportFactory = (_, _) => new Transport { @@ -456,7 +869,12 @@ class SubscriptionConnectionSpec extends munit.FunSuite { def close(): Unit = () } val connection = - new SubscriptionConnection(factory, Vector(Connection.ping()), new ManualScheduler, fixedBackoff, noWatchdog, 1000L, 16, () => true) + new SubscriptionConnection( + factory, + new ManualScheduler, + SageConfig(reconnect = fixedBackoff, watchdog = noWatchdog, connectTimeout = 1000.millis, pubsub = PubSubConfig(bufferSize = 16)), + () => true + ) intercept[sage.SageException.ConnectionFailed](connection.subscribeChannels(Vector("news"))) } @@ -473,7 +891,7 @@ class SubscriptionConnectionSpec extends munit.FunSuite { } var healthy = true val respond: Bytes => Seq[Frame] = payload => - if (!healthy && payload.asUtf8String.contains("\r\nPING\r\n")) Seq(Frame.SimpleError("WRONGPASS invalid password")) + if (!healthy && payload.asUtf8String.contains("\r\nHELLO\r\n")) Seq(Frame.SimpleError("WRONGPASS invalid password")) else serverResponder(payload) val transports = mutable.ArrayBuffer.empty[FakeTransport] val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => { @@ -483,7 +901,13 @@ class SubscriptionConnectionSpec extends munit.FunSuite { } val scheduler = new ManualScheduler val connection = - new SubscriptionConnection(factory, Vector(Connection.ping()), scheduler, fixedBackoff, noWatchdog, 1000L, 16, () => true, events = events) + new SubscriptionConnection( + factory, + scheduler, + SageConfig(reconnect = fixedBackoff, watchdog = noWatchdog, connectTimeout = 1000.millis, pubsub = PubSubConfig(bufferSize = 16)), + () => true, + SubscriptionConnection.OnLoss.Reconnect(() => (), events) + ) connection.subscribeChannels(Vector("news")) healthy = false @@ -498,33 +922,47 @@ class SubscriptionConnectionSpec extends munit.FunSuite { ) } - test("detach swallows an interrupted unsubscribe write so a cluster close can still terminate the sink") { - final class InterruptingUnsubscribe(onFrame: Frame => Unit, onClosed: () => Unit) extends Transport { - def start(): Unit = () - def send(item: Transport.Item): Unit = { - if (item.payload.asUtf8String.contains("\r\nUNSUBSCRIBE\r\n")) throw new InterruptedException("interrupted") - item.writeAttempted() - serverResponder(item.payload).foreach(onFrame) - } - def close(): Unit = onClosed() + test("a failed master-replica subscription reconnect names the node it tried, or none when there is no master") { + val recorded = new java.util.concurrent.ConcurrentLinkedQueue[sage.SageEvent]() + val events = new Events { + def enabled: Boolean = true + def emitsEvents: Boolean = true + def tracer: Option[sage.CommandTracer] = None + def serverNode: Option[sage.cluster.Node] = None + def emit(event: sage.SageEvent): Unit = recorded.add(event): Unit + def close(): Unit = () } - val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => new InterruptingUnsubscribe(onFrame, onClosed) - val connection = - new SubscriptionConnection( - factory, - Vector(Connection.ping()), - new ManualScheduler, - fixedBackoff, - noWatchdog, - 1000L, - 16, - () => true, - cluster = true - ) + val (first, second) = (Node("first", 1), Node("second", 2)) + @volatile var master = Option(first) + val transports = mutable.ArrayBuffer.empty[FakeTransport] + val nodeFactory: Node => MultiplexedConnection.TransportFactory = node => + (onFrame, onClosed) => { + if (node == second) throw new java.net.ConnectException("connection refused") + val t = new FakeTransport(onFrame, onClosed, serverResponder) + transports += t + t + } + val following = new SubscriptionConnection.Following(nodeFactory, () => master) + val scheduler = new ManualScheduler + val connection = new SubscriptionConnection( + following.factory, + scheduler, + SageConfig(reconnect = fixedBackoff, watchdog = noWatchdog, connectTimeout = 1000.millis, pubsub = PubSubConfig(bufferSize = 16)), + () => true, + SubscriptionConnection.OnLoss.Reconnect(() => (), events, () => following.node) + ) + def failedOn(): Vector[Option[Node]] = + recorded.toArray.toVector.collect { case sage.SageEvent.Connection.ReconnectFailed(node, _) => node } - val sink = new SubscriptionConnection.Sink(Vector("news"), SubscriptionConnection.Kind.Channel, 16) - connection.attach(sink, Vector("news"), SubscriptionConnection.Kind.Channel) - assert(connection.detach(sink, Vector("news"), SubscriptionConnection.Kind.Channel)) // must return, not throw the interrupt + connection.subscribeChannels(Vector("news")) + master = Some(second) + transports.head.close() + scheduler.advance(1.milli) + assertEquals(failedOn(), Vector(Some(second))) + master = None + scheduler.advance(1.milli) + assertEquals(failedOn(), Vector(Some(second), None)) + connection.close() } test("close aborts a subscription connection still blocked in the connect (start) phase") { @@ -535,7 +973,12 @@ class SubscriptionConnectionSpec extends munit.FunSuite { transport } val connection = - new SubscriptionConnection(factory, Vector(Connection.ping()), new ManualScheduler, fixedBackoff, noWatchdog, 1000L, 16, () => true) + new SubscriptionConnection( + factory, + new ManualScheduler, + SageConfig(reconnect = fixedBackoff, watchdog = noWatchdog, connectTimeout = 1000.millis, pubsub = PubSubConfig(bufferSize = 16)), + () => true + ) val subscribing = new Thread(() => try connection.subscribeChannels(Vector("news")): Unit @@ -556,26 +999,31 @@ class SubscriptionConnectionSpec extends munit.FunSuite { } test("close unblocks a subscriber parked waiting for the bootstrap reply") { - val pinged = new CountDownLatch(1) + val greeted = new CountDownLatch(1) val factory: MultiplexedConnection.TransportFactory = (onFrame, onClosed) => new FakeTransport( onFrame, onClosed, payload => { - if (payload.asUtf8String.contains("PING")) pinged.countDown() + if (payload.asUtf8String.contains("HELLO")) greeted.countDown() Nil } ) val connection = - new SubscriptionConnection(factory, Vector(Connection.ping()), new ManualScheduler, fixedBackoff, noWatchdog, 60000L, 16, () => true) + new SubscriptionConnection( + factory, + new ManualScheduler, + SageConfig(reconnect = fixedBackoff, watchdog = noWatchdog, connectTimeout = 60000.millis, pubsub = PubSubConfig(bufferSize = 16)), + () => true + ) val subscribing = new Thread(() => try connection.subscribeChannels(Vector("news")): Unit catch { case _: Throwable => () } ) subscribing.start() - assert(pinged.await(2, TimeUnit.SECONDS), "the bootstrap PING was never sent") + assert(greeted.await(2, TimeUnit.SECONDS), "the bootstrap HELLO was never sent") connection.close() subscribing.join(2000) @@ -599,7 +1047,12 @@ class SubscriptionConnectionSpec extends munit.FunSuite { t } val connection = - new SubscriptionConnection(factory, Vector(Connection.ping()), scheduler, fixedBackoff, noWatchdog, 5000L, 16, () => true) + new SubscriptionConnection( + factory, + scheduler, + SageConfig(reconnect = fixedBackoff, watchdog = noWatchdog, connectTimeout = 5000.millis, pubsub = PubSubConfig(bufferSize = 16)), + () => true + ) def awaitReached(i: Int): ConnectingTransport = { val deadline = System.currentTimeMillis() + 2000 @@ -634,4 +1087,65 @@ class SubscriptionConnectionSpec extends munit.FunSuite { assert(conn1.wasClosed, "close must abort the reconnect establishment") assert(conn2.wasClosed, "close must abort the concurrent attach establishment") } + + test("a shard connection that reports its loss late does not unregister the connection that replaced it") { + val (n, m, mid) = (Node("n", 1), Node("m", 2), Slot.Count / 2) + def channelOn(node: Node) = Iterator.from(0).map(i => s"c$i").find(c => (Slot.of(Bytes.utf8(c)).value < mid) == (node == n)).get + def owned(from: Int, to: Int, node: Node) = SlotRange(Slot.at(from).get, Slot.at(to).get, node, Vector.empty) + val topology = ClusterTopology.from(Vector(owned(0, mid - 1, n), owned(mid, Slot.Count - 1, m))) + val transports = mutable.Map.empty[Node, mutable.ArrayBuffer[FakeTransport]] + var whileOpeningM: () => Unit = () => () + val factory: Node => MultiplexedConnection.TransportFactory = node => { + if (node == m) whileOpeningM() // runs while the manager holds its lock + (onFrame, onClosed) => { + val transport = new FakeTransport(onFrame, onClosed, serverResponder) + transports.getOrElseUpdate(node, mutable.ArrayBuffer.empty) += transport + transport + } + } + val idle = new Scheduler { + def nowMillis: Long = 0L + def jitterMillis(boundExclusive: Long): Long = 0L + def after(delay: FiniteDuration)(task: => Unit): Unit = () + def every(interval: FiniteDuration)(task: => Unit): Scheduler.Cancelable = () => () + } + val manager = + new ClusterSubscriptions(factory, idle, SageConfig(watchdog = noWatchdog), () => topology, () => (), () => Some(n), Events.disabled) + val first = manager.subscribeShard(Vector(channelOn(n))) + whileOpeningM = () => { + whileOpeningM = () => () + // the first connection dies, and its report waits for the manager lock that this thread holds + val dying = new Thread(() => transports(n).head.close()) + dying.start() + while (dying.getState != Thread.State.WAITING) Thread.onSpinWait() + first.close() // evicts the dead connection, which is now empty + manager.subscribeShard(Vector(channelOn(n))): Unit + } + manager.subscribeShard(Vector(channelOn(m))) + manager.close() + assertEquals(transports(n).map(_.closeCount).toVector, Vector(1, 1), "close must reach the replacement connection") + } + + test("a shard owner whose connection cannot be created, as with unusable trust material, is retried instead of failing the subscribe") { + val node = Node("n", 1) + val topology = ClusterTopology.from(Vector(SlotRange(Slot.at(0).get, Slot.at(Slot.Count - 1).get, node, Vector.empty))) + @volatile var unusable = true + val transports = mutable.ArrayBuffer.empty[FakeTransport] + val factory: Node => MultiplexedConnection.TransportFactory = _ => + if (unusable) throw TlsError("unusable TLS trust material") + else + (onFrame, onClosed) => { + val transport = new FakeTransport(onFrame, onClosed, serverResponder) + transports += transport + transport + } + val scheduler = new ManualScheduler + val manager = + new ClusterSubscriptions(factory, scheduler, SageConfig(watchdog = noWatchdog), () => topology, () => (), () => Some(node), Events.disabled) + manager.subscribeShard(Vector("orders")) + unusable = false + scheduler.advance(1.second) + assert(transports.exists(wrote(_, "SSUBSCRIBE")), "the retry never placed the channel") + manager.close() + } } diff --git a/sage-client/shared/src/test/scala/sage/client/internal/TlsSpec.scala b/sage-client/shared/src/test/scala/sage/client/internal/TlsSpec.scala index f12a8ea5..2f84a764 100644 --- a/sage-client/shared/src/test/scala/sage/client/internal/TlsSpec.scala +++ b/sage-client/shared/src/test/scala/sage/client/internal/TlsSpec.scala @@ -3,10 +3,12 @@ package sage.client.internal import java.net.{InetAddress, InetSocketAddress, Socket} import java.nio.file.{Files, Path} import java.security.KeyStore +import java.security.cert.Certificate import javax.net.ssl.{KeyManagerFactory, SSLContext, SSLServerSocket, TrustManagerFactory} import sage.SageException.TlsError -import sage.client.{TlsConfig, TrustSource} +import sage.client.{SageConfig, TlsConfig, TrustSource} +import sage.cluster.Node class TlsSpec extends munit.FunSuite { @@ -14,7 +16,9 @@ class TlsSpec extends munit.FunSuite { // each test changes only the connection host. private lazy val material = certMaterial() - private def certMaterial(): (SSLContext, SSLContext) = { + private def certificate: Certificate = material._3 + + private def certMaterial(): (SSLContext, SSLContext, Certificate) = { val dir = Files.createTempDirectory("sage-tls-unit") val store = dir.resolve("server.p12") val pass = "changeit".toCharArray @@ -61,7 +65,7 @@ class TlsSpec extends munit.FunSuite { tmf.init(trust) val client = SSLContext.getInstance("TLS") client.init(null, tmf.getTrustManagers, null) - (server, client) + (server, client, ks.getCertificate("server")) } private def withServer(body: Int => Unit): Unit = { @@ -100,4 +104,23 @@ class TlsSpec extends munit.FunSuite { () } } + + test("a node's connections share the TLS context built for the node, so connecting after its trust file is gone succeeds") { + val pem = Files.createTempFile("sage-tls-trust", ".pem") + val encoded = java.util.Base64.getMimeEncoder.encodeToString(certificate.getEncoded) + Files.writeString(pem, s"-----BEGIN CERTIFICATE-----\n$encoded\n-----END CERTIFICATE-----\n") + withServer { port => + val factory = Client.transports(SageConfig(tls = Some(TlsConfig(TrustSource.Pem(pem)))))(Node("localhost", port)) + Files.delete(pem) + val transport = factory(_ => (), () => ()) + try transport.start() + finally transport.close() + } + } + + test("unusable trust material fails with TlsError before connecting, even to an unreachable node") { + val missing = Files.createTempFile("sage-tls-trust", ".pem") + Files.delete(missing) + intercept[TlsError](Client.transports(SageConfig(tls = Some(TlsConfig(TrustSource.Pem(missing)))))(Node("localhost", 1))) + } } diff --git a/sage-client/zio/src/main/scala/sage/backend/SageClient.scala b/sage-client/zio/src/main/scala/sage/backend/SageClient.scala index a088c3b9..ee502a99 100644 --- a/sage-client/zio/src/main/scala/sage/backend/SageClient.scala +++ b/sage-client/zio/src/main/scala/sage/backend/SageClient.scala @@ -9,7 +9,7 @@ import zio.stream.ZStream import sage.{Message, PatternMessage, SageException} import sage.client.SageConfig -import sage.client.internal.{Client, LoweredClient, Paged, ScanStep, ScanTarget, Subscription} +import sage.client.internal.{Client, LoweredClient, Paged, Subscription} import sage.codec.{KeyCodec, ValueCodec} import sage.commands.* @@ -36,7 +36,7 @@ extension [K](client: Client[IO[SageException, *], K])(using @unused ev: KeyCode count: Option[Long] = None, ofType: Option[RedisType] = None ): ZStream[Any, SageException, K] = - scanStreamAll(target => cursor => client.runOn(target, Keys.scan[K](cursor, pattern, count, ofType))) + paged(Paged.scanAll[K](client.runner, pattern, count, ofType)) /** * Iterates over all HSCAN field/value pairs until the server returns a zero cursor. An empty page with a non-zero cursor continues the scan. @@ -46,7 +46,7 @@ extension [K](client: Client[IO[SageException, *], K])(using @unused ev: KeyCode pattern: Option[String] = None, count: Option[Long] = None ): ZStream[Any, SageException, (F, V)] = - scanStream(cursor => client.run(Hashes.hScan[K, F, V](key, cursor, pattern, count))) + paged(Paged.scanKey(client.runner)(cursor => Hashes.hScan[K, F, V](key, cursor, pattern, count))) /** * Iterates over all SSCAN members until the server returns a zero cursor. An empty page with a non-zero cursor continues the scan. @@ -56,7 +56,7 @@ extension [K](client: Client[IO[SageException, *], K])(using @unused ev: KeyCode pattern: Option[String] = None, count: Option[Long] = None ): ZStream[Any, SageException, V] = - scanStream(cursor => client.run(Sets.sScan[K, V](key, cursor, pattern, count))) + paged(Paged.scanKey(client.runner)(cursor => Sets.sScan[K, V](key, cursor, pattern, count))) /** * Iterates over all ZSCAN member/score pairs until the server returns a zero cursor. An empty page with a non-zero cursor continues the scan. @@ -66,18 +66,11 @@ extension [K](client: Client[IO[SageException, *], K])(using @unused ev: KeyCode pattern: Option[String] = None, count: Option[Long] = None ): ZStream[Any, SageException, (V, Double)] = - scanStream(cursor => client.run(SortedSets.zScan[K, V](key, cursor, pattern, count))) + paged(Paged.scanKey(client.runner)(cursor => SortedSets.zScan[K, V](key, cursor, pattern, count))) // convert pages from the shared Paged helper into individual ZStream elements - private def paged[S, A](init: S)(step: Paged.Step[S, A]): ZStream[Any, SageException, A] = - CStream.unfold[S, Vector[A]](init)(step).flatMap(items => CStream.init(items)).lower.refineToOrDie[SageException] - - private def scanStream[A](fetch: ScanCursor => IO[SageException, ScanPage[A]]): ZStream[Any, SageException, A] = - paged[Option[ScanCursor], A](Some(ScanCursor.start))(Paged.byCursor(cursor => CIO.lift(fetch(cursor)))) - - // scan each target in sequence with its own node-local cursor. A cluster scan visits every master that owns slots. - private def scanStreamAll[A](fetch: ScanTarget => ScanCursor => IO[SageException, ScanPage[A]]): ZStream[Any, SageException, A] = - paged[ScanStep, A](ScanStep.Begin)(Paged.acrossTargets(CIO.lift(client.scanTargets))(target => cursor => CIO.lift(fetch(target)(cursor)))) + private def paged[S, A](pages: Paged.Pages[S, A]): ZStream[Any, SageException, A] = + CStream.unfold[S, Vector[A]](pages.init)(pages.step).flatMap(items => CStream.init(items)).lower.refineToOrDie[SageException] /** * Lazily pages an entire stream by range, batching `XRANGE` and advancing past the last id each page. Stops when a page comes back empty. @@ -88,9 +81,7 @@ extension [K](client: Client[IO[SageException, *], K])(using @unused ev: KeyCode end: StreamRangeId = StreamRangeId.Max, batch: Long = 100L ): ZStream[Any, SageException, StreamEntry[F, V]] = - paged[Option[StreamRangeId], StreamEntry[F, V]](Some(start))( - Paged.byRange(batch)(from => CIO.lift(client.run(Streams.xRange[K, F, V](key, from, end, Some(batch))))) - ) + paged(Paged.xRangeAll[K, F, V](client.runner, key, start, end, batch)) /** * Auto-claims idle pending entries for `consumer`, advancing the `XAUTOCLAIM` cursor until it returns to the start. Entries whose data @@ -104,9 +95,7 @@ extension [K](client: Client[IO[SageException, *], K])(using @unused ev: KeyCode start: StreamId = StreamId.Zero, count: Option[Long] = None ): ZStream[Any, SageException, StreamEntry[F, V]] = - paged[Option[StreamId], StreamEntry[F, V]](Some(start))( - Paged.byAutoClaim(from => CIO.lift(client.run(Streams.xAutoClaim[K, F, V](key, group, consumer, minIdle, from, count)))) - ) + paged(Paged.xAutoClaimAll[K, F, V](client.runner, key, group, consumer, minIdle, start, count)) /** * Follows a stream without a consumer group. It first reads every entry after `from`, then waits for new entries. The explicit entry ID @@ -119,11 +108,7 @@ extension [K](client: Client[IO[SageException, *], K])(using @unused ev: KeyCode count: Option[Long] = None, block: BlockTimeout = Paged.defaultPoll ): ZStream[Any, SageException, StreamEntry[F, V]] = - paged[StreamId, StreamEntry[F, V]](from)( - Paged.tail(last => - CIO.lift(client.run(Streams.xRead[K, F, V]((key, ReadId.After(last)))(count = count, block = Some(block)))).map(_.flatMap(_._2)) - ) - ) + paged(Paged.xTail[K, F, V](client.runner, key, from, count, block)) /** * Follows a stream as part of a consumer group. It processes this consumer's pending entries first, then waits for new entries. Each @@ -136,27 +121,10 @@ extension [K](client: Client[IO[SageException, *], K])(using @unused ev: KeyCode count: Option[Long] = None, block: BlockTimeout = Paged.defaultPoll )(handle: StreamEntry[F, V] => IO[SageException, Unit]): IO[SageException, Unit] = - consumeStream[F, V](group, consumer, key, count, block) + paged(Paged.xConsume[K, F, V](client.runner, group, consumer, key, count, block)) .mapZIO(entry => handle(entry) *> client.run(Streams.xAck(key, group)(entry.id)).unit) .runDrain - private def consumeStream[F: KeyCodec, V: ValueCodec]( - group: String, - consumer: String, - key: K, - count: Option[Long], - block: BlockTimeout - ): ZStream[Any, SageException, StreamEntry[F, V]] = - paged[Either[StreamId, Unit], StreamEntry[F, V]](Left(StreamId.Zero))( - Paged.consume( - drainPending = after => - CIO.lift(client.run(Streams.xReadGroup[K, F, V](group, consumer)((key, GroupReadId.After(after)))(count = count))).map(_.flatMap(_._2)), - tailNew = CIO - .lift(client.run(Streams.xReadGroup[K, F, V](group, consumer)((key, GroupReadId.New))(count = count, block = Some(block)))) - .map(_.flatMap(_._2)) - ) - ) - /** * Subscribes to one or more channels. Closing the stream's scope unsubscribes. Sage resubscribes after reconnecting, but messages * published while the connection is down are lost. @@ -256,6 +224,6 @@ object SageClient { timeout: FiniteDuration, replicaAcknowledgement: Boolean ): CIO[Boolean] = - CIO.lift(underlying.lockWrite(command, timeout, replicaAcknowledgement).lower.interruptible) + CIO.lift(underlying.runner.lockWrite(command, timeout, replicaAcknowledgement).lower.interruptible) } } diff --git a/sage-client/zio/src/test/scala/sage/client/ClusterClientSpec.scala b/sage-client/zio/src/test/scala/sage/client/ClusterClientSpec.scala index 8b282524..161e5078 100644 --- a/sage-client/zio/src/test/scala/sage/client/ClusterClientSpec.scala +++ b/sage-client/zio/src/test/scala/sage/client/ClusterClientSpec.scala @@ -6,21 +6,23 @@ import scala.concurrent.duration.* import kyo.compat.* -import sage.Bytes +import sage.{Bytes, SageEvent, SageListener} import sage.SageException.{ConnectionLost, CrossSlot, DecodeError, InvalidArgument, NotConnected, ServerError, TimedOut, UnsupportedServer} import sage.client.internal.{ ClusterLive, CountingScheduler, + Events, FakeTransport, ManualScheduler, MultiplexedConnection, + RecordingTracer, Replies, Scheduler, StaggeringScheduler } import sage.client.internal.Replies.bulk import sage.cluster.{Node, Slot} -import sage.commands.{BroadcastReduce, Command, Connection, Json, JsonPath, Keys, Scripting, Server, Strings} +import sage.commands.{BlockTimeout, BroadcastReduce, Command, Connection, Json, JsonPath, Keys, Lists, Scripting, Server, Sets, Strings} import sage.protocol.Frame class ClusterClientSpec extends munit.FunSuite { @@ -46,7 +48,9 @@ class ClusterClientSpec extends munit.FunSuite { connectGate: (Node, Int) => Unit = (_, _) => (), // blocks a node's nth transport while it is being opened scheduler: Scheduler = Scheduler.real, caching: Boolean = false, - cluster: ClusterConfig = ClusterConfig() + cluster: ClusterConfig = ClusterConfig(), + events: Events = Events.disabled, + onTransport: FakeTransport => Unit = _ => () ) { // Accumulate every transport per node (a node has both a Multiplexed and, once a transaction pins, a Dedicated connection) so a refresh @@ -56,6 +60,9 @@ class ClusterClientSpec extends munit.FunSuite { // make the next HELLO fail once to verify that subscription reassignment retries after a transient connection error. val flakyHello = mutable.Set.empty[Node] + // nodes whose HELLO fails until removed, as a crashed master does until failover replaces it + val down = java.util.concurrent.ConcurrentHashMap.newKeySet[Node]() + private def transportsOf(node: Node): Vector[FakeTransport] = transports.synchronized(transports.get(node).map(_.toVector).getOrElse(Vector.empty)) @@ -67,12 +74,14 @@ class ClusterClientSpec extends munit.FunSuite { val respond: Bytes => Seq[Frame] = payload => { val text = payload.asUtf8String if (text.contains("HELLO")) - if (unreachable(node) || flakyHello.remove(node)) Seq(Frame.SimpleError("ERR node is down")) else Seq(Replies.hello) - else if (text.contains("TRACKING")) Seq(Frame.SimpleString("OK")) + if (unreachable(node) || down.contains(node) || flakyHello.remove(node)) Seq(Frame.SimpleError("ERR node is down")) + else Seq(Replies.hello) + else if (text.contains("TRACKING") || Replies.isSetup(payload)) Seq(Frame.SimpleString("OK")) else behaviour(node, text) } val transport = new FakeTransport(onFrame, onClosed, respond) transports.synchronized(transports.getOrElseUpdate(node, mutable.ArrayBuffer.empty) += transport) + onTransport(transport) transport } @@ -80,30 +89,29 @@ class ClusterClientSpec extends munit.FunSuite { new ClusterLive( factory, scheduler, - Vector(Connection.hello(None)), - BackoffConfig(), - WatchdogConfig(enabled = false), - 1.second, - Duration.Zero, - DedicatedPoolConfig(), + SageConfig( + watchdog = WatchdogConfig(enabled = false), + connectTimeout = 1.second, + closeTimeout = Duration.Zero, + pubsub = PubSubConfig(bufferSize = 1024), + readFrom = readFrom, + clientCache = CacheConfig(enabled = caching, maxBytes = 1L << 20) + ), cluster, - 1024, seeds, - readFrom, - cachingEnabled = caching, - cacheMaxBytes = if (caching) 1L << 20 else 0L + events ) - live.bootstrapTopology() + live.start() def written(node: Node): Vector[String] = transportsOf(node).flatMap(_.written.map(_.asUtf8String)) def clusterSlotsCount(node: Node): Int = written(node).count(_.contains("CLUSTER")) - // Simulate the disconnect after a slot migration by closing the connection used to send SSUBSCRIBE. The manager must then assign the - // subscription again. - def dropShardConn(node: Node): Unit = - transportsOf(node).find(_.written.exists(_.asUtf8String.contains("SSUBSCRIBE"))).foreach(_.close()) + // the connection a node's SSUBSCRIBE was written to + def shardConn(node: Node): Option[FakeTransport] = transportsOf(node).findLast(_.written.exists(_.asUtf8String.contains("SSUBSCRIBE"))) + + def dropShardConn(node: Node): Unit = shardConn(node).foreach(_.close()) // Close the master's classic subscription connection so the manager assigns those subscriptions again. def dropClassicConn(): Unit = @@ -165,6 +173,63 @@ class ClusterClientSpec extends munit.FunSuite { .andThen { case _ => fixture.live.close.unsafeRun } } + test("an interrupted discovery closes the seed connection it opened") { + val opened = new java.util.concurrent.ConcurrentLinkedQueue[FakeTransport]() + @volatile var failure: Throwable = null + // CLUSTER SLOTS is never answered; the discovering thread is interrupted instead + Thread + .ofVirtual() + .start { () => + try + new Fixture( + (_, cmd) => { if (cmd.contains("CLUSTER")) Thread.currentThread().interrupt(); Nil }, + Vector(nodeA), + onTransport = opened.add(_): Unit + ): Unit + catch { case e: Throwable => failure = e } + } + .join() + assert(failure.isInstanceOf[InterruptedException], String.valueOf(failure)) + assertEquals(opened.size, 1) + assertEquals(opened.peek().closeCount, 1) + } + + test("cancelling a blocking command requests a topology refresh") { + val fixture = new Fixture( + (_, text) => if (text.contains("CLUSTER")) Seq(wholeClusterOn(nodeA)) else Nil, + Vector(nodeA), + cluster = ClusterConfig(minRefreshInterval = 10.millis) + ) + val cancel = for { + fiber <- fixture.live.run(Lists.blPop[String, String]("k")(BlockTimeout.Forever)).lower.fork + _ <- zio.ZIO.attemptBlocking(awaitWritten(fixture, nodeA, "BLPOP")) + _ <- fiber.interrupt + } yield () + CIO + .lift(cancel) + .unsafeRun + .map(_ => awaitRefreshed(fixture, nodeA)) + .andThen { case _ => fixture.live.close.unsafeRun } + } + + test("a lock write that times out requests a topology refresh") { + val fixture = new Fixture( + (_, text) => if (text.contains("CLUSTER")) Seq(wholeClusterOn(nodeA)) else Nil, + Vector(nodeA), + cluster = ClusterConfig(minRefreshInterval = 10.millis) + ) + val stalled = Command[Boolean]("STALL", Vector(0), Vector(Bytes.utf8("k")), _ => Right(true)) + fixture.live + .lockWrite(stalled, 100.millis, replicaAcknowledgement = false) + .liftToTry + .unsafeRun + .map { result => + assert(result.failed.get.isInstanceOf[TimedOut], result.toString) + awaitRefreshed(fixture, nodeA) + } + .andThen { case _ => fixture.live.close.unsafeRun } + } + test("cluster locks retry after the per-command redirect limit is exhausted") { val key = "lock-key" val slot = Slot.of(Bytes.utf8(s"4:lock:$key")).value @@ -381,6 +446,107 @@ class ClusterClientSpec extends munit.FunSuite { } } + test("a cached read that follows an ASK has one span, routed to the importing node") { + val slot = Slot.of(Bytes.utf8("foo")).value + val behaviour = (node: Node, text: String) => + if (text.contains("CLUSTER")) Seq(wholeClusterOn(nodeA)) + else if (node == nodeB && text.contains("ASKING")) Seq(Frame.SimpleString("OK"), bulk("v")) + else if (text.contains("GET")) Seq(Frame.SimpleString("OK"), Frame.SimpleError(s"ASK $slot b:6379")) + else Seq(Frame.Null) + val tracer = new RecordingTracer + val fixture = new Fixture(behaviour, Vector(nodeA), caching = true, events = Events(Vector.empty, Some(tracer))) + + fixture.live.cached(Strings.get[String, String]("foo"), 1.minute).unsafeRun.map { result => + assertEquals(result, Some("v")) + val log = tracer.log.synchronized(tracer.log.toVector) + assertEquals(log.count(_ == "start:GET"), 1, log.toString) + assert(log.contains(s"routed:${nodeB.host}:${nodeB.port}") && log.contains("settled:Succeeded"), log.toString) + } + } + + test("a cached miss redirected by MOVED has one span covering both attempts") { + val behaviour = (node: Node, text: String) => + if (text.contains("CLUSTER")) Seq(wholeClusterOn(nodeA)) + else if (node == nodeB && text.contains("GET")) Seq(Frame.SimpleString("OK"), bulk("v")) + else if (text.contains("GET")) Seq(Frame.SimpleString("OK"), Frame.SimpleError("MOVED 0 b:6379")) + else Seq(Frame.Null) + val tracer = new RecordingTracer + val fixture = new Fixture(behaviour, Vector(nodeA), caching = true, events = Events(Vector.empty, Some(tracer))) + + fixture.live.cached(Strings.get[String, String]("foo"), 1.minute).unsafeRun.map { result => + assertEquals(result, Some("v")) + val log = tracer.log.synchronized(tracer.log.toVector) + assertEquals(log.count(_ == "start:GET"), 1, log.toString) + assertEquals(log.count(_.startsWith("settled:")), 1, log.toString) + assert(log.contains("settled:Succeeded"), log.toString) + } + } + + test("a cached cross-slot read settles a failed span, as an uncached one does") { + val tracer = new RecordingTracer + val fixture = new Fixture((_, _) => Seq(wholeClusterOn(nodeA)), Vector(nodeA), caching = true, events = Events(Vector.empty, Some(tracer))) + fixture.live.cached(Sets.sInter[String, String]("{a}x", "{b}y"), 1.minute).unsafeRun.failed.map { error => + assert(error.isInstanceOf[CrossSlot], error.toString) + assertEquals(tracer.log.synchronized(tracer.log.toVector), Vector("start:SINTER", s"settled:Failed($error)")) + } + } + + test("a cached cross-slot MGET miss split across two owners has one span, as an unsplit miss does") { + val behaviour = (node: Node, text: String) => + if (text.contains("CLUSTER")) Seq(splitOn(Slot.of(Bytes.utf8("{b}")).value)) + else if (text.contains("MGET")) Seq(Frame.SimpleString("OK"), Frame.Array(Vector(bulk(node.host)))) + else Seq(Frame.Null) + val tracer = new RecordingTracer + val fixture = new Fixture(behaviour, Vector(nodeA), caching = true, events = Events(Vector.empty, Some(tracer))) + + fixture.live.cached(Strings.mGet[String, String]("{a}", "{b}"), 1.minute).unsafeRun.map { result => + assertEquals(result, Vector(Some(nodeA.host), Some(nodeB.host))) + val log = tracer.log.synchronized(tracer.log.toVector) + assertEquals(log.count(_ == "start:MGET"), 1, log.toString) + assertEquals(log.filter(_.startsWith("settled:")), Vector("settled:Succeeded"), log.toString) + } + } + + test("a cached read on a closed cluster client settles a failed span, as standalone and master-replica do") { + val tracer = new RecordingTracer + val fixture = new Fixture((_, _) => Seq(wholeClusterOn(nodeA)), Vector(nodeA), caching = true, events = Events(Vector.empty, Some(tracer))) + fixture.live.close.unsafeRun.flatMap { _ => + fixture.live.cached(Strings.get[String, String]("foo"), 1.minute).unsafeRun.failed.map { error => + assert(error.isInstanceOf[NotConnected], error.toString) + assertEquals(tracer.log.synchronized(tracer.log.toVector), Vector("start:GET", s"settled:Failed($error)")) + } + } + } + + test("a cached hit that fails to decode reports no CommandCompleted, as standalone does") { + val behaviour = (_: Node, text: String) => + if (text.contains("CLUSTER")) Seq(wholeClusterOn(nodeA)) + else if (text.contains("GET")) Seq(Frame.SimpleString("OK"), bulk("v")) + else Seq(Frame.Null) + val completed = new java.util.concurrent.LinkedBlockingQueue[String]() + val listener = new SageListener { + def onEvent(event: SageEvent): Unit = event match { + case SageEvent.CommandCompleted(name, _, _, outcome) => completed.add(s"$name:$outcome"): Unit + case _ => () + } + } + val fixture = new Fixture(behaviour, Vector(nodeA), caching = true, events = Events(Vector(listener))) + val get = Strings.get[String, String]("foo") + val undecoded = get.copy(decode = _ => Left(DecodeError("a number", "a string"))) + val afterwards = Strings.get[String, String]("bar") + + for { + _ <- fixture.live.cached(get, 1.minute).unsafeRun + error <- fixture.live.cached(undecoded, 1.minute).unsafeRun.failed + _ <- fixture.live.cached(afterwards, 1.minute).unsafeRun + } yield { + assert(error.isInstanceOf[DecodeError], error.toString) + // events arrive in order, so the miss of `afterwards` comes after any event of the failed hit + val seen = Iterator.continually(completed.poll(2, java.util.concurrent.TimeUnit.SECONDS)).take(2).toVector + assertEquals(seen, Vector("GET:Succeeded", "GET:Succeeded")) + } + } + test("a MOVED retires the source master's whole cache, so a previously cached key can no longer serve a stale local hit") { val moved = new java.util.concurrent.atomic.AtomicBoolean(false) val behaviour = (node: Node, text: String) => @@ -783,6 +949,37 @@ class ClusterClientSpec extends munit.FunSuite { } } + private def masterRefusesFirstGet(replica: Node): (Node, String) => Seq[Frame] = { + val refused = new java.util.concurrent.atomic.AtomicBoolean(false) + (node: Node, text: String) => + if (text.contains("CLUSTER")) Seq(Replies.clusterShard(nodeA, replica)) + else if (node == replica && text.contains("GET")) Seq(bulk("from-replica")) + else if (text.contains("GET")) { + val get = if (refused.compareAndSet(false, true)) Frame.SimpleError("TRYAGAIN rehashing") else bulk("from-master") + if (text.contains("SET")) Seq(Frame.SimpleString("OK"), get) else Seq(get) + } else Seq(Frame.SimpleString("OK")) + } + + test("an uncached cluster cached() read retried after TRYAGAIN stays on the master") { + val nodeR = Node("r", 6379) + val fixture = new Fixture(masterRefusesFirstGet(nodeR), Vector(nodeA), readFrom = ReadFrom.ReplicaPreferred) + + fixture.live.cached(Strings.get[String, String]("foo"), 1.minute).unsafeRun.map { result => + assertEquals(result, Some("from-master")) + assert(!fixture.written(nodeR).exists(_.contains("GET")), "a cached read must never reach a replica") + } + } + + test("a mixed pipeline position retried after TRYAGAIN stays on the master") { + val nodeR = Node("r", 6379) + val fixture = new Fixture(masterRefusesFirstGet(nodeR), Vector(nodeA), readFrom = ReadFrom.ReplicaPreferred) + + fixture.live.pipeline((Strings.set("foo", "v"), Strings.get[String, String]("foo"))).unsafeRun.map { case (_, read) => + assertEquals(read, Some("from-master")) + assert(!fixture.written(nodeR).exists(_.contains("GET")), "a read after a write in the same pipeline must stay on the master") + } + } + /** * A shard whose replica's socket dies with the read already written, so its reply fails as `ConnectionLost(mayHaveExecuted = true)`. */ @@ -840,6 +1037,27 @@ class ClusterClientSpec extends munit.FunSuite { } } + test("under a strict Replica policy a replica answering ASK fails the read and its pipeline instead of reaching the importing master") { + val nodeR = Node("r", 6379) + val ask = Frame.SimpleError(s"ASK ${Slot.of(Bytes.utf8("foo")).value} ${nodeB.host}:${nodeB.port}") + val behaviour = (node: Node, text: String) => + if (text.contains("CLUSTER")) Seq(Replies.clusterShard(nodeA, nodeR)) + else if (text.contains("READONLY")) Seq(Frame.SimpleString("OK")) + else if (node == nodeR) Seq.fill(text.split("\r\nGET\r\n", -1).length - 1)(ask) + else Seq(bulk("from-master")) + val fixture = new Fixture(behaviour, Vector(nodeA), readFrom = ReadFrom.Replica) + val get = Strings.get[String, String]("foo") + + fixture.live.run(get).unsafeRun.failed.flatMap { error => + assert(error.isInstanceOf[NotConnected], s"expected NotConnected, got $error") + fixture.live.pipeline((get, get)).unsafeRun.failed.map { error => + assert(error.isInstanceOf[NotConnected], s"expected NotConnected, got $error") + assert(!fixture.written(nodeB).exists(_.contains("GET")), "a strict Replica read must not follow ASK to a master") + assert(!fixture.written(nodeA).exists(_.contains("GET")), "a strict Replica read must not fall back to the master") + } + } + } + test("under a Replica policy an eligible read routes to the shard's replica, which gets READONLY at setup") { val nodeR = Node("r", 6379) // CLUSTER SLOTS lists nodeR as nodeA's replica for the whole keyspace @@ -1089,6 +1307,21 @@ class ClusterClientSpec extends munit.FunSuite { } } + test("a command that exhausts its TRYAGAIN retries is routed once, to the node that refused it") { + val behaviour = (_: Node, text: String) => + if (text.contains("CLUSTER")) Seq(wholeClusterOn(nodeA)) + else Seq(Frame.SimpleError("TRYAGAIN Multiple keys request during rehashing of slot")) + val tracer = new RecordingTracer + val fixture = new Fixture(behaviour, Vector(nodeA), events = Events(Vector.empty, Some(tracer))) + + fixture.live.run(Strings.get[String, String]("foo")).unsafeRun.failed.map { error => + assertEquals( + tracer.log.synchronized(tracer.log.toVector), + Vector("start:GET", s"routed:${nodeA.host}:${nodeA.port}", s"settled:Failed($error)") + ) + } + } + test("an unsupported multi-key command whose keys span slots fails CrossSlot") { val behaviour = (_: Node, text: String) => if (text.contains("CLUSTER")) Seq(wholeClusterOn(nodeA)) else Seq(Frame.Integer(1)) val fixture = new Fixture(behaviour, Vector(nodeA)) @@ -1257,6 +1490,26 @@ class ClusterClientSpec extends munit.FunSuite { } } + test("a command that ran on several nodes reports no node: a cross-slot MGET and a broadcast leave the span unrouted") { + val keyA = "{a}" + val keyB = "{b}" + val behaviour = (node: Node, text: String) => + if (text.contains("CLUSTER")) Seq(splitOn(Slot.of(Bytes.utf8(keyB)).value)) + else if (text.contains("DBSIZE")) Seq(Frame.Integer(1)) + else Seq(Frame.Array(Vector(bulk(node.host)))) + val tracer = new RecordingTracer + val fixture = new Fixture(behaviour, Vector(nodeA), events = Events(Vector.empty, Some(tracer))) + + for { + _ <- fixture.live.run(Strings.mGet[String, String](keyA, keyB)).unsafeRun + _ <- fixture.live.run(Server.dbSize).unsafeRun + } yield { + val log = tracer.log.synchronized(tracer.log.toVector) + assertEquals(log.count(_.startsWith("start:")), 2, log.toString) + assert(!log.exists(_.startsWith("routed:")), log.toString) + } + } + test("a cross-slot MGET fails as a whole when one slot group fails") { val keyA = "{a}" val keyB = "{b}" @@ -1762,6 +2015,18 @@ class ClusterClientSpec extends munit.FunSuite { } } + test("a sharded subscribe the owner rejects fails with the server's error instead of retrying") { + val behaviour = (_: Node, text: String) => + if (text.contains("CLUSTER")) Seq(splitOn(slotB)) + else if (text.contains("SSUBSCRIBE")) Seq(Frame.SimpleError("NOPERM no permissions to access the channel")) + else Seq(Frame.SimpleString("OK")) + val fixture = new Fixture(behaviour, Vector(nodeA)) + + fixture.live.subscribeShardChannels[String](keyB).unsafeRun.failed.map { error => + assertEquals(Option(error).collect { case ServerError(code, _) => code }, Some("NOPERM")) + } + } + test("sPublish routes by slot to the channel's owner") { val behaviour = (_: Node, text: String) => if (text.contains("CLUSTER")) Seq(splitOn(slotB)) @@ -1789,7 +2054,7 @@ class ClusterClientSpec extends munit.FunSuite { } } - test("a sharded subscription re-homes to the new owner after its connection drops") { + test("a sharded subscription re-homes to the new owner when the server drops its channel after a slot migration") { @volatile var migrated = false val behaviour = (_: Node, text: String) => if (text.contains("CLUSTER")) Seq(if (migrated) wholeClusterOn(nodeA) else splitOn(slotB)) @@ -1799,20 +2064,241 @@ class ClusterClientSpec extends munit.FunSuite { fixture.live.subscribeShardChannels[String](keyB).unsafeRun.map { _ => assert(fixture.written(nodeB).exists(_.contains("SSUBSCRIBE")), "initial subscribe did not reach the owner nodeB") - // slotB migrates to nodeA and nodeB drops the subscriber connection (the server's post-migration disconnect) + // slotB migrates to nodeA, and nodeB unsubscribes the channel but keeps the connection open + migrated = true + fixture.shardConn(nodeB).foreach(_.emit(Frame.Push(Vector(bulk("sunsubscribe"), bulk(keyB), Frame.Integer(0))))) + awaitWritten(fixture, nodeA, "SSUBSCRIBE") + } + } + + test("a sharded subscription re-homes to the new owner at once when its connection drops while the old owner is still reachable") { + val scheduler = new ManualScheduler + @volatile var migrated = false + val behaviour = (_: Node, text: String) => + if (text.contains("CLUSTER")) Seq(if (migrated) wholeClusterOn(nodeA) else splitOn(slotB)) + else if (text.contains("SSUBSCRIBE")) Seq(subscribed("ssubscribe", keyB)) + else Seq(Frame.SimpleString("OK")) + val fixture = new Fixture(behaviour, Vector(nodeA), scheduler = scheduler) + + fixture.live.subscribeShardChannels[String](keyB).unsafeRun.flatMap { sub => + assert(fixture.written(nodeB).exists(_.contains("SSUBSCRIBE")), "initial subscribe did not reach the owner nodeB") migrated = true fixture.dropShardConn(nodeB) - Thread.sleep(300) // re-homing is offloaded: force a refresh, reconcile, and re-SSUBSCRIBE on the new owner + scheduler.advance(Duration.Zero) assert(fixture.written(nodeA).exists(_.contains("SSUBSCRIBE")), "subscription did not re-home to the new owner nodeA") + fixture.shardConn(nodeA).foreach(_.emit(Frame.Push(Vector(bulk("smessage"), bulk(keyB), bulk("after-rehome"))))) + CIO.timeout(5.seconds)(sub.next).unsafeRun.map(message => assertEquals(message, Some(Some(sage.Message(keyB, "after-rehome"))))) + } + } + + test("a sharded subscription the new owner refuses after a slot migration ends with the server's error instead of retrying silently") { + @volatile var migrated = false + val behaviour = (node: Node, text: String) => + if (text.contains("CLUSTER")) Seq(if (migrated) wholeClusterOn(nodeA) else splitOn(slotB)) + else if (text.contains("SSUBSCRIBE")) + Seq(if (node == nodeA) Frame.SimpleError("NOPERM no permissions to access the channel") else subscribed("ssubscribe", keyB)) + else Seq(Frame.SimpleString("OK")) + val fixture = new Fixture(behaviour, Vector(nodeA)) + + fixture.live.subscribeShardChannels[String](keyB).unsafeRun.flatMap { sub => + migrated = true + fixture.shardConn(nodeB).foreach(_.emit(Frame.Push(Vector(bulk("sunsubscribe"), bulk(keyB), Frame.Integer(0))))) + sub.next.unsafeRun.failed.map(error => assertEquals(Option(error).collect { case ServerError(code, _) => code }, Some("NOPERM"))) + } + } + + test("a sharded subscribe redirected after its confirmation wait ended is placed on the new owner") { + @volatile var migrated = false + val behaviour = (node: Node, text: String) => + if (text.contains("CLUSTER")) Seq(if (migrated) wholeClusterOn(nodeA) else splitOn(slotB)) + else if (text.contains("SSUBSCRIBE")) if (node == nodeA) Seq(subscribed("ssubscribe", keyB)) else Nil + else Seq(Frame.SimpleString("OK")) + val fixture = new Fixture(behaviour, Vector(nodeA)) + + fixture.live.subscribeShardChannels[String](keyB).unsafeRun.map { _ => + migrated = true + fixture.shardConn(nodeB).foreach(_.emit(Frame.SimpleError(s"MOVED ${Slot.of(Bytes.utf8(keyB)).value} ${nodeA.host}:${nodeA.port}"))) + awaitWritten(fixture, nodeA, "SSUBSCRIBE") + } + } + + test("shard subscriptions waiting for an unreachable owner share one retry, which refreshes at most once per 200ms") { + val scheduler = new ManualScheduler + val behaviour = (_: Node, text: String) => + if (text.contains("CLUSTER")) Seq(splitOn(slotB)) + else if (text.contains("SSUBSCRIBE")) Seq(subscribed("ssubscribe", keyB)) + else Seq(Frame.SimpleString("OK")) + val fixture = new Fixture(behaviour, Vector(nodeA), unreachable = Set(nodeB), scheduler = scheduler) + val subscribe = fixture.live.subscribeShardChannels[String](keyB) + + subscribe.unsafeRun.flatMap(_ => subscribe.unsafeRun).flatMap(_ => subscribe.unsafeRun).map { _ => + val before = fixture.clusterSlotsCount(nodeA) + // one retry chain retries at 50, 150, 350, 550, 750 and 950ms, and refreshes at most once per 200ms, so not at 150ms + scheduler.advance(1.second) + assertEquals(fixture.clusterSlotsCount(nodeA), before + 5) + } + } + + test("a shard connection the server closes right after each subscribe backs off further on each loss") { + val scheduler = new ManualScheduler + val behaviour = (_: Node, text: String) => + if (text.contains("CLUSTER")) Seq(splitOn(slotB)) + else if (text.contains("SSUBSCRIBE")) Seq(subscribed("ssubscribe", keyB)) + else Seq(Frame.SimpleString("OK")) + val fixture = new Fixture(behaviour, Vector(nodeA), scheduler = scheduler) + + def subscribes = fixture.written(nodeB).count(_.contains("SSUBSCRIBE")) + + fixture.live.subscribeShardChannels[String](keyB).unsafeRun.map { _ => + fixture.dropShardConn(nodeB) + scheduler.advance(Duration.Zero) + assertEquals(subscribes, 2, "the first loss must re-home at once") + fixture.dropShardConn(nodeB) + scheduler.advance(99.millis) + assertEquals(subscribes, 2, "the second loss must wait longer than the first") + scheduler.advance(1.milli) + assertEquals(subscribes, 3) + fixture.dropShardConn(nodeB) + scheduler.advance(199.millis) + assertEquals(subscribes, 3, "the third loss must wait longer than the second") + scheduler.advance(1.milli) + assertEquals(subscribes, 4) + val refreshes = fixture.clusterSlotsCount(nodeA) + fixture.dropShardConn(nodeB) + scheduler.advance(199.millis) + assertEquals(subscribes, 4, "the fourth loss waits as long as the third") + scheduler.advance(1.milli) + assertEquals(subscribes, 5, "the delay stops growing at four initial delays") + assertEquals(fixture.clusterSlotsCount(nodeA), refreshes + 1, "a retry refreshes at most once per 200ms") + } + } + + test("a sharded subscribe whose slot has no owner yet refreshes the topology and places the channel at once") { + val scheduler = new ManualScheduler + @volatile var owned = false + val behaviour = (_: Node, text: String) => + if (text.contains("CLUSTER")) + Seq(if (owned) splitOn(slotB) else Replies.clusterSlots((nodeA, 0, slotB - 1), (nodeA, slotB + 1, Slot.Count - 1))) + else if (text.contains("SSUBSCRIBE")) Seq(subscribed("ssubscribe", keyB)) + else Seq(Frame.SimpleString("OK")) + val fixture = new Fixture(behaviour, Vector(nodeA), scheduler = scheduler) + owned = true + + fixture.live.subscribeShardChannels[String](keyB).unsafeRun.map { _ => + assert(fixture.written(nodeB).exists(_.contains("SSUBSCRIBE")), "the subscribe waited for a retry instead of refreshing first") + } + } + + test("a sharded subscription resumes within four initial delays when failover names a new owner after a long outage") { + val scheduler = new ManualScheduler + @volatile var promoted = false + val behaviour = (_: Node, text: String) => + if (text.contains("CLUSTER")) Seq(if (promoted) wholeClusterOn(nodeA) else splitOn(slotB)) + else if (text.contains("SSUBSCRIBE")) Seq(subscribed("ssubscribe", keyB)) + else Seq(Frame.SimpleString("OK")) + val fixture = new Fixture(behaviour, Vector(nodeA), scheduler = scheduler) + + fixture.live.subscribeShardChannels[String](keyB).unsafeRun.map { _ => + fixture.down.add(nodeB) + fixture.dropShardConn(nodeB) + scheduler.advance(30.seconds) + promoted = true + scheduler.advance(200.millis) + assert(fixture.written(nodeA).exists(_.contains("SSUBSCRIBE")), "the subscription did not resume on the new owner within 200ms") + } + } + + test("a classic subscription reconnects at once after its connection drops") { + val scheduler = new ManualScheduler + val behaviour = (_: Node, text: String) => + if (text.contains("CLUSTER")) Seq(wholeClusterOn(nodeA)) + else if (text.contains("SUBSCRIBE")) Seq(subscribed("subscribe", "news")) + else Seq(Frame.SimpleString("OK")) + val fixture = new Fixture(behaviour, Vector(nodeA), scheduler = scheduler) + def classicSubscribes = fixture.written(nodeA).count(_.contains("\r\nSUBSCRIBE\r\n")) + + fixture.live.subscribeChannels[String]("news").unsafeRun.map { _ => + fixture.dropClassicConn() + scheduler.advance(Duration.Zero) + assertEquals(classicSubscribes, 2, "the first reconnect waited") + } + } + + test("a classic subscription resumes within four initial delays when its master answers again after a long outage") { + val scheduler = new ManualScheduler + val behaviour = (_: Node, text: String) => + if (text.contains("CLUSTER")) Seq(wholeClusterOn(nodeA)) + else if (text.contains("SUBSCRIBE")) Seq(subscribed("subscribe", "news")) + else Seq(Frame.SimpleString("OK")) + val fixture = new Fixture(behaviour, Vector(nodeA), scheduler = scheduler) + def classicSubscribes = fixture.written(nodeA).count(_.contains("\r\nSUBSCRIBE\r\n")) + + fixture.live.subscribeChannels[String]("news").unsafeRun.map { _ => + fixture.down.add(nodeA) + fixture.dropClassicConn() + scheduler.advance(30.seconds) + fixture.down.remove(nodeA) + scheduler.advance(200.millis) + assertEquals(classicSubscribes, 2, "the subscription did not resume within 200ms") + } + } + + test("a sharded subscription whose owner is down keeps refreshing until failover names a new owner") { + @volatile var promoted = false + val behaviour = (_: Node, text: String) => + if (text.contains("CLUSTER")) Seq(if (promoted) wholeClusterOn(nodeA) else splitOn(slotB)) + else if (text.contains("SSUBSCRIBE")) Seq(subscribed("ssubscribe", keyB)) + else Seq(Frame.SimpleString("OK")) + val fixture = new Fixture(behaviour, Vector(nodeA)) + + fixture.live.subscribeShardChannels[String](keyB).unsafeRun.map { _ => + assert(fixture.written(nodeB).exists(_.contains("SSUBSCRIBE")), "initial subscribe did not reach the owner nodeB") + // nodeB crashes; until failover completes, CLUSTER SLOTS still names it as the owner + fixture.down.add(nodeB) + fixture.dropShardConn(nodeB) + await("the subscription did not retry its unreachable owner")(fixture.written(nodeB).count(_.contains("HELLO")) >= 3) + promoted = true + awaitWritten(fixture, nodeA, "SSUBSCRIBE") } } - test("a classic subscription recovers when a re-home's establish fails, rather than stranding") { + test("a classic subscription moves to another master when its node leaves the cluster but keeps running") { + val scheduler = new ManualScheduler + @volatile var left = false + val behaviour = (_: Node, text: String) => + if (text.contains("CLUSTER")) Seq(if (left) wholeClusterOn(nodeB) else splitOn(slotB)) + else if (text.contains("SUBSCRIBE")) Seq(subscribed("subscribe", "news")) + else Seq(Frame.SimpleString("OK")) + val fixture = + new Fixture(behaviour, Vector(nodeA), scheduler = scheduler, cluster = ClusterConfig(topologyRefreshInterval = Some(1.minute))) + def subscribes(node: Node) = fixture.written(node).count(_.contains("\r\nSUBSCRIBE\r\n")) + + fixture.live.subscribeChannels[String]("news").unsafeRun.map { _ => + assertEquals(subscribes(nodeA), 1, "the classic subscription did not start on nodeA") + // nodeA hands its slots to nodeB and is removed (redis-cli --cluster del-node resets it), so PUBLISH no longer reaches it, but its + // connections stay open + left = true + scheduler.advance(1.minute) + await("the classic subscription stayed on the node that left the cluster") { + scheduler.advance(1.second) + subscribes(nodeB) >= 1 + } + } + } + + test("a classic subscription recovers when a re-home's establish fails, rather than stranding, and reports the failure with its node") { val behaviour = (_: Node, text: String) => if (text.contains("CLUSTER")) Seq(wholeClusterOn(nodeA)) else if (text.contains("SUBSCRIBE")) Seq(subscribed("subscribe", "news")) else Seq(Frame.SimpleString("OK")) - val fixture = new Fixture(behaviour, Vector(nodeA)) + val failures = new java.util.concurrent.atomic.AtomicInteger + val listener = new SageListener { + def onEvent(event: SageEvent): Unit = event match { + case SageEvent.Connection.ReconnectFailed(Some(`nodeA`), _) => failures.incrementAndGet(): Unit + case _ => () + } + } + val fixture = new Fixture(behaviour, Vector(nodeA), events = Events(Vector(listener))) def classicSubscribes = fixture.written(nodeA).count(_.contains("\r\nSUBSCRIBE\r\n")) @@ -1822,8 +2308,9 @@ class ClusterClientSpec extends munit.FunSuite { // Earlier behavior ignored this failure and left the subscription inactive. The manager must retry until it attaches again. fixture.flakyHello += nodeA fixture.dropClassicConn() - Thread.sleep(300) // the failed establish retries after the 50ms backoff, once HELLO succeeds again - assert(classicSubscribes >= 2, "classic subscription did not recover after the failed re-home") + // the failed establish retries after a backoff, once HELLO succeeds again + await("classic subscription did not recover after the failed re-home")(classicSubscribes >= 2) + await("the failed re-home emitted no ReconnectFailed naming its node")(failures.get() >= 1) } } } diff --git a/sage-client/zio/src/test/scala/sage/client/ConnectSpec.scala b/sage-client/zio/src/test/scala/sage/client/ConnectSpec.scala index f485fafd..9febcdde 100644 --- a/sage-client/zio/src/test/scala/sage/client/ConnectSpec.scala +++ b/sage-client/zio/src/test/scala/sage/client/ConnectSpec.scala @@ -45,6 +45,14 @@ class ConnectSpec extends munit.FunSuite { } } + test("a server that accepts the socket but never answers HELLO fails connect with ConnectionFailed") { + val (factory, transport) = ScriptedTransport(_ => Nil) + Client.connectWith(factory, config = SageConfig(connectTimeout = 10.millis)).unsafeRun.failed.map { error => + assert(error.isInstanceOf[ConnectionFailed], s"expected ConnectionFailed, got $error") + assertEquals(transport().closeCount, 1) + } + } + test("a server without RESP3 is rejected with UnsupportedServer and the connection is released") { val (factory, transport) = ScriptedTransport(_ => Seq(Frame.SimpleError("ERR unknown command 'HELLO'"))) Client.connectWith(factory).unsafeRun.failed.map { error => diff --git a/sage-client/zio/src/test/scala/sage/client/MasterReplicaPipelineSpec.scala b/sage-client/zio/src/test/scala/sage/client/MasterReplicaPipelineSpec.scala index af5b0814..9585a920 100644 --- a/sage-client/zio/src/test/scala/sage/client/MasterReplicaPipelineSpec.scala +++ b/sage-client/zio/src/test/scala/sage/client/MasterReplicaPipelineSpec.scala @@ -13,7 +13,7 @@ import sage.{Bytes, CommandSpan, CommandTracer, Outcome, SageEvent, SageListener import sage.client.internal.{CountingScheduler, Events, FakeTransport, MasterReplicaLive, MultiplexedConnection, Replies, Scheduler} import sage.client.internal.Replies.ok import sage.cluster.Node -import sage.commands.{Command, Connection} +import sage.commands.Command import sage.protocol.Frame class MasterReplicaPipelineSpec extends munit.FunSuite { @@ -52,18 +52,21 @@ class MasterReplicaPipelineSpec extends munit.FunSuite { val readBatches = new ConcurrentLinkedQueue[(Node, Int)]() val readFailure = new AtomicReference(failure) val readReplies = new AtomicReference[Option[(Node, Vector[Frame])]](None) + val opened = new ConcurrentLinkedQueue[(Node, FakeTransport)]() + val refuseHello = new java.util.concurrent.atomic.AtomicBoolean(false) val factory: Node => MultiplexedConnection.TransportFactory = node => (onFrame, onClosed) => { var transport: FakeTransport = null transport = new FakeTransport(onFrame, onClosed, respondFor(node, () => transport.close())) + opened.add(node -> transport) transport } private def respondFor(node: Node, disconnect: () => Unit): Bytes => Seq[Frame] = payload => { val s = payload.asUtf8String - if (s.contains("HELLO")) Seq(Replies.hello) + if (s.contains("HELLO")) Seq(if (refuseHello.get()) Frame.SimpleError("ERR node is down") else Replies.hello) else if (s.contains("ROLE")) if (node == master) Seq(role) else Nil else { val reads = occurrences(s, "PREAD") @@ -123,13 +126,12 @@ class MasterReplicaPipelineSpec extends munit.FunSuite { new MasterReplicaLive( script.factory, scheduler, - Vector(Connection.hello(None)), SageConfig(readFrom = readFrom), Vector(master), MasterReplicaConfig(1.second), Events(Vector(listener), Some(tracer)) ) - live.bootstrapRoles() + live.start() Fixture(live, completions, tracer, latch, script) } @@ -277,6 +279,19 @@ class MasterReplicaPipelineSpec extends munit.FunSuite { } } + test("a pipeline whose master connection is down when it is sent fails unsent, attributing no node or span") { + val f = build(ReadFrom.Master) + f.script.refuseHello.set(true) + f.script.opened.asScala.foreach { case (_, transport) => transport.close() } + f.live.pipeline(Seq(readCmd, writeCmd)).unsafeRun.failed.map { error => + assert(error.isInstanceOf[sage.SageException.NotConnected], s"an unsent batch must fail with NotConnected, got $error") + assert(f.latch.await(2, TimeUnit.SECONDS), "expected a completion per pipeline position") + assertEquals(f.completions.asScala.toVector.map(c => (c.name, c.node)), Vector("PREAD" -> None, "PWRITE" -> None)) + assertEquals(f.tracer.routed.asScala.toVector.filter(r => r._1 == "PREAD" || r._1 == "PWRITE"), Vector.empty) + f.live.close.unsafeRun + } + } + test("a ReplicaPreferred read-only pipeline falls back when the replica disconnects after submission") { val f = build(ReadFrom.ReplicaPreferred, disconnectOnRead = Some(replica)) f.live diff --git a/sage-client/zio/src/test/scala/sage/client/MasterReplicaTopologySpec.scala b/sage-client/zio/src/test/scala/sage/client/MasterReplicaTopologySpec.scala index 10b2e7da..0946cdd3 100644 --- a/sage-client/zio/src/test/scala/sage/client/MasterReplicaTopologySpec.scala +++ b/sage-client/zio/src/test/scala/sage/client/MasterReplicaTopologySpec.scala @@ -12,11 +12,11 @@ import scala.util.Try import kyo.compat.* import sage.{Bytes, SageEvent} -import sage.SageException.{ConnectionFailed, NotConnected, TimedOut} +import sage.SageException.{ConnectionFailed, TimedOut} import sage.client.internal.{ConnectFailureRecorder, Events, FakeTransport, MasterReplicaLive, MultiplexedConnection, Replies, Scheduler} import sage.client.internal.Replies.{masterRole, replicaRole} import sage.cluster.Node -import sage.commands.{Command, Connection} +import sage.commands.{BlockTimeout, Command, Lists} import sage.protocol.Frame class MasterReplicaTopologySpec extends munit.FunSuite { @@ -89,6 +89,7 @@ class MasterReplicaTopologySpec extends munit.FunSuite { roleRequests.add(node) roles.get(node).toSeq } else if (text.contains("EVALSHA")) Seq(Frame.Integer(1)) + else if (text.contains("STALL") || text.contains("BLPOP")) Nil else if (text.contains("WAIT")) Seq(Frame.Integer(lockAcknowledgements)) else if (reads > 0 && node == diesOnRead) { kill(node) @@ -107,7 +108,6 @@ class MasterReplicaTopologySpec extends munit.FunSuite { val live = new MasterReplicaLive( factory, scheduler, - Vector(Connection.hello(None)), SageConfig(readFrom = ReadFrom.ReplicaPreferred, connectTimeout = 500.millis, closeTimeout = Duration.Zero), seeds, MasterReplicaConfig(minRefreshInterval, topologyRefreshInterval), @@ -125,6 +125,9 @@ class MasterReplicaTopologySpec extends munit.FunSuite { transport.written.count(_.asUtf8String.contains("ZSCORE")) }.sum + def wrote(node: Node, token: String): Int = + transports.asScala.toVector.collect { case (`node`, transport) => transport.written.count(_.asUtf8String.contains(token)) }.sum + def lockConfirmations(node: Node): Int = transports.asScala.toVector.collect { case (`node`, transport) => transport.written.count(_.asUtf8String.contains("WAIT")) @@ -154,7 +157,7 @@ class MasterReplicaTopologySpec extends munit.FunSuite { test("master-replica locks confirm acquisition and renewal on the master with replica reads") { val fixture = new Fixture(Vector(primary), Map(primary -> masterRole(reader), reader -> replicaRole(primary))) - fixture.live.bootstrapRoles() + fixture.live.start() try { val result = scala.concurrent.Await.result( fixture.live @@ -173,7 +176,7 @@ class MasterReplicaTopologySpec extends munit.FunSuite { test("master-replica locks reject a missing replica and recover after refreshing its removal") { val fixture = new Fixture(Vector(primary), Map(primary -> masterRole(reader), reader -> replicaRole(primary))) - fixture.live.bootstrapRoles() + fixture.live.start() fixture.roles(primary) = masterRole() fixture.lockAcknowledgements = 0 var evaluated = false @@ -201,6 +204,32 @@ class MasterReplicaTopologySpec extends munit.FunSuite { } finally fixture.close() } + test("a master-replica lock write that times out requests role discovery") { + val fixture = new Fixture(Vector(primary), Map(primary -> masterRole())) + fixture.live.start() + val stalled = Command[Boolean]("STALL", Vector(0), Vector(Bytes.utf8("k")), _ => Right(true)) + try { + val result = + Try(scala.concurrent.Await.result(fixture.live.lockWrite(stalled, 100.millis, replicaAcknowledgement = false).unsafeRun, 5.seconds)) + assert(result.failed.get.isInstanceOf[TimedOut], result.toString) + fixture.awaitTrue(fixture.roleRequestCount(primary) >= 2, "the lock write timeout did not request role discovery") + } finally fixture.close() + } + + test("cancelling a blocking command requests role discovery") { + val fixture = new Fixture(Vector(primary), Map(primary -> masterRole())) + fixture.live.start() + val cancel = for { + fiber <- fixture.live.run(Lists.blPop[String, String]("k")(BlockTimeout.Forever)).lower.fork + _ <- zio.ZIO.attemptBlocking(fixture.awaitTrue(fixture.wrote(primary, "BLPOP") > 0, "BLPOP was never written")) + _ <- fiber.interrupt + } yield () + try { + scala.concurrent.Await.result(CIO.lift(cancel).unsafeRun, 5.seconds) + fixture.awaitTrue(fixture.roleRequestCount(primary) >= 2, "the cancelled blocking command did not request role discovery") + } finally fixture.close() + } + test("several seeds keep the supplied addresses and never dial a ROLE-advertised address") { val fixture = new Fixture( seeds = Vector(primary, reader), @@ -211,7 +240,7 @@ class MasterReplicaTopologySpec extends munit.FunSuite { ), unreachable = Set(advertisedReplica) ) - fixture.live.bootstrapRoles() + fixture.live.start() assertEquals(fixture.read(), 1L) assertEquals(fixture.readsServedBy(reader), 1) @@ -227,7 +256,7 @@ class MasterReplicaTopologySpec extends munit.FunSuite { initialRoles = Map(primary -> masterRole()), minRefreshInterval = Duration.Zero ) - fixture.live.bootstrapRoles() + fixture.live.start() assertEquals(fixture.read(), 1L) fixture.awaitTrue(fixture.roleRequestCount(primary) >= 2, "the first read did not trigger a re-discovery") @@ -251,7 +280,7 @@ class MasterReplicaTopologySpec extends munit.FunSuite { seeds = Vector(primary), initialRoles = Map(primary -> masterRole(advertisedReplica), advertisedReplica -> replicaRole(primary)) ) - fixture.live.bootstrapRoles() + fixture.live.start() fixture.diesOnRead = advertisedReplica assertEquals(fixture.read(), 1L) @@ -265,7 +294,7 @@ class MasterReplicaTopologySpec extends munit.FunSuite { seeds = Vector(primary), initialRoles = Map(primary -> masterRole(advertisedReplica), advertisedReplica -> replicaRole(primary)) ) - fixture.live.bootstrapRoles() + fixture.live.start() assertEquals(fixture.read(), 1L) assertEquals(fixture.readsServedBy(advertisedReplica), 1) @@ -291,7 +320,7 @@ class MasterReplicaTopologySpec extends munit.FunSuite { events = events ) - fixture.live.bootstrapRoles() + fixture.live.start() fixture.close() assertEquals(connected.asScala.toVector, Vector.empty) @@ -305,7 +334,7 @@ class MasterReplicaTopologySpec extends munit.FunSuite { events = recorder.events ) - intercept[NotConnected](fixture.live.bootstrapRoles()) + intercept[TimedOut](fixture.live.start()) assert(recorder.await(), "the singleton ROLE timeout was not reported") val failure = recorder.failures.head @@ -313,6 +342,12 @@ class MasterReplicaTopologySpec extends munit.FunSuite { assert(failure.error.isInstanceOf[TimedOut], s"unexpected cause: ${failure.error}") } + test("supplied endpoints that all time out on ROLE return the timeout") { + val fixture = new Fixture(seeds = Vector(primary, reader), initialRoles = Map.empty) + + intercept[TimedOut](fixture.live.start()) + } + test("an unreachable supplied endpoint is omitted and reported while the available topology connects") { val recorder = new ConnectFailureRecorder val fixture = new Fixture( @@ -321,7 +356,7 @@ class MasterReplicaTopologySpec extends munit.FunSuite { unreachable = Set(reader), events = recorder.events ) - fixture.live.bootstrapRoles() + fixture.live.start() assertEquals(fixture.read(), 1L) assert(recorder.await(), "ConnectFailed was not delivered") @@ -340,7 +375,7 @@ class MasterReplicaTopologySpec extends munit.FunSuite { unreachable = down, minRefreshInterval = 10.millis ) - fixture.live.bootstrapRoles() + fixture.live.start() assertEquals(fixture.read(), 1L) assertEquals(fixture.readsServedBy(primary), 1) @@ -365,7 +400,7 @@ class MasterReplicaTopologySpec extends munit.FunSuite { initialRoles = Map(primary -> masterRole()), events = recorder.events ) - fixture.live.bootstrapRoles() + fixture.live.start() assert(recorder.await(), "the ROLE timeout was not reported") val failure = recorder.failures.head @@ -384,7 +419,7 @@ class MasterReplicaTopologySpec extends munit.FunSuite { unreachable = down, events = recorder.events ) - fixture.live.bootstrapRoles() + fixture.live.start() assert(fixture.write().isSuccess, "the master pool should be established before discovery connections fail") down += primary @@ -405,7 +440,7 @@ class MasterReplicaTopologySpec extends munit.FunSuite { seeds = Vector(primary, reader), initialRoles = Map(primary -> masterRole(reader), reader -> replicaRole(primary, state = "sync")) ) - fixture.live.bootstrapRoles() + fixture.live.start() assertEquals(fixture.read(), 1L) assertEquals(fixture.readsServedBy(primary), 1) @@ -434,7 +469,7 @@ class MasterReplicaTopologySpec extends munit.FunSuite { ), topologyRefreshInterval = Some(50.millis) ) - fixture.live.bootstrapRoles() + fixture.live.start() assertEquals(fixture.read(), 1L) assertEquals(fixture.readsServedBy(reader), 1) @@ -452,6 +487,26 @@ class MasterReplicaTopologySpec extends munit.FunSuite { fixture.close() } + test("a subscription moves to the current master when its former master leaves replication but keeps running") { + val fixture = new Fixture( + seeds = Vector(primary), + initialRoles = Map(primary -> masterRole(reader), reader -> replicaRole(primary)), + topologyRefreshInterval = Some(50.millis) + ) + fixture.live.start() + val subscription = scala.concurrent.Await.result(fixture.live.subscribeChannels[String]("news").unsafeRun, 10.seconds) + assertEquals(fixture.wrote(primary, "\r\nSUBSCRIBE\r\n"), 1) + + // a failover promotes the reader, and the former master replicates it, so PUBLISH still reaches the subscription + fixture.roles ++= Map(reader -> masterRole(primary), primary -> replicaRole(reader)) + fixture.awaitTrue({ fixture.write(); fixture.wrote(reader, "ZADD") > 0 }, "the failover was not discovered") + // REPLICAOF NO ONE on the former master keeps its connections open, but PUBLISH on the current master no longer reaches it + fixture.roles ++= Map(reader -> masterRole(), primary -> masterRole()) + fixture.awaitTrue(fixture.wrote(reader, "\r\nSUBSCRIBE\r\n") > 0, "the subscription stayed on the former master") + scala.concurrent.Await.result(subscription.close.unsafeRun, 10.seconds) + fixture.close() + } + test("a re-discovery the poll queued before close does not probe after it") { val scheduler = new DeferringScheduler val fixture = new Fixture( @@ -460,7 +515,7 @@ class MasterReplicaTopologySpec extends munit.FunSuite { topologyRefreshInterval = Some(1.minute), scheduler = scheduler ) - fixture.live.bootstrapRoles() + fixture.live.start() val dialledAtBootstrap = fixture.dialled.size scheduler.tick() @@ -476,7 +531,7 @@ class MasterReplicaTopologySpec extends munit.FunSuite { initialRoles = Map(primary -> masterRole(reader), reader -> replicaRole(primary, state = "sync")), minRefreshInterval = 10.millis ) - fixture.live.bootstrapRoles() + fixture.live.start() fixture.readPipeline() assertEquals(fixture.readsServedBy(primary), 1) @@ -500,7 +555,7 @@ class MasterReplicaTopologySpec extends munit.FunSuite { initialRoles = Map(primary -> masterRole(advertisedReplica)), failDial = (node, attempt) => node == primary && attempt == 1 ) - fixture.live.bootstrapRoles() + fixture.live.start() assertEquals(fixture.roleRequestCount(primary), 1) assert(fixture.writePipeline().isFailure, "the master connection establishment should fail") @@ -517,7 +572,7 @@ class MasterReplicaTopologySpec extends munit.FunSuite { initialRoles = Map(reader -> replicaRole(primary), advertisedReplica -> replicaRole(primary)) ) - val error = intercept[ConnectionFailed](fixture.live.bootstrapRoles()) + val error = intercept[ConnectionFailed](fixture.live.start()) assert(error.getMessage.contains("no supplied endpoint reports the master role"), error.getMessage) } @@ -527,7 +582,7 @@ class MasterReplicaTopologySpec extends munit.FunSuite { initialRoles = Map(primary -> masterRole(advertisedReplica), reader -> replicaRole(primary)), unreachable = Set(advertisedReplica) ) - fixture.live.bootstrapRoles() + fixture.live.start() assertEquals(fixture.read(), 1L) assertEquals(fixture.readsServedBy(reader), 1) diff --git a/sage-client/zio/src/test/scala/sage/client/internal/TxScopeFaultSpec.scala b/sage-client/zio/src/test/scala/sage/client/internal/TxScopeFaultSpec.scala index cef8899c..d6e84b23 100644 --- a/sage-client/zio/src/test/scala/sage/client/internal/TxScopeFaultSpec.scala +++ b/sage-client/zio/src/test/scala/sage/client/internal/TxScopeFaultSpec.scala @@ -6,7 +6,7 @@ import scala.concurrent.ExecutionContext import kyo.compat.* import sage.Bytes -import sage.SageException.{ConnectionLost, ServerError} +import sage.SageException.{ConnectionLost, TransactionDiscarded} import sage.client.DedicatedPoolConfig import sage.commands.{Connection, Strings} import sage.protocol.Frame @@ -17,26 +17,19 @@ class TxScopeFaultSpec extends munit.FunSuite { private val readonly = Frame.SimpleError("READONLY You can't write against a read only replica.") - private def txScope(respond: Bytes => Seq[Frame]): (Client.TxScope, mutable.ArrayBuffer[Throwable]) = { + private def txScope(respond: Bytes => Seq[Frame]): (Client.TxScope, mutable.ArrayBuffer[RefreshPolicy]) = { val scheduler = new ManualScheduler - val gen = MultiplexedConnection.Generation.initial val factory: MultiplexedConnection.TransportFactory = ScriptedTransport.factory(respond) val pool = - new DedicatedPool(factory, Vector(Connection.hello()), scheduler, () => true, () => Some(gen), _ == gen, DedicatedPoolConfig(), 1000L) - val faults = mutable.ArrayBuffer.empty[Throwable] - (new Client.TxScope(pool.acquireForTransaction(), faults += _), faults) + new DedicatedPool(factory, Vector(Connection.hello()), scheduler, () => true, DedicatedPoolConfig(), 1000L) + val faults = mutable.ArrayBuffer.empty[RefreshPolicy] + (new Client.TxScope(pool.acquireForTransaction(), pool.releaseTransaction, faults += _), faults) } - private def isOwnershipFault(error: Throwable): Boolean = error match { - case e: ServerError => e.code == "READONLY" - case _: ConnectionLost => true - case _ => false - } - - test("a READONLY command fault invokes the onFault hook") { + test("a READONLY command fault requests a forced refresh") { val (scope, faults) = txScope(p => if (p.asUtf8String.contains("HELLO")) Seq(Replies.hello) else Seq(readonly)) scope.run(Strings.set("k", "v")).unsafeRun.failed.map { _ => - assert(faults.exists(isOwnershipFault), s"expected an ownership fault, got $faults") + assert(faults.contains(RefreshPolicy.Forced), s"expected a forced refresh, got $faults") } } @@ -45,14 +38,15 @@ class TxScopeFaultSpec extends munit.FunSuite { scope.run(Strings.set("k", "v")).unsafeRun.failed.map(e => assert(e.isInstanceOf[ConnectionLost], s"expected ConnectionLost, got $e")) } - test("a transaction whose EXEC hits a READONLY invokes the onFault hook") { + test("a transaction whose queued command hits a READONLY is discarded and requests a forced refresh") { val respond: Bytes => Seq[Frame] = p => if (p.asUtf8String.contains("HELLO")) Seq(Replies.hello) else if (p.asUtf8String.contains("MULTI")) Seq(Frame.SimpleString("OK"), readonly, Frame.SimpleError("EXECABORT discarded")) else Seq(Frame.SimpleString("OK")) val (scope, faults) = txScope(respond) - scope.exec(Vector(Strings.set("k", "v"))).unsafeRun.failed.map { _ => - assert(faults.exists(isOwnershipFault), s"expected an ownership fault, got $faults") + scope.exec(Vector(Strings.set("k", "v"))).unsafeRun.failed.map { error => + assert(error.isInstanceOf[TransactionDiscarded] && error.getMessage.contains("READONLY"), s"expected TransactionDiscarded, got $error") + assert(faults.contains(RefreshPolicy.Forced), s"expected a forced refresh, got $faults") } } } diff --git a/sage-client/zio/src/test/scala/sage/client/internal/ZioLockCancellationSpec.scala b/sage-client/zio/src/test/scala/sage/client/internal/ZioLockCancellationSpec.scala index c61e113e..d43fd34d 100644 --- a/sage-client/zio/src/test/scala/sage/client/internal/ZioLockCancellationSpec.scala +++ b/sage-client/zio/src/test/scala/sage/client/internal/ZioLockCancellationSpec.scala @@ -11,10 +11,10 @@ import sage.SageException import sage.backend.SageClient class ZioLockCancellationSpec extends LockCancellationSpec { - override protected def tryWithLock[A](commands: CommandRunner[CIO, String], lease: FiniteDuration)(body: CIO[A]): CIO[Option[A]] = + override protected def tryWithLock[A](commands: SharedRunner, lease: FiniteDuration)(body: CIO[A]): CIO[Option[A]] = CIO.lift(new SageClient.Lowered(new LockTestClient(commands)).lock[String](lease).tryWithLock("key")(body.lower.refineToOrDie[SageException])) - override protected def withLock[A](commands: CommandRunner[CIO, String], lease: FiniteDuration, wait: FiniteDuration)(body: CIO[A]): CIO[A] = + override protected def withLock[A](commands: SharedRunner, lease: FiniteDuration, wait: FiniteDuration)(body: CIO[A]): CIO[A] = CIO.lift(new SageClient.Lowered(new LockTestClient(commands)).lock[String](lease).withLock("key", wait)(body.lower.refineToOrDie[SageException])) test("a body defect ends the scope and releases ownership") { diff --git a/sage-compat-ce/.conformance/ce/src/test/scala/kyo/compat/CeBindingTest.scala b/sage-compat-ce/.conformance/ce/src/test/scala/kyo/compat/CeBindingTest.scala new file mode 100644 index 00000000..b9570af2 --- /dev/null +++ b/sage-compat-ce/.conformance/ce/src/test/scala/kyo/compat/CeBindingTest.scala @@ -0,0 +1,54 @@ +package kyo.compat + +import java.util.concurrent.atomic.AtomicInteger + +import scala.concurrent.duration.* +import scala.util.Failure + +class CeBindingTest extends CompatTest { + + private def failsWhileBuilding(): CIO[Int] = throw TestError("built") + + "ensure runs the cleanup when building the computation throws" in run { + val ran = new AtomicInteger(0) + CIO.ensure(CIO.defer { val _ = ran.incrementAndGet() })(failsWhileBuilding()).liftToTry.map { + case Failure(TestError("built")) => assert(ran.get == 1) + case other => fail(s"expected Failure(TestError(built)), got $other") + } + } + + "foreachIndexed numbers a Set's elements in iteration order" in run { + val set = (1 to 20).map(i => s"e$i").toSet + CIO.foreachIndexed(set)((i, a) => CIO.value(i -> a)).map(out => assert(out.toSeq == set.toSeq.zipWithIndex.map(_.swap))) + } + + "timeout reports a null result as completed" in run { + CIO.timeout(1.minute)(CIO.value(null: String)).map(out => assert(out == Some(null))) + } + + "atomic compareAndSet compares values outside the boxing cache" in run { + for { + i <- CAtomicInt.init(1000) + okInt <- i.compareAndSet(Integer.valueOf(1000).intValue, 2000) + l <- CAtomicLong.init(1000L) + okLng <- l.compareAndSet(java.lang.Long.valueOf(1000L).longValue, 2000L) + ra <- CAtomicRef.init[Any](1000L) + okAny <- ra.compareAndSet(java.lang.Long.valueOf(1000L), 2000L) + vi <- i.get + vl <- l.get + va <- ra.get + } yield assert(okInt && okLng && okAny && vi == 2000 && vl == 2000L && va == 2000L) + } + + "atomic arithmetic wraps on overflow" in run { + for { + i <- CAtomicInt.init(Int.MaxValue) + vi <- i.incrementAndGet + di <- i.getAndDecrement + l <- CAtomicLong.init(Long.MinValue) + vl <- l.decrementAndGet + al <- l.getAndAdd(2L) + nl <- l.get + } yield assert(vi == Int.MinValue && di == Int.MinValue && vl == Long.MaxValue && al == Long.MaxValue && nl == Long.MinValue + 1) + } +} diff --git a/sage-compat-ce/README.md b/sage-compat-ce/README.md index ea7b54d2..4dc4ab1f 100644 --- a/sage-compat-ce/README.md +++ b/sage-compat-ce/README.md @@ -6,7 +6,7 @@ This module implements the full [kyo-compat](https://github.com/getkyo/kyo/tree/ Kyo removed its Cats Effect integrations in 1.0.0-RC6 in [kyo#1779](https://github.com/getkyo/kyo/pull/1779) and moved community bindings outside the project in [kyo#1840](https://github.com/getkyo/kyo/pull/1840). This module vendors the last upstream Cats Effect binding from commit [`eae31e1d`](https://github.com/getkyo/kyo/tree/eae31e1d39d4b8ff2df168272e60b38d9e9dd502/kyo-compat/bindings/ce). Those sources match `io.getkyo:kyo-compat-ce_3:1.0.0-RC5`. -The vendored copy merges the `shared` and `jvm` source trees because Sage supports only the JVM. It also uses this repository's brace-based formatting and corrects several comments. +The vendored copy merges the `shared` and `jvm` source trees because Sage supports only the JVM. It also uses this repository's brace-based formatting and corrects several comments. `CAtomicInt`, `CAtomicLong`, and `CAtomicBoolean` are aliases of `CAtomicRef`, which defines their operations once. Sage changes `CIO.async` from `IO.async_` to `IO.async` with a cancellation token so a timeout or cancellation can stop waiting for a callback. This applies to every CE command wait, including distributed lock acquisition, renewal, and release. Cancellation does not stop the external operation or retract a command already sent to the server. The transport still consumes its reply in order. The conformance suite checks compatibility with the shared API, and lock regression tests cover delayed callbacks. diff --git a/sage-compat-ce/src/main/scala/kyo/compat/CAtomicBoolean.scala b/sage-compat-ce/src/main/scala/kyo/compat/CAtomicBoolean.scala index c1471ee4..5509ddc2 100644 --- a/sage-compat-ce/src/main/scala/kyo/compat/CAtomicBoolean.scala +++ b/sage-compat-ce/src/main/scala/kyo/compat/CAtomicBoolean.scala @@ -4,52 +4,19 @@ import cats.effect.IO import cats.effect.kernel.Ref /** - * Uses `cats.effect.kernel.Ref[IO, Boolean]`. Cats Effect has no `Frame` or `Trace` to propagate. `lift` and `lower` return the existing Cats - * Effect ref. `compareAndSet` uses `Ref.modify` because Cats Effect does not provide a native compare-and-set operation. + * Uses `cats.effect.kernel.Ref[IO, Boolean]` through [[CAtomicRef]], which provides every operation. */ -opaque type CAtomicBoolean = Ref[IO, Boolean] +type CAtomicBoolean = CAtomicRef[Boolean] object CAtomicBoolean { /** * Allocates a fresh atomic boolean initialized to `v`. */ - inline def init(inline v: Boolean): CIO[CAtomicBoolean] = - CIO.lift(Ref.of[IO, Boolean](v)) + inline def init(inline v: Boolean): CIO[CAtomicBoolean] = CAtomicRef.init(v) /** * Wraps a native `cats.effect.kernel.Ref[IO, Boolean]` as a `CAtomicBoolean`. The conversion is the identity on the carrier. */ - inline def lift(inline u: Ref[IO, Boolean]): CAtomicBoolean = u - - extension (inline self: CAtomicBoolean) { - - /** - * Unwraps to the native `cats.effect.kernel.Ref[IO, Boolean]`. The conversion is the identity on the carrier. - */ - inline def lower: Ref[IO, Boolean] = self - - /** - * Reads the current value. - */ - inline def get: CIO[Boolean] = CIO.lift(self.get) - - /** - * Atomically sets the value to `v`. - */ - inline def set(inline v: Boolean): CIO[Unit] = CIO.lift(self.set(v)) - - /** - * Atomically sets the value to `v` and returns the previous value. - */ - inline def getAndSet(inline v: Boolean): CIO[Boolean] = CIO.lift(self.getAndSet(v)) - - /** - * Atomic compare-and-set: replaces `expected` with `updated` iff the current value equals `expected`. - */ - inline def compareAndSet(inline expected: Boolean, inline updated: Boolean): CIO[Boolean] = - CIO.lift(self.modify(cur => if (cur == expected) (updated, true) else (cur, false))) - - } - + inline def lift(inline u: Ref[IO, Boolean]): CAtomicBoolean = CAtomicRef.lift(u) } diff --git a/sage-compat-ce/src/main/scala/kyo/compat/CAtomicInt.scala b/sage-compat-ce/src/main/scala/kyo/compat/CAtomicInt.scala index 62e9aa14..1fc1b3ff 100644 --- a/sage-compat-ce/src/main/scala/kyo/compat/CAtomicInt.scala +++ b/sage-compat-ce/src/main/scala/kyo/compat/CAtomicInt.scala @@ -4,83 +4,19 @@ import cats.effect.IO import cats.effect.kernel.Ref /** - * Uses `cats.effect.kernel.Ref[IO, Int]`. Cats Effect has no `Frame` or `Trace` to propagate. `lift` and `lower` return the existing Cats - * Effect ref. Cats Effect does not provide a specialized atomic integer. Arithmetic uses `Ref[Int]` updates, and `compareAndSet` uses - * `Ref.modify`. + * Uses `cats.effect.kernel.Ref[IO, Int]` through [[CAtomicRef]], which provides every operation. */ -opaque type CAtomicInt = Ref[IO, Int] +type CAtomicInt = CAtomicRef[Int] object CAtomicInt { /** * Allocates a fresh atomic int initialized to `v`. */ - inline def init(inline v: Int): CIO[CAtomicInt] = - CIO.lift(Ref.of[IO, Int](v)) + inline def init(inline v: Int): CIO[CAtomicInt] = CAtomicRef.init(v) /** * Wraps a native `cats.effect.kernel.Ref[IO, Int]` as a `CAtomicInt`. The conversion is the identity on the carrier. */ - inline def lift(inline u: Ref[IO, Int]): CAtomicInt = u - - extension (inline self: CAtomicInt) { - - /** - * Unwraps to the native `cats.effect.kernel.Ref[IO, Int]`. The conversion is the identity on the carrier. - */ - inline def lower: Ref[IO, Int] = self - - /** - * Reads the current value. - */ - inline def get: CIO[Int] = CIO.lift(self.get) - - /** - * Atomically sets the value to `v`. - */ - inline def set(inline v: Int): CIO[Unit] = CIO.lift(self.set(v)) - - /** - * Atomically sets the value to `v` and returns the previous value. - */ - inline def getAndSet(inline v: Int): CIO[Int] = CIO.lift(self.getAndSet(v)) - - /** - * Atomically increments by 1 and returns the new value. - */ - inline def incrementAndGet: CIO[Int] = CIO.lift(self.updateAndGet(_ + 1)) - - /** - * Atomically increments by 1 and returns the previous value. - */ - inline def getAndIncrement: CIO[Int] = CIO.lift(self.getAndUpdate(_ + 1)) - - /** - * Atomically decrements by 1 and returns the new value. - */ - inline def decrementAndGet: CIO[Int] = CIO.lift(self.updateAndGet(_ - 1)) - - /** - * Atomically decrements by 1 and returns the previous value. - */ - inline def getAndDecrement: CIO[Int] = CIO.lift(self.getAndUpdate(_ - 1)) - - /** - * Atomically adds `delta` and returns the new value. - */ - inline def addAndGet(inline delta: Int): CIO[Int] = CIO.lift(self.updateAndGet(_ + delta)) - - /** - * Atomically adds `delta` and returns the previous value. - */ - inline def getAndAdd(inline delta: Int): CIO[Int] = CIO.lift(self.getAndUpdate(_ + delta)) - - /** - * Atomic compare-and-set: replaces `expected` with `updated` iff the current value equals `expected`. - */ - inline def compareAndSet(inline expected: Int, inline updated: Int): CIO[Boolean] = - CIO.lift(self.modify(cur => if (cur == expected) (updated, true) else (cur, false))) - - } - + inline def lift(inline u: Ref[IO, Int]): CAtomicInt = CAtomicRef.lift(u) } diff --git a/sage-compat-ce/src/main/scala/kyo/compat/CAtomicLong.scala b/sage-compat-ce/src/main/scala/kyo/compat/CAtomicLong.scala index 651103ad..01474cb6 100644 --- a/sage-compat-ce/src/main/scala/kyo/compat/CAtomicLong.scala +++ b/sage-compat-ce/src/main/scala/kyo/compat/CAtomicLong.scala @@ -4,83 +4,19 @@ import cats.effect.IO import cats.effect.kernel.Ref /** - * Uses `cats.effect.kernel.Ref[IO, Long]`. Cats Effect has no `Frame` or `Trace` to propagate. `lift` and `lower` return the existing Cats - * Effect ref. Cats Effect does not provide a specialized atomic long. Arithmetic uses `Ref[Long]` updates, and `compareAndSet` uses - * `Ref.modify`. + * Uses `cats.effect.kernel.Ref[IO, Long]` through [[CAtomicRef]], which provides every operation. */ -opaque type CAtomicLong = Ref[IO, Long] +type CAtomicLong = CAtomicRef[Long] object CAtomicLong { /** * Allocates a fresh atomic long initialized to `v`. */ - inline def init(inline v: Long): CIO[CAtomicLong] = - CIO.lift(Ref.of[IO, Long](v)) + inline def init(inline v: Long): CIO[CAtomicLong] = CAtomicRef.init(v) /** * Wraps a native `cats.effect.kernel.Ref[IO, Long]` as a `CAtomicLong`. The conversion is the identity on the carrier. */ - inline def lift(inline u: Ref[IO, Long]): CAtomicLong = u - - extension (inline self: CAtomicLong) { - - /** - * Unwraps to the native `cats.effect.kernel.Ref[IO, Long]`. The conversion is the identity on the carrier. - */ - inline def lower: Ref[IO, Long] = self - - /** - * Reads the current value. - */ - inline def get: CIO[Long] = CIO.lift(self.get) - - /** - * Atomically sets the value to `v`. - */ - inline def set(inline v: Long): CIO[Unit] = CIO.lift(self.set(v)) - - /** - * Atomically sets the value to `v` and returns the previous value. - */ - inline def getAndSet(inline v: Long): CIO[Long] = CIO.lift(self.getAndSet(v)) - - /** - * Atomically increments by 1 and returns the new value. - */ - inline def incrementAndGet: CIO[Long] = CIO.lift(self.updateAndGet(_ + 1L)) - - /** - * Atomically increments by 1 and returns the previous value. - */ - inline def getAndIncrement: CIO[Long] = CIO.lift(self.getAndUpdate(_ + 1L)) - - /** - * Atomically decrements by 1 and returns the new value. - */ - inline def decrementAndGet: CIO[Long] = CIO.lift(self.updateAndGet(_ - 1L)) - - /** - * Atomically decrements by 1 and returns the previous value. - */ - inline def getAndDecrement: CIO[Long] = CIO.lift(self.getAndUpdate(_ - 1L)) - - /** - * Atomically adds `delta` and returns the new value. - */ - inline def addAndGet(inline delta: Long): CIO[Long] = CIO.lift(self.updateAndGet(_ + delta)) - - /** - * Atomically adds `delta` and returns the previous value. - */ - inline def getAndAdd(inline delta: Long): CIO[Long] = CIO.lift(self.getAndUpdate(_ + delta)) - - /** - * Atomic compare-and-set: replaces `expected` with `updated` iff the current value equals `expected`. - */ - inline def compareAndSet(inline expected: Long, inline updated: Long): CIO[Boolean] = - CIO.lift(self.modify(cur => if (cur == expected) (updated, true) else (cur, false))) - - } - + inline def lift(inline u: Ref[IO, Long]): CAtomicLong = CAtomicRef.lift(u) } diff --git a/sage-compat-ce/src/main/scala/kyo/compat/CAtomicRef.scala b/sage-compat-ce/src/main/scala/kyo/compat/CAtomicRef.scala index a334556c..4afdf772 100644 --- a/sage-compat-ce/src/main/scala/kyo/compat/CAtomicRef.scala +++ b/sage-compat-ce/src/main/scala/kyo/compat/CAtomicRef.scala @@ -65,4 +65,38 @@ object CAtomicRef { } + extension [A](inline self: CAtomicRef[A])(using n: Numeric[A]) { + + /** + * Atomically increments by 1 and returns the new value. + */ + inline def incrementAndGet: CIO[A] = CIO.lift(self.updateAndGet(n.plus(_, n.one))) + + /** + * Atomically increments by 1 and returns the previous value. + */ + inline def getAndIncrement: CIO[A] = CIO.lift(self.getAndUpdate(n.plus(_, n.one))) + + /** + * Atomically decrements by 1 and returns the new value. + */ + inline def decrementAndGet: CIO[A] = CIO.lift(self.updateAndGet(n.minus(_, n.one))) + + /** + * Atomically decrements by 1 and returns the previous value. + */ + inline def getAndDecrement: CIO[A] = CIO.lift(self.getAndUpdate(n.minus(_, n.one))) + + /** + * Atomically adds `delta` and returns the new value. + */ + inline def addAndGet(inline delta: A): CIO[A] = CIO.lift(self.updateAndGet(n.plus(_, delta))) + + /** + * Atomically adds `delta` and returns the previous value. + */ + inline def getAndAdd(inline delta: A): CIO[A] = CIO.lift(self.getAndUpdate(n.plus(_, delta))) + + } + } diff --git a/sage-compat-ce/src/main/scala/kyo/compat/CFiber.scala b/sage-compat-ce/src/main/scala/kyo/compat/CFiber.scala index 96fece96..af13ae2e 100644 --- a/sage-compat-ce/src/main/scala/kyo/compat/CFiber.scala +++ b/sage-compat-ce/src/main/scala/kyo/compat/CFiber.scala @@ -4,6 +4,7 @@ import java.util.concurrent.CancellationException import cats.effect.FiberIO import cats.effect.IO +import cats.effect.Outcome /** * Underlying carrier is `cats.effect.FiberIO[A]`. Cats Effect has no `Frame` / `Trace` to propagate. `lift` and `lower` are identity since @@ -38,28 +39,17 @@ object CFiber { * Joins the fiber and returns its result. Cancellation fails with `CancellationException`. */ inline def get: CIO[A] = - CIO.lift( - self.join.flatMap { - case cats.effect.Outcome.Succeeded(ioa) => ioa - case cats.effect.Outcome.Errored(t) => IO.raiseError(t) - case cats.effect.Outcome.Canceled() => IO.raiseError(new CancellationException("CFiber.interrupt")) - } - ) + CIO.lift(self.join.flatMap { + case Outcome.Succeeded(ioa) => ioa + case Outcome.Errored(t) => IO.raiseError(t) + case Outcome.Canceled() => IO.raiseError(new CancellationException("CFiber.interrupt")) + }) /** * Registers `cb` to fire when the fiber completes; success and failure are reified as `scala.util.Try`, and `Outcome.Canceled` is * translated to `Failure(CancellationException)` before the callback runs. */ inline def onComplete(cb: scala.util.Try[A] => CIO[Unit]): CIO[Unit] = - CIO.lift( - self.join - .flatMap { - case cats.effect.Outcome.Succeeded(ioa) => ioa.flatMap(a => cb(scala.util.Success(a)).lower) - case cats.effect.Outcome.Errored(t) => cb(scala.util.Failure(t)).lower - case cats.effect.Outcome.Canceled() => cb(scala.util.Failure(new CancellationException("CFiber.interrupt"))).lower - } - .start - .void - ) + CIO.lift(get.lower.attempt.flatMap(r => cb(r.toTry).lower).start.void) } } diff --git a/sage-compat-ce/src/main/scala/kyo/compat/CIO.scala b/sage-compat-ce/src/main/scala/kyo/compat/CIO.scala index b32755c7..602588a9 100644 --- a/sage-compat-ce/src/main/scala/kyo/compat/CIO.scala +++ b/sage-compat-ce/src/main/scala/kyo/compat/CIO.scala @@ -1,9 +1,7 @@ package kyo.compat -import scala.annotation.nowarn import scala.concurrent.Future as ScalaFuture import scala.concurrent.duration.FiniteDuration -import scala.concurrent.duration.NANOSECONDS import cats.effect.IO import cats.syntax.parallel.* @@ -124,10 +122,7 @@ object CIO { * Reifies failure as `Try`; the resulting `CIO` always succeeds. */ inline def liftToTry: CIO[scala.util.Try[A]] = - lift(self.lower.attempt.map { - case Right(a) => scala.util.Success(a) - case Left(t) => scala.util.Failure(t) - }) + lift(self.lower.attempt.map(_.toTry)) /** * Discards the success value; failure propagates. @@ -144,11 +139,8 @@ object CIO { /** * Rewrites the error value through `f`. */ - @nowarn("msg=anonymous") inline def mapError(inline f: Throwable => Throwable): CIO[A] = - lift(self.lower.adaptError { case t => - f(t) - }) + lift(self.lower.handleErrorWith(t => IO.raiseError(f(t)))) /** * Transforms the success value with a pure function. @@ -185,13 +177,13 @@ object CIO { * Reads a monotonic timestamp expressed as a `FiniteDuration` since a backend-defined origin, suitable for measuring intervals. */ inline def nowMonotonic: CIO[FiniteDuration] = - lift(IO.monotonic.map(d => FiniteDuration(d.toNanos, NANOSECONDS))) + lift(IO.monotonic) /** * Runs `c` with a deadline; resolves to `None` if `d` elapses first. */ inline def timeout[A](inline d: FiniteDuration)(inline c: CIO[A]): CIO[Option[A]] = - lift(c.lower.map(Option(_)).timeoutTo(d, IO.none[A])) + lift(c.lower.map(Some(_)).timeoutTo(d, IO.none[A])) /** * Runs `c` with a deadline; fails with `e` if `d` elapses first. @@ -216,6 +208,12 @@ object CIO { ): CIO[A] = lift(IO.race(a.lower, b.lower).map(_.merge)) + private inline def parTraverse[A, B](list: List[A], inline concurrency: Int)(f: A => IO[B]): IO[List[B]] = + if (concurrency == Int.MaxValue) list.parTraverse(f) else IO.parTraverseN(concurrency)(list)(f) + + private inline def parTraverseDiscard[A](list: List[A], inline concurrency: Int)(f: A => IO[Any]): IO[Unit] = + if (concurrency == Int.MaxValue) list.parTraverse_(f) else IO.parTraverseN_(concurrency)(list)(f) + /** * Parallel map. `concurrency` caps the number of in-flight elements, is unbounded by default, and must be positive when bounded. */ @@ -223,13 +221,7 @@ object CIO { inline coll: Iterable[A], inline concurrency: Int = Int.MaxValue )(f: A => CIO[B]): CIO[CChunk[B]] = - lift { - val list = coll.toList - val io = - if (concurrency == Int.MaxValue) list.parTraverse(a => f(a).lower) - else IO.parTraverseN(concurrency)(list)(a => f(a).lower) - io.map(lst => CChunk.lift(lst.toVector)) - } + lift(parTraverse(coll.toList, concurrency)(a => f(a).lower).map(lst => CChunk.lift(lst.toVector))) /** * Parallel map that passes the element index to `f`; same concurrency semantics as `foreach`. @@ -238,13 +230,7 @@ object CIO { inline coll: Iterable[A], inline concurrency: Int = Int.MaxValue )(f: (Int, A) => CIO[B]): CIO[CChunk[B]] = - lift { - val list = coll.toList.zipWithIndex - val io = - if (concurrency == Int.MaxValue) list.parTraverse { case (a, i) => f(i, a).lower } - else IO.parTraverseN(concurrency)(list) { case (a, i) => f(i, a).lower } - io.map(lst => CChunk.lift(lst.toVector)) - } + foreach(coll.toList.zipWithIndex, concurrency)((a, i) => f(i, a)) /** * Runs `f` for its effects on each element and discards the results; same concurrency semantics as `foreach`. @@ -253,27 +239,16 @@ object CIO { inline coll: Iterable[A], inline concurrency: Int = Int.MaxValue )(f: A => CIO[Any]): CIO[Unit] = - lift { - val list = coll.toList - if (concurrency == Int.MaxValue) list.parTraverse_(a => f(a).lower) - else IO.parTraverseN_(concurrency)(list)(a => f(a).lower) - } + lift(parTraverseDiscard(coll.toList, concurrency)(a => f(a).lower)) /** * Filters the collection with an effectful predicate; same concurrency semantics as `foreach`. */ - @nowarn("msg=anonymous") inline def filter[A]( inline coll: Iterable[A], inline concurrency: Int = Int.MaxValue )(p: A => CIO[Boolean]): CIO[CChunk[A]] = - lift { - val list = coll.toList - if (concurrency == Int.MaxValue) list.parFilterA(a => p(a).lower).map(lst => CChunk.lift(lst.toVector)) - else - IO.parTraverseN(concurrency)(list)(a => p(a).lower.map(b => a -> b)) - .map(pairs => CChunk.lift(pairs.collect { case (a, true) => a }.toVector)) - } + lift(parTraverse(coll.toList, concurrency)(a => p(a).lower.map(a -> _)).map(pairs => CChunk.lift(pairs.filter(_._2).map(_._1).toVector))) /** * Sequences an `Iterable[CIO[A]]`; same concurrency semantics as `foreach`. @@ -282,11 +257,7 @@ object CIO { inline coll: Iterable[CIO[A]], inline concurrency: Int = Int.MaxValue ): CIO[CChunk[A]] = - lift { - val list = coll.toList.map(_.lower) - if (concurrency == Int.MaxValue) list.parSequence.map(lst => CChunk.lift(lst.toVector)) - else IO.parSequenceN(concurrency)(list).map(lst => CChunk.lift(lst.toVector)) - } + foreach(coll, concurrency)(identity) /** * Sequences and discards the results; same concurrency semantics as `foreach`. @@ -295,11 +266,7 @@ object CIO { inline coll: Iterable[CIO[Any]], inline concurrency: Int = Int.MaxValue ): CIO[Unit] = - lift { - val list = coll.toList.map(_.lower) - if (concurrency == Int.MaxValue) list.parSequence_ - else IO.parSequenceN_(concurrency)(list) - } + foreachDiscard(coll, concurrency)(identity) /** * Bridges a one-shot completion callback into `CIO`; `register` receives a `Try[A] => Unit`. @@ -307,11 +274,7 @@ object CIO { inline def async[A](inline register: ((scala.util.Try[A] => Unit) => Unit)): CIO[A] = lift(IO.async[A] { k => IO.delay { - val cb: scala.util.Try[A] => Unit = { - case scala.util.Success(a) => k(Right(a)) - case scala.util.Failure(e) => k(Left(e)) - } - register(cb) + register(t => k(t.toEither)) // A cancel token makes callback waiting cancelable even when the external operation cannot be stopped. Some(IO.unit) } diff --git a/sage-compat-ce/src/main/scala/kyo/compat/CLatch.scala b/sage-compat-ce/src/main/scala/kyo/compat/CLatch.scala index fda2d841..8a6954ce 100644 --- a/sage-compat-ce/src/main/scala/kyo/compat/CLatch.scala +++ b/sage-compat-ce/src/main/scala/kyo/compat/CLatch.scala @@ -16,17 +16,16 @@ object CLatch { * Allocates a latch with counter `n`; `n <= 0` is normalized to "already released". */ inline def init(inline n: Int): CIO[CLatch] = - CIO.lift(initImpl(n)) + CIO.lift( + if (n <= 0) CountDownLatch[IO](1).flatMap(l => l.release.as(l)) + else CountDownLatch[IO](n) + ) /** * Wraps a native `cats.effect.std.CountDownLatch` as a `CLatch`. The conversion is the identity on the carrier. */ inline def lift(inline u: CountDownLatch[IO]): CLatch = u - private inline def initImpl(inline n: Int): IO[CountDownLatch[IO]] = - if (n <= 0) CountDownLatch[IO](1).flatMap(l => l.release.as(l)) - else CountDownLatch[IO](n) - extension (inline self: CLatch) { /** diff --git a/sage-compat-ce/src/main/scala/kyo/compat/CLocal.scala b/sage-compat-ce/src/main/scala/kyo/compat/CLocal.scala index d81b0401..428ab91d 100644 --- a/sage-compat-ce/src/main/scala/kyo/compat/CLocal.scala +++ b/sage-compat-ce/src/main/scala/kyo/compat/CLocal.scala @@ -45,9 +45,7 @@ object CLocal { * Reads the current value, applies `f`, installs the result for the duration of `c`, then reverts. */ inline def update[B](inline f: A => A)(inline c: CIO[B]): CIO[B] = - CIO.lift(self.get.flatMap { cur => - self.set(f(cur)).bracket(_ => c.lower)(_ => self.set(cur)) - }) + CIO.lift(self.get.flatMap(cur => self.let(f(cur))(c).lower)) } } diff --git a/sage-compat-ce/src/main/scala/kyo/compat/CPromise.scala b/sage-compat-ce/src/main/scala/kyo/compat/CPromise.scala index 48f7176e..5456b75d 100644 --- a/sage-compat-ce/src/main/scala/kyo/compat/CPromise.scala +++ b/sage-compat-ce/src/main/scala/kyo/compat/CPromise.scala @@ -50,10 +50,7 @@ object CPromise { * Suspends until the promise is completed and returns its value. */ inline def get: CIO[A] = - CIO.lift(self.get.flatMap { - case Success(a) => IO.pure(a) - case Failure(e) => IO.raiseError(e) - }) + CIO.lift(self.get.flatMap(IO.fromTry)) /** * Returns the current state without blocking: `None` if pending, `Some(Try)` if completed. diff --git a/sage-core/src/main/scala/sage/SageEvent.scala b/sage-core/src/main/scala/sage/SageEvent.scala index e1495c2e..7e064480 100644 --- a/sage-core/src/main/scala/sage/SageEvent.scala +++ b/sage-core/src/main/scala/sage/SageEvent.scala @@ -15,7 +15,7 @@ sealed trait SageEvent object SageEvent { /** - * Reports a completed user command with its name, final node (`None` for standalone and all-master commands), client-observed duration, and + * Reports a completed user command with its name, final node (`None` for standalone, all-master and cross-slot split commands), client-observed duration, and * outcome. The duration includes cluster redirects and retries. A locally served cached read produces only [[Cache.Hit]]. A cache miss * produces [[Cache.Miss]] and a `CommandCompleted` after the server request finishes. */ diff --git a/sage-core/src/main/scala/sage/cluster/Redirect.scala b/sage-core/src/main/scala/sage/cluster/Redirect.scala index e881a943..311032f5 100644 --- a/sage-core/src/main/scala/sage/cluster/Redirect.scala +++ b/sage-core/src/main/scala/sage/cluster/Redirect.scala @@ -9,28 +9,20 @@ private[sage] enum RedirectKind { } /** - * A parsed `MOVED`/`ASK` reply. An empty target host means the IP of the current connection, which the runtime substitutes. + * A parsed `MOVED`/`ASK` reply. */ -final private[sage] case class Redirect(kind: RedirectKind, slot: Slot, target: Node) +final private[sage] case class Redirect(kind: RedirectKind, private val announced: Node) { -private[sage] object Redirect { + // an empty announced host means the node that sent the redirect (e.g. `MOVED 3999 :6381`) + def target(from: Node): Node = if (announced.host.isEmpty) Node(from.host, announced.port) else announced +} - def parse(error: String): Option[Redirect] = { - val parts = error.split(' ') - if (parts.length == 3) - for { - kind <- kindOf(parts(0)) - slot <- parts(1).toIntOption.flatMap(Slot.at) - target <- addressOf(parts(2)) - } yield Redirect(kind, slot, target) - else None - } +private[sage] object Redirect { - private def kindOf(token: String): Option[RedirectKind] = - token match { - case "MOVED" => Some(RedirectKind.Moved) - case "ASK" => Some(RedirectKind.Ask) - case _ => None + def parse(kind: RedirectKind, detail: String): Option[Redirect] = + detail.split(' ') match { + case Array(_, address) => addressOf(address).map(Redirect(kind, _)) + case _ => None } private def addressOf(address: String): Option[Node] = { diff --git a/sage-core/src/main/scala/sage/cluster/Route.scala b/sage-core/src/main/scala/sage/cluster/Route.scala index 8adf67d4..ca2f7d7e 100644 --- a/sage-core/src/main/scala/sage/cluster/Route.scala +++ b/sage-core/src/main/scala/sage/cluster/Route.scala @@ -6,9 +6,9 @@ package sage.cluster * its arguments. */ private[sage] enum Route { - case ToNode(node: Node, slot: Slot) + case ToNode(shard: Shard, slot: Slot) case Keyless case Unowned(slot: Slot) - case CrossSlot(slots: Set[Slot]) + case CrossSlot case Malformed } diff --git a/sage-core/src/main/scala/sage/cluster/Slot.scala b/sage-core/src/main/scala/sage/cluster/Slot.scala index 19ca9dc2..8496849b 100644 --- a/sage-core/src/main/scala/sage/cluster/Slot.scala +++ b/sage-core/src/main/scala/sage/cluster/Slot.scala @@ -16,11 +16,6 @@ private[sage] object Slot { */ def at(index: Int): Option[Slot] = if (index >= 0 && index < Count) Some(index) else None - /** - * Wrap an index already known in range (a hash modulo, a validated bound). Unchecked: out-of-range breaks topology lookup. - */ - private[sage] def unsafe(index: Int): Slot = index - def of(key: Bytes): Slot = { val bytes = key.unsafeArray // read-only; CRC16 never mutates val open = indexOf(bytes, '{'.toByte, 0) diff --git a/sage-core/src/main/scala/sage/cluster/SplitPlan.scala b/sage-core/src/main/scala/sage/cluster/SplitPlan.scala index 189406ab..b4b59998 100644 --- a/sage-core/src/main/scala/sage/cluster/SplitPlan.scala +++ b/sage-core/src/main/scala/sage/cluster/SplitPlan.scala @@ -1,20 +1,9 @@ package sage.cluster -final private[sage] case class NodeGroup(node: Node, positions: Vector[Int]) - -private[sage] enum Rejected { - case CrossSlot(slots: Set[Slot]) - case Unowned(slot: Slot) - case Malformed -} +final private[sage] case class NodeGroup(shard: Shard, positions: Vector[Int]) /** - * Describes how to run a pipeline across a cluster. `perNode` groups commands by target node and keeps their original positions. `keyless` - * contains commands that can be added to any node group. `rejected` contains commands that cannot be planned. Each pipeline position appears - * in exactly one of these collections. + * Describes how to run a pipeline across a cluster. `routes` holds the route of every pipeline position. `perNode` groups the positions + * routed to a node in their original order; the first group also holds every keyless position. With no group, keyless positions run alone. */ -final private[sage] case class SplitPlan( - perNode: Vector[NodeGroup], - keyless: Vector[Int], - rejected: Vector[(Int, Rejected)] -) +final private[sage] case class SplitPlan(routes: Vector[Route], perNode: Vector[NodeGroup]) diff --git a/sage-core/src/main/scala/sage/cluster/Topology.scala b/sage-core/src/main/scala/sage/cluster/Topology.scala index b1f89841..441a9a11 100644 --- a/sage-core/src/main/scala/sage/cluster/Topology.scala +++ b/sage-core/src/main/scala/sage/cluster/Topology.scala @@ -2,7 +2,7 @@ package sage.cluster import scala.collection.mutable -import sage.commands.{Command, Pipeline} +import sage.commands.Command /** * One server process in a cluster or master-replica deployment, addressed by host and port. [[sage.SageEvent]] values expose a `Node` to @@ -11,26 +11,32 @@ import sage.commands.{Command, Pipeline} final case class Node(host: String, port: Int) /** - * Inclusive on both ends. + * One `CLUSTER SLOTS` row: the slots from `start` to `end`, inclusive on both ends, and the nodes that serve them. */ -final private[sage] case class SlotRange(start: Slot, end: Slot) +final private[sage] case class SlotRange(start: Slot, end: Slot, master: Node, replicas: Vector[Node]) -final private[sage] case class Shard(master: Node, replicas: Vector[Node], slots: Vector[SlotRange]) +final private[sage] case class Shard(master: Node, replicas: Vector[Node]) /** * Records the master and replicas for each slot. Routing and pipeline splitting return a classification for every command. The runtime uses * that result to choose connections and decide whether to retry. */ -final private[sage] class ClusterTopology private (val shards: Vector[Shard], private val owners: Array[Node], shardOwners: Array[Shard]) { +final private[sage] class ClusterTopology private (private val owners: Array[Shard]) { - def nodeForSlot(slot: Slot): Option[Node] = Option(owners(slot.value)) + // derived from the owner array, in slot order, so broadcasts, scans and topology events use the same ownership as route + val shards: Vector[Shard] = owners.iterator.filter(_ != null).distinct.toVector - // compares the master for every slot so shard subscriptions stay in place when a refresh finds the same topology - def sameOwnership(other: ClusterTopology): Boolean = - java.util.Arrays.equals(owners.asInstanceOf[Array[AnyRef]], other.owners.asInstanceOf[Array[AnyRef]]) + val masters: Vector[Node] = shards.map(_.master) + + private def masterAt(slot: Int): Node = { + val shard = owners(slot) + if (shard == null) null else shard.master + } - // the core only locates the owning shard; selecting a live replica and applying the read policy is the runtime's job - def shardForSlot(slot: Slot): Option[Shard] = Option(shardOwners(slot.value)) + def nodeForSlot(slot: Slot): Option[Node] = Option(masterAt(slot.value)) + + // compares the master for every slot so shard subscriptions stay in place when a refresh finds the same topology + def sameOwnership(other: ClusterTopology): Boolean = (0 until Slot.Count).forall(slot => masterAt(slot) == other.masterAt(slot)) def replicasForMaster(master: Node): Vector[Node] = shards.find(_.master == master).fold(Vector.empty[Node])(_.replicas) @@ -39,8 +45,8 @@ final private[sage] class ClusterTopology private (val shards: Vector[Shard], pr val losing = mutable.Set.empty[Node] var slot = 0 while (slot < Slot.Count) { - val before = previous.owners(slot) - if (before != null && !before.equals(owners(slot))) losing += before + val before = previous.masterAt(slot) + if (before != null && !before.equals(masterAt(slot))) losing += before slot += 1 } losing.toSet @@ -48,72 +54,54 @@ final private[sage] class ClusterTopology private (val shards: Vector[Shard], pr def route(command: Command[?]): Route = if (command.hasMalformedKeys) Route.Malformed - else - command.keyIndices.length match { - case 0 => Route.Keyless - case 1 => routeSlot(Slot.of(command.args(command.keyIndices.head))) - case _ => - val keyIndices = command.keyIndices - val first = Slot.of(command.args(keyIndices.head)) - var i = 1 - var crossed = false - while (i < keyIndices.length && !crossed) { - if (Slot.of(command.args(keyIndices(i))) != first) crossed = true - i += 1 - } - if (crossed) Route.CrossSlot(slotsOf(command)) else routeSlot(first) - } - - def split(pipeline: Pipeline[?, ?]): SplitPlan = { - val perNode = mutable.LinkedHashMap.empty[Node, mutable.ArrayBuffer[Int]] - val keyless = mutable.ArrayBuffer.empty[Int] - val rejected = mutable.ArrayBuffer.empty[(Int, Rejected)] - pipeline.commands.iterator.zipWithIndex.foreach { case (command, index) => - route(command) match { - case Route.ToNode(node, _) => perNode.getOrElseUpdate(node, mutable.ArrayBuffer.empty) += index - case Route.Keyless => keyless += index - case Route.Unowned(slot) => rejected += ((index, Rejected.Unowned(slot))) - case Route.CrossSlot(slots) => rejected += ((index, Rejected.CrossSlot(slots))) - case Route.Malformed => rejected += ((index, Rejected.Malformed)) - } + else if (command.keyIndices.isEmpty) Route.Keyless + else { + val keyIndices = command.keyIndices + val first = Slot.of(command.args(keyIndices(0))) + var i = 1 + while (i < keyIndices.length && Slot.of(command.args(keyIndices(i))) == first) i += 1 + if (i < keyIndices.length) Route.CrossSlot else routeSlot(first) } - SplitPlan( - perNode.iterator.map { case (node, indices) => NodeGroup(node, indices.toVector) }.toVector, - keyless.toVector, - rejected.toVector - ) - } - private def routeSlot(slot: Slot): Route = - nodeForSlot(slot) match { - case Some(node) => Route.ToNode(node, slot) - case None => Route.Unowned(slot) + def split(commands: Vector[Command[?]]): SplitPlan = { + val routes = commands.map(route) + val perNode = mutable.LinkedHashMap.empty[Node, (Shard, mutable.ArrayBuffer[Int])] + // keyless positions join the first node's batch, which adopts this buffer so its positions stay ascending + val first = mutable.ArrayBuffer.empty[Int] + routes.iterator.zipWithIndex.foreach { + case (Route.ToNode(shard, _), index) => + perNode.getOrElseUpdate(shard.master, (shard, if (perNode.isEmpty) first else mutable.ArrayBuffer.empty))._2 += index + case (Route.Keyless, index) => first += index + case _ => () } + SplitPlan(routes, perNode.valuesIterator.map { case (shard, indices) => NodeGroup(shard, indices.toVector) }.toVector) + } - // safe to index args directly: route rejects out-of-range keyIndices as Malformed before calling this - private def slotsOf(command: Command[?]): Set[Slot] = - command.keyIndices.iterator.map(index => Slot.of(command.args(index))).toSet + private def routeSlot(slot: Slot): Route = { + val shard = owners(slot.value) + if (shard == null) Route.Unowned(slot) else Route.ToNode(shard, slot) + } } private[sage] object ClusterTopology { /** * Builds a topology even when slot ranges are incomplete or overlap. Uncovered slots remain unowned and trigger a refresh when routed. - * For overlapping ranges, the last listed shard owns the slot. The core preserves the topology reported by the server. + * For overlapping ranges, the last listed range owns the slot. A master's shard lists the replicas of all its ranges. */ - def from(shards: Vector[Shard]): ClusterTopology = { - val owners = new Array[Node](Slot.Count) - val shardOwners = new Array[Shard](Slot.Count) - shards.foreach { shard => - shard.slots.foreach { range => - var slot = range.start.value - while (slot <= range.end.value) { - owners(slot) = shard.master - shardOwners(slot) = shard - slot += 1 - } + def from(ranges: Vector[SlotRange]): ClusterTopology = { + val shardOf = ranges.groupMapReduce(_.master)(_.replicas)(_ ++ _).map((master, replicas) => master -> Shard(master, replicas.distinct)) + val owners = new Array[Shard](Slot.Count) + for { + range <- ranges + shard <- shardOf.get(range.master) + } { + var slot = range.start.value + while (slot <= range.end.value) { + owners(slot) = shard + slot += 1 } } - new ClusterTopology(shards, owners, shardOwners) + new ClusterTopology(owners) } } diff --git a/sage-core/src/main/scala/sage/codec/Doubles.scala b/sage-core/src/main/scala/sage/codec/Doubles.scala index e07c38a2..e923c797 100644 --- a/sage-core/src/main/scala/sage/codec/Doubles.scala +++ b/sage-core/src/main/scala/sage/codec/Doubles.scala @@ -8,10 +8,9 @@ package sage.codec private[sage] object Doubles { def format(value: Double): String = - special(value == Double.PositiveInfinity, value == Double.NegativeInfinity, value.isNaN).getOrElse(value.toString) + if (value == Double.PositiveInfinity) "inf" else if (value == Double.NegativeInfinity) "-inf" else if (value.isNaN) "nan" else value.toString - def formatFloat(value: Float): String = - special(value == Float.PositiveInfinity, value == Float.NegativeInfinity, value.isNaN).getOrElse(value.toString) + def formatFloat(value: Float): String = if (value.isNaN || value.isInfinite) format(value.toDouble) else value.toString def parse(text: String): Option[Double] = parseWith(text)(Double.PositiveInfinity, Double.NegativeInfinity, Double.NaN, _.toDoubleOption) @@ -19,9 +18,6 @@ private[sage] object Doubles { def parseFloat(text: String): Option[Float] = parseWith(text)(Float.PositiveInfinity, Float.NegativeInfinity, Float.NaN, _.toFloatOption) - private def special(posInf: Boolean, negInf: Boolean, nan: Boolean): Option[String] = - if (posInf) Some("inf") else if (negInf) Some("-inf") else if (nan) Some("nan") else None - private def parseWith[A](text: String)(posInf: A, negInf: A, nan: A, fallback: String => Option[A]): Option[A] = text match { case "inf" | "+inf" => Some(posInf) diff --git a/sage-core/src/main/scala/sage/codec/KeyCodec.scala b/sage-core/src/main/scala/sage/codec/KeyCodec.scala index bb60df33..c988f361 100644 --- a/sage-core/src/main/scala/sage/codec/KeyCodec.scala +++ b/sage-core/src/main/scala/sage/codec/KeyCodec.scala @@ -44,38 +44,36 @@ object KeyCodec { /** * Builds a key codec from encode and decode functions. The decoder returns `Either` for invalid input. */ - def from[A](enc: A => Bytes)(dec: Bytes => Either[DecodeError, A]): KeyCodec[A] = instance(enc, dec) + def from[A](enc: A => Bytes)(dec: Bytes => Either[DecodeError, A]): KeyCodec[A] = + new KeyCodec[A] { + + def encode(value: A): Bytes = enc(value) + + def decode(bytes: Bytes): Either[DecodeError, A] = dec(bytes) + } /** * UTF-8 text; decoding rejects malformed UTF-8. */ - given string: KeyCodec[String] = instance(Bytes.utf8, Primitives.decodeUtf8) + given string: KeyCodec[String] = from(Bytes.utf8)(Primitives.decodeUtf8) /** * Decimal `Int`; decoding rejects non-numeric or out-of-range input. */ - given int: KeyCodec[Int] = instance(Primitives.encodeInt, Primitives.decodeNumber("Int", Primitives.parseInt)) + given int: KeyCodec[Int] = from(Primitives.encodeInt)(Primitives.decodeLong("Int", Int.MinValue, Int.MaxValue)(_).map(_.toInt)) /** * Decimal `Long`; decoding rejects non-numeric or out-of-range input. */ - given long: KeyCodec[Long] = instance(Primitives.encodeLong, Primitives.decodeNumber("Long", Primitives.parseLong)) + given long: KeyCodec[Long] = from(Primitives.encodeLong)(Primitives.decodeLong("Long", Long.MinValue, Long.MaxValue)) /** * Raw [[sage.Bytes]], passed through unchanged in both directions. */ - given bytes: KeyCodec[Bytes] = instance(identity, Right(_)) + given bytes: KeyCodec[Bytes] = from[Bytes](identity)(Right(_)) /** * Raw `Array[Byte]`, copied at the boundary in both directions (see [[sage.Bytes.fromArray]]/[[sage.Bytes.toArray]]). */ - given byteArray: KeyCodec[Array[Byte]] = instance(Bytes.fromArray, raw => Right(raw.toArray)) - - private def instance[A](enc: A => Bytes, dec: Bytes => Either[DecodeError, A]): KeyCodec[A] = - new KeyCodec[A] { - - def encode(value: A): Bytes = enc(value) - - def decode(bytes: Bytes): Either[DecodeError, A] = dec(bytes) - } + given byteArray: KeyCodec[Array[Byte]] = from(Bytes.fromArray)(raw => Right(raw.toArray)) } diff --git a/sage-core/src/main/scala/sage/codec/Primitives.scala b/sage-core/src/main/scala/sage/codec/Primitives.scala index 23685ecb..1c6b6508 100644 --- a/sage-core/src/main/scala/sage/codec/Primitives.scala +++ b/sage-core/src/main/scala/sage/codec/Primitives.scala @@ -7,45 +7,53 @@ import java.util.Arrays import sage.Bytes import sage.SageException.DecodeError -private[codec] object Primitives { +private[sage] object Primitives { - private val True = Bytes.utf8("1") - private val False = Bytes.utf8("0") private val Zero = Bytes.utf8("0") + private val One = Bytes.utf8("1") private val MaxPreview = 64 def encodeInt(value: Int): Bytes = encodeLong(value.toLong) - // digits come from the negative magnitude (`value % 10` is <= 0), so Long.MinValue, which has no positive counterpart, is safe def encodeLong(value: Long): Bytes = if (value == 0L) Zero + else if (value == 1L) One else { - val negative = value < 0 - var counter = value - var digits = 0 - while (counter != 0L) { - counter /= 10 - digits += 1 - } - val start = if (negative) 1 else 0 - val out = new Array[Byte](start + digits) - var i = out.length - 1 - var remaining = value - while (i >= start) { - val digit = if (negative) -(remaining % 10) else remaining % 10 - out(i) = ('0' + digit).toByte - remaining /= 10 - i -= 1 - } - if (negative) out(0) = '-' + val start = if (value < 0) 1 else 0 + val digits = digitCount(value) + val out = new Array[Byte](start + digits) + writeDigits(out, start, digits, value) + if (start == 1) out(0) = '-' Bytes.wrap(IArray.unsafeFromArray(out)) } - def encodeBoolean(value: Boolean): Bytes = if (value) True else False + // digits come from the negative magnitude, so Long.MinValue, which has no positive counterpart, is safe + def digitCount(value: Long): Int = { + val magnitude = if (value < 0) value else -value + var digits = 1 + var floor = -10L + while (digits < 19 && magnitude <= floor) { + digits += 1 + floor *= 10 + } + digits + } + + def writeDigits(out: Array[Byte], start: Int, digits: Int, value: Long): Unit = { + var magnitude = if (value < 0) value else -value + var i = start + digits - 1 + while (i >= start) { + out(i) = ('0' - magnitude % 10).toByte + magnitude /= 10 + i -= 1 + } + } + + def encodeBoolean(value: Boolean): Bytes = if (value) One else Zero def decodeBoolean(bytes: Bytes): Either[DecodeError, Boolean] = - if (bytes.sameBytes(True)) Right(true) - else if (bytes.sameBytes(False)) Right(false) + if (bytes.sameBytes(One)) Right(true) + else if (bytes.sameBytes(Zero)) Right(false) else Left(DecodeError("boolean (1 or 0)", preview(bytes))) def decodeUtf8(bytes: Bytes): Either[DecodeError, String] = { @@ -68,9 +76,23 @@ private[codec] object Primitives { def decodeNumber[A](expected: String, parse: String => Option[A])(bytes: Bytes): Either[DecodeError, A] = parse(bytes.asUtf8String).toRight(DecodeError(expected, preview(bytes))) - // reject a leading '+' so "5"/"+5" don't decode to the same key - def parseInt(text: String): Option[Int] = if (text.startsWith("+")) None else text.toIntOption - def parseLong(text: String): Option[Long] = if (text.startsWith("+")) None else text.toLongOption + // ASCII digits with an optional '-'; a leading '+' is rejected so "5" and "+5" don't decode to the same key + def decodeLong(expected: String, min: Long, max: Long)(bytes: Bytes): Either[DecodeError, Long] = { + val a = bytes.unsafeArray + val negative = a.length > 1 && a(0) == '-' + var i = if (negative) 1 else 0 + var ok = i < a.length + var acc = 0L // accumulates the negated value so Long.MinValue fits + while (ok && i < a.length) { + val digit = a(i) - '0' + ok = digit >= 0 && digit <= 9 && acc >= (Long.MinValue + digit) / 10 + if (ok) acc = acc * 10 - digit + i += 1 + } + val value = if (negative) acc else -acc + if (ok && (negative || acc != Long.MinValue) && value >= min && value <= max) Right(value) + else Left(DecodeError(expected, preview(bytes))) + } def preview(bytes: Bytes): String = { // MaxPreview code points never need more than 4 bytes each, so decoding a bounded window avoids diff --git a/sage-core/src/main/scala/sage/codec/ValueCodec.scala b/sage-core/src/main/scala/sage/codec/ValueCodec.scala index f29468a3..fabe841d 100644 --- a/sage-core/src/main/scala/sage/codec/ValueCodec.scala +++ b/sage-core/src/main/scala/sage/codec/ValueCodec.scala @@ -44,55 +44,53 @@ object ValueCodec { /** * Builds a codec from encode and decode functions. The decoder returns `Either` for invalid input. */ - def from[A](enc: A => Bytes)(dec: Bytes => Either[DecodeError, A]): ValueCodec[A] = instance(enc, dec) + def from[A](enc: A => Bytes)(dec: Bytes => Either[DecodeError, A]): ValueCodec[A] = + new ValueCodec[A] { + + def encode(value: A): Bytes = enc(value) + + def decode(bytes: Bytes): Either[DecodeError, A] = dec(bytes) + } /** * UTF-8 text; decoding rejects malformed UTF-8. */ - given string: ValueCodec[String] = instance(Bytes.utf8, Primitives.decodeUtf8) + given string: ValueCodec[String] = from(Bytes.utf8)(Primitives.decodeUtf8) /** * Decimal `Int`; decoding rejects non-numeric or out-of-range input. */ - given int: ValueCodec[Int] = instance(Primitives.encodeInt, Primitives.decodeNumber("Int", Primitives.parseInt)) + given int: ValueCodec[Int] = from(Primitives.encodeInt)(Primitives.decodeLong("Int", Int.MinValue, Int.MaxValue)(_).map(_.toInt)) /** * Decimal `Long`; decoding rejects non-numeric or out-of-range input. */ - given long: ValueCodec[Long] = instance(Primitives.encodeLong, Primitives.decodeNumber("Long", Primitives.parseLong)) + given long: ValueCodec[Long] = from(Primitives.encodeLong)(Primitives.decodeLong("Long", Long.MinValue, Long.MaxValue)) /** * `Double` in Redis's number format, including `inf`/`-inf`/`nan`. */ given double: ValueCodec[Double] = - instance(d => Bytes.utf8(Doubles.format(d)), Primitives.decodeNumber("Double", Doubles.parse)) + from[Double](d => Bytes.utf8(Doubles.format(d)))(Primitives.decodeNumber("Double", Doubles.parse)) /** * `Float` in Redis's number format, including `inf`/`-inf`/`nan`. */ given float: ValueCodec[Float] = - instance(f => Bytes.utf8(Doubles.formatFloat(f)), Primitives.decodeNumber("Float", Doubles.parseFloat)) + from[Float](f => Bytes.utf8(Doubles.formatFloat(f)))(Primitives.decodeNumber("Float", Doubles.parseFloat)) /** * `1`/`0` on the wire; decoding accepts only those two tokens. */ - given boolean: ValueCodec[Boolean] = instance(Primitives.encodeBoolean, Primitives.decodeBoolean) + given boolean: ValueCodec[Boolean] = from(Primitives.encodeBoolean)(Primitives.decodeBoolean) /** * Raw [[sage.Bytes]], passed through unchanged in both directions. */ - given bytes: ValueCodec[Bytes] = instance(identity, Right(_)) + given bytes: ValueCodec[Bytes] = from[Bytes](identity)(Right(_)) /** * Raw `Array[Byte]`, copied at the boundary in both directions (see [[sage.Bytes.fromArray]]/[[sage.Bytes.toArray]]). */ - given byteArray: ValueCodec[Array[Byte]] = instance(Bytes.fromArray, raw => Right(raw.toArray)) - - private def instance[A](enc: A => Bytes, dec: Bytes => Either[DecodeError, A]): ValueCodec[A] = - new ValueCodec[A] { - - def encode(value: A): Bytes = enc(value) - - def decode(bytes: Bytes): Either[DecodeError, A] = dec(bytes) - } + given byteArray: ValueCodec[Array[Byte]] = from(Bytes.fromArray)(raw => Right(raw.toArray)) } diff --git a/sage-core/src/main/scala/sage/commands/Acl.scala b/sage-core/src/main/scala/sage/commands/Acl.scala index 4433719c..bdd595c0 100644 --- a/sage-core/src/main/scala/sage/commands/Acl.scala +++ b/sage-core/src/main/scala/sage/commands/Acl.scala @@ -55,27 +55,24 @@ private[sage] object Acl { Command("ACL", Command.NoKeys, Vector(GetUser, Bytes.utf8(username)), decodeUser) def aclLog(count: Option[Long] = None): Command[Vector[AclLogEntry]] = - Command("ACL", Command.NoKeys, Log +: count.map(n => Bytes.utf8(n.toString)).toVector, Decode.vector(decodeLogEntry)) + Command("ACL", Command.NoKeys, Log +: count.map(n => Args.long(n)).toVector, Decode.vector(decodeLogEntry)) - private val decodeUser: Frame => Either[DecodeError, Option[AclUser]] = { - case Frame.Null => Right(None) - case frame => + private val decodeUser: Frame => Either[DecodeError, Option[AclUser]] = + Decode.nullable { frame => Decode.fieldMap(frame).map { fields => - Some( - AclUser( - flags = strings(fields.get("flags")), - passwords = strings(fields.get("passwords")), - commands = string(fields, "commands"), - keys = string(fields, "keys"), - channels = string(fields, "channels"), - selectors = fields.get("selectors") match { - case Some(Frame.Array(rows)) => rows.flatMap(selectorOf) - case _ => Vector.empty - } - ) + AclUser( + flags = strings(fields.get("flags")), + passwords = strings(fields.get("passwords")), + commands = string(fields, "commands"), + keys = string(fields, "keys"), + channels = string(fields, "channels"), + selectors = fields.get("selectors") match { + case Some(Frame.Array(rows)) => rows.flatMap(selectorOf) + case _ => Vector.empty + } ) } - } + } private def selectorOf(frame: Frame): Option[Map[String, String]] = Decode.fieldMap(frame).toOption.map(_.collect { case (k, Frame.BulkString(v)) => k -> v.asUtf8String }) @@ -88,21 +85,14 @@ private[sage] object Acl { context = string(fields, "context"), obj = string(fields, "object"), username = string(fields, "username"), - ageSeconds = fields.get("age-seconds").flatMap(asDouble).getOrElse(0.0), + ageSeconds = fields.get("age-seconds").collect { case Frame.Double(v) => v }.getOrElse(0.0), clientInfo = string(fields, "client-info"), entryId = long(fields, "entry-id") ) } private def string(fields: Map[String, Frame], key: String): String = - fields - .get(key) - .flatMap { - case Frame.BulkString(b) => Some(b.asUtf8String) - case Frame.SimpleString(s) => Some(s) - case _ => None - } - .getOrElse("") + fields.get(key).collect { case Decode.Text(s) => s }.getOrElse("") private def long(fields: Map[String, Frame], key: String): Long = fields.get(key).collect { case Frame.Integer(n) => n }.getOrElse(0L) @@ -113,12 +103,4 @@ private[sage] object Acl { case Some(Frame.Set(elements)) => elements.collect { case Frame.BulkString(b) => b.asUtf8String } case _ => Vector.empty } - - private def asDouble(frame: Frame): Option[Double] = - frame match { - case Frame.Double(v) => Some(v) - case Frame.Integer(v) => Some(v.toDouble) - case Frame.BulkString(b) => b.asUtf8String.toDoubleOption - case _ => None - } } diff --git a/sage-core/src/main/scala/sage/commands/Args.scala b/sage-core/src/main/scala/sage/commands/Args.scala new file mode 100644 index 00000000..90b1373c --- /dev/null +++ b/sage-core/src/main/scala/sage/commands/Args.scala @@ -0,0 +1,65 @@ +package sage.commands + +import sage.Bytes +import sage.codec.{Doubles, KeyCodec, Primitives, ValueCodec} + +// argument keywords and option encoders shared by the command families +private[commands] object Args { + + val Match: Bytes = Bytes.utf8("MATCH") + val Count: Bytes = Bytes.utf8("COUNT") + val LimitWord: Bytes = Bytes.utf8("LIMIT") + val Nx: Bytes = Bytes.utf8("NX") + val Xx: Bytes = Bytes.utf8("XX") + val Gt: Bytes = Bytes.utf8("GT") + val Lt: Bytes = Bytes.utf8("LT") + val Ch: Bytes = Bytes.utf8("CH") + val Get: Bytes = Bytes.utf8("GET") + val Replace: Bytes = Bytes.utf8("REPLACE") + val Rev: Bytes = Bytes.utf8("REV") + val Min: Bytes = Bytes.utf8("MIN") + val Max: Bytes = Bytes.utf8("MAX") + val Desc: Bytes = Bytes.utf8("DESC") + val WithValues: Bytes = Bytes.utf8("WITHVALUES") + + def long(value: Long): Bytes = Primitives.encodeLong(value) + + def double(value: Double): Bytes = Bytes.utf8(Doubles.format(value)) + + // the keyword of each case of a fieldless enum is its upper-cased name + def keywords[E <: scala.reflect.Enum](values: Array[E]): E => Bytes = { + val words = values.toVector.map(value => Bytes.utf8(value.toString.toUpperCase(java.util.Locale.ROOT))) + value => words(value.ordinal) + } + + def pairs[A, B](entries: Vector[(A, B)])(using a: KeyCodec[A], b: ValueCodec[B]): Vector[Bytes] = + entries.flatMap { case (x, y) => Vector(a.encode(x), b.encode(y)) } + + def keyThen[K, A](key: K, first: A, rest: Seq[A])(encode: A => Bytes)(using keyCodec: KeyCodec[K]): Vector[Bytes] = { + val args = Vector.newBuilder[Bytes] + args.sizeHint(rest.length + 2) + args += keyCodec.encode(key) += encode(first) + rest.foreach(a => args += encode(a)) + args.result() + } + + def flag(on: Boolean, word: Bytes): Vector[Bytes] = if (on) Vector(word) else Vector.empty + + def opt[A](word: Bytes, value: Option[A])(encode: A => Bytes): Vector[Bytes] = + value match { + case Some(a) => Vector(word, encode(a)) + case None => Vector.empty + } + + def optLong(word: Bytes, value: Option[Long]): Vector[Bytes] = opt(word, value)(long) + + def optText(word: Bytes, value: Option[String]): Vector[Bytes] = opt(word, value)(Bytes.utf8) + + def limit(value: Option[Limit]): Vector[Bytes] = + value match { + case Some(l) => Vector(LimitWord, Args.long(l.offset), Args.long(l.count)) + case None => Vector.empty + } + + def scanOptions(pattern: Option[String], count: Option[Long]): Vector[Bytes] = optText(Match, pattern) ++ optLong(Count, count) +} diff --git a/sage-core/src/main/scala/sage/commands/Arrays.scala b/sage-core/src/main/scala/sage/commands/Arrays.scala index b03e2df3..b5b1c701 100644 --- a/sage-core/src/main/scala/sage/commands/Arrays.scala +++ b/sage-core/src/main/scala/sage/commands/Arrays.scala @@ -3,6 +3,7 @@ package sage.commands import sage.Bytes import sage.SageException.DecodeError import sage.codec.{KeyCodec, ValueCodec} +import sage.commands.Args.{LimitWord, Match, Max, Min, Rev, WithValues} import sage.protocol.Frame /** @@ -62,24 +63,19 @@ final case class ArrayInfoFull( */ private[sage] object Arrays { - private val Rev = Bytes.utf8("REV") - private val WithValues = Bytes.utf8("WITHVALUES") - private val LimitWord = Bytes.utf8("LIMIT") private val NoCaseWord = Bytes.utf8("NOCASE") private val And = Bytes.utf8("AND") private val Or = Bytes.utf8("OR") + private val combineArg = Args.keywords(ArGrepCombine.values) private val Full = Bytes.utf8("FULL") private val Exact = Bytes.utf8("EXACT") - private val MatchWord = Bytes.utf8("MATCH") private val Glob = Bytes.utf8("GLOB") private val Re = Bytes.utf8("RE") private val Sum = Bytes.utf8("SUM") - private val Min = Bytes.utf8("MIN") - private val Max = Bytes.utf8("MAX") private val Xor = Bytes.utf8("XOR") private val Used = Bytes.utf8("USED") - private def idx(i: Long): Bytes = Bytes.utf8(i.toString) + private def idx(i: Long): Bytes = Args.long(i) def arSet[K, V](key: K, index: Long, first: V, rest: V*)(using k: KeyCodec[K], v: ValueCodec[V]): Command[Long] = Command("ARSET", Command.FirstKey, k.encode(key) +: idx(index) +: (first +: rest.toVector).map(v.encode), Decode.long) @@ -96,7 +92,7 @@ private[sage] object Arrays { Command.read("ARGET", Command.FirstKey, Vector(k.encode(key), idx(index)), Decode.optionalValue) def arMGet[K, V](key: K, first: Long, rest: Long*)(using k: KeyCodec[K], v: ValueCodec[V]): Command[Vector[Option[V]]] = - Command.read("ARMGET", Command.FirstKey, k.encode(key) +: (first +: rest.toVector).map(idx), Decode.vector(Decode.optionalValue)) + Command.read("ARMGET", Command.FirstKey, Args.keyThen(key, first, rest)(idx), Decode.vector(Decode.optionalValue)) def arLen[K](key: K)(using k: KeyCodec[K]): Command[Long] = Command.read("ARLEN", Command.FirstKey, Vector(k.encode(key)), Decode.long) @@ -114,12 +110,12 @@ private[sage] object Arrays { Command.read( "ARLASTITEMS", Command.FirstKey, - Vector(k.encode(key), idx(count)) ++ (if (rev) Vector(Rev) else Vector.empty), + Vector(k.encode(key), idx(count)) ++ Args.flag(rev, Rev), Decode.vector(Decode.value) ) def arDel[K](key: K, first: Long, rest: Long*)(using k: KeyCodec[K]): Command[Long] = - Command("ARDEL", Command.FirstKey, k.encode(key) +: (first +: rest.toVector).map(idx), Decode.long) + Command("ARDEL", Command.FirstKey, Args.keyThen(key, first, rest)(idx), Decode.long) def arDelRange[K](key: K, first: (Long, Long), rest: (Long, Long)*)(using k: KeyCodec[K]): Command[Long] = Command( @@ -130,7 +126,7 @@ private[sage] object Arrays { ) def arInsert[K, V](key: K, first: V, rest: V*)(using k: KeyCodec[K], v: ValueCodec[V]): Command[Long] = - Command("ARINSERT", Command.FirstKey, k.encode(key) +: (first +: rest.toVector).map(v.encode), Decode.long) + Command("ARINSERT", Command.FirstKey, Args.keyThen(key, first, rest)(v.encode), Decode.long) def arNext[K](key: K)(using k: KeyCodec[K]): Command[Option[Long]] = Command.read("ARNEXT", Command.FirstKey, Vector(k.encode(key)), Decode.optionalLong) @@ -142,7 +138,7 @@ private[sage] object Arrays { Command.read( "ARSCAN", Command.FirstKey, - Vector(k.encode(key), idx(start), idx(end)) ++ limit.toVector.flatMap(l => Vector(LimitWord, idx(l))), + Vector(k.encode(key), idx(start), idx(end)) ++ Args.optLong(LimitWord, limit), indexValuePairs ) @@ -153,7 +149,7 @@ private[sage] object Arrays { Command.read( "ARGREP", Command.FirstKey, - grepArgs(k.encode(key), start, end, first +: rest.toVector, combine, limit, noCase, withValues = false), + grepArgs(k.encode(key), start, end, first, rest, combine, limit, noCase, withValues = false), Decode.vector(Decode.long) ) @@ -168,22 +164,21 @@ private[sage] object Arrays { Command.read( "ARGREP", Command.FirstKey, - grepArgs(k.encode(key), start, end, first +: rest.toVector, combine, limit, noCase, withValues = true), + grepArgs(k.encode(key), start, end, first, rest, combine, limit, noCase, withValues = true), indexValuePairs ) - def arOpSum[K](key: K, start: Long, end: Long)(using KeyCodec[K]): Command[Option[Double]] = aropDouble(key, start, end, Sum) - def arOpMin[K](key: K, start: Long, end: Long)(using KeyCodec[K]): Command[Option[Double]] = aropDouble(key, start, end, Min) - def arOpMax[K](key: K, start: Long, end: Long)(using KeyCodec[K]): Command[Option[Double]] = aropDouble(key, start, end, Max) - def arOpAnd[K](key: K, start: Long, end: Long)(using KeyCodec[K]): Command[Option[Long]] = aropLong(key, start, end, And) - def arOpOr[K](key: K, start: Long, end: Long)(using KeyCodec[K]): Command[Option[Long]] = aropLong(key, start, end, Or) - def arOpXor[K](key: K, start: Long, end: Long)(using KeyCodec[K]): Command[Option[Long]] = aropLong(key, start, end, Xor) + def arOpSum[K](key: K, start: Long, end: Long)(using KeyCodec[K]): Command[Option[Double]] = arop(key, start, end, Sum, Decode.optionalDouble) + def arOpMin[K](key: K, start: Long, end: Long)(using KeyCodec[K]): Command[Option[Double]] = arop(key, start, end, Min, Decode.optionalDouble) + def arOpMax[K](key: K, start: Long, end: Long)(using KeyCodec[K]): Command[Option[Double]] = arop(key, start, end, Max, Decode.optionalDouble) + def arOpAnd[K](key: K, start: Long, end: Long)(using KeyCodec[K]): Command[Option[Long]] = arop(key, start, end, And, Decode.optionalLong) + def arOpOr[K](key: K, start: Long, end: Long)(using KeyCodec[K]): Command[Option[Long]] = arop(key, start, end, Or, Decode.optionalLong) + def arOpXor[K](key: K, start: Long, end: Long)(using KeyCodec[K]): Command[Option[Long]] = arop(key, start, end, Xor, Decode.optionalLong) - def arOpUsed[K](key: K, start: Long, end: Long)(using k: KeyCodec[K]): Command[Long] = - Command.read("AROP", Command.FirstKey, Vector(k.encode(key), idx(start), idx(end), Used), Decode.long) + def arOpUsed[K](key: K, start: Long, end: Long)(using KeyCodec[K]): Command[Long] = arop(key, start, end, Used, Decode.long) def arOpMatch[K, V](key: K, start: Long, end: Long, value: V)(using k: KeyCodec[K], v: ValueCodec[V]): Command[Long] = - Command.read("AROP", Command.FirstKey, Vector(k.encode(key), idx(start), idx(end), MatchWord, v.encode(value)), Decode.long) + Command.read("AROP", Command.FirstKey, Vector(k.encode(key), idx(start), idx(end), Match, v.encode(value)), Decode.long) def arInfo[K](key: K)(using k: KeyCodec[K]): Command[ArrayInfo] = Command.read("ARINFO", Command.FirstKey, Vector(k.encode(key)), decodeInfo) @@ -193,37 +188,31 @@ private[sage] object Arrays { // --- helpers --------------------------------------------------------------------------------------------------------------------------- - private def aropDouble[K](key: K, start: Long, end: Long, op: Bytes)(using k: KeyCodec[K]): Command[Option[Double]] = - Command.read("AROP", Command.FirstKey, Vector(k.encode(key), idx(start), idx(end), op), Decode.optionalDouble) - - private def aropLong[K](key: K, start: Long, end: Long, op: Bytes)(using k: KeyCodec[K]): Command[Option[Long]] = - Command.read("AROP", Command.FirstKey, Vector(k.encode(key), idx(start), idx(end), op), Decode.optionalLong) + private def arop[K, A](key: K, start: Long, end: Long, op: Bytes, decode: Frame => Either[DecodeError, A])(using k: KeyCodec[K]): Command[A] = + Command.read("AROP", Command.FirstKey, Vector(k.encode(key), idx(start), idx(end), op), decode) private def grepArgs( key: Bytes, start: Long, end: Long, - predicates: Vector[ArMatch], + first: ArMatch, + rest: Seq[ArMatch], combine: ArGrepCombine, limit: Option[Long], noCase: Boolean, withValues: Boolean ): Vector[Bytes] = { - val combineToken = combine match { - case ArGrepCombine.And => And - case ArGrepCombine.Or => Or - } - val predTokens = predicates.map(predicateArgs).reduce((a, b) => (a :+ combineToken) ++ b) + val predTokens = rest.foldLeft(predicateArgs(first))((acc, p) => (acc :+ combineArg(combine)) ++ predicateArgs(p)) Vector(key, idx(start), idx(end)) ++ predTokens ++ - limit.toVector.flatMap(l => Vector(LimitWord, idx(l))) ++ - (if (withValues) Vector(WithValues) else Vector.empty) ++ - (if (noCase) Vector(NoCaseWord) else Vector.empty) + Args.optLong(LimitWord, limit) ++ + Args.flag(withValues, WithValues) ++ + Args.flag(noCase, NoCaseWord) } private def predicateArgs(predicate: ArMatch): Vector[Bytes] = predicate match { case ArMatch.Exact(value) => Vector(Exact, Bytes.utf8(value)) - case ArMatch.Match(substring) => Vector(MatchWord, Bytes.utf8(substring)) + case ArMatch.Match(substring) => Vector(Match, Bytes.utf8(substring)) case ArMatch.Glob(pattern) => Vector(Glob, Bytes.utf8(pattern)) case ArMatch.Re(pattern) => Vector(Re, Bytes.utf8(pattern)) } @@ -235,58 +224,35 @@ private[sage] object Arrays { case other => Left(DecodeError("[index, value] pair", Frame.describe(other))) } - private val decodeInfo: Frame => Either[DecodeError, ArrayInfo] = - frame => - Decode.fieldMap(frame).flatMap { fields => - for { - count <- longField(fields, "count") - len <- longField(fields, "len") - next <- longField(fields, "next-insert-index") - } yield ArrayInfo( - count, - len, - next, - optLong(fields, "slices"), - optLong(fields, "directory-size"), - optLong(fields, "super-dir-entries"), - optLong(fields, "slice-size") - ) - } - private val decodeInfoFull: Frame => Either[DecodeError, ArrayInfoFull] = - frame => - Decode.fieldMap(frame).flatMap { fields => - for { - count <- longField(fields, "count") - len <- longField(fields, "len") - next <- longField(fields, "next-insert-index") - } yield ArrayInfoFull( - count, - len, - next, - optLong(fields, "slices"), - optLong(fields, "directory-size"), - optLong(fields, "super-dir-entries"), - optLong(fields, "slice-size"), - optLong(fields, "dense-slices"), - optLong(fields, "sparse-slices"), - optDouble(fields, "avg-dense-size"), - optDouble(fields, "avg-dense-fill"), - optDouble(fields, "avg-sparse-size") - ) - } - - private def longField(fields: Map[String, Frame], name: String): Either[DecodeError, Long] = - fields.get(name) match { - case Some(Frame.Integer(n)) => Right(n) - case Some(other) => Left(DecodeError(s"ARINFO $name integer", Frame.describe(other))) - case None => Left(DecodeError(s"ARINFO $name", "absent")) + Decode.fields { fields => + for { + count <- fields.required("count", Decode.long) + len <- fields.required("len", Decode.long) + next <- fields.required("next-insert-index", Decode.long) + } yield ArrayInfoFull( + count, + len, + next, + optLong(fields, "slices"), + optLong(fields, "directory-size"), + optLong(fields, "super-dir-entries"), + optLong(fields, "slice-size"), + optLong(fields, "dense-slices"), + optLong(fields, "sparse-slices"), + optDouble(fields, "avg-dense-size"), + optDouble(fields, "avg-dense-fill"), + optDouble(fields, "avg-sparse-size") + ) } - private def optLong(fields: Map[String, Frame], name: String): Option[Long] = + private val decodeInfo: Frame => Either[DecodeError, ArrayInfo] = + frame => decodeInfoFull(frame).map(f => ArrayInfo(f.count, f.len, f.nextInsertIndex, f.slices, f.directorySize, f.superDirEntries, f.sliceSize)) + + private def optLong(fields: Decode.Fields, name: String): Option[Long] = fields.get(name).collect { case Frame.Integer(n) => n } - private def optDouble(fields: Map[String, Frame], name: String): Option[Double] = + private def optDouble(fields: Decode.Fields, name: String): Option[Double] = fields.get(name).collect { case Frame.Double(d) => d case Frame.Integer(n) => n.toDouble diff --git a/sage-core/src/main/scala/sage/commands/Bitmaps.scala b/sage-core/src/main/scala/sage/commands/Bitmaps.scala index 6daf8458..e01fd937 100644 --- a/sage-core/src/main/scala/sage/commands/Bitmaps.scala +++ b/sage-core/src/main/scala/sage/commands/Bitmaps.scala @@ -1,7 +1,8 @@ package sage.commands import sage.Bytes -import sage.codec.KeyCodec +import sage.codec.{KeyCodec, Primitives} +import sage.commands.Args.Get /** * Whether a `BITCOUNT`/`BITPOS` range is measured in `Byte`s or `Bit`s. @@ -60,33 +61,27 @@ enum BitFieldOp { private[sage] object Bitmaps { - private val Zero = Bytes.utf8("0") - private val One = Bytes.utf8("1") private val And = Bytes.utf8("AND") private val Or = Bytes.utf8("OR") private val Xor = Bytes.utf8("XOR") private val Not = Bytes.utf8("NOT") - private val GetWord = Bytes.utf8("GET") private val SetWord = Bytes.utf8("SET") private val IncrByWord = Bytes.utf8("INCRBY") private val OverflowWord = Bytes.utf8("OVERFLOW") - private val WrapWord = Bytes.utf8("WRAP") - private val SatWord = Bytes.utf8("SAT") - private val FailWord = Bytes.utf8("FAIL") - private val ByteWord = Bytes.utf8("BYTE") - private val BitWord = Bytes.utf8("BIT") + private val overflowArg = Args.keywords(BitFieldOverflow.values) + private val unitArg = Args.keywords(BitUnit.values) def setBit[K](key: K, offset: Long, value: Boolean)(using keyCodec: KeyCodec[K]): Command[Boolean] = - Command("SETBIT", Command.FirstKey, Vector(keyCodec.encode(key), Bytes.utf8(offset.toString), bitToken(value)), Decode.flag) + Command("SETBIT", Command.FirstKey, Vector(keyCodec.encode(key), Args.long(offset), Primitives.encodeBoolean(value)), Decode.flag) def getBit[K](key: K, offset: Long)(using keyCodec: KeyCodec[K]): Command[Boolean] = - Command.read("GETBIT", Command.FirstKey, Vector(keyCodec.encode(key), Bytes.utf8(offset.toString)), Decode.flag) + Command.read("GETBIT", Command.FirstKey, Vector(keyCodec.encode(key), Args.long(offset)), Decode.flag) def bitCount[K](key: K, range: Option[BitRange] = None)(using keyCodec: KeyCodec[K]): Command[Long] = Command.read("BITCOUNT", Command.FirstKey, keyCodec.encode(key) +: rangeArgs(range), Decode.long) def bitPos[K](key: K, bit: Boolean, range: Option[BitPosRange] = None)(using keyCodec: KeyCodec[K]): Command[Long] = - Command.read("BITPOS", Command.FirstKey, Vector(keyCodec.encode(key), bitToken(bit)) ++ posRangeArgs(range), Decode.long) + Command.read("BITPOS", Command.FirstKey, Vector(keyCodec.encode(key), Primitives.encodeBoolean(bit)) ++ posRangeArgs(range), Decode.long) def bitOpAnd[K](destination: K, first: K, rest: K*)(using KeyCodec[K]): Command[Long] = bitOp(And, destination, first +: rest.toVector) @@ -108,20 +103,20 @@ private[sage] object Bitmaps { } private def rangeArgs(range: Option[BitRange]): Vector[Bytes] = - range.toVector.flatMap(r => Vector(Bytes.utf8(r.start.toString), Bytes.utf8(r.end.toString), unitArg(r.unit))) + range.toVector.flatMap(r => Vector(Args.long(r.start), Args.long(r.end), unitArg(r.unit))) private def posRangeArgs(range: Option[BitPosRange]): Vector[Bytes] = range match { case None => Vector.empty - case Some(BitPosRange.FromStart(start)) => Vector(Bytes.utf8(start.toString)) - case Some(BitPosRange.Within(start, end, unit)) => Vector(Bytes.utf8(start.toString), Bytes.utf8(end.toString), unitArg(unit)) + case Some(BitPosRange.FromStart(start)) => Vector(Args.long(start)) + case Some(BitPosRange.Within(start, end, unit)) => Vector(Args.long(start), Args.long(end), unitArg(unit)) } private def opArgs(op: BitFieldOp): Vector[Bytes] = op match { - case BitFieldOp.Get(fieldType, offset) => Vector(GetWord, typeArg(fieldType), offsetArg(offset)) - case BitFieldOp.Set(fieldType, offset, value) => Vector(SetWord, typeArg(fieldType), offsetArg(offset), Bytes.utf8(value.toString)) - case BitFieldOp.IncrBy(fieldType, offset, incr) => Vector(IncrByWord, typeArg(fieldType), offsetArg(offset), Bytes.utf8(incr.toString)) + case BitFieldOp.Get(fieldType, offset) => Vector(Get, typeArg(fieldType), offsetArg(offset)) + case BitFieldOp.Set(fieldType, offset, value) => Vector(SetWord, typeArg(fieldType), offsetArg(offset), Args.long(value)) + case BitFieldOp.IncrBy(fieldType, offset, incr) => Vector(IncrByWord, typeArg(fieldType), offsetArg(offset), Args.long(incr)) case BitFieldOp.Overflow(behavior) => Vector(OverflowWord, overflowArg(behavior)) } @@ -133,22 +128,8 @@ private[sage] object Bitmaps { private def offsetArg(offset: BitFieldOffset): Bytes = offset match { - case BitFieldOffset.Absolute(value) => Bytes.utf8(value.toString) + case BitFieldOffset.Absolute(value) => Args.long(value) case BitFieldOffset.TypeWidth(factor) => Bytes.utf8("#" + factor.toString) } - private def overflowArg(behavior: BitFieldOverflow): Bytes = - behavior match { - case BitFieldOverflow.Wrap => WrapWord - case BitFieldOverflow.Sat => SatWord - case BitFieldOverflow.Fail => FailWord - } - - private def unitArg(unit: BitUnit): Bytes = - unit match { - case BitUnit.Byte => ByteWord - case BitUnit.Bit => BitWord - } - - private def bitToken(value: Boolean): Bytes = if (value) One else Zero } diff --git a/sage-core/src/main/scala/sage/commands/BlockTimeout.scala b/sage-core/src/main/scala/sage/commands/BlockTimeout.scala index 3938d761..7cdc3bbe 100644 --- a/sage-core/src/main/scala/sage/commands/BlockTimeout.scala +++ b/sage-core/src/main/scala/sage/commands/BlockTimeout.scala @@ -20,9 +20,9 @@ object BlockTimeout { // server does not wait for less time than requested. Use at least one millisecond because the wire value 0 means to wait forever. private[commands] def wire(timeout: BlockTimeout): Bytes = timeout match { - case Forever => Zero + case Forever => Args.long(0) case After(duration) => - val millis = Math.max(1L, Math.ceilDiv(duration.toNanos, 1000000L)) + val millis = TimeArgs.positiveMillis(duration) val text = if (millis % 1000L == 0L) (millis / 1000L).toString else java.math.BigDecimal.valueOf(millis, 3).stripTrailingZeros.toPlainString @@ -32,9 +32,7 @@ object BlockTimeout { // the millisecond form `XREAD`/`XREADGROUP` `BLOCK` requires (the SECONDS [[wire]] form would silently shorten the wait 1000x) private[commands] def millisWire(timeout: BlockTimeout): Bytes = timeout match { - case Forever => Zero - case After(duration) => Bytes.utf8(Math.max(1L, Math.ceilDiv(duration.toNanos, 1000000L)).toString) + case Forever => Args.long(0) + case After(duration) => Args.long(TimeArgs.positiveMillis(duration)) } - - private val Zero = Bytes.utf8("0") } diff --git a/sage-core/src/main/scala/sage/commands/Cluster.scala b/sage-core/src/main/scala/sage/commands/Cluster.scala index 770a8c1b..7cce9cae 100644 --- a/sage-core/src/main/scala/sage/commands/Cluster.scala +++ b/sage-core/src/main/scala/sage/commands/Cluster.scala @@ -2,65 +2,48 @@ package sage.commands import sage.Bytes import sage.SageException.DecodeError -import sage.cluster.{Node, Shard, Slot, SlotRange} +import sage.cluster.{Node, Slot, SlotRange} import sage.protocol.Frame private[sage] object Cluster { - val slots: Command[Vector[Shard]] = - Command("CLUSTER", keyIndices = Command.NoKeys, args = Vector(Bytes.utf8("SLOTS")), decode = decodeSlots) - - final private case class Range(master: Node, replicas: Vector[Node], slots: SlotRange) - // Each range has the form [start, end, master, replica*], and each node has the form [ip, port, id, meta?]. The first node is the master. If // any entry cannot be decoded, reject the complete topology reply to avoid routing with incomplete information. - private def decodeSlots(frame: Frame): Either[DecodeError, Vector[Shard]] = - Decode.vector(decodeRange)(frame).map(merge) - - private def decodeRange(frame: Frame): Either[DecodeError, Range] = - frame match { - case Frame.Array(elements) if elements.length >= 3 => - for { - start <- slotOf(elements(0)) - end <- slotOf(elements(1)) - nodes <- Decode.vector(decodeNode)(Frame.Array(elements.drop(2))) - } yield Range(nodes.head, nodes.tail, SlotRange(start, end)) - case other => Left(DecodeError("array of [start, end, master, replicas...]", Frame.describe(other))) - } - - private def decodeNode(frame: Frame): Either[DecodeError, Node] = - frame match { - case Frame.Array(elements) if elements.length >= 2 => - for { - host <- endpointOf(elements(0)) - port <- intOf(elements(1)) - } yield Node(host, port) - case other => Left(DecodeError("node array [ip, port, id, ...]", Frame.describe(other))) + def slots(queried: Node): Command[Vector[SlotRange]] = + Command("CLUSTER", keyIndices = Command.NoKeys, args = Vector(Bytes.utf8("SLOTS")), decode = Decode.vector(decodeRange(queried))) + + private def decodeRange(queried: Node): Frame => Either[DecodeError, SlotRange] = { + val node = decodeNode(queried) + Decode.shape("array of [start, end, master, replicas...]") { case Frame.Array(startFrame +: endFrame +: masterFrame +: replicaFrames) => + for { + start <- slotOf(startFrame) + end <- slotOf(endFrame) + master <- node(masterFrame) + replicas <- Decode.each(replicaFrames)(node) + } yield SlotRange(start, end, master, replicas) } + } - private def merge(ranges: Vector[Range]): Vector[Shard] = - ranges.map(_.master).distinct.map { master => - val forMaster = ranges.filter(_.master == master) - Shard(master, forMaster.flatMap(_.replicas).distinct, forMaster.map(_.slots)) - } + // an empty or null endpoint means the node that answered CLUSTER SLOTS + private def decodeNode(queried: Node): Frame => Either[DecodeError, Node] = Decode.shape("node array [ip, port, id, ...]") { + case Frame.Array(hostFrame +: portFrame +: _) => + for { + host <- endpointOf(hostFrame) + port <- intOf(portFrame) + } yield Node(if (host.isEmpty) queried.host else host, port) + } private def slotOf(frame: Frame): Either[DecodeError, Slot] = intOf(frame).flatMap(index => Slot.at(index).toRight(DecodeError("slot in [0, 16384)", index.toString))) - private def intOf(frame: Frame): Either[DecodeError, Int] = - frame match { - case Frame.Integer(value) if value.isValidInt => - Right(value.toInt) // guard the narrowing: a Long outside Int range must not wrap into a valid slot or port - case Frame.Integer(value) => Left(DecodeError("integer within Int range", value.toString)) - case other => Left(DecodeError("integer", Frame.describe(other))) - } + private val intOf: Frame => Either[DecodeError, Int] = Decode.shape("integer") { + case Frame.Integer(value) if value.isValidInt => + Right(value.toInt) // guard the narrowing: a Long outside Int range must not wrap into a valid slot or port + case Frame.Integer(value) => Left(DecodeError("integer within Int range", value.toString)) + } - // a null endpoint means "the node you queried", which the empty host already denotes downstream - private def endpointOf(frame: Frame): Either[DecodeError, String] = - frame match { - case Frame.BulkString(bytes) => Right(bytes.asUtf8String) - case Frame.SimpleString(text) => Right(text) - case Frame.Null => Right("") - case other => Left(DecodeError("string or null endpoint", Frame.describe(other))) - } + private val endpointOf: Frame => Either[DecodeError, String] = Decode.shape("string or null endpoint") { + case Decode.Text(text) => Right(text) + case Frame.Null => Right("") + } } diff --git a/sage-core/src/main/scala/sage/commands/Command.scala b/sage-core/src/main/scala/sage/commands/Command.scala index 30967dea..347ad336 100644 --- a/sage-core/src/main/scala/sage/commands/Command.scala +++ b/sage-core/src/main/scala/sage/commands/Command.scala @@ -1,5 +1,7 @@ package sage.commands +import scala.util.Try + import sage.Bytes import sage.SageException.DecodeError import sage.protocol.{Frame, RespWriter} @@ -15,18 +17,34 @@ enum Execution { /** * Selects how to combine replies from an `allMasters` command. `First` returns the first reply when every node should return the same value, - * such as the SHA from `SCRIPT LOAD`. `Concat` joins the array elements returned by each node, as required by `KEYS`. `Fold` combines replies - * two at a time with the supplied function. This setting is used only when `allMasters` is true. + * such as the SHA from `SCRIPT LOAD`, and fails if any other reply does not decode. `Concat` joins the array elements returned by each node, + * as required by `KEYS`. `Fold` combines replies two at a time with the supplied function. This setting is used only when `allMasters` is + * true. */ enum BroadcastReduce { case First case Concat case Fold(combine: (Frame, Frame) => Frame) + + // takes only the wrapped decoder, so a codec that throws on a dropped First reply fails as a DecodeError + private[sage] def reduce(first: Frame, rest: Vector[Frame], decode: Frame => Try[Any]): Frame = + this match { + case First => + rest.foreach(decode(_).get) + first + case Concat => + Frame.Array((first +: rest).flatMap { + case Frame.Array(elements) => elements + case Frame.Set(elements) => elements + case other => throw DecodeError("array replies", Frame.describe(other)) + }) + case Fold(combine) => rest.foldLeft(first)(combine) + } } /** * Stores one server command, including its encoded arguments and reply decoder. `keyIndices` marks the argument positions used as keys for - * cluster routing. [[Reply.run]] handles top-level error frames before calling `decode`. + * cluster routing. [[Reply.decode]] handles top-level error frames before calling `decode`. * * `isReadOnly` marks side-effect-free reads for replica routing. `cacheable` is narrower: the result must depend only on the named keys' * current state, allowing server invalidations to cover every change. Time-varying reads (`TTL`, `OBJECT IDLETIME`) and non-deterministic @@ -67,29 +85,13 @@ final case class Command[+Out]( /** * Transforms the decoded result, leaving the wire encoding and routing metadata untouched. */ - def map[B](f: Out => B): Command[B] = withDecode(frame => decode(frame).map(f)) + def map[B](f: Out => B): Command[B] = copy[B](decode = frame => decode(frame).map(f)) /** * Whether this command requires a dedicated connection because it blocks the connection it uses. */ def isBlocking: Boolean = execution == Execution.Blocking - // rebuild by named fields to preserve their order and ensure future fields are copied explicitly. - private def withDecode[B](decode: Frame => Either[DecodeError, B]): Command[B] = - Command( - name = name, - keyIndices = keyIndices, - args = args, - decode = decode, - execution = execution, - isReadOnly = isReadOnly, - cacheable = cacheable, - allMasters = allMasters, - cursorBound = cursorBound, - broadcast = broadcast, - requiresClusterWideTxResult = requiresClusterWideTxResult - ) - /** * The command's key bytes in argument order. Client-side cache invalidations use these bytes to remove affected results. */ @@ -107,9 +109,12 @@ final case class Command[+Out]( /** * This command's wire form with decoding replaced by the raw reply frame, preserving routing metadata. The cluster runtime uses it to - * collect a broadcast's per-node replies and fold them before decoding once. + * collect a broadcast's per-node replies and combine them before decoding the result. */ - def rawFrame: Command[Frame] = withDecode(frame => Right(frame)) + def rawFrame: Command[Frame] = copy[Frame](decode = Decode.frame) + + // the caller decodes the result, so First decodes only the replies it drops + private[sage] def reduceReplies(first: Frame, rest: Vector[Frame]): Frame = broadcast.reduce(first, rest, Reply.decode(this, _)) } object Command { diff --git a/sage-core/src/main/scala/sage/commands/Connection.scala b/sage-core/src/main/scala/sage/commands/Connection.scala index e23a48b2..af46e569 100644 --- a/sage-core/src/main/scala/sage/commands/Connection.scala +++ b/sage-core/src/main/scala/sage/commands/Connection.scala @@ -10,7 +10,7 @@ private[sage] object Connection { val multi: Command[Unit] = Command("MULTI", keyIndices = Command.NoKeys, args = Vector.empty, decode = Decode.ok) // the reply (an array of per-command results, or a null array on WATCH abort) is interpreted by the runtime, not this passthrough decoder - val exec: Command[Frame] = Command("EXEC", keyIndices = Command.NoKeys, args = Vector.empty, decode = frame => Right(frame)) + val exec: Command[Frame] = Command("EXEC", keyIndices = Command.NoKeys, args = Vector.empty, decode = Decode.frame) val unwatch: Command[Unit] = Command("UNWATCH", keyIndices = Command.NoKeys, args = Vector.empty, decode = Decode.ok) @@ -33,39 +33,39 @@ private[sage] object Connection { def isClientTracking(command: Command[?]): Boolean = command.name == "CLIENT" && command.args.headOption.exists(_.asUtf8String == "TRACKING") - def watch[K](first: K, rest: K*)(using keyCodec: KeyCodec[K]): Command[Unit] = { - val keys = (first +: rest.toVector).map(keyCodec.encode) - Command("WATCH", keyIndices = Vector.range(0, keys.length), args = keys, decode = Decode.ok) - } + def watch[K](first: K, rest: K*)(using KeyCodec[K]): Command[Unit] = KeyArgs.allKeys("WATCH", first +: rest.toVector, Decode.ok) def ping(message: Option[String] = None): Command[String] = Command( "PING", keyIndices = Command.NoKeys, args = message.map(Bytes.utf8).toVector, - decode = { - case Frame.SimpleString(value) => Right(value) - case Frame.BulkString(value) => Right(value.asUtf8String) - case other => Left(DecodeError("simple or bulk string", Frame.describe(other))) - } + decode = pingReply ) + private val pingReply = Decode.shape("simple or bulk string") { case Decode.Text(value) => Right(value) } + /** - * The protocol handshake. Unknown reply entries are ignored for forward compatibility. + * The protocol handshake. The reply must confirm RESP3, the only supported protocol. Other reply entries are ignored. */ - def hello(auth: Option[(String, String)] = None): Command[HelloReply] = + def hello(auth: Option[(String, String)] = None): Command[Unit] = Command( "HELLO", keyIndices = Command.NoKeys, args = Bytes.utf8("3") +: auth.toVector.flatMap { case (username, password) => Vector("AUTH", username, password).map(Bytes.utf8) }, - decode = HelloReply.decode + decode = Decode.fields(_.required("proto", proto3)) ) + private val proto3: Frame => Either[DecodeError, Unit] = Decode.shape("integer for 'proto'") { + case Frame.Integer(3) => Right(()) + case Frame.Integer(value) => Left(DecodeError("proto 3", s"proto $value")) + } + // These commands configure connection state during the HELLO setup and run again after reconnecting. They are not exposed as ordinary // operations because concurrent users of a shared connection would all observe the changed state. def select(database: Int): Command[Unit] = - Command("SELECT", Command.NoKeys, Vector(Bytes.utf8(database.toString)), Decode.ok) + Command("SELECT", Command.NoKeys, Vector(Args.long(database)), Decode.ok) def clientSetName(name: String): Command[Unit] = Command("CLIENT", Command.NoKeys, Vector("SETNAME", name).map(Bytes.utf8), Decode.ok) @@ -83,10 +83,9 @@ private[sage] object Connection { "CLIENT", Command.NoKeys, Vector(Bytes.utf8("GETNAME")), - { + Decode.shape("bulk string or null") { case Frame.Null => Right("") case Frame.BulkString(bytes) => Right(bytes.asUtf8String) - case other => Left(DecodeError("bulk string or null", Frame.describe(other))) } ) val clientInfo: Command[String] = Command("CLIENT", Command.NoKeys, Vector(Bytes.utf8("INFO")), Decode.text) diff --git a/sage-core/src/main/scala/sage/commands/Decode.scala b/sage-core/src/main/scala/sage/commands/Decode.scala index f20dcc09..df31a8b8 100644 --- a/sage-core/src/main/scala/sage/commands/Decode.scala +++ b/sage-core/src/main/scala/sage/commands/Decode.scala @@ -3,38 +3,48 @@ package sage.commands import java.time.Instant import scala.collection.mutable -import scala.concurrent.duration.FiniteDuration +import scala.concurrent.duration.{FiniteDuration, MILLISECONDS} import sage.Bytes import sage.SageException.DecodeError -import sage.codec.{Doubles, KeyCodec, ValueCodec} +import sage.codec.{Doubles, KeyCodec, Primitives, ValueCodec} import sage.protocol.Frame private[commands] object Decode { + // a decoder for the frames `accept` handles; any other frame fails with a mismatch against `expected` + def shape[A](expected: String)(accept: PartialFunction[Frame, Either[DecodeError, A]]): Frame => Either[DecodeError, A] = { + val mismatch: Frame => Either[DecodeError, A] = other => Left(DecodeError(expected, Frame.describe(other))) + frame => accept.applyOrElse(frame, mismatch) + } + val frame: Frame => Either[DecodeError, Frame] = Right(_) - val long: Frame => Either[DecodeError, Long] = { - case Frame.Integer(value) => Right(value) - case other => Left(DecodeError("integer", Frame.describe(other))) + val long: Frame => Either[DecodeError, Long] = shape("integer") { case Frame.Integer(value) => + Right(value) + } + + val millisDuration: Frame => Either[DecodeError, FiniteDuration] = shape("integer") { case Frame.Integer(ms) => + Right(FiniteDuration(ms, MILLISECONDS)) } - val flag: Frame => Either[DecodeError, Boolean] = { + val millisInstant: Frame => Either[DecodeError, Instant] = shape("integer") { case Frame.Integer(ms) => Right(Instant.ofEpochMilli(ms)) } + + def decimal(expected: String): Frame => Either[DecodeError, Long] = + shape(expected) { case Frame.BulkString(text) => Primitives.decodeLong(expected, Long.MinValue, Long.MaxValue)(text) } + + val flag: Frame => Either[DecodeError, Boolean] = shape("integer 0 or 1") { case Frame.Integer(0) => Right(false) case Frame.Integer(1) => Right(true) - case other => Left(DecodeError("integer 0 or 1", Frame.describe(other))) } - val ok: Frame => Either[DecodeError, Unit] = { - case Frame.SimpleString("OK") => Right(()) - case other => Left(DecodeError("simple string 'OK'", Frame.describe(other))) + val ok: Frame => Either[DecodeError, Unit] = shape("simple string 'OK'") { case Frame.SimpleString("OK") => + Right(()) } - val double: Frame => Either[DecodeError, Double] = { - case Frame.BulkString(bytes) => - val text = bytes.asUtf8String - Doubles.parse(text).toRight(DecodeError("double bulk string", s"bulk string '$text'")) - case other => Left(DecodeError("double bulk string", Frame.describe(other))) + val okOrNull: Frame => Either[DecodeError, Boolean] = shape("simple string 'OK' or null") { + case Frame.SimpleString("OK") => Right(true) + case Frame.Null => Right(false) } // Decode the integer format shared by TTL, EXPIRETIME, HTTL, and related commands. -2 means absent, -1 means no expiry, and a non-negative @@ -57,34 +67,29 @@ private[commands] object Decode { case other => Left(DecodeError("bulk string or null", Frame.describe(other))) } - val utf8String: Frame => Either[DecodeError, String] = { - case Frame.BulkString(bytes) => Right(bytes.asUtf8String) - case other => Left(DecodeError("bulk string", Frame.describe(other))) + val utf8String: Frame => Either[DecodeError, String] = shape("bulk string") { case Frame.BulkString(bytes) => + Right(bytes.asUtf8String) } // text however the server framed it: simple, bulk, or the RESP3 verbatim form INFO/CLIENT INFO/CLUSTER NODES use - val text: Frame => Either[DecodeError, String] = { + val text: Frame => Either[DecodeError, String] = shape("string") { case Frame.SimpleString(value) => Right(value) case Frame.BulkString(bytes) => Right(bytes.asUtf8String) case Frame.VerbatimString(_, bytes) => Right(bytes.asUtf8String) - case other => Left(DecodeError("string", Frame.describe(other))) } - val optionalUtf8String: Frame => Either[DecodeError, Option[String]] = { + val optionalUtf8String: Frame => Either[DecodeError, Option[String]] = shape("bulk string or null") { case Frame.Null => Right(None) case Frame.BulkString(bytes) => Right(Some(bytes.asUtf8String)) - case other => Left(DecodeError("bulk string or null", Frame.describe(other))) } - val bytes: Frame => Either[DecodeError, Bytes] = { - case Frame.BulkString(value) => Right(value) - case other => Left(DecodeError("bulk string", Frame.describe(other))) + val bytes: Frame => Either[DecodeError, Bytes] = shape("bulk string") { case Frame.BulkString(value) => + Right(value) } - val optionalBytes: Frame => Either[DecodeError, Option[Bytes]] = { + val optionalBytes: Frame => Either[DecodeError, Option[Bytes]] = shape("bulk string or null") { case Frame.Null => Right(None) case Frame.BulkString(value) => Right(Some(value)) - case other => Left(DecodeError("bulk string or null", Frame.describe(other))) } def key[K](using codec: KeyCodec[K]): Frame => Either[DecodeError, K] = { @@ -98,25 +103,25 @@ private[commands] object Decode { case other => Left(DecodeError("bulk string or null", Frame.describe(other))) } - val optionalLong: Frame => Either[DecodeError, Option[Long]] = { + val optionalLong: Frame => Either[DecodeError, Option[Long]] = shape("integer or null") { case Frame.Null => Right(None) case Frame.Integer(value) => Right(Some(value)) - case other => Left(DecodeError("integer or null", Frame.describe(other))) } - // a double however the server framed it: a RESP3 Double, or a bulk string under RESP2 (geo coordinates, distances) - val lenientDouble: Frame => Either[DecodeError, Double] = { + // a double however the server framed it: a RESP3 Double (scores), or a bulk string (INCRBYFLOAT, ZSCAN scores, geo under RESP2) + val double: Frame => Either[DecodeError, Double] = shape("double") { case Frame.Double(value) => Right(value) case Frame.BulkString(bytes) => val text = bytes.asUtf8String Doubles.parse(text).toRight(DecodeError("double", s"bulk string '$text'")) - case other => Left(DecodeError("double", Frame.describe(other))) } // GEODIST replies the distance as a double, or null when a member is absent - val optionalDouble: Frame => Either[DecodeError, Option[Double]] = { + val optionalDouble: Frame => Either[DecodeError, Option[Double]] = nullable(double) + + def nullable[A](decode: Frame => Either[DecodeError, A]): Frame => Either[DecodeError, Option[A]] = { case Frame.Null => Right(None) - case other => lenientDouble(other).map(Some(_)) + case other => decode(other).map(Some(_)) } private def buildEach[A, B, C](items: IterableOnce[A], builder: mutable.Builder[B, C])( @@ -133,15 +138,24 @@ private[commands] object Decode { Right(builder.result()) } + def mapEntries[A, K, V](items: IterableOnce[A])(f: A => Either[DecodeError, (K, V)]): Either[DecodeError, Map[K, V]] = + buildEach(items, Map.newBuilder[K, V])(f) + def each[A, B](items: IterableOnce[A])(f: A => Either[DecodeError, B]): Either[DecodeError, Vector[B]] = buildEach(items, Vector.newBuilder[B])(f) - // steps a flat alternating array two elements at a time, without grouped(2)'s throwaway 2-element Vector per pair; caller guarantees even length - private def buildPairs[B, C](elements: Vector[Frame], builder: mutable.Builder[B, C])( - f: (Frame, Frame) => Either[DecodeError, B] - ): Either[DecodeError, C] = { + def flatPairsOf[B](label: String)(f: (Frame, Frame) => Either[DecodeError, B]): Frame => Either[DecodeError, Vector[B]] = { + case array: Frame.Array => buildPairs(array, label)(f) + case other => Left(DecodeError(label, Frame.describe(other))) + } + + // steps a flat alternating array two elements at a time, without grouped(2)'s throwaway 2-element Vector per pair + private def buildPairs[B](array: Frame.Array, label: String)(f: (Frame, Frame) => Either[DecodeError, B]): Either[DecodeError, Vector[B]] = { + val elements = array.elements + if (elements.length % 2 != 0) return Left(DecodeError(label, Frame.describe(array))) + val builder = Vector.newBuilder[B] builder.sizeHint(elements.length / 2) - var i = 0 + var i = 0 while (i < elements.length) { f(elements(i), elements(i + 1)) match { case Right(value) => builder += value @@ -196,75 +210,84 @@ private[commands] object Decode { case other => Left(DecodeError(label, Frame.describe(other))) } - def vector[A](element: Frame => Either[DecodeError, A]): Frame => Either[DecodeError, Vector[A]] = { + def byLowerName[E](cases: E*): Map[String, E] = cases.iterator.map(value => value.toString.toLowerCase(java.util.Locale.ROOT) -> value).toMap + + def vector[A](element: Frame => Either[DecodeError, A], expected: String = "array"): Frame => Either[DecodeError, Vector[A]] = { case Frame.Array(elements) => buildEach(elements, Vector.newBuilder[A])(element) - case other => Left(DecodeError("array", Frame.describe(other))) + case other => Left(DecodeError(expected, Frame.describe(other))) } // a missing list replies null where a present one replies an array; a stored list is never empty, so null collapses to an empty vector - def vectorOrEmpty[A](element: Frame => Either[DecodeError, A]): Frame => Either[DecodeError, Vector[A]] = { - val decodeVector = vector(element) - frame => - frame match { - case Frame.Null => Right(Vector.empty) - case other => decodeVector(other) - } + def orEmpty[A](decode: Frame => Either[DecodeError, Vector[A]]): Frame => Either[DecodeError, Vector[A]] = { + case Frame.Null => Right(Vector.empty) + case other => decode(other) } // a RESP3 map, or the flat RESP2 array of alternating key/value some introspection replies still use; non-string keys are dropped - val fieldMap: Frame => Either[DecodeError, Map[String, Frame]] = { - case Frame.Map(entries) => Right(entries.collect { case (Frame.BulkString(k), v) => k.asUtf8String -> v }.toMap) - case Frame.Array(elements) if elements.length % 2 == 0 => - val builder = Map.newBuilder[String, Frame] - builder.sizeHint(elements.length / 2) - var i = 0 - while (i < elements.length) { - elements(i) match { - case Frame.BulkString(k) => builder += (k.asUtf8String -> elements(i + 1)) - case _ => () - } - i += 2 - } - Right(builder.result()) - case other => Left(DecodeError("map", Frame.describe(other))) + val fieldMap: Frame => Either[DecodeError, Map[String, Frame]] = shape("map") { + case Frame.Map(entries) => Right(entries.collect { case (Text(k), v) => k -> v }.toMap) + case array: Frame.Array => buildPairs(array, "map")((k, v) => Right(Text.unapply(k).map(_ -> v))).map(_.flatten.toMap) } - def map[K, V](using KeyCodec[K], ValueCodec[V]): Frame => Either[DecodeError, Map[K, V]] = { - case Frame.Map(entries) => - buildEach(entries, Map.newBuilder[K, V]) { case (fieldFrame, valueFrame) => - for { - field <- key(fieldFrame) - value <- this.value(valueFrame) - } yield field -> value + def fields[A](read: Fields => Either[DecodeError, A]): Frame => Either[DecodeError, A] = frame => fieldMap(frame).flatMap(m => read(new Fields(m))) + + def fieldValues[A](value: Frame => Either[DecodeError, A]): Frame => Either[DecodeError, Map[String, A]] = + frame => fieldMap(frame).flatMap(mapEntries(_) { case (name, field) => value(field).map(name -> _) }) + + // text framed as a bulk or simple string + object Text { + def unapply(frame: Frame): Option[String] = + frame match { + case Frame.BulkString(bytes) => Some(bytes.asUtf8String) + case Frame.SimpleString(name) => Some(name) + case _ => None } - case other => Left(DecodeError("map", Frame.describe(other))) } - // HSCAN's items are a flat field, value, field, value, … array; HRANDFIELD WITHVALUES nests each pair in its own array. - def flatPairs[K, V](using KeyCodec[K], ValueCodec[V]): Frame => Either[DecodeError, Vector[(K, V)]] = { - case Frame.Array(elements) if elements.length % 2 == 0 => - buildPairs(elements, Vector.newBuilder[(K, V)]) { (fieldFrame, valueFrame) => - for { - field <- key(fieldFrame) - value <- this.value(valueFrame) - } yield field -> value + /** + * A lenient view over an introspection reply map: read fields by known name, ignore the rest. + */ + final class Fields(table: Map[String, Frame]) { + + def get(name: String): Option[Frame] = table.get(name) + + def required[A](name: String, decode: Frame => Either[DecodeError, A]): Either[DecodeError, A] = + table.get(name).toRight(DecodeError(s"field '$name'", "absent")).flatMap(decode) + + // a field that is core to the reply but whose absence on some server we tolerate with a default rather than failing the whole decode + def requiredOr[A](name: String, decode: Frame => Either[DecodeError, A], fallback: A): Either[DecodeError, A] = + table.get(name) match { + case None | Some(Frame.Null) => Right(fallback) + case Some(frame) => decode(frame) } - case other => Left(DecodeError("array of field/value pairs", Frame.describe(other))) - } - def nestedPairs[K, V](using KeyCodec[K], ValueCodec[V]): Frame => Either[DecodeError, Vector[(K, V)]] = { - val pair = array2(key[K], value[V], "field/value pair")(_ -> _) - frame => - frame match { - case Frame.Array(rows) => each(rows)(pair) - case other => Left(DecodeError("array of field/value pairs", Frame.describe(other))) + def optional[A](name: String, decode: Frame => Either[DecodeError, A]): Either[DecodeError, Option[A]] = + table.get(name) match { + case None | Some(Frame.Null) => Right(None) + case Some(frame) => decode(frame).map(Some(_)) } + + def optionalVector[A](name: String, element: Frame => Either[DecodeError, A]): Either[DecodeError, Vector[A]] = + requiredOr(name, vector(element), Vector.empty) } - private val scanCursor: Frame => Either[DecodeError, Option[ScanCursor]] = { - case Frame.BulkString(bytes) => - Right(if (bytes.sameBytes(ScanCursor.bytes(ScanCursor.start))) None else Some(ScanCursor.wrap(bytes))) - case other => Left(DecodeError("cursor bulk string", Frame.describe(other))) + def pair[A, B](a: Frame => Either[DecodeError, A], b: Frame => Either[DecodeError, B]): (Frame, Frame) => Either[DecodeError, (A, B)] = + (x, y) => a(x).flatMap(first => b(y).map(first -> _)) + + def map[K, V](using KeyCodec[K], ValueCodec[V]): Frame => Either[DecodeError, Map[K, V]] = { + val entry = pair(key[K], value[V]).tupled + shape("map") { case Frame.Map(entries) => mapEntries(entries)(entry) } + } + + // HSCAN's items are a flat field, value, field, value, … array; HRANDFIELD WITHVALUES nests each pair in its own array. + def flatPairs[K, V](using KeyCodec[K], ValueCodec[V]): Frame => Either[DecodeError, Vector[(K, V)]] = + flatPairsOf("array of field/value pairs")(pair(key[K], value[V])) + + def nestedPairs[K, V](using KeyCodec[K], ValueCodec[V]): Frame => Either[DecodeError, Vector[(K, V)]] = + vector(array2(key[K], value[V], "field/value pair")(_ -> _), "array of field/value pairs") + + private val scanCursor: Frame => Either[DecodeError, Option[ScanCursor]] = shape("cursor bulk string") { case Frame.BulkString(bytes) => + Right(if (bytes.sameBytes(ScanCursor.bytes(ScanCursor.start))) None else Some(ScanCursor.wrap(bytes))) } def scanPage[A](items: Frame => Either[DecodeError, Vector[A]]): Frame => Either[DecodeError, ScanPage[A]] = @@ -276,25 +299,17 @@ private[commands] object Decode { case other => Left(DecodeError("set", Frame.describe(other))) } - val score: Frame => Either[DecodeError, Double] = { - case Frame.Double(value) => Right(value) - case other => Left(DecodeError("double", Frame.describe(other))) - } - - val optionalScore: Frame => Either[DecodeError, Option[Double]] = { + val optionalScore: Frame => Either[DecodeError, Option[Double]] = shape("double or null") { case Frame.Null => Right(None) case Frame.Double(value) => Right(Some(value)) - case other => Left(DecodeError("double or null", Frame.describe(other))) } private def memberScore[V](using ValueCodec[V]): Frame => Either[DecodeError, (V, Double)] = - array2(value[V], score, "member/score pair")(_ -> _) + array2(value[V], double, "member/score pair")(_ -> _) // RESP3 nests each member with its Double score in a two-element array (ZRANGE WITHSCORES, ZPOPMIN count, …) - def scoredMembers[V](using ValueCodec[V]): Frame => Either[DecodeError, Vector[(V, Double)]] = { - case Frame.Array(rows) => each(rows)(memberScore[V]) - case other => Left(DecodeError("array of member/score pairs", Frame.describe(other))) - } + def scoredMembers[V](using ValueCodec[V]): Frame => Either[DecodeError, Vector[(V, Double)]] = + vector(memberScore[V], "array of member/score pairs") // ZPOPMIN/ZPOPMAX without a count: a flat [member, score], or an empty array when the key is absent def optionalScoredMember[V](using ValueCodec[V]): Frame => Either[DecodeError, Option[(V, Double)]] = { @@ -305,16 +320,8 @@ private[commands] object Decode { } // ZSCAN's items are a flat member, score, member, score array with scores as bulk strings, not RESP3 doubles - def scoredMembersFlat[V](using ValueCodec[V]): Frame => Either[DecodeError, Vector[(V, Double)]] = { - case Frame.Array(elements) if elements.length % 2 == 0 => - buildPairs(elements, Vector.newBuilder[(V, Double)]) { (memberFrame, scoreFrame) => - for { - member <- value(memberFrame) - s <- double(scoreFrame) - } yield member -> s - } - case other => Left(DecodeError("array of member/score pairs", Frame.describe(other))) - } + def scoredMembersFlat[V](using ValueCodec[V]): Frame => Either[DecodeError, Vector[(V, Double)]] = + flatPairsOf("array of member/score pairs")(pair(value[V], double)) } private[commands] object TimeArgs { @@ -327,6 +334,11 @@ private[commands] object TimeArgs { // requested time. Rounding down would encode a sub-millisecond duration as 0, which expires immediately. def millis(duration: FiniteDuration): Long = Math.ceilDiv(duration.toNanos, 1000000L) + // for arguments where the wire value 0 means "no timeout" or "no expiry" + def positiveMillis(duration: FiniteDuration): Long = Math.max(1L, millis(duration)) + + def positiveMillis(timestamp: Instant): Long = Math.max(1L, millis(timestamp)) + // Use saturating arithmetic for an Instant outside the millisecond range. This avoids an overflow while building the command and preserves // upward rounding at the maximum value. def millis(timestamp: Instant): Long = @@ -341,12 +353,12 @@ private[commands] object TimeArgs { catch { case _: ArithmeticException => if (a < 0L) Long.MinValue else Long.MaxValue } def relative(duration: FiniteDuration): Vector[Bytes] = - if (wholeSeconds(duration)) Vector(Ex, Bytes.utf8(duration.toSeconds.toString)) - else Vector(Px, Bytes.utf8(millis(duration).toString)) + if (wholeSeconds(duration)) Vector(Ex, Args.long(duration.toSeconds)) + else Vector(Px, Args.long(millis(duration))) def absolute(timestamp: Instant): Vector[Bytes] = - if (wholeSeconds(timestamp)) Vector(ExAt, Bytes.utf8(timestamp.getEpochSecond.toString)) - else Vector(PxAt, Bytes.utf8(millis(timestamp).toString)) + if (wholeSeconds(timestamp)) Vector(ExAt, Args.long(timestamp.getEpochSecond)) + else Vector(PxAt, Args.long(millis(timestamp))) def expireCommand(secName: String, msName: String, duration: FiniteDuration): (String, Long) = if (wholeSeconds(duration)) (secName, duration.toSeconds) else (msName, millis(duration)) diff --git a/sage-core/src/main/scala/sage/commands/FrameDecode.scala b/sage-core/src/main/scala/sage/commands/FrameDecode.scala index 41bb4ff7..01a3d32d 100644 --- a/sage-core/src/main/scala/sage/commands/FrameDecode.scala +++ b/sage-core/src/main/scala/sage/commands/FrameDecode.scala @@ -19,7 +19,7 @@ extension (frame: Frame) { /** * Decodes an array/set/push frame into a `Vector[A]`, decoding each element through its [[sage.codec.ValueCodec]]. */ - def asArrayOf[A](using ValueCodec[A]): Either[DecodeError, Vector[A]] = Decode.vector(Decode.value[A]).apply(frame) + def asArrayOf[A](using ValueCodec[A]): Either[DecodeError, Vector[A]] = asArray.flatMap(Decode.each(_)(Decode.value[A])) /** * The raw, undecoded elements of an array/set/push frame, for replies whose elements are heterogeneous. diff --git a/sage-core/src/main/scala/sage/commands/Functions.scala b/sage-core/src/main/scala/sage/commands/Functions.scala index 62d5357e..f9915531 100644 --- a/sage-core/src/main/scala/sage/commands/Functions.scala +++ b/sage-core/src/main/scala/sage/commands/Functions.scala @@ -5,6 +5,8 @@ import scala.concurrent.duration.* import sage.Bytes import sage.SageException.DecodeError import sage.codec.{KeyCodec, ValueCodec} +import sage.commands.Args.Replace +import sage.commands.KeyArgs.ScriptVerb import sage.protocol.Frame /** @@ -49,7 +51,6 @@ final case class FunctionStats(runningScript: Option[RunningScript], engines: Ma private[sage] object Functions { private val Load = Bytes.utf8("LOAD") - private val Replace = Bytes.utf8("REPLACE") private val Delete = Bytes.utf8("DELETE") private val Flush = Bytes.utf8("FLUSH") private val Kill = Bytes.utf8("KILL") @@ -59,34 +60,29 @@ private[sage] object Functions { private val LibraryName = Bytes.utf8("LIBRARYNAME") private val WithCode = Bytes.utf8("WITHCODE") private val Stats = Bytes.utf8("STATS") + private val policyArg = Args.keywords(RestorePolicy.values) - def fCall(function: String): Command[Frame] = fCallCommand("FCALL", function, Vector.empty, Vector.empty, readOnly = false) + def fCall(function: String): Command[Frame] = KeyArgs.scriptCall(ScriptVerb.FCall, function, Seq.empty[Bytes], Seq.empty[Bytes]) - def fCall[K](function: String, keys: Seq[K])(using keyCodec: KeyCodec[K]): Command[Frame] = - fCallCommand("FCALL", function, keys.iterator.map(keyCodec.encode).toVector, Vector.empty, readOnly = false) + def fCall[K](function: String, keys: Seq[K])(using KeyCodec[K]): Command[Frame] = + KeyArgs.scriptCall(ScriptVerb.FCall, function, keys, Seq.empty[Bytes]) - def fCall[K, V](function: String, keys: Seq[K], args: Seq[V])(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Frame] = - fCallCommand("FCALL", function, keys.iterator.map(keyCodec.encode).toVector, args.iterator.map(valueCodec.encode).toVector, readOnly = false) + def fCall[K, V](function: String, keys: Seq[K], args: Seq[V])(using KeyCodec[K], ValueCodec[V]): Command[Frame] = + KeyArgs.scriptCall(ScriptVerb.FCall, function, keys, args) - def fCallRo(function: String): Command[Frame] = fCallCommand("FCALL_RO", function, Vector.empty, Vector.empty, readOnly = true) + def fCallRo(function: String): Command[Frame] = KeyArgs.scriptCall(ScriptVerb.FCallRo, function, Seq.empty[Bytes], Seq.empty[Bytes]) - def fCallRo[K](function: String, keys: Seq[K])(using keyCodec: KeyCodec[K]): Command[Frame] = - fCallCommand("FCALL_RO", function, keys.iterator.map(keyCodec.encode).toVector, Vector.empty, readOnly = true) + def fCallRo[K](function: String, keys: Seq[K])(using KeyCodec[K]): Command[Frame] = + KeyArgs.scriptCall(ScriptVerb.FCallRo, function, keys, Seq.empty[Bytes]) - def fCallRo[K, V](function: String, keys: Seq[K], args: Seq[V])(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Frame] = - fCallCommand("FCALL_RO", function, keys.iterator.map(keyCodec.encode).toVector, args.iterator.map(valueCodec.encode).toVector, readOnly = true) - - private def fCallCommand(name: String, function: String, keys: Vector[Bytes], args: Vector[Bytes], readOnly: Boolean): Command[Frame] = { - val allArgs = (Bytes.utf8(function) +: Bytes.utf8(keys.length.toString) +: keys) ++ args - val keyIndices = Vector.range(2, 2 + keys.length) - Command(name, keyIndices, allArgs, Decode.frame, Execution.Ordinary, isReadOnly = readOnly, cacheable = false) - } + def fCallRo[K, V](function: String, keys: Seq[K], args: Seq[V])(using KeyCodec[K], ValueCodec[V]): Command[Frame] = + KeyArgs.scriptCall(ScriptVerb.FCallRo, function, keys, args) def functionLoad(code: String, replace: Boolean = false): Command[String] = Command( "FUNCTION", Command.NoKeys, - (Load +: (if (replace) Vector(Replace) else Vector.empty)) :+ Bytes.utf8(code), + (Load +: Args.flag(replace, Replace)) :+ Bytes.utf8(code), Decode.utf8String, allMasters = true ) @@ -105,7 +101,7 @@ private[sage] object Functions { Command( "FUNCTION", Command.NoKeys, - Vector(Restore, payload) ++ policy.map(p => Bytes.utf8(p.toString.toUpperCase)).toVector, + Vector(Restore, payload) ++ policy.map(policyArg).toVector, Decode.ok, allMasters = true ) @@ -114,85 +110,59 @@ private[sage] object Functions { Command( "FUNCTION", Command.NoKeys, - List +: (libraryName.toVector.flatMap(name => Vector(LibraryName, Bytes.utf8(name))) ++ (if (withCode) Vector(WithCode) else Vector.empty)), + List +: (Args.optText(LibraryName, libraryName) ++ Args.flag(withCode, WithCode)), Decode.vector(decodeLibrary) ) private val decodeStats: Frame => Either[DecodeError, FunctionStats] = - frame => - Decode.fieldMap(frame).flatMap { fields => - for { - running <- fields.get("running_script") match { - case None | Some(Frame.Null) => Right(None) - case Some(scriptFrame) => decodeRunningScript(scriptFrame).map(Some(_)) - } - engines <- fields.get("engines").fold[Either[DecodeError, Map[String, EngineStats]]](Right(Map.empty))(decodeEngines) - } yield FunctionStats(running, engines) - } + Decode.fields { f => + for { + running <- f.optional("running_script", decodeRunningScript) + engines <- f.requiredOr("engines", decodeEngines, Map.empty) + } yield FunctionStats(running, engines) + } val functionStats: Command[FunctionStats] = Command("FUNCTION", Command.NoKeys, Vector(Stats), decodeStats) - private def decodeLibrary(frame: Frame): Either[DecodeError, LibraryInfo] = - Decode.fieldMap(frame).flatMap { fields => + private val decodeLibrary: Frame => Either[DecodeError, LibraryInfo] = + Decode.fields { f => for { - name <- requireString(fields, "library_name") - engine <- requireString(fields, "engine") - functions <- fields.get("functions").fold[Either[DecodeError, Vector[FunctionInfo]]](Right(Vector.empty))(Decode.vector(decodeFunction)) - } yield LibraryInfo(name, engine, functions, optionalString(fields, "library_code")) + name <- f.required("library_name", Decode.text) + engine <- f.required("engine", Decode.text) + functions <- f.optionalVector("functions", decodeFunction) + code <- f.optional("library_code", Decode.text) + } yield LibraryInfo(name, engine, functions, code) } - private def decodeFunction(frame: Frame): Either[DecodeError, FunctionInfo] = - Decode.fieldMap(frame).flatMap { fields => - requireString(fields, "name").map { name => - FunctionInfo(name, optionalString(fields, "description"), flagSet(fields.get("flags"))) - } + private val decodeFunction: Frame => Either[DecodeError, FunctionInfo] = + Decode.fields { f => + for { + name <- f.required("name", Decode.text) + description <- f.optional("description", Decode.text) + } yield FunctionInfo(name, description, flagSet(f.get("flags"))) } - private def decodeRunningScript(frame: Frame): Either[DecodeError, RunningScript] = - Decode.fieldMap(frame).flatMap { fields => + private val decodeRunningScript: Frame => Either[DecodeError, RunningScript] = + Decode.fields { f => for { - name <- requireString(fields, "name") - command <- fields.get("command").fold[Either[DecodeError, Vector[String]]](Right(Vector.empty))(Decode.vector(Decode.utf8String)) - duration <- fields.get("duration_ms").fold[Either[DecodeError, Long]](Right(0L))(Decode.long) + name <- f.required("name", Decode.text) + command <- f.optionalVector("command", Decode.utf8String) + duration <- f.requiredOr("duration_ms", Decode.long, 0L) } yield RunningScript(name, command, duration.millis) } - private def decodeEngines(frame: Frame): Either[DecodeError, Map[String, EngineStats]] = - Decode.fieldMap(frame).flatMap { engines => - engines.foldLeft[Either[DecodeError, Map[String, EngineStats]]](Right(Map.empty)) { case (acc, (name, statsFrame)) => - for { - map <- acc - fields <- Decode.fieldMap(statsFrame) - } yield map + (name -> EngineStats( - fields.get("libraries_count").flatMap(asLong).getOrElse(0L), - fields.get("functions_count").flatMap(asLong).getOrElse(0L) - )) - } - } - - private def requireString(fields: Map[String, Frame], key: String): Either[DecodeError, String] = - fields.get(key).flatMap(asString).toRight(DecodeError(s"'$key' field", s"map without '$key'")) - - private def optionalString(fields: Map[String, Frame], key: String): Option[String] = - fields.get(key).flatMap(asString) - - private def asString(frame: Frame): Option[String] = - frame match { - case Frame.BulkString(b) => Some(b.asUtf8String) - case Frame.SimpleString(s) => Some(s) - case _ => None - } - - private def asLong(frame: Frame): Option[Long] = - frame match { - case Frame.Integer(v) => Some(v) - case _ => None - } + private val decodeEngines: Frame => Either[DecodeError, Map[String, EngineStats]] = + Decode.fieldValues(Decode.fields { stats => + for { + libraries <- stats.requiredOr("libraries_count", Decode.long, 0L) + functions <- stats.requiredOr("functions_count", Decode.long, 0L) + } yield EngineStats(libraries, functions) + }) private def flagSet(frame: Option[Frame]): Set[String] = frame match { - case Some(Frame.Set(elements)) => elements.flatMap(asString).toSet - case Some(Frame.Array(elements)) => elements.flatMap(asString).toSet + case Some(Frame.Set(elements)) => elements.flatMap(Decode.text(_).toOption).toSet + case Some(Frame.Array(elements)) => elements.flatMap(Decode.text(_).toOption).toSet case _ => Set.empty } } diff --git a/sage-core/src/main/scala/sage/commands/Geo.scala b/sage-core/src/main/scala/sage/commands/Geo.scala index ffbde61c..369fee1d 100644 --- a/sage-core/src/main/scala/sage/commands/Geo.scala +++ b/sage-core/src/main/scala/sage/commands/Geo.scala @@ -2,7 +2,8 @@ package sage.commands import sage.Bytes import sage.SageException.DecodeError -import sage.codec.{Doubles, KeyCodec, ValueCodec} +import sage.codec.{KeyCodec, ValueCodec} +import sage.commands.Args.{Ch, Count, Nx, Xx} import sage.protocol.Frame /** @@ -59,16 +60,11 @@ final case class GeoSearchResult[V](member: V, distance: Option[Double], hash: O private[sage] object Geo { - private val Nx = Bytes.utf8("NX") - private val Xx = Bytes.utf8("XX") - private val Ch = Bytes.utf8("CH") private val FromMember = Bytes.utf8("FROMMEMBER") private val FromLonLat = Bytes.utf8("FROMLONLAT") private val ByRadius = Bytes.utf8("BYRADIUS") private val ByBox = Bytes.utf8("BYBOX") - private val Asc = Bytes.utf8("ASC") - private val Desc = Bytes.utf8("DESC") - private val CountWord = Bytes.utf8("COUNT") + private val sortWord = Args.keywords(GeoSort.values) private val AnyWord = Bytes.utf8("ANY") private val WithCoord = Bytes.utf8("WITHCOORD") private val WithDist = Bytes.utf8("WITHDIST") @@ -86,7 +82,7 @@ private[sage] object Geo { Command( "GEOADD", Command.FirstKey, - (keyCodec.encode(key) +: conditionArgs(condition)) ++ (if (changed) Vector(Ch) else Vector.empty) ++ memberCoordArgs(first +: rest.toVector), + (keyCodec.encode(key) +: conditionArgs(condition)) ++ Args.flag(changed, Ch) ++ memberCoordArgs(first +: rest.toVector), Decode.long ) @@ -102,20 +98,10 @@ private[sage] object Geo { ) def geoHash[K, V](key: K, first: V, rest: V*)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Vector[Option[String]]] = - Command.read( - "GEOHASH", - Command.FirstKey, - keyCodec.encode(key) +: (first +: rest.toVector).map(valueCodec.encode), - Decode.vector(Decode.optionalUtf8String) - ) + Command.read("GEOHASH", Command.FirstKey, Args.keyThen(key, first, rest)(valueCodec.encode), Decode.vector(Decode.optionalUtf8String)) def geoPos[K, V](key: K, first: V, rest: V*)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Vector[Option[GeoCoordinates]]] = - Command.read( - "GEOPOS", - Command.FirstKey, - keyCodec.encode(key) +: (first +: rest.toVector).map(valueCodec.encode), - Decode.vector(optionalCoordinates) - ) + Command.read("GEOPOS", Command.FirstKey, Args.keyThen(key, first, rest)(valueCodec.encode), Decode.vector(optionalCoordinates)) def geoSearch[K, V]( key: K, @@ -127,7 +113,7 @@ private[sage] object Geo { Command.read( "GEOSEARCH", Command.FirstKey, - keyCodec.encode(key) +: (originArgs(origin) ++ shapeArgs(shape) ++ sortArgs(sort) ++ countArgs(count)), + keyCodec.encode(key) +: searchArgs(origin, shape, sort, count), Decode.vector(Decode.value[V]) ) @@ -144,8 +130,7 @@ private[sage] object Geo { Command.read( "GEOSEARCH", Command.FirstKey, - keyCodec - .encode(key) +: (originArgs(origin) ++ shapeArgs(shape) ++ sortArgs(sort) ++ countArgs(count) ++ withArgs(withCoord, withDist, withHash)), + keyCodec.encode(key) +: (searchArgs(origin, shape, sort, count) ++ withArgs(withCoord, withDist, withHash)), searchReply[V](withCoord, withDist, withHash) ) @@ -161,10 +146,7 @@ private[sage] object Geo { Command( "GEOSEARCHSTORE", Vector(0, 1), - (Vector(keyCodec.encode(destination), keyCodec.encode(source)) ++ originArgs(origin) ++ shapeArgs(shape) ++ sortArgs(sort) ++ countArgs( - count - )) ++ - (if (storeDist) Vector(StoreDist) else Vector.empty), + (Vector(keyCodec.encode(destination), keyCodec.encode(source)) ++ searchArgs(origin, shape, sort, count)) ++ Args.flag(storeDist, StoreDist), Decode.long ) @@ -176,33 +158,28 @@ private[sage] object Geo { } private def memberCoordArgs[V](pairs: Vector[(V, GeoCoordinates)])(using valueCodec: ValueCodec[V]): Vector[Bytes] = - pairs.flatMap { case (member, coords) => Vector(coordArg(coords.longitude), coordArg(coords.latitude), valueCodec.encode(member)) } + pairs.flatMap { case (member, coords) => Vector(Args.double(coords.longitude), Args.double(coords.latitude), valueCodec.encode(member)) } private def originArgs[V](origin: GeoOrigin[V])(using valueCodec: ValueCodec[V]): Vector[Bytes] = origin match { case GeoOrigin.FromMember(member) => Vector(FromMember, valueCodec.encode(member)) - case GeoOrigin.FromLonLat(coords) => Vector(FromLonLat, coordArg(coords.longitude), coordArg(coords.latitude)) + case GeoOrigin.FromLonLat(coords) => Vector(FromLonLat, Args.double(coords.longitude), Args.double(coords.latitude)) } private def shapeArgs(shape: GeoShape): Vector[Bytes] = shape match { - case GeoShape.ByRadius(radius, unit) => Vector(ByRadius, Bytes.utf8(Doubles.format(radius)), unitArg(unit)) - case GeoShape.ByBox(width, height, unit) => Vector(ByBox, Bytes.utf8(Doubles.format(width)), Bytes.utf8(Doubles.format(height)), unitArg(unit)) - } - - private def sortArgs(sort: Option[GeoSort]): Vector[Bytes] = - sort.toVector.map { - case GeoSort.Asc => Asc - case GeoSort.Desc => Desc + case GeoShape.ByRadius(radius, unit) => Vector(ByRadius, Args.double(radius), unitArg(unit)) + case GeoShape.ByBox(width, height, unit) => Vector(ByBox, Args.double(width), Args.double(height), unitArg(unit)) } - private def countArgs(count: Option[GeoCount]): Vector[Bytes] = - count.toVector.flatMap(c => Vector(CountWord, Bytes.utf8(c.count.toString)) ++ (if (c.any) Vector(AnyWord) else Vector.empty)) + private def searchArgs[V: ValueCodec](origin: GeoOrigin[V], shape: GeoShape, sort: Option[GeoSort], count: Option[GeoCount]): Vector[Bytes] = + originArgs(origin) ++ shapeArgs(shape) ++ sort.toVector.map(sortWord) ++ + count.toVector.flatMap(c => Vector(Count, Args.long(c.count)) ++ Args.flag(c.any, AnyWord)) private def withArgs(withCoord: Boolean, withDist: Boolean, withHash: Boolean): Vector[Bytes] = - (if (withCoord) Vector(WithCoord) else Vector.empty) ++ - (if (withDist) Vector(WithDist) else Vector.empty) ++ - (if (withHash) Vector(WithHash) else Vector.empty) + Args.flag(withCoord, WithCoord) ++ + Args.flag(withDist, WithDist) ++ + Args.flag(withHash, WithHash) private def unitArg(unit: GeoUnit): Bytes = unit match { @@ -212,15 +189,10 @@ private[sage] object Geo { case GeoUnit.Feet => UnitFeet } - private def coordArg(value: Double): Bytes = Bytes.utf8(Doubles.format(value)) - private val coordinates: Frame => Either[DecodeError, GeoCoordinates] = - Decode.array2(Decode.lenientDouble, Decode.lenientDouble, "longitude/latitude pair")(GeoCoordinates(_, _)) + Decode.array2(Decode.double, Decode.double, "longitude/latitude pair")(GeoCoordinates(_, _)) - private val optionalCoordinates: Frame => Either[DecodeError, Option[GeoCoordinates]] = { - case Frame.Null => Right(None) - case other => coordinates(other).map(Some(_)) - } + private val optionalCoordinates: Frame => Either[DecodeError, Option[GeoCoordinates]] = Decode.nullable(coordinates) // Without any WITH flag GEOSEARCH replies a flat array of member bulk strings; any flag turns each row into [member, dist?, hash?, coord?] // in that fixed field order regardless of the order the flags were requested @@ -238,7 +210,7 @@ private[sage] object Geo { case Frame.Array(fields) if fields.length == width => for { member <- Decode.value[V](fields(0)) - distance <- if (withDist) Decode.lenientDouble(fields(distIdx)).map(Some(_)) else Right(None) + distance <- if (withDist) Decode.double(fields(distIdx)).map(Some(_)) else Right(None) hash <- if (withHash) Decode.long(fields(hashIdx)).map(Some(_)) else Right(None) coords <- if (withCoord) coordinates(fields(coordIdx)).map(Some(_)) else Right(None) } yield GeoSearchResult(member, distance, hash, coords) diff --git a/sage-core/src/main/scala/sage/commands/Hashes.scala b/sage-core/src/main/scala/sage/commands/Hashes.scala index 23dd05e3..94c8cda2 100644 --- a/sage-core/src/main/scala/sage/commands/Hashes.scala +++ b/sage-core/src/main/scala/sage/commands/Hashes.scala @@ -6,7 +6,8 @@ import scala.concurrent.duration.{FiniteDuration, MILLISECONDS, SECONDS, TimeUni import sage.Bytes import sage.SageException.DecodeError -import sage.codec.{Doubles, KeyCodec, ValueCodec} +import sage.codec.{KeyCodec, ValueCodec} +import sage.commands.Args.WithValues import sage.protocol.Frame /** @@ -63,18 +64,17 @@ enum HSetExCondition { */ private[sage] object Hashes { - private val WithValues = Bytes.utf8("WITHVALUES") - private val NoValues = Bytes.utf8("NOVALUES") - private val Fields = Bytes.utf8("FIELDS") - private val Fnx = Bytes.utf8("FNX") - private val Fxx = Bytes.utf8("FXX") + private val NoValuesTail = Vector(Bytes.utf8("NOVALUES")) + private val Fields = Bytes.utf8("FIELDS") + private val Fnx = Bytes.utf8("FNX") + private val Fxx = Bytes.utf8("FXX") def hSet[K, F, V](key: K, first: (F, V), rest: (F, V)*)( using keyCodec: KeyCodec[K], fieldCodec: KeyCodec[F], valueCodec: ValueCodec[V] ): Command[Long] = - Command("HSET", Command.FirstKey, keyCodec.encode(key) +: fieldValueArgs(first +: rest.toVector), Decode.long) + Command("HSET", Command.FirstKey, keyCodec.encode(key) +: Args.pairs(first +: rest.toVector), Decode.long) def hSetNx[K, F, V](key: K, field: F, value: V)(using keyCodec: KeyCodec[K], fieldCodec: KeyCodec[F], valueCodec: ValueCodec[V]): Command[Boolean] = Command("HSETNX", Command.FirstKey, Vector(keyCodec.encode(key), fieldCodec.encode(field), valueCodec.encode(value)), Decode.flag) @@ -87,15 +87,10 @@ private[sage] object Hashes { fieldCodec: KeyCodec[F], valueCodec: ValueCodec[V] ): Command[Vector[Option[V]]] = - Command.read( - "HMGET", - Command.FirstKey, - keyCodec.encode(key) +: (first +: rest.toVector).map(fieldCodec.encode), - Decode.vector(Decode.optionalValue) - ) + Command.read("HMGET", Command.FirstKey, Args.keyThen(key, first, rest)(fieldCodec.encode), Decode.vector(Decode.optionalValue)) def hDel[K, F](key: K, first: F, rest: F*)(using keyCodec: KeyCodec[K], fieldCodec: KeyCodec[F]): Command[Long] = - Command("HDEL", Command.FirstKey, keyCodec.encode(key) +: (first +: rest.toVector).map(fieldCodec.encode), Decode.long) + Command("HDEL", Command.FirstKey, Args.keyThen(key, first, rest)(fieldCodec.encode), Decode.long) def hExists[K, F](key: K, field: F)(using keyCodec: KeyCodec[K], fieldCodec: KeyCodec[F]): Command[Boolean] = Command.read("HEXISTS", Command.FirstKey, Vector(keyCodec.encode(key), fieldCodec.encode(field)), Decode.flag) @@ -116,13 +111,13 @@ private[sage] object Hashes { Command.read("HGETALL", Command.FirstKey, Vector(keyCodec.encode(key)), Decode.map[F, V]) def hIncrBy[K, F](key: K, field: F, increment: Long)(using keyCodec: KeyCodec[K], fieldCodec: KeyCodec[F]): Command[Long] = - Command("HINCRBY", Command.FirstKey, Vector(keyCodec.encode(key), fieldCodec.encode(field), Bytes.utf8(increment.toString)), Decode.long) + Command("HINCRBY", Command.FirstKey, Vector(keyCodec.encode(key), fieldCodec.encode(field), Args.long(increment)), Decode.long) def hIncrByFloat[K, F](key: K, field: F, increment: Double)(using keyCodec: KeyCodec[K], fieldCodec: KeyCodec[F]): Command[Double] = Command( "HINCRBYFLOAT", Command.FirstKey, - Vector(keyCodec.encode(key), fieldCodec.encode(field), Bytes.utf8(Doubles.format(increment))), + Vector(keyCodec.encode(key), fieldCodec.encode(field), Args.double(increment)), Decode.double ) @@ -130,7 +125,7 @@ private[sage] object Hashes { Command.readUncacheable("HRANDFIELD", Command.FirstKey, Vector(keyCodec.encode(key)), Decode.optionalKey[F]) def hRandField[K, F](key: K, count: Long)(using keyCodec: KeyCodec[K], fieldCodec: KeyCodec[F]): Command[Vector[F]] = - Command.readUncacheable("HRANDFIELD", Command.FirstKey, Vector(keyCodec.encode(key), Bytes.utf8(count.toString)), Decode.vector(Decode.key[F])) + Command.readUncacheable("HRANDFIELD", Command.FirstKey, Vector(keyCodec.encode(key), Args.long(count)), Decode.vector(Decode.key[F])) def hRandFieldWithValues[K, F, V]( key: K, @@ -139,7 +134,7 @@ private[sage] object Hashes { Command.readUncacheable( "HRANDFIELD", Command.FirstKey, - Vector(keyCodec.encode(key), Bytes.utf8(count.toString), WithValues), + Vector(keyCodec.encode(key), Args.long(count), WithValues), Decode.nestedPairs[F, V] ) @@ -148,86 +143,50 @@ private[sage] object Hashes { fieldCodec: KeyCodec[F], valueCodec: ValueCodec[V] ): Command[ScanPage[(F, V)]] = - Command.readCursor( - "HSCAN", - Command.FirstKey, - Vector(keyCodec.encode(key), ScanCursor.bytes(cursor)) ++ ScanArgs.options(pattern, count), - Decode.scanPage(Decode.flatPairs[F, V]) - ) + KeyArgs.keyScan("HSCAN", key, cursor, pattern, count)(Decode.flatPairs[F, V]) def hScanNoValues[K, F](key: K, cursor: ScanCursor, pattern: Option[String] = None, count: Option[Long] = None)( using keyCodec: KeyCodec[K], fieldCodec: KeyCodec[F] ): Command[ScanPage[F]] = - Command.readCursor( - "HSCAN", - Command.FirstKey, - (Vector(keyCodec.encode(key), ScanCursor.bytes(cursor)) ++ ScanArgs.options(pattern, count)) :+ NoValues, - Decode.scanPage(Decode.vector(Decode.key[F])) - ) + KeyArgs.keyScan("HSCAN", key, cursor, pattern, count, NoValuesTail)(Decode.vector(Decode.key[F])) def hExpire[K, F](key: K, ttl: FiniteDuration, condition: ExpireCondition = ExpireCondition.Always)(first: F, rest: F*)( using keyCodec: KeyCodec[K], fieldCodec: KeyCodec[F] - ): Command[Vector[FieldExpiry]] = { - val (name, amount) = TimeArgs.expireCommand("HEXPIRE", "HPEXPIRE", ttl) - Command( - name, - Command.FirstKey, - Vector(keyCodec.encode(key), Bytes.utf8(amount.toString)) ++ Keys.conditionArgs(condition) ++ fieldsArgs(first +: rest.toVector), - Decode.vector(fieldExpiry) - ) - } + ): Command[Vector[FieldExpiry]] = + expireFields(TimeArgs.expireCommand("HEXPIRE", "HPEXPIRE", ttl), keyCodec.encode(key), condition, first +: rest.toVector) def hExpireAt[K, F](key: K, at: Instant, condition: ExpireCondition = ExpireCondition.Always)(first: F, rest: F*)( using keyCodec: KeyCodec[K], fieldCodec: KeyCodec[F] - ): Command[Vector[FieldExpiry]] = { - val (name, amount) = TimeArgs.expireCommand("HEXPIREAT", "HPEXPIREAT", at) - Command( - name, - Command.FirstKey, - Vector(keyCodec.encode(key), Bytes.utf8(amount.toString)) ++ Keys.conditionArgs(condition) ++ fieldsArgs(first +: rest.toVector), - Decode.vector(fieldExpiry) - ) - } + ): Command[Vector[FieldExpiry]] = + expireFields(TimeArgs.expireCommand("HEXPIREAT", "HPEXPIREAT", at), keyCodec.encode(key), condition, first +: rest.toVector) + + private def expireFields[F: KeyCodec](command: (String, Long), key: Bytes, condition: ExpireCondition, fields: Vector[F]) = + Command(command._1, Command.FirstKey, Vector(key, Args.long(command._2)) ++ Keys.conditionArgs(condition) ++ fieldsArgs(fields), fieldExpiries) def hExpireTime[K, F](key: K)(first: F, rest: F*)(using keyCodec: KeyCodec[K], fieldCodec: KeyCodec[F]): Command[Vector[FieldExpiryTime]] = - Command.readUncacheable( - "HEXPIRETIME", - Command.FirstKey, - keyCodec.encode(key) +: fieldsArgs(first +: rest.toVector), - Decode.vector(fieldExpiryTime(Instant.ofEpochSecond)) - ) + Command.readUncacheable("HEXPIRETIME", Command.FirstKey, keyFields(key, first, rest), Decode.vector(fieldExpiryTime(Instant.ofEpochSecond))) def hpExpireTime[K, F](key: K)(first: F, rest: F*)(using keyCodec: KeyCodec[K], fieldCodec: KeyCodec[F]): Command[Vector[FieldExpiryTime]] = - Command.readUncacheable( - "HPEXPIRETIME", - Command.FirstKey, - keyCodec.encode(key) +: fieldsArgs(first +: rest.toVector), - Decode.vector(fieldExpiryTime(Instant.ofEpochMilli)) - ) + Command.readUncacheable("HPEXPIRETIME", Command.FirstKey, keyFields(key, first, rest), Decode.vector(fieldExpiryTime(Instant.ofEpochMilli))) def hTtl[K, F](key: K)(first: F, rest: F*)(using keyCodec: KeyCodec[K], fieldCodec: KeyCodec[F]): Command[Vector[FieldTtl]] = - Command.readUncacheable("HTTL", Command.FirstKey, keyCodec.encode(key) +: fieldsArgs(first +: rest.toVector), Decode.vector(fieldTtl(SECONDS))) + Command.readUncacheable("HTTL", Command.FirstKey, keyFields(key, first, rest), Decode.vector(fieldTtl(SECONDS))) def hpTtl[K, F](key: K)(first: F, rest: F*)(using keyCodec: KeyCodec[K], fieldCodec: KeyCodec[F]): Command[Vector[FieldTtl]] = - Command.readUncacheable( - "HPTTL", - Command.FirstKey, - keyCodec.encode(key) +: fieldsArgs(first +: rest.toVector), - Decode.vector(fieldTtl(MILLISECONDS)) - ) + Command.readUncacheable("HPTTL", Command.FirstKey, keyFields(key, first, rest), Decode.vector(fieldTtl(MILLISECONDS))) def hPersist[K, F](key: K)(first: F, rest: F*)(using keyCodec: KeyCodec[K], fieldCodec: KeyCodec[F]): Command[Vector[FieldPersist]] = - Command("HPERSIST", Command.FirstKey, keyCodec.encode(key) +: fieldsArgs(first +: rest.toVector), Decode.vector(fieldPersist)) + Command("HPERSIST", Command.FirstKey, keyFields(key, first, rest), Decode.vector(fieldPersist)) def hGetDel[K, F, V](key: K)(first: F, rest: F*)( using keyCodec: KeyCodec[K], fieldCodec: KeyCodec[F], valueCodec: ValueCodec[V] ): Command[Vector[Option[V]]] = - Command("HGETDEL", Command.FirstKey, keyCodec.encode(key) +: fieldsArgs(first +: rest.toVector), Decode.vector(Decode.optionalValue)) + Command("HGETDEL", Command.FirstKey, keyFields(key, first, rest), Decode.vector(Decode.optionalValue)) def hGetEx[K, F, V](key: K, expiry: GetExpiry = GetExpiry.Keep)(first: F, rest: F*)( using keyCodec: KeyCodec[K], @@ -253,14 +212,14 @@ private[sage] object Hashes { Decode.flag ) - private def fieldValueArgs[F, V](pairs: Vector[(F, V)])(using fieldCodec: KeyCodec[F], valueCodec: ValueCodec[V]): Vector[Bytes] = - pairs.flatMap { case (field, value) => Vector(fieldCodec.encode(field), valueCodec.encode(value)) } + private def keyFields[K, F](key: K, first: F, rest: Seq[F])(using keyCodec: KeyCodec[K], fieldCodec: KeyCodec[F]): Vector[Bytes] = + keyCodec.encode(key) +: fieldsArgs(first +: rest.toVector) private def fieldsArgs[F](fields: Vector[F])(using fieldCodec: KeyCodec[F]): Vector[Bytes] = - Fields +: Bytes.utf8(fields.size.toString) +: fields.map(fieldCodec.encode) + Fields +: Args.long(fields.size) +: fields.map(fieldCodec.encode) private def fieldValuePairsArgs[F, V](pairs: Vector[(F, V)])(using fieldCodec: KeyCodec[F], valueCodec: ValueCodec[V]): Vector[Bytes] = - Fields +: Bytes.utf8(pairs.size.toString) +: fieldValueArgs(pairs) + Fields +: Args.long(pairs.size) +: Args.pairs(pairs) private def setExConditionArgs(condition: HSetExCondition): Vector[Bytes] = condition match { @@ -269,13 +228,12 @@ private[sage] object Hashes { case HSetExCondition.IfAllExist => Vector(Fxx) } - private val fieldExpiry: Frame => Either[DecodeError, FieldExpiry] = { + private val fieldExpiries: Frame => Either[DecodeError, Vector[FieldExpiry]] = Decode.vector(Decode.shape("field expiry integer") { case Frame.Integer(-2) => Right(FieldExpiry.NoField) case Frame.Integer(0) => Right(FieldExpiry.ConditionNotMet) case Frame.Integer(1) => Right(FieldExpiry.Updated) case Frame.Integer(2) => Right(FieldExpiry.Deleted) - case other => Left(DecodeError("field expiry integer", Frame.describe(other))) - } + }) private def fieldTtl(unit: TimeUnit): Frame => Either[DecodeError, FieldTtl] = Decode.expiryInteger(FieldTtl.NoField, FieldTtl.NoExpiry, "field ttl integer")(amount => FieldTtl.Expires(FiniteDuration(amount, unit))) @@ -285,10 +243,9 @@ private[sage] object Hashes { FieldExpiryTime.At(toInstant(amount)) ) - private val fieldPersist: Frame => Either[DecodeError, FieldPersist] = { + private val fieldPersist: Frame => Either[DecodeError, FieldPersist] = Decode.shape("field persist integer") { case Frame.Integer(-2) => Right(FieldPersist.NoField) case Frame.Integer(-1) => Right(FieldPersist.NoExpiry) case Frame.Integer(1) => Right(FieldPersist.Persisted) - case other => Left(DecodeError("field persist integer", Frame.describe(other))) } } diff --git a/sage-core/src/main/scala/sage/commands/HelloReply.scala b/sage-core/src/main/scala/sage/commands/HelloReply.scala deleted file mode 100644 index 51e9fda2..00000000 --- a/sage-core/src/main/scala/sage/commands/HelloReply.scala +++ /dev/null @@ -1,43 +0,0 @@ -package sage.commands - -import sage.SageException.DecodeError -import sage.protocol.Frame - -final private[sage] case class HelloReply(server: String, version: String, proto: Int, role: String) - -private[sage] object HelloReply { - - private[commands] def decode(frame: Frame): Either[DecodeError, HelloReply] = - frame match { - case Frame.Map(entries) => - val fields = entries.collect { - case (Frame.BulkString(key), value) => key.asUtf8String -> value - case (Frame.SimpleString(key), value) => key -> value - }.toMap - for { - server <- string(fields, "server") - version <- string(fields, "version") - proto <- proto3(fields) - role <- string(fields, "role") - } yield HelloReply(server, version, proto, role) - case other => - Left(DecodeError("map", Frame.describe(other))) - } - - private def string(fields: Map[String, Frame], name: String): Either[DecodeError, String] = - fields.get(name) match { - case Some(Frame.BulkString(value)) => Right(value.asUtf8String) - case Some(Frame.SimpleString(value)) => Right(value) - case Some(other) => Left(DecodeError(s"string for '$name'", Frame.describe(other))) - case None => Left(DecodeError(s"map entry '$name'", "absent")) - } - - // RESP3 is the only supported protocol, so any other proto value is a decode failure. - private def proto3(fields: Map[String, Frame]): Either[DecodeError, Int] = - fields.get("proto") match { - case Some(Frame.Integer(3)) => Right(3) - case Some(Frame.Integer(value)) => Left(DecodeError("proto 3", s"proto $value")) - case Some(other) => Left(DecodeError("integer for 'proto'", Frame.describe(other))) - case None => Left(DecodeError("map entry 'proto'", "absent")) - } -} diff --git a/sage-core/src/main/scala/sage/commands/HyperLogLog.scala b/sage-core/src/main/scala/sage/commands/HyperLogLog.scala index db65a474..17b4af56 100644 --- a/sage-core/src/main/scala/sage/commands/HyperLogLog.scala +++ b/sage-core/src/main/scala/sage/commands/HyperLogLog.scala @@ -9,13 +9,9 @@ private[sage] object HyperLogLog { Command("PFADD", Command.FirstKey, keyCodec.encode(key) +: elements.toVector.map(valueCodec.encode), Decode.flag) // PFCOUNT is cacheable. Its documented internal register-cache write does not change the estimate or send an invalidation. - def pfCount[K](first: K, rest: K*)(using keyCodec: KeyCodec[K]): Command[Long] = { - val keys = (first +: rest).iterator.map(keyCodec.encode).toVector - Command.read("PFCOUNT", keys.indices.toVector, keys, Decode.long) - } + def pfCount[K](first: K, rest: K*)(using KeyCodec[K]): Command[Long] = + KeyArgs.allKeys("PFCOUNT", first +: rest.toVector, Decode.long, readOnly = true) - def pfMerge[K](destination: K, sources: K*)(using keyCodec: KeyCodec[K]): Command[Unit] = { - val keys = (destination +: sources).iterator.map(keyCodec.encode).toVector - Command("PFMERGE", keys.indices.toVector, keys, Decode.ok) - } + def pfMerge[K](destination: K, sources: K*)(using KeyCodec[K]): Command[Unit] = + KeyArgs.allKeys("PFMERGE", destination +: sources.toVector, Decode.ok) } diff --git a/sage-core/src/main/scala/sage/commands/Invalidation.scala b/sage-core/src/main/scala/sage/commands/Invalidation.scala index fa0f8edd..a1fca940 100644 --- a/sage-core/src/main/scala/sage/commands/Invalidation.scala +++ b/sage-core/src/main/scala/sage/commands/Invalidation.scala @@ -15,32 +15,14 @@ private[sage] enum Invalidation { private[sage] object Invalidation { /** - * Decodes the elements of an invalidation push frame. Returns `None` for malformed frames and other push kinds. A null key list becomes - * [[FlushAll]]. Pub/sub push frames are handled separately by [[Pubsub.decode]]. + * Decodes the elements of an invalidation push frame. Returns `None` for other push kinds. A null or undecodable key list becomes + * [[FlushAll]], so no stale entry survives. Pub/sub push frames are handled separately by [[Pubsub.decode]]. */ def decode(elements: Vector[Frame]): Option[Invalidation] = elements match { - case Vector(kind, keys) if isInvalidate(kind) => - keys match { - case Frame.Null => Some(FlushAll) - case Frame.Array(elements) => - val builder = Vector.newBuilder[Bytes] - val it = elements.iterator - while (it.hasNext) - it.next() match { - case Frame.BulkString(key) => builder += key - case _ => return None - } - Some(Evict(builder.result())) - case _ => None - } - case _ => None + case Vector(Decode.Text("invalidate"), keys) => Some(evictedKeys(keys).fold(_ => FlushAll, Evict(_))) + case _ => None } - private def isInvalidate(frame: Frame): Boolean = - frame match { - case Frame.BulkString(b) => b.asUtf8String == "invalidate" - case Frame.SimpleString(s) => s == "invalidate" - case _ => false - } + private val evictedKeys = Decode.vector(Decode.bytes) } diff --git a/sage-core/src/main/scala/sage/commands/Json.scala b/sage-core/src/main/scala/sage/commands/Json.scala index 5dad79fb..cb166514 100644 --- a/sage-core/src/main/scala/sage/commands/Json.scala +++ b/sage-core/src/main/scala/sage/commands/Json.scala @@ -3,6 +3,7 @@ package sage.commands import sage.Bytes import sage.SageException.DecodeError import sage.codec.{Doubles, KeyCodec, ValueCodec} +import sage.commands.Args.{Nx, Xx} import sage.protocol.Frame /** @@ -39,24 +40,13 @@ enum JsonType { object JsonType { - private[commands] def fromWireName(name: java.lang.String): JsonType = - name match { - case "object" => Object - case "array" => Array - case "string" => String - case "number" => Number - case "integer" => Integer - case "boolean" => Boolean - case "null" => Null - case other => Other(other) - } + private val byWireName = Decode.byLowerName(Object, Array, String, Number, Integer, Boolean, Null) + + private[commands] def fromWireName(name: java.lang.String): JsonType = byWireName.getOrElse(name, Other(name)) } private[sage] object Json { - private val Nx = Bytes.utf8("NX") - private val Xx = Bytes.utf8("XX") - def jsonSet[K, V](key: K, path: JsonPath, value: V, condition: JsonSetCondition = JsonSetCondition.Always)( using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V] @@ -65,11 +55,7 @@ private[sage] object Json { "JSON.SET", Command.FirstKey, Vector(keyCodec.encode(key), JsonPath.encode(path), valueCodec.encode(value)) ++ conditionArgs(condition), - decode = { - case Frame.SimpleString("OK") => Right(true) - case Frame.Null => Right(false) - case other => Left(DecodeError("simple string 'OK' or null", Frame.describe(other))) - } + Decode.okOrNull ) def jsonGet[K, V](key: K, paths: JsonPath*)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Option[V]] = @@ -126,7 +112,7 @@ private[sage] object Json { Command( "JSON.NUMINCRBY", Command.FirstKey, - Vector(keyCodec.encode(key), JsonPath.encode(path), Bytes.utf8(Doubles.format(increment))), + Vector(keyCodec.encode(key), JsonPath.encode(path), Args.double(increment)), numResult ) @@ -134,7 +120,7 @@ private[sage] object Json { Command( "JSON.NUMMULTBY", Command.FirstKey, - Vector(keyCodec.encode(key), JsonPath.encode(path), Bytes.utf8(Doubles.format(multiplier))), + Vector(keyCodec.encode(key), JsonPath.encode(path), Args.double(multiplier)), numResult ) @@ -156,7 +142,7 @@ private[sage] object Json { Command.read( "JSON.ARRINDEX", Command.FirstKey, - Vector(keyCodec.encode(key), JsonPath.encode(path), valueCodec.encode(value), Bytes.utf8(start.toString), Bytes.utf8(stop.toString)), + Vector(keyCodec.encode(key), JsonPath.encode(path), valueCodec.encode(value), Args.long(start), Args.long(stop)), pathMulti(Decode.optionalLong) ) @@ -167,7 +153,7 @@ private[sage] object Json { Command( "JSON.ARRINSERT", Command.FirstKey, - Vector(keyCodec.encode(key), JsonPath.encode(path), Bytes.utf8(index.toString)) ++ (first +: rest.toVector).map(valueCodec.encode), + Vector(keyCodec.encode(key), JsonPath.encode(path), Args.long(index)) ++ (first +: rest.toVector).map(valueCodec.encode), pathMulti(Decode.optionalLong) ) @@ -181,7 +167,7 @@ private[sage] object Json { Command( "JSON.ARRPOP", Command.FirstKey, - Vector(keyCodec.encode(key), JsonPath.encode(path), Bytes.utf8(index.toString)), + Vector(keyCodec.encode(key), JsonPath.encode(path), Args.long(index)), pathMulti(Decode.optionalValue) ) @@ -189,7 +175,7 @@ private[sage] object Json { Command( "JSON.ARRTRIM", Command.FirstKey, - Vector(keyCodec.encode(key), JsonPath.encode(path), Bytes.utf8(start.toString), Bytes.utf8(stop.toString)), + Vector(keyCodec.encode(key), JsonPath.encode(path), Args.long(start), Args.long(stop)), pathMulti(Decode.optionalLong) ) @@ -211,10 +197,8 @@ private[sage] object Json { Command.readUncacheable("JSON.RESP", Command.FirstKey, Vector(keyCodec.encode(key), JsonPath.encode(path)), Decode.frame) // JSONPath commands return an array with one value for each match. A legacy path without $ returns a single value and is rejected here. - private def pathMulti[A](element: Frame => Either[DecodeError, A]): Frame => Either[DecodeError, Vector[A]] = { - case array @ Frame.Array(_) => Decode.vector(element)(array) - case other => Left(DecodeError("a JSONPath ($) reply array (legacy '.' paths are unsupported)", Frame.describe(other))) - } + private def pathMulti[A](element: Frame => Either[DecodeError, A]): Frame => Either[DecodeError, Vector[A]] = + Decode.vector(element, "a JSONPath ($) reply array (legacy '.' paths are unsupported)") private def conditionArgs(condition: JsonSetCondition): Vector[Bytes] = condition match { @@ -226,45 +210,37 @@ private[sage] object Json { private def tripleArgs[K, V](triple: (K, JsonPath, V))(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Vector[Bytes] = Vector(keyCodec.encode(triple._1), JsonPath.encode(triple._2), valueCodec.encode(triple._3)) - private val optionalFlag: Frame => Either[DecodeError, Option[Boolean]] = { + private val optionalFlag: Frame => Either[DecodeError, Option[Boolean]] = Decode.shape("integer 0 or 1, or null") { case Frame.Null => Right(None) case Frame.Integer(0) => Right(Some(false)) case Frame.Integer(1) => Right(Some(true)) - case other => Left(DecodeError("integer 0 or 1, or null", Frame.describe(other))) } - private val optionalKeys: Frame => Either[DecodeError, Option[Vector[String]]] = { - case Frame.Null => Right(None) - case array @ Frame.Array(_) => Decode.vector(Decode.utf8String)(array).map(Some(_)) - case other => Left(DecodeError("array of keys, or null", Frame.describe(other))) - } + private val optionalKeys = Decode.nullable(Decode.vector(Decode.utf8String, "array of keys, or null")) - private val optionalTypeName: Frame => Either[DecodeError, Option[JsonType]] = { + private val optionalTypeName: Frame => Either[DecodeError, Option[JsonType]] = Decode.shape("type name string or null") { case Frame.Null => Right(None) case Frame.SimpleString(v) => Right(Some(JsonType.fromWireName(v))) case Frame.BulkString(b) => Right(Some(JsonType.fromWireName(b.asUtf8String))) - case other => Left(DecodeError("type name string or null", Frame.describe(other))) } // JSON.TYPE: Redis wraps the whole type list in one outer array, Valkey replies it flat; unify to one type name per match - private val typeReply: Frame => Either[DecodeError, Vector[Option[JsonType]]] = { - case Frame.Array(Vector(Frame.Array(inner))) => Decode.each(inner)(optionalTypeName) - case Frame.Array(elements) => Decode.each(elements)(optionalTypeName) - case other => Left(DecodeError("a JSONPath ($) reply array of type names (legacy '.' paths are unsupported)", Frame.describe(other))) - } + private val typeReply: Frame => Either[DecodeError, Vector[Option[JsonType]]] = + Decode.shape("a JSONPath ($) reply array of type names (legacy '.' paths are unsupported)") { + case Frame.Array(Vector(Frame.Array(inner))) => Decode.each(inner)(optionalTypeName) + case Frame.Array(elements) => Decode.each(elements)(optionalTypeName) + } - private val optionalNumber: Frame => Either[DecodeError, Option[Double]] = { + private val optionalNumber: Frame => Either[DecodeError, Option[Double]] = Decode.shape("number or null") { case Frame.Null => Right(None) case Frame.Integer(v) => Right(Some(v.toDouble)) case Frame.Double(v) => Right(Some(v)) - case other => Left(DecodeError("number or null", Frame.describe(other))) } // JSON.NUMINCRBY: Redis replies a RESP3 number array, Valkey a JSON-array bulk string; unify to one new value per match - private val numResult: Frame => Either[DecodeError, Vector[Option[Double]]] = { + private val numResult: Frame => Either[DecodeError, Vector[Option[Double]]] = Decode.shape("array of numbers or a JSON array string") { case Frame.Array(elements) => Decode.each(elements)(optionalNumber) case Frame.BulkString(bytes) => parseNumberArray(bytes.asUtf8String) - case other => Left(DecodeError("array of numbers or a JSON array string", Frame.describe(other))) } private def parseNumberArray(text: String): Either[DecodeError, Vector[Option[Double]]] = { diff --git a/sage-core/src/main/scala/sage/commands/KeyArgs.scala b/sage-core/src/main/scala/sage/commands/KeyArgs.scala index 78f9b2b0..80888406 100644 --- a/sage-core/src/main/scala/sage/commands/KeyArgs.scala +++ b/sage-core/src/main/scala/sage/commands/KeyArgs.scala @@ -1,13 +1,92 @@ package sage.commands import sage.Bytes -import sage.codec.KeyCodec +import sage.SageException.DecodeError +import sage.codec.{KeyCodec, ValueCodec} +import sage.protocol.Frame private[commands] object KeyArgs { + def allKeys[K, Out](name: String, keys: Vector[K], decode: Frame => Either[DecodeError, Out], readOnly: Boolean = false)( + using keyCodec: KeyCodec[K] + ): Command[Out] = { + val args = keys.map(keyCodec.encode) + if (readOnly) Command.read(name, args.indices.toVector, args, decode) + else Command(name, args.indices.toVector, args, decode) + } + // the `numkeys key…` block: encoded keys prefixed by their count, with 1-based key positions def numKeyed[K](keys: Vector[K])(using keyCodec: KeyCodec[K]): (Vector[Int], Vector[Bytes]) = { val encoded = keys.map(keyCodec.encode) - (Vector.tabulate(encoded.size)(_ + 1), Bytes.utf8(encoded.size.toString) +: encoded) + (Vector.tabulate(encoded.size)(_ + 1), Args.long(encoded.size) +: encoded) + } + + // `leading numkeys key…`, the layout of BLMPOP, BZMPOP, the sorted-set *STORE commands, EVAL and FCALL + def numKeyedAfter[K](leading: Bytes, keys: Seq[K])(using keyCodec: KeyCodec[K]): (Vector[Int], Vector[Bytes]) = { + val encoded = keys.iterator.map(keyCodec.encode).toVector + (Vector.range(2, 2 + encoded.length), leading +: Args.long(encoded.length) +: encoded) + } + + enum ScriptVerb(val wire: String, val readOnly: Boolean) { + case Eval extends ScriptVerb("EVAL", false) + case EvalRo extends ScriptVerb("EVAL_RO", true) + case EvalSha extends ScriptVerb("EVALSHA", false) + case EvalShaRo extends ScriptVerb("EVALSHA_RO", true) + case FCall extends ScriptVerb("FCALL", false) + case FCallRo extends ScriptVerb("FCALL_RO", true) + } + + def scriptCall[K, V](verb: ScriptVerb, target: String, keys: Seq[K], args: Seq[V])( + using keyCodec: KeyCodec[K], + valueCodec: ValueCodec[V] + ): Command[Frame] = { + val (indices, prefix) = numKeyedAfter(Bytes.utf8(target), keys) + Command(verb.wire, indices, prefix ++ args.iterator.map(valueCodec.encode), Decode.frame, isReadOnly = verb.readOnly) + } + + // `[timeout] numkeys key… where [COUNT count]`, the layout of LMPOP, ZMPOP and their blocking forms + def multiPop[K, A]( + name: String, + timeout: Option[BlockTimeout], + keys: Vector[K], + where: Bytes, + count: Option[Long], + items: Frame => Either[DecodeError, A], + label: String + )(using keyCodec: KeyCodec[K]): Command[Option[(K, A)]] = { + val (indices, prefix) = timeout match { + case None => numKeyed(keys) + case Some(timeout) => numKeyedAfter(BlockTimeout.wire(timeout), keys) + } + val decode = Decode.nullable(Decode.array2(Decode.key[K], items, label)(_ -> _)) + Command( + name, + indices, + (prefix :+ where) ++ Args.optLong(Args.Count, count), + decode, + if (timeout.isEmpty) Execution.Ordinary else Execution.Blocking + ) + } + + def interCard[K](name: String, keys: Vector[K], limit: Option[Long])(using keyCodec: KeyCodec[K]): Command[Long] = { + val (indices, prefix) = numKeyed(keys) + Command.read(name, indices, prefix ++ Args.optLong(Args.LimitWord, limit), Decode.long) + } + + def keyScan[K, A](name: String, key: K, cursor: ScanCursor, pattern: Option[String], count: Option[Long], tail: Vector[Bytes] = Vector.empty)( + items: Frame => Either[DecodeError, Vector[A]] + )(using keyCodec: KeyCodec[K]): Command[ScanPage[A]] = + Command.readCursor( + name, + Command.FirstKey, + (Vector(keyCodec.encode(key), ScanCursor.bytes(cursor)) ++ Args.scanOptions(pattern, count)) ++ tail, + Decode.scanPage(items) + ) + + def blockingPop[K, A](name: String, keys: Vector[K], timeout: BlockTimeout, decode: Frame => Either[DecodeError, A])( + using keyCodec: KeyCodec[K] + ): Command[Option[A]] = { + val encoded = keys.map(keyCodec.encode) + Command(name, Vector.range(0, encoded.size), encoded :+ BlockTimeout.wire(timeout), Decode.nullable(decode), Execution.Blocking) } } diff --git a/sage-core/src/main/scala/sage/commands/Keys.scala b/sage-core/src/main/scala/sage/commands/Keys.scala index 6190d115..307808ba 100644 --- a/sage-core/src/main/scala/sage/commands/Keys.scala +++ b/sage-core/src/main/scala/sage/commands/Keys.scala @@ -7,6 +7,7 @@ import scala.concurrent.duration.{FiniteDuration, MILLISECONDS, SECONDS} import sage.Bytes import sage.SageException.DecodeError import sage.codec.{KeyCodec, ValueCodec} +import sage.commands.Args.{Desc, Get, Gt, Lt, Nx, Replace, Xx} import sage.protocol.Frame /** @@ -34,25 +35,13 @@ object RedisType { private[commands] def wireName(tpe: RedisType): java.lang.String = tpe match { - case String => "string" - case List => "list" - case Set => "set" - case ZSet => "zset" - case Hash => "hash" - case Stream => "stream" case Other(name) => name + case simple => simple.toString.toLowerCase(java.util.Locale.ROOT) } - private[commands] def fromWireName(name: java.lang.String): RedisType = - name match { - case "string" => String - case "list" => List - case "set" => Set - case "zset" => ZSet - case "hash" => Hash - case "stream" => Stream - case other => Other(other) - } + private val byWireName = Decode.byLowerName(String, List, Set, ZSet, Hash, Stream) + + private[commands] def fromWireName(name: java.lang.String): RedisType = byWireName.getOrElse(name, Other(name)) } /** @@ -129,54 +118,46 @@ enum MigrateResult { private[sage] object Keys { - private val Replace = Bytes.utf8("REPLACE") - private val Type = Bytes.utf8("TYPE") - private val Nx = Bytes.utf8("NX") - private val Xx = Bytes.utf8("XX") - private val Gt = Bytes.utf8("GT") - private val Lt = Bytes.utf8("LT") - private val By = Bytes.utf8("BY") - private val Get = Bytes.utf8("GET") - private val LimitWord = Bytes.utf8("LIMIT") - private val Desc = Bytes.utf8("DESC") - private val Alpha = Bytes.utf8("ALPHA") - private val Store = Bytes.utf8("STORE") - private val AbsTtl = Bytes.utf8("ABSTTL") - private val IdleTime = Bytes.utf8("IDLETIME") - private val Freq = Bytes.utf8("FREQ") - private val Copy = Bytes.utf8("COPY") - private val Auth = Bytes.utf8("AUTH") - private val Auth2 = Bytes.utf8("AUTH2") - private val KeysWord = Bytes.utf8("KEYS") - private val Encoding = Bytes.utf8("ENCODING") - private val RefCount = Bytes.utf8("REFCOUNT") + private val Type = Bytes.utf8("TYPE") + private val By = Bytes.utf8("BY") + private val Alpha = Bytes.utf8("ALPHA") + private val Store = Bytes.utf8("STORE") + private val AbsTtl = Bytes.utf8("ABSTTL") + private val IdleTime = Bytes.utf8("IDLETIME") + private val Freq = Bytes.utf8("FREQ") + private val Copy = Bytes.utf8("COPY") + private val Auth = Bytes.utf8("AUTH") + private val Auth2 = Bytes.utf8("AUTH2") + private val KeysWord = Bytes.utf8("KEYS") + private val Encoding = Bytes.utf8("ENCODING") + private val RefCount = Bytes.utf8("REFCOUNT") def copy[K](source: K, destination: K, replace: Boolean = false)(using keyCodec: KeyCodec[K]): Command[Boolean] = Command( "COPY", keyIndices = Vector(0, 1), - args = Vector(keyCodec.encode(source), keyCodec.encode(destination)) ++ (if (replace) Vector(Replace) else Vector.empty), + args = Vector(keyCodec.encode(source), keyCodec.encode(destination)) ++ Args.flag(replace, Replace), Decode.flag ) def del[K](first: K, rest: K*)(using keyCodec: KeyCodec[K]): Command[Long] = - allKeys("DEL", first +: rest.toVector, Decode.long) + KeyArgs.allKeys("DEL", first +: rest.toVector, Decode.long) // Valkey's atomic compare-and-delete: removes the key only if its current string value equals `value` def delIfEq[K, V](key: K, value: V)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Boolean] = Command("DELIFEQ", Command.FirstKey, Vector(keyCodec.encode(key), valueCodec.encode(value)), Decode.flag) def exists[K](first: K, rest: K*)(using keyCodec: KeyCodec[K]): Command[Long] = - allKeys("EXISTS", first +: rest.toVector, Decode.long, readOnly = true) + KeyArgs.allKeys("EXISTS", first +: rest.toVector, Decode.long, readOnly = true) def expire[K](key: K, in: FiniteDuration, condition: ExpireCondition = ExpireCondition.Always)(using keyCodec: KeyCodec[K]): Command[Boolean] = { val (name, amount) = TimeArgs.expireCommand("EXPIRE", "PEXPIRE", in) - Command(name, Command.FirstKey, Vector(keyCodec.encode(key), Bytes.utf8(amount.toString)) ++ conditionArgs(condition), Decode.flag) + Command(name, Command.FirstKey, Vector(keyCodec.encode(key), Args.long(amount)) ++ conditionArgs(condition), Decode.flag) } def expireAt[K](key: K, at: Instant, condition: ExpireCondition = ExpireCondition.Always)(using keyCodec: KeyCodec[K]): Command[Boolean] = { val (name, amount) = TimeArgs.expireCommand("EXPIREAT", "PEXPIREAT", at) - Command(name, Command.FirstKey, Vector(keyCodec.encode(key), Bytes.utf8(amount.toString)) ++ conditionArgs(condition), Decode.flag) + Command(name, Command.FirstKey, Vector(keyCodec.encode(key), Args.long(amount)) ++ conditionArgs(condition), Decode.flag) } def expireTime[K](key: K)(using keyCodec: KeyCodec[K]): Command[ExpiryTime] = @@ -223,13 +204,13 @@ private[sage] object Keys { "SCAN", Command.NoKeys, ScanCursor.bytes(cursor) +: - (ScanArgs.options(pattern, count) ++ - ofType.toVector.flatMap(t => Vector(Type, Bytes.utf8(RedisType.wireName(t))))), + (Args.scanOptions(pattern, count) ++ + Args.opt(Type, ofType)(t => Bytes.utf8(RedisType.wireName(t)))), Decode.scanPage(Decode.vector(Decode.key[K])) ) def touch[K](first: K, rest: K*)(using keyCodec: KeyCodec[K]): Command[Long] = - allKeys("TOUCH", first +: rest.toVector, Decode.long) + KeyArgs.allKeys("TOUCH", first +: rest.toVector, Decode.long) def ttl[K](key: K)(using keyCodec: KeyCodec[K]): Command[Ttl] = Command.readUncacheable("TTL", Command.FirstKey, Vector(keyCodec.encode(key)), ttlDecode(SECONDS)) @@ -247,7 +228,7 @@ private[sage] object Keys { ) def unlink[K](first: K, rest: K*)(using keyCodec: KeyCodec[K]): Command[Long] = - allKeys("UNLINK", first +: rest.toVector, Decode.long) + KeyArgs.allKeys("UNLINK", first +: rest.toVector, Decode.long) def sort[K, V]( key: K, @@ -289,7 +270,7 @@ private[sage] object Keys { } def move[K](key: K, db: Int)(using keyCodec: KeyCodec[K]): Command[Boolean] = - Command("MOVE", Command.FirstKey, Vector(keyCodec.encode(key), Bytes.utf8(db.toString)), Decode.flag) + Command("MOVE", Command.FirstKey, Vector(keyCodec.encode(key), Args.long(db)), Decode.flag) def dump[K](key: K)(using keyCodec: KeyCodec[K]): Command[Option[Bytes]] = Command.read("DUMP", Command.FirstKey, Vector(keyCodec.encode(key)), Decode.optionalBytes) @@ -304,17 +285,17 @@ private[sage] object Keys { )(using keyCodec: KeyCodec[K]): Command[Unit] = { val (ttl, absTtl) = expiry match { case RestoreExpiry.NoExpiry => (0L, false) - case RestoreExpiry.In(duration) => (TimeArgs.millis(duration), false) - case RestoreExpiry.At(at) => (TimeArgs.millis(at), true) + case RestoreExpiry.In(duration) => (TimeArgs.positiveMillis(duration), false) + case RestoreExpiry.At(at) => (TimeArgs.positiveMillis(at), true) } Command( "RESTORE", Command.FirstKey, - Vector(keyCodec.encode(key), Bytes.utf8(ttl.toString), payload) ++ - (if (replace) Vector(Replace) else Vector.empty) ++ - (if (absTtl) Vector(AbsTtl) else Vector.empty) ++ - idleTime.toVector.flatMap(d => Vector(IdleTime, Bytes.utf8(d.toSeconds.toString))) ++ - freq.toVector.flatMap(f => Vector(Freq, Bytes.utf8(f.toString))), + Vector(keyCodec.encode(key), Args.long(ttl), payload) ++ + Args.flag(replace, Replace) ++ + Args.flag(absTtl, AbsTtl) ++ + Args.opt(IdleTime, idleTime)(d => Args.long(d.toSeconds)) ++ + Args.optLong(Freq, freq), Decode.ok ) } @@ -330,9 +311,15 @@ private[sage] object Keys { )(first: K, rest: K*)(using keyCodec: KeyCodec[K]): Command[MigrateResult] = { val keys = (first +: rest.toVector).map(keyCodec.encode) val args = - Vector(Bytes.utf8(host), Bytes.utf8(port.toString), Bytes.empty, Bytes.utf8(destinationDb.toString), Bytes.utf8(timeout.toMillis.toString)) ++ - (if (copy) Vector(Copy) else Vector.empty) ++ - (if (replace) Vector(Replace) else Vector.empty) ++ + Vector( + Bytes.utf8(host), + Args.long(port), + Bytes.empty, + Args.long(destinationDb), + Args.long(TimeArgs.positiveMillis(timeout)) + ) ++ + Args.flag(copy, Copy) ++ + Args.flag(replace, Replace) ++ authArgs(auth) ++ (KeysWord +: keys) Command("MIGRATE", Vector.range(args.size - keys.size, args.size), args, migrateResult) @@ -350,23 +337,12 @@ private[sage] object Keys { def objectIdleTime[K](key: K)(using keyCodec: KeyCodec[K]): Command[Option[FiniteDuration]] = Command.readUncacheable("OBJECT", Vector(1), Vector(IdleTime, keyCodec.encode(key)), Decode.optionalLong).map(_.map(FiniteDuration(_, SECONDS))) - private def allKeys[K, Out](name: String, keys: Vector[K], decode: Frame => Either[DecodeError, Out], readOnly: Boolean = false)( - using keyCodec: KeyCodec[K] - ): Command[Out] = { - val args = keys.map(keyCodec.encode) - if (readOnly) Command.read(name, args.indices.toVector, args, decode) - else Command(name, args.indices.toVector, args, decode) - } - private def sortOptionArgs(by: Option[String], limit: Option[Limit], get: Vector[String], order: SortOrder, alpha: Boolean): Vector[Bytes] = - by.toVector.flatMap(p => Vector(By, Bytes.utf8(p))) ++ - limit.toVector.flatMap(l => Vector(LimitWord, Bytes.utf8(l.offset.toString), Bytes.utf8(l.count.toString))) ++ + Args.optText(By, by) ++ + Args.limit(limit) ++ get.flatMap(p => Vector(Get, Bytes.utf8(p))) ++ - (order match { - case SortOrder.Asc => Vector.empty - case SortOrder.Desc => Vector(Desc) - }) ++ - (if (alpha) Vector(Alpha) else Vector.empty) + Args.flag(order == SortOrder.Desc, Desc) ++ + Args.flag(alpha, Alpha) private def authArgs(auth: MigrateAuth): Vector[Bytes] = auth match { @@ -375,10 +351,9 @@ private[sage] object Keys { case MigrateAuth.UserPassword(user, password) => Vector(Auth2, Bytes.utf8(user), Bytes.utf8(password)) } - private val migrateResult: Frame => Either[DecodeError, MigrateResult] = { + private val migrateResult: Frame => Either[DecodeError, MigrateResult] = Decode.shape("simple string 'OK' or 'NOKEY'") { case Frame.SimpleString("OK") => Right(MigrateResult.Ok) case Frame.SimpleString("NOKEY") => Right(MigrateResult.NoKey) - case other => Left(DecodeError("simple string 'OK' or 'NOKEY'", Frame.describe(other))) } private[commands] def conditionArgs(condition: ExpireCondition): Vector[Bytes] = diff --git a/sage-core/src/main/scala/sage/commands/Lists.scala b/sage-core/src/main/scala/sage/commands/Lists.scala index 2b1b2348..ed63b430 100644 --- a/sage-core/src/main/scala/sage/commands/Lists.scala +++ b/sage-core/src/main/scala/sage/commands/Lists.scala @@ -3,6 +3,7 @@ package sage.commands import sage.Bytes import sage.SageException.DecodeError import sage.codec.{KeyCodec, ValueCodec} +import sage.commands.Args.Count import sage.protocol.Frame /** @@ -12,17 +13,6 @@ enum ListSide { case Left, Right } -object ListSide { - - private[commands] def wire(side: ListSide): Bytes = - side match { - case ListSide.Left => LeftWord - case ListSide.Right => RightWord - } - private val LeftWord = Bytes.utf8("LEFT") - private val RightWord = Bytes.utf8("RIGHT") -} - /** * Whether `LINSERT` places the new element `Before` or `After` the pivot. */ @@ -32,23 +22,22 @@ enum InsertPosition { private[sage] object Lists { - private val Before = Bytes.utf8("BEFORE") - private val After = Bytes.utf8("AFTER") - private val Rank = Bytes.utf8("RANK") - private val Count = Bytes.utf8("COUNT") - private val MaxLen = Bytes.utf8("MAXLEN") + private val sideArg = Args.keywords(ListSide.values) + private val positionArg = Args.keywords(InsertPosition.values) + private val Rank = Bytes.utf8("RANK") + private val MaxLen = Bytes.utf8("MAXLEN") def lPush[K, V](key: K, first: V, rest: V*)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Long] = - push("LPUSH", key, first +: rest.toVector) + Command("LPUSH", Command.FirstKey, Args.keyThen(key, first, rest)(valueCodec.encode), Decode.long) def rPush[K, V](key: K, first: V, rest: V*)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Long] = - push("RPUSH", key, first +: rest.toVector) + Command("RPUSH", Command.FirstKey, Args.keyThen(key, first, rest)(valueCodec.encode), Decode.long) def lPushX[K, V](key: K, first: V, rest: V*)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Long] = - push("LPUSHX", key, first +: rest.toVector) + Command("LPUSHX", Command.FirstKey, Args.keyThen(key, first, rest)(valueCodec.encode), Decode.long) def rPushX[K, V](key: K, first: V, rest: V*)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Long] = - push("RPUSHX", key, first +: rest.toVector) + Command("RPUSHX", Command.FirstKey, Args.keyThen(key, first, rest)(valueCodec.encode), Decode.long) def lPop[K, V](key: K)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Option[V]] = Command("LPOP", Command.FirstKey, Vector(keyCodec.encode(key)), Decode.optionalValue) @@ -57,10 +46,10 @@ private[sage] object Lists { Command("RPOP", Command.FirstKey, Vector(keyCodec.encode(key)), Decode.optionalValue) def lPopCount[K, V](key: K, count: Long)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Vector[V]] = - Command("LPOP", Command.FirstKey, Vector(keyCodec.encode(key), Bytes.utf8(count.toString)), Decode.vectorOrEmpty(Decode.value[V])) + Command("LPOP", Command.FirstKey, Vector(keyCodec.encode(key), Args.long(count)), Decode.orEmpty(Decode.vector(Decode.value[V]))) def rPopCount[K, V](key: K, count: Long)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Vector[V]] = - Command("RPOP", Command.FirstKey, Vector(keyCodec.encode(key), Bytes.utf8(count.toString)), Decode.vectorOrEmpty(Decode.value[V])) + Command("RPOP", Command.FirstKey, Vector(keyCodec.encode(key), Args.long(count)), Decode.orEmpty(Decode.vector(Decode.value[V]))) def lLen[K](key: K)(using keyCodec: KeyCodec[K]): Command[Long] = Command.read("LLEN", Command.FirstKey, Vector(keyCodec.encode(key)), Decode.long) @@ -69,15 +58,15 @@ private[sage] object Lists { Command.read( "LRANGE", Command.FirstKey, - Vector(keyCodec.encode(key), Bytes.utf8(start.toString), Bytes.utf8(stop.toString)), + Vector(keyCodec.encode(key), Args.long(start), Args.long(stop)), Decode.vector(Decode.value[V]) ) def lIndex[K, V](key: K, index: Long)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Option[V]] = - Command.read("LINDEX", Command.FirstKey, Vector(keyCodec.encode(key), Bytes.utf8(index.toString)), Decode.optionalValue) + Command.read("LINDEX", Command.FirstKey, Vector(keyCodec.encode(key), Args.long(index)), Decode.optionalValue) def lSet[K, V](key: K, index: Long, value: V)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Unit] = - Command("LSET", Command.FirstKey, Vector(keyCodec.encode(key), Bytes.utf8(index.toString), valueCodec.encode(value)), Decode.ok) + Command("LSET", Command.FirstKey, Vector(keyCodec.encode(key), Args.long(index), valueCodec.encode(value)), Decode.ok) // the list length after the insert; 0 if the key is absent, -1 if the pivot is not found def lInsert[K, V](key: K, position: InsertPosition, pivot: V, value: V)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Long] = @@ -89,35 +78,28 @@ private[sage] object Lists { ) def lRem[K, V](key: K, count: Long, value: V)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Long] = - Command("LREM", Command.FirstKey, Vector(keyCodec.encode(key), Bytes.utf8(count.toString), valueCodec.encode(value)), Decode.long) + Command("LREM", Command.FirstKey, Vector(keyCodec.encode(key), Args.long(count), valueCodec.encode(value)), Decode.long) def lTrim[K](key: K, start: Long, stop: Long)(using keyCodec: KeyCodec[K]): Command[Unit] = - Command("LTRIM", Command.FirstKey, Vector(keyCodec.encode(key), Bytes.utf8(start.toString), Bytes.utf8(stop.toString)), Decode.ok) + Command("LTRIM", Command.FirstKey, Vector(keyCodec.encode(key), Args.long(start), Args.long(stop)), Decode.ok) def lPos[K, V](key: K, element: V, rank: Option[Long] = None, maxLen: Option[Long] = None)( using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V] ): Command[Option[Long]] = - Command.read( - "LPOS", - Command.FirstKey, - Vector(keyCodec.encode(key), valueCodec.encode(element)) ++ longArg(Rank, rank) ++ longArg(MaxLen, maxLen), - Decode.optionalLong - ) + Command.read("LPOS", Command.FirstKey, posArgs(key, element, rank, Vector.empty, maxLen), Decode.optionalLong) def lPosCount[K, V](key: K, element: V, count: Long, rank: Option[Long] = None, maxLen: Option[Long] = None)( using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V] ): Command[Vector[Long]] = - Command.read( - "LPOS", - Command.FirstKey, - Vector(keyCodec.encode(key), valueCodec.encode(element)) ++ longArg(Rank, rank) ++ Vector(Count, Bytes.utf8(count.toString)) ++ longArg( - MaxLen, - maxLen - ), - Decode.vector(Decode.long) - ) + Command.read("LPOS", Command.FirstKey, posArgs(key, element, rank, Vector(Count, Args.long(count)), maxLen), Decode.vector(Decode.long)) + + private def posArgs[K, V](key: K, element: V, rank: Option[Long], count: Vector[Bytes], maxLen: Option[Long])( + using keyCodec: KeyCodec[K], + valueCodec: ValueCodec[V] + ): Vector[Bytes] = + Vector(keyCodec.encode(key), valueCodec.encode(element)) ++ Args.optLong(Rank, rank) ++ count ++ Args.optLong(MaxLen, maxLen) def lMove[K, V](source: K, destination: K, from: ListSide, to: ListSide)( using keyCodec: KeyCodec[K], @@ -126,28 +108,21 @@ private[sage] object Lists { Command( "LMOVE", Vector(0, 1), - Vector(keyCodec.encode(source), keyCodec.encode(destination), ListSide.wire(from), ListSide.wire(to)), + Vector(keyCodec.encode(source), keyCodec.encode(destination), sideArg(from), sideArg(to)), Decode.optionalValue ) def lMpop[K, V]( first: K, rest: K* - )(side: ListSide, count: Option[Long] = None)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Option[(K, Vector[V])]] = { - val (keyIndices, prefix) = KeyArgs.numKeyed(first +: rest.toVector) - Command( - "LMPOP", - keyIndices, - args = (prefix :+ ListSide.wire(side)) ++ longArg(Count, count), - decode = mpopReply[K, V] - ) - } + )(side: ListSide, count: Option[Long] = None)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Option[(K, Vector[V])]] = + KeyArgs.multiPop("LMPOP", None, first +: rest.toVector, sideArg(side), count, Decode.vector(Decode.value[V]), MpopLabel) def blPop[K, V](first: K, rest: K*)(timeout: BlockTimeout)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Option[(K, V)]] = - blockingPop("BLPOP", first, rest.toVector, timeout) + KeyArgs.blockingPop("BLPOP", first +: rest.toVector, timeout, poppedPair[K, V]) def brPop[K, V](first: K, rest: K*)(timeout: BlockTimeout)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Option[(K, V)]] = - blockingPop("BRPOP", first, rest.toVector, timeout) + KeyArgs.blockingPop("BRPOP", first +: rest.toVector, timeout, poppedPair[K, V]) def blMove[K, V](source: K, destination: K, from: ListSide, to: ListSide, timeout: BlockTimeout)( using keyCodec: KeyCodec[K], @@ -156,7 +131,7 @@ private[sage] object Lists { Command( "BLMOVE", Vector(0, 1), - Vector(keyCodec.encode(source), keyCodec.encode(destination), ListSide.wire(from), ListSide.wire(to), BlockTimeout.wire(timeout)), + Vector(keyCodec.encode(source), keyCodec.encode(destination), sideArg(from), sideArg(to), BlockTimeout.wire(timeout)), Decode.optionalValue, Execution.Blocking ) @@ -164,49 +139,11 @@ private[sage] object Lists { def blMpop[K, V](first: K, rest: K*)(side: ListSide, timeout: BlockTimeout, count: Option[Long] = None)( using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V] - ): Command[Option[(K, Vector[V])]] = { - val keys = (first +: rest.toVector).map(keyCodec.encode) - Command( - "BLMPOP", - keyIndices = Vector.tabulate(keys.size)(_ + 2), - args = (BlockTimeout.wire(timeout) +: Bytes.utf8(keys.size.toString) +: keys :+ ListSide.wire(side)) ++ longArg(Count, count), - decode = mpopReply[K, V], - execution = Execution.Blocking - ) - } - - private def blockingPop[K, V](name: String, first: K, rest: Vector[K], timeout: BlockTimeout)( - using keyCodec: KeyCodec[K], - valueCodec: ValueCodec[V] - ): Command[Option[(K, V)]] = { - val keys = (first +: rest).map(keyCodec.encode) - Command( - name, - keyIndices = Vector.tabulate(keys.size)(identity), - args = keys :+ BlockTimeout.wire(timeout), - decode = { - case Frame.Null => Right(None) - case other => Decode.array2(Decode.key[K], Decode.value[V], "array of key and value or null")(_ -> _)(other).map(Some(_)) - }, - execution = Execution.Blocking - ) - } - - private def mpopReply[K, V](using KeyCodec[K], ValueCodec[V]): Frame => Either[DecodeError, Option[(K, Vector[V])]] = { - case Frame.Null => Right(None) - case other => - Decode.array2(Decode.key[K], Decode.vector(Decode.value[V]), "array of key and values or null")(_ -> _)(other).map(Some(_)) - } - - private def longArg(keyword: Bytes, value: Option[Long]): Vector[Bytes] = - value.toVector.flatMap(v => Vector(keyword, Bytes.utf8(v.toString))) + ): Command[Option[(K, Vector[V])]] = + KeyArgs.multiPop("BLMPOP", Some(timeout), first +: rest.toVector, sideArg(side), count, Decode.vector(Decode.value[V]), MpopLabel) - private def push[K, V](name: String, key: K, values: Vector[V])(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Long] = - Command(name, Command.FirstKey, keyCodec.encode(key) +: values.map(valueCodec.encode), Decode.long) + private def poppedPair[K, V](using KeyCodec[K], ValueCodec[V]): Frame => Either[DecodeError, (K, V)] = + Decode.array2(Decode.key[K], Decode.value[V], "array of key and value or null")(_ -> _) - private def positionArg(position: InsertPosition): Bytes = - position match { - case InsertPosition.Before => Before - case InsertPosition.After => After - } + private val MpopLabel = "array of key and values or null" } diff --git a/sage-core/src/main/scala/sage/commands/Merge.scala b/sage-core/src/main/scala/sage/commands/Merge.scala index 4bca5dda..f8f25e53 100644 --- a/sage-core/src/main/scala/sage/commands/Merge.scala +++ b/sage-core/src/main/scala/sage/commands/Merge.scala @@ -1,68 +1,23 @@ package sage.commands -import scala.collection.mutable - +import sage.SageException.DecodeError import sage.protocol.Frame /** - * Pairwise reply combiners used by [[BroadcastReduce.Fold]]. If either reply has an unexpected shape, the combiner returns that reply for - * the command's decoder to reject. Command-specific validation remains with the command. For example, `SCRIPT EXISTS` must also verify + * Pairwise reply combiners used by [[BroadcastReduce.Fold]]. If either reply has an unexpected shape, the combiner throws a + * `DecodeError`. Command-specific validation remains with the command. For example, `SCRIPT EXISTS` must also verify * that both arrays describe the same SHAs, which a general array combiner cannot do. */ private[commands] object Merge { - val sum: (Frame, Frame) => Frame = integers((x, y) => math.addExact(x, y)) - - val min: (Frame, Frame) => Frame = integers((x, y) => math.min(x, y)) - - /** - * Appends two arrays, dropping repeats: a classic channel can hold subscribers on several masters, so more than one reports it. - */ - val distinctChannels: (Frame, Frame) => Frame = (a, b) => - (a, b) match { - case (Frame.Array(x), Frame.Array(y)) => Frame.Array((x ++ y).distinct) - case (Frame.Array(_), bad) => bad - case (bad, _) => bad - } + val sum: (Frame, Frame) => Frame = typed(Decode.long, Frame.Integer(_))((x, y) => math.addExact(x, y)) - /** - * Sums a flat `[channel, count, …]` reply by channel and preserves the order in which channels first appear. Appending replies - * would repeat channels, and converting the result to a `Map` would discard all but the last count for each one. - */ - val sumByChannel: (Frame, Frame) => Frame = (a, b) => - (channelCounts(a), channelCounts(b)) match { - case (Some(x), Some(y)) => - val totals = mutable.LinkedHashMap.empty[Frame.BulkString, Long] - (x ++ y).foreach { case (channel, count) => totals.update(channel, math.addExact(totals.getOrElse(channel, 0L), count)) } - Frame.Array(totals.iterator.flatMap { case (channel, total) => Vector(channel, Frame.Integer(total)) }.toVector) - case (None, _) => a - case _ => b - } + val min: (Frame, Frame) => Frame = typed(Decode.long, Frame.Integer(_))((x, y) => math.min(x, y)) - /** - * Extracts channel/count pairs from a flat `[channel, count, …]` reply. Returns `None` for any other shape. The merge and the `PUBSUB - * NUMSUB` decoder share this function so they apply the same validation. - */ - def channelCounts(frame: Frame): Option[Vector[(Frame.BulkString, Long)]] = - frame match { - case Frame.Array(elements) if elements.length % 2 == 0 => - val builder = Vector.newBuilder[(Frame.BulkString, Long)] - var i = 0 - while (i < elements.length) { - (elements(i), elements(i + 1)) match { - case (channel: Frame.BulkString, Frame.Integer(count)) => builder += ((channel, count)) - case _ => return None - } - i += 2 - } - Some(builder.result()) - case _ => None - } + // appends two arrays, dropping repeats: a classic channel can hold subscribers on several masters, so more than one reports it + val distinct: (Frame, Frame) => Frame = typed(Decode.vector(Decode.frame), Frame.Array(_))((x, y) => (x ++ y).distinct) - private def integers(combine: (Long, Long) => Long): (Frame, Frame) => Frame = (a, b) => - (a, b) match { - case (Frame.Integer(x), Frame.Integer(y)) => Frame.Integer(combine(x, y)) - case (Frame.Integer(_), bad) => bad - case (bad, _) => bad - } + // decoding both replies with the command's decoder makes the merge accept exactly what the decoder accepts + def typed[A](decode: Frame => Either[DecodeError, A], encode: A => Frame)(combine: (A, A) => A): (Frame, Frame) => Frame = (a, b) => + decode(a).flatMap(x => decode(b).map(y => encode(combine(x, y)))).fold(error => throw error, identity) } diff --git a/sage-core/src/main/scala/sage/commands/Pipeline.scala b/sage-core/src/main/scala/sage/commands/Pipeline.scala index 3ad9f7bc..f5ff95ba 100644 --- a/sage-core/src/main/scala/sage/commands/Pipeline.scala +++ b/sage-core/src/main/scala/sage/commands/Pipeline.scala @@ -1,5 +1,7 @@ package sage.commands +import scala.util.boundary + import sage.SageException /** @@ -12,17 +14,16 @@ type Attempt[A] = Either[SageException, A] * Assembles the commands passed to `pipeline` and `exec`. The runtime batches commands by target connection. A cluster pipeline normally * sends one batch per target node and routes any command whose target cannot be resolved individually. Each command produces one typed * result. A pipeline does not provide transaction atomicity. The public methods accept either a tuple of commands with different result types or a - * `Seq[Command[A]]` with one result type. `Out` contains the all-success result, while `Results` contains an [[Attempt]] for each command. - * The runtime decodes positions independently into a `Vector[Either[SageException, Any]]`, then converts it to `Out` or `Results`. The - * internal `Any` values do not appear in the public result. + * `Seq[Command[A]]` with one result type. The runtime decodes positions independently into a `Vector[Either[SageException, Any]]`, then + * `finish` converts it to `R`: the strict factories fail with the first error, and the `*Attempt` factories keep an [[Attempt]] for each + * command. The internal `Any` values do not appear in the public result. */ -final private[sage] class Pipeline[Out, Results] private[commands] ( +final private[sage] class Pipeline[R] private ( /** * The composed commands, in send order. */ val commands: Vector[Command[?]], - private[sage] val toOut: Vector[Any] => Out, - private[sage] val toResults: Vector[Either[SageException, Any]] => Results + val finish: Vector[Attempt[Any]] => Attempt[R] ) private[sage] object Pipeline { @@ -30,25 +31,27 @@ private[sage] object Pipeline { /** * A dynamic, homogeneous pipeline. An empty sequence is a no-op that yields an empty result without touching the socket. */ - def sequence[A](commands: Seq[Command[A]]): Pipeline[Vector[A], Vector[Attempt[A]]] = - new Pipeline( - commands.toVector, - values => values.asInstanceOf[Vector[A]], - results => results.asInstanceOf[Vector[Attempt[A]]] - ) + def sequence[A](commands: Seq[Command[A]]): Pipeline[Vector[A]] = strict(commands.toVector, _.asInstanceOf[Vector[A]]) + + def sequenceAttempt[A](commands: Seq[Command[A]]): Pipeline[Vector[Attempt[A]]] = + new Pipeline(commands.toVector, results => Right(results.asInstanceOf[Vector[Attempt[A]]])) /** * A fixed-arity pipeline from a tuple of `Command`s, whose result tuple mirrors it element-for-element. `(get, incr)` yields - * `Pipeline[(Option[V], Long), (Attempt[Option[V]], Attempt[Long])]`. + * `Pipeline[(Option[V], Long)]`, and `fromTupleAttempt` yields `Pipeline[(Attempt[Option[V]], Attempt[Long])]`. */ - def fromTuple[T <: NonEmptyTuple]( + def fromTuple[T <: NonEmptyTuple](commands: T)(using Tuple.IsMappedBy[Command][T]): Pipeline[Tuple.InverseMap[T, Command]] = + strict(listed(commands), tuple) + + def fromTupleAttempt[T <: NonEmptyTuple]( commands: T - )(using Tuple.IsMappedBy[Command][T]): Pipeline[Tuple.InverseMap[T, Command], Tuple.Map[Tuple.InverseMap[T, Command], Attempt]] = { - val cmds = commands.toList.asInstanceOf[List[Command[?]]].toVector - new Pipeline( - cmds, - values => Tuple.fromArray(values.toArray).asInstanceOf[Tuple.InverseMap[T, Command]], - results => Tuple.fromArray(results.toArray).asInstanceOf[Tuple.Map[Tuple.InverseMap[T, Command], Attempt]] - ) - } + )(using Tuple.IsMappedBy[Command][T]): Pipeline[Tuple.Map[Tuple.InverseMap[T, Command], Attempt]] = + new Pipeline(listed(commands), results => Right(tuple(results))) + + private def listed(commands: NonEmptyTuple): Vector[Command[?]] = commands.toList.asInstanceOf[List[Command[?]]].toVector + + private def tuple[R](values: Vector[Any]): R = Tuple.fromArray(values.toArray).asInstanceOf[R] + + private def strict[R](commands: Vector[Command[?]], assemble: Vector[Any] => R): Pipeline[R] = + new Pipeline(commands, results => boundary(Right(assemble(results.map(_.fold(error => boundary.break(Left(error)), identity)))))) } diff --git a/sage-core/src/main/scala/sage/commands/Pubsub.scala b/sage-core/src/main/scala/sage/commands/Pubsub.scala index ff05ac0a..58d884aa 100644 --- a/sage-core/src/main/scala/sage/commands/Pubsub.scala +++ b/sage-core/src/main/scala/sage/commands/Pubsub.scala @@ -1,9 +1,9 @@ package sage.commands -import sage.Bytes +import sage.{Bytes, Message, PatternMessage} import sage.SageException.DecodeError import sage.codec.ValueCodec -import sage.protocol.{Frame, RespWriter} +import sage.protocol.Frame /** * Pub/sub command definitions. `PUBLISH`, `SPUBLISH`, and `PUBSUB` use the usual request/reply flow. In a cluster, `SPUBLISH` uses its @@ -20,16 +20,16 @@ private[sage] object Pubsub { Command("SPUBLISH", Command.FirstKey, Vector(Bytes.utf8(channel), codec.encode(message)), Decode.long) def pubsubChannels(pattern: Option[String] = None): Command[Vector[String]] = - introspect(Bytes.utf8("CHANNELS") +: pattern.map(Bytes.utf8).toVector, decodeStrings, Merge.distinctChannels) + introspect(Bytes.utf8("CHANNELS") +: pattern.map(Bytes.utf8).toVector, Decode.vector(Decode.utf8String), Merge.distinct) def pubsubShardChannels(pattern: Option[String] = None): Command[Vector[String]] = - introspect(Bytes.utf8("SHARDCHANNELS") +: pattern.map(Bytes.utf8).toVector, decodeStrings, Merge.distinctChannels) + introspect(Bytes.utf8("SHARDCHANNELS") +: pattern.map(Bytes.utf8).toVector, Decode.vector(Decode.utf8String), Merge.distinct) def pubsubNumSub(channels: String*): Command[Map[String, Long]] = - introspect(Bytes.utf8("NUMSUB") +: channels.toVector.map(Bytes.utf8), decodeNumSub, Merge.sumByChannel) + introspect(Bytes.utf8("NUMSUB") +: channels.toVector.map(Bytes.utf8), numSub, sumByChannel) def pubsubShardNumSub(channels: String*): Command[Map[String, Long]] = - introspect(Bytes.utf8("SHARDNUMSUB") +: channels.toVector.map(Bytes.utf8), decodeNumSub, Merge.sumByChannel) + introspect(Bytes.utf8("SHARDNUMSUB") +: channels.toVector.map(Bytes.utf8), numSub, sumByChannel) val pubsubNumPat: Command[Long] = introspect(Vector(Bytes.utf8("NUMPAT")), Decode.long, Merge.sum) @@ -46,24 +46,37 @@ private[sage] object Pubsub { ): Command[Out] = Command("PUBSUB", Command.NoKeys, args, decode, allMasters = true, broadcast = BroadcastReduce.Fold(merge)) - def subscribe(channels: Vector[String]): Bytes = RespWriter.writeCommand("SUBSCRIBE", channels.map(Bytes.utf8)) - def unsubscribe(channels: Vector[String]): Bytes = RespWriter.writeCommand("UNSUBSCRIBE", channels.map(Bytes.utf8)) - def psubscribe(patterns: Vector[String]): Bytes = RespWriter.writeCommand("PSUBSCRIBE", patterns.map(Bytes.utf8)) - def punsubscribe(patterns: Vector[String]): Bytes = RespWriter.writeCommand("PUNSUBSCRIBE", patterns.map(Bytes.utf8)) - def ssubscribe(channels: Vector[String]): Bytes = RespWriter.writeCommand("SSUBSCRIBE", channels.map(Bytes.utf8)) - def sunsubscribe(channels: Vector[String]): Bytes = RespWriter.writeCommand("SUNSUBSCRIBE", channels.map(Bytes.utf8)) + /** + * The three subscription kinds, each with its wire encoders: classic channels (`SUBSCRIBE`), glob patterns (`PSUBSCRIBE`), and shard + * channels (`SSUBSCRIBE`). + */ + enum Kind(subscribeVerb: String, unsubscribeVerb: String) { + case Channel extends Kind("SUBSCRIBE", "UNSUBSCRIBE") + case Pattern extends Kind("PSUBSCRIBE", "PUNSUBSCRIBE") + case Shard extends Kind("SSUBSCRIBE", "SUNSUBSCRIBE") + + // one name per command, so the server answers it with exactly one confirmation push or one error reply + def subscribe(name: String): Command[Unit] = Command(subscribeVerb, Command.NoKeys, Vector(Bytes.utf8(name)), _ => Right(())) + def unsubscribe(name: String): Command[Frame] = Command(unsubscribeVerb, Command.NoKeys, Vector(Bytes.utf8(name)), Decode.frame) + } + + // HELLO without arguments returns the connection's server information. It runs while the server is loading, stale or busy, and every + // connection may run it, since its setup did. + val helloInfo: Command[Frame] = Command("HELLO", Command.NoKeys, Vector.empty, Decode.frame) + + type Delivery = Message[Bytes] | PatternMessage[Bytes] /** - * A classified pub/sub push frame. Confirmations include the current subscription count. Deliveries contain raw payload bytes, which are - * decoded to the subscriber's value type at the stream boundary. + * A classified pub/sub push frame: a confirmation, a shard channel unsubscription, or a delivery. Deliveries contain raw payload bytes, + * which are decoded to the subscriber's value type at the stream boundary. */ enum Event { - case Subscribed(channel: String, count: Long) - case Unsubscribed(channel: String, count: Long) - case Message(channel: String, payload: Bytes) - // kept separate from Message so a connection can route classic and sharded deliveries to different subscribers. - case ShardMessage(channel: String, payload: Bytes) - case PatternMessage(pattern: String, channel: String, payload: Bytes) + // the reply to a SUBSCRIBE, PSUBSCRIBE, SSUBSCRIBE, UNSUBSCRIBE or PUNSUBSCRIBE, which the server never sends on its own + case Confirmed + // the server sends this push to confirm an SUNSUBSCRIBE and also when it drops a shard channel whose slot moved + case ShardUnsubscribed(channel: String) + // `subscription` is the channel or pattern under which the subscribers of `kind` are registered + case Delivered(kind: Kind, subscription: String, delivery: Delivery) } /** @@ -72,79 +85,30 @@ private[sage] object Pubsub { */ def decode(elements: Vector[Frame]): Option[Event] = elements match { - case Vector(kind, a, b) => - text(kind).flatMap { - case "message" => - for { - ch <- text(a) - p <- bytes(b) - } yield Event.Message(ch, p) - case "smessage" => - for { - ch <- text(a) - p <- bytes(b) - } yield Event.ShardMessage(ch, p) - case "subscribe" | "psubscribe" | "ssubscribe" => - for { - ch <- text(a) - c <- int(b) - } yield Event.Subscribed(ch, c) - case "unsubscribe" | "punsubscribe" | "sunsubscribe" => - for { - ch <- text(a) - c <- int(b) - } yield Event.Unsubscribed(ch, c) - case _ => None - } - case Vector(kind, p, ch, payload) => - text(kind).flatMap { - case "pmessage" => - for { - pat <- text(p) - c <- text(ch) - pl <- bytes(payload) - } yield Event.PatternMessage(pat, c, pl) - case _ => None - } - case _ => None - } - - private def text(frame: Frame): Option[String] = - frame match { - case Frame.BulkString(b) => Some(b.asUtf8String) - case Frame.SimpleString(s) => Some(s) - case _ => None + case Vector(Decode.Text("message"), Decode.Text(channel), Frame.BulkString(payload)) => + Some(Event.Delivered(Kind.Channel, channel, Message(channel, payload))) + case Vector(Decode.Text("smessage"), Decode.Text(channel), Frame.BulkString(payload)) => + Some(Event.Delivered(Kind.Shard, channel, Message(channel, payload))) + case Vector(Decode.Text("pmessage"), Decode.Text(pattern), Decode.Text(channel), Frame.BulkString(payload)) => + Some(Event.Delivered(Kind.Pattern, pattern, PatternMessage(pattern, channel, payload))) + case Vector(Decode.Text("subscribe" | "psubscribe" | "ssubscribe" | "unsubscribe" | "punsubscribe"), _, _) => Some(Event.Confirmed) + case Vector(Decode.Text("sunsubscribe"), Decode.Text(channel), _) => Some(Event.ShardUnsubscribed(channel)) + case _ => None } - private def bytes(frame: Frame): Option[Bytes] = - frame match { - case Frame.BulkString(b) => Some(b) - case _ => None - } - - private def int(frame: Frame): Option[Long] = - frame match { - case Frame.Integer(i) => Some(i) - case _ => None - } - - private def decodeStrings(frame: Frame): Either[DecodeError, Vector[String]] = - frame match { - case Frame.Array(elements) => - val builder = Vector.newBuilder[String] - val it = elements.iterator - while (it.hasNext) - it.next() match { - case Frame.BulkString(b) => builder += b.asUtf8String - case other => return Left(DecodeError("bulk string", Frame.describe(other))) - } - Right(builder.result()) - case other => Left(DecodeError("array", Frame.describe(other))) + private val numSub: Frame => Either[DecodeError, Map[String, Long]] = { + val pairs = Decode.flatPairsOf("array of channel/count pairs") { + case (Frame.BulkString(channel), Frame.Integer(count)) => Right(channel.asUtf8String -> count) + case (channel, count) => Left(DecodeError("channel/count pair", s"${Frame.describe(channel)} and ${Frame.describe(count)}")) } + pairs(_).map(_.toMap) + } - private def decodeNumSub(frame: Frame): Either[DecodeError, Map[String, Long]] = - Merge - .channelCounts(frame) - .map(_.iterator.map { case (channel, count) => channel.value.asUtf8String -> count }.toMap) - .toRight(DecodeError("array of channel/count pairs", Frame.describe(frame))) + // sums the decoded maps, so a channel a master reports twice still counts once for that master + private val sumByChannel = Merge.typed( + numSub, + counts => Frame.Array(counts.iterator.flatMap((c, n) => Iterator(Frame.BulkString(Bytes.utf8(c)), Frame.Integer(n))).toVector) + ) { (x, y) => + y.foldLeft(x) { case (acc, (channel, count)) => acc.updated(channel, math.addExact(acc.getOrElse(channel, 0L), count)) } + } } diff --git a/sage-core/src/main/scala/sage/commands/Reply.scala b/sage-core/src/main/scala/sage/commands/Reply.scala index f6500a6c..4eef42c6 100644 --- a/sage-core/src/main/scala/sage/commands/Reply.scala +++ b/sage-core/src/main/scala/sage/commands/Reply.scala @@ -1,6 +1,6 @@ package sage.commands -import scala.util.{Failure, Success, Try} +import scala.util.{Failure, Try} import scala.util.control.NonFatal import sage.SageException @@ -13,22 +13,16 @@ import sage.protocol.Frame */ private[sage] object Reply { - def run[Out](command: Command[Out], frame: Frame): Either[SageException, Out] = - frame match { - case Frame.SimpleError(message) => Left(ServerError.of(message)) - case Frame.BulkError(message) => Left(ServerError.of(message.asUtf8String)) - case other => command.decode(other) - } - /** * The decode boundary every transport uses: a throwing codec is caught and wrapped as a [[DecodeError]] (keeping the cause) rather than * escaping as a raw throwable. */ def decode[Out](command: Command[Out], frame: Frame): Try[Out] = try - run(command, frame) match { - case Right(value) => Success(value) - case Left(error) => Failure(error) + frame match { + case Frame.SimpleError(message) => Failure(ServerError.of(message)) + case Frame.BulkError(message) => Failure(ServerError.of(message.asUtf8String)) + case other => command.decode(other).toTry } catch { case error: SageException => Failure(error) diff --git a/sage-core/src/main/scala/sage/commands/ScanArgs.scala b/sage-core/src/main/scala/sage/commands/ScanArgs.scala deleted file mode 100644 index 39beb564..00000000 --- a/sage-core/src/main/scala/sage/commands/ScanArgs.scala +++ /dev/null @@ -1,13 +0,0 @@ -package sage.commands - -import sage.Bytes - -private[commands] object ScanArgs { - - val Match: Bytes = Bytes.utf8("MATCH") - val Count: Bytes = Bytes.utf8("COUNT") - - def options(pattern: Option[String], count: Option[Long]): Vector[Bytes] = - pattern.toVector.flatMap(p => Vector(Match, Bytes.utf8(p))) ++ - count.toVector.flatMap(n => Vector(Count, Bytes.utf8(n.toString))) -} diff --git a/sage-core/src/main/scala/sage/commands/Scripting.scala b/sage-core/src/main/scala/sage/commands/Scripting.scala index 0983e672..dd9a1f75 100644 --- a/sage-core/src/main/scala/sage/commands/Scripting.scala +++ b/sage-core/src/main/scala/sage/commands/Scripting.scala @@ -3,6 +3,7 @@ package sage.commands import sage.Bytes import sage.SageException.DecodeError import sage.codec.{KeyCodec, ValueCodec} +import sage.commands.KeyArgs.ScriptVerb import sage.protocol.Frame /** @@ -19,65 +20,47 @@ private[sage] object Scripting { private val Kill = Bytes.utf8("KILL") private val Show = Bytes.utf8("SHOW") - def eval(script: String): Command[Frame] = evalCommand("EVAL", script, Vector.empty, Vector.empty, readOnly = false) + def eval(script: String): Command[Frame] = KeyArgs.scriptCall(ScriptVerb.Eval, script, Seq.empty[Bytes], Seq.empty[Bytes]) - def eval[K](script: String, keys: Seq[K])(using keyCodec: KeyCodec[K]): Command[Frame] = - evalCommand("EVAL", script, keys.iterator.map(keyCodec.encode).toVector, Vector.empty, readOnly = false) + def eval[K](script: String, keys: Seq[K])(using KeyCodec[K]): Command[Frame] = + KeyArgs.scriptCall(ScriptVerb.Eval, script, keys, Seq.empty[Bytes]) - def eval[K, V](script: String, keys: Seq[K], args: Seq[V])(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Frame] = - evalCommand("EVAL", script, keys.iterator.map(keyCodec.encode).toVector, args.iterator.map(valueCodec.encode).toVector, readOnly = false) + def eval[K, V](script: String, keys: Seq[K], args: Seq[V])(using KeyCodec[K], ValueCodec[V]): Command[Frame] = + KeyArgs.scriptCall(ScriptVerb.Eval, script, keys, args) - def evalRo(script: String): Command[Frame] = evalCommand("EVAL_RO", script, Vector.empty, Vector.empty, readOnly = true) + def evalRo(script: String): Command[Frame] = KeyArgs.scriptCall(ScriptVerb.EvalRo, script, Seq.empty[Bytes], Seq.empty[Bytes]) - def evalRo[K](script: String, keys: Seq[K])(using keyCodec: KeyCodec[K]): Command[Frame] = - evalCommand("EVAL_RO", script, keys.iterator.map(keyCodec.encode).toVector, Vector.empty, readOnly = true) + def evalRo[K](script: String, keys: Seq[K])(using KeyCodec[K]): Command[Frame] = + KeyArgs.scriptCall(ScriptVerb.EvalRo, script, keys, Seq.empty[Bytes]) - def evalRo[K, V](script: String, keys: Seq[K], args: Seq[V])(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Frame] = - evalCommand("EVAL_RO", script, keys.iterator.map(keyCodec.encode).toVector, args.iterator.map(valueCodec.encode).toVector, readOnly = true) + def evalRo[K, V](script: String, keys: Seq[K], args: Seq[V])(using KeyCodec[K], ValueCodec[V]): Command[Frame] = + KeyArgs.scriptCall(ScriptVerb.EvalRo, script, keys, args) - def evalSha(sha: String): Command[Frame] = evalCommand("EVALSHA", sha, Vector.empty, Vector.empty, readOnly = false) + def evalSha(sha: String): Command[Frame] = KeyArgs.scriptCall(ScriptVerb.EvalSha, sha, Seq.empty[Bytes], Seq.empty[Bytes]) - def evalSha[K](sha: String, keys: Seq[K])(using keyCodec: KeyCodec[K]): Command[Frame] = - evalCommand("EVALSHA", sha, keys.iterator.map(keyCodec.encode).toVector, Vector.empty, readOnly = false) + def evalSha[K](sha: String, keys: Seq[K])(using KeyCodec[K]): Command[Frame] = + KeyArgs.scriptCall(ScriptVerb.EvalSha, sha, keys, Seq.empty[Bytes]) - def evalSha[K, V](sha: String, keys: Seq[K], args: Seq[V])(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Frame] = - evalCommand("EVALSHA", sha, keys.iterator.map(keyCodec.encode).toVector, args.iterator.map(valueCodec.encode).toVector, readOnly = false) + def evalSha[K, V](sha: String, keys: Seq[K], args: Seq[V])(using KeyCodec[K], ValueCodec[V]): Command[Frame] = + KeyArgs.scriptCall(ScriptVerb.EvalSha, sha, keys, args) - def evalShaRo(sha: String): Command[Frame] = evalCommand("EVALSHA_RO", sha, Vector.empty, Vector.empty, readOnly = true) + def evalShaRo(sha: String): Command[Frame] = KeyArgs.scriptCall(ScriptVerb.EvalShaRo, sha, Seq.empty[Bytes], Seq.empty[Bytes]) - def evalShaRo[K](sha: String, keys: Seq[K])(using keyCodec: KeyCodec[K]): Command[Frame] = - evalCommand("EVALSHA_RO", sha, keys.iterator.map(keyCodec.encode).toVector, Vector.empty, readOnly = true) + def evalShaRo[K](sha: String, keys: Seq[K])(using KeyCodec[K]): Command[Frame] = + KeyArgs.scriptCall(ScriptVerb.EvalShaRo, sha, keys, Seq.empty[Bytes]) - def evalShaRo[K, V](sha: String, keys: Seq[K], args: Seq[V])(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Frame] = - evalCommand("EVALSHA_RO", sha, keys.iterator.map(keyCodec.encode).toVector, args.iterator.map(valueCodec.encode).toVector, readOnly = true) - - private def evalCommand(name: String, script: String, keys: Vector[Bytes], args: Vector[Bytes], readOnly: Boolean): Command[Frame] = { - val allArgs = (Bytes.utf8(script) +: Bytes.utf8(keys.length.toString) +: keys) ++ args - val keyIndices = Vector.range(2, 2 + keys.length) - Command(name, keyIndices, allArgs, Decode.frame, Execution.Ordinary, isReadOnly = readOnly, cacheable = false) - } + def evalShaRo[K, V](sha: String, keys: Seq[K], args: Seq[V])(using KeyCodec[K], ValueCodec[V]): Command[Frame] = + KeyArgs.scriptCall(ScriptVerb.EvalShaRo, sha, keys, args) def scriptLoad(script: String): Command[String] = Command("SCRIPT", Command.NoKeys, Vector(Load, Bytes.utf8(script)), Decode.utf8String, allMasters = true) - private def isFlag(n: Long): Boolean = n == 0L || n == 1L - - private def flagFrame(f: Frame): Boolean = f match { - case Frame.Integer(n) => isFlag(n) - case _ => false - } + private val existsFlags = Decode.vector(Decode.flag) - private def andFlag(a: Frame, b: Frame): Frame = - (a, b) match { - case (Frame.Integer(x), Frame.Integer(y)) if isFlag(x) && isFlag(y) => Frame.Integer(if (x == 1L && y == 1L) 1L else 0L) - case (bad, _) if !flagFrame(bad) => bad - case (_, bad) => bad - } - - private val existsAnd: (Frame, Frame) => Frame = (a, b) => - (a, b) match { - case (Frame.Array(xs), Frame.Array(ys)) if xs.length == ys.length => Frame.Array(xs.lazyZip(ys).map(andFlag)) - case _ => throw DecodeError("SCRIPT EXISTS per-master flag arrays of equal length", s"${Frame.describe(a)} vs ${Frame.describe(b)}") + private val existsAnd = + Merge.typed[Vector[Boolean]](existsFlags, flags => Frame.Array(flags.map(f => Frame.Integer(if (f) 1L else 0L)))) { (xs, ys) => + if (xs.length == ys.length) xs.lazyZip(ys).map(_ && _) + else throw DecodeError("SCRIPT EXISTS per-master flag arrays of equal length", s"${xs.length} and ${ys.length} flags") } def scriptExists(first: String, rest: String*): Command[Vector[Boolean]] = @@ -85,7 +68,7 @@ private[sage] object Scripting { "SCRIPT", Command.NoKeys, Exists +: (first +: rest).iterator.map(Bytes.utf8).toVector, - Decode.vector(Decode.flag), + existsFlags, allMasters = true, broadcast = BroadcastReduce.Fold(existsAnd) ) @@ -99,3 +82,26 @@ private[sage] object Scripting { def scriptShow(sha: String): Command[String] = Command("SCRIPT", Command.NoKeys, Vector(Show, Bytes.utf8(sha)), Decode.utf8String) } + +// a Lua script that declares exactly one key, sent by digest (EVALSHA) or by body (EVAL) +final private[sage] class SingleKeyScript(source: String) { + private val body = Bytes.utf8(source) + val sha: String = java.security.MessageDigest.getInstance("SHA-1").digest(body.toArray).iterator.map(b => f"${b & 0xff}%02x").mkString + private val digest = Bytes.utf8(sha) + + def verb(cached: Boolean): String = if (cached) "EVALSHA" else "EVAL" + def reference(cached: Boolean): Bytes = if (cached) digest else body +} + +private[sage] object SingleKeyScript { + // the arguments are the script reference, numkeys = 1, the key, then ARGV + val NumKeys: Bytes = Bytes.utf8("1") + val KeyIndices: Vector[Int] = Vector(2) + + // length framing distinguishes namespace `a` with key `b:c` from namespace `a:b` with key `c` + def namespaced(namespace: String): Bytes => Bytes = { + val ns = Bytes.utf8(namespace) + val prefix = Bytes.concat(Vector(Bytes.utf8(s"${ns.length}:"), ns, Bytes.utf8(":"))) + key => Bytes.concat(Vector(prefix, key)) + } +} diff --git a/sage-core/src/main/scala/sage/commands/Server.scala b/sage-core/src/main/scala/sage/commands/Server.scala index 265839fd..3474b707 100644 --- a/sage-core/src/main/scala/sage/commands/Server.scala +++ b/sage-core/src/main/scala/sage/commands/Server.scala @@ -6,6 +6,8 @@ import scala.concurrent.duration.* import sage.Bytes import sage.SageException.DecodeError +import sage.codec.Primitives +import sage.commands.Args.{Count, Get} import sage.protocol.Frame /** @@ -17,7 +19,8 @@ enum FlushMode { } object FlushMode { - private[commands] def args(mode: Option[FlushMode]): Vector[Bytes] = mode.map(m => Bytes.utf8(m.toString.toUpperCase)).toVector + private val word = Args.keywords(FlushMode.values) + private[commands] def args(mode: Option[FlushMode]): Vector[Bytes] = mode.map(word).toVector } /** @@ -106,7 +109,6 @@ enum CommandFilterBy { */ private[sage] object Server { - private val Get = Bytes.utf8("GET") private val Set = Bytes.utf8("SET") private val Usage = Bytes.utf8("USAGE") private val Samples = Bytes.utf8("SAMPLES") @@ -116,7 +118,6 @@ private[sage] object Server { private val Latest = Bytes.utf8("LATEST") private val Reset = Bytes.utf8("RESET") private val Histogram = Bytes.utf8("HISTOGRAM") - private val Count = Bytes.utf8("COUNT") private val ListCmd = Bytes.utf8("LIST") private val GetKeys = Bytes.utf8("GETKEYS") private val GetKeysAndFlags = Bytes.utf8("GETKEYSANDFLAGS") @@ -125,17 +126,16 @@ private[sage] object Server { private val Module = Bytes.utf8("MODULE") private val AclCat = Bytes.utf8("ACLCAT") private val PatternWord = Bytes.utf8("PATTERN") - private val ClInfo = Bytes.utf8("INFO") private val ClNodes = Bytes.utf8("NODES") private val ClMyId = Bytes.utf8("MYID") private val ClKeySlot = Bytes.utf8("KEYSLOT") private val ClCountKeys = Bytes.utf8("COUNTKEYSINSLOT") def configGet(parameter: String, rest: String*): Command[Map[String, String]] = - Command("CONFIG", Command.NoKeys, Get +: (parameter +: rest).iterator.map(Bytes.utf8).toVector, decodeStringMap) + Command("CONFIG", Command.NoKeys, Get +: (parameter +: rest).iterator.map(Bytes.utf8).toVector, Decode.fieldValues(Decode.text)) def configSet(setting: (String, String), rest: (String, String)*): Command[Unit] = - Command("CONFIG", Command.NoKeys, Set +: (setting +: rest).flatMap { case (k, v) => Vector(Bytes.utf8(k), Bytes.utf8(v)) }.toVector, Decode.ok) + Command("CONFIG", Command.NoKeys, Set +: Args.pairs(setting +: rest.toVector), Decode.ok) def info(sections: String*): Command[String] = Command("INFO", Command.NoKeys, sections.iterator.map(Bytes.utf8).toVector, Decode.text) @@ -156,39 +156,32 @@ private[sage] object Server { "TIME", Command.NoKeys, Vector.empty, - { - case Frame.Array(Vector(Frame.BulkString(sec), Frame.BulkString(micros))) => - for { - s <- sec.asUtf8String.toLongOption.toRight(DecodeError("epoch seconds", sec.asUtf8String)) - u <- micros.asUtf8String.toLongOption.toRight(DecodeError("microseconds", micros.asUtf8String)) - } yield Instant.ofEpochSecond(s, u * 1000L) - case other => Left(DecodeError("TIME [seconds, microseconds]", Frame.describe(other))) - } + Decode.array2(Decode.decimal("epoch seconds"), Decode.decimal("microseconds"), "TIME [seconds, microseconds]")((s, u) => + Instant.ofEpochSecond(s, u * 1000L) + ) ) - private val decodeRole: Frame => Either[DecodeError, Role] = { - case Frame.Array(Frame.BulkString(kind) +: rest) => - kind.asUtf8String match { - case "master" => - rest match { - case Frame.Integer(offset) +: replicasFrame +: _ => decodeReplicas(replicasFrame).map(Role.Master(offset, _)) - case other => Left(DecodeError("master role [offset, replicas]", other.map(Frame.describe).mkString(", "))) - } - case "slave" => - rest match { - case Vector(Frame.BulkString(host), Frame.Integer(port), Frame.BulkString(state), Frame.Integer(offset)) => - decodePort(port).map(Role.Replica(host.asUtf8String, _, state.asUtf8String, offset)) - case other => - Left(DecodeError("replica role [host, port, state, offset]", other.map(Frame.describe).mkString(", "))) - } - case "sentinel" => - rest.headOption match { - case Some(masters) => Decode.vector(Decode.utf8String)(masters).map(Role.Sentinel(_)) - case None => Left(DecodeError("sentinel role [masterNames]", "empty")) - } - case other => Left(DecodeError("role master|slave|sentinel", other)) - } - case other => Left(DecodeError("ROLE array", Frame.describe(other))) + private val decodeRole: Frame => Either[DecodeError, Role] = Decode.shape("ROLE array") { case Frame.Array(Frame.BulkString(kind) +: rest) => + kind.asUtf8String match { + case "master" => + rest match { + case Frame.Integer(offset) +: replicasFrame +: _ => decodeReplicas(replicasFrame).map(Role.Master(offset, _)) + case other => Left(DecodeError("master role [offset, replicas]", other.map(Frame.describe).mkString(", "))) + } + case "slave" => + rest match { + case Vector(Frame.BulkString(host), Frame.Integer(port), Frame.BulkString(state), Frame.Integer(offset)) => + decodePort(port).map(Role.Replica(host.asUtf8String, _, state.asUtf8String, offset)) + case other => + Left(DecodeError("replica role [host, port, state, offset]", other.map(Frame.describe).mkString(", "))) + } + case "sentinel" => + rest.headOption match { + case Some(masters) => Decode.vector(Decode.utf8String)(masters).map(Role.Sentinel(_)) + case None => Left(DecodeError("sentinel role [masterNames]", "empty")) + } + case other => Left(DecodeError("role master|slave|sentinel", other)) + } } val role: Command[Role] = Command("ROLE", Command.NoKeys, Vector.empty, decodeRole) @@ -199,38 +192,35 @@ private[sage] object Server { def flushDb(mode: Option[FlushMode] = None): Command[Unit] = Command("FLUSHDB", Command.NoKeys, FlushMode.args(mode), Decode.ok, allMasters = true) - private val waitAofMin: (Frame, Frame) => Frame = (a, b) => - (a, b) match { - case (Frame.Array(Vector(Frame.Integer(l1), Frame.Integer(r1))), Frame.Array(Vector(Frame.Integer(l2), Frame.Integer(r2)))) => - Frame.Array(Vector(Frame.Integer(math.min(l1, l2)), Frame.Integer(math.min(r1, r2)))) - case (Frame.Array(Vector(Frame.Integer(_), Frame.Integer(_))), bad) => bad - case (bad, _) => bad - } - // WAIT and WAITAOF interpret an encoded timeout of 0 as an unlimited wait. Encode Duration.Zero as 0. Round every other duration, including - // a negative one, up to at least 1 ms so a sub-millisecond timeout remains finite. This matches BlockTimeout.millisWire. - private def waitTimeoutMillis(timeout: FiniteDuration): Long = - if (timeout == Duration.Zero) 0L else Math.max(1L, Math.ceilDiv(timeout.toNanos, 1000000L)) + // a negative one, up to at least 1 ms so a sub-millisecond timeout remains finite. + private def waitTimeout(timeout: FiniteDuration): Bytes = + BlockTimeout.millisWire(if (timeout == Duration.Zero) BlockTimeout.Forever else BlockTimeout.After(timeout)) def waitReplicas(numReplicas: Long, timeout: FiniteDuration): Command[Long] = Command( "WAIT", Command.NoKeys, - Vector(Bytes.utf8(numReplicas.toString), Bytes.utf8(waitTimeoutMillis(timeout).toString)), + Vector(Args.long(numReplicas), waitTimeout(timeout)), Decode.long, allMasters = true, broadcast = BroadcastReduce.Fold(Merge.min) ) + private val waitAofReply = Decode.shape("WAITAOF [numlocal, numreplicas]") { + case Frame.Array(Vector(Frame.Integer(local), Frame.Integer(replicas))) => Right((local, replicas)) + } + + private val waitAofMin = Merge.typed[(Long, Long)](waitAofReply, (l, r) => Frame.Array(Vector(Frame.Integer(l), Frame.Integer(r)))) { + case ((l1, r1), (l2, r2)) => (math.min(l1, l2), math.min(r1, r2)) + } + def waitAof(numLocal: Long, numReplicas: Long, timeout: FiniteDuration): Command[(Long, Long)] = Command( "WAITAOF", Command.NoKeys, - Vector(numLocal, numReplicas, waitTimeoutMillis(timeout)).map(n => Bytes.utf8(n.toString)), - { - case Frame.Array(Vector(Frame.Integer(local), Frame.Integer(replicas))) => Right((local, replicas)) - case other => Left(DecodeError("WAITAOF [numlocal, numreplicas]", Frame.describe(other))) - }, + Vector(Args.long(numLocal), Args.long(numReplicas), waitTimeout(timeout)), + waitAofReply, allMasters = true, broadcast = BroadcastReduce.Fold(waitAofMin) ) @@ -239,21 +229,21 @@ private[sage] object Server { Command( "MEMORY", Vector(1), - Vector(Usage, keyCodec.encode(key)) ++ samples.toVector.flatMap(n => Vector(Samples, Bytes.utf8(n.toString))), + Vector(Usage, keyCodec.encode(key)) ++ Args.optLong(Samples, samples), Decode.optionalLong ) val memoryPurge: Command[Unit] = Command("MEMORY", Command.NoKeys, Vector(Purge), Decode.ok, allMasters = true) def slowLogGet(count: Option[Long] = None): Command[Vector[SlowLogEntry]] = - Command("SLOWLOG", Command.NoKeys, Get +: count.map(n => Bytes.utf8(n.toString)).toVector, Decode.vector(decodeSlowLog)) + Command("SLOWLOG", Command.NoKeys, Get +: count.map(n => Args.long(n)).toVector, Decode.vector(decodeSlowLog)) val slowLogLen: Command[Long] = Command("SLOWLOG", Command.NoKeys, Vector(SlowLen), Decode.long) val slowLogReset: Command[Unit] = Command("SLOWLOG", Command.NoKeys, Vector(Reset), Decode.ok) // `count` of -1 returns every entry of the type def commandLogGet(count: Long, logType: CommandLogType): Command[Vector[CommandLogEntry]] = - Command("COMMANDLOG", Command.NoKeys, Vector(Get, Bytes.utf8(count.toString), CommandLogType.wire(logType)), Decode.vector(decodeCommandLog)) + Command("COMMANDLOG", Command.NoKeys, Vector(Get, Args.long(count), CommandLogType.wire(logType)), Decode.vector(decodeCommandLog)) def commandLogLen(logType: CommandLogType): Command[Long] = Command("COMMANDLOG", Command.NoKeys, Vector(SlowLen, CommandLogType.wire(logType)), Decode.long) @@ -264,6 +254,11 @@ private[sage] object Server { def latencyHistory(event: String): Command[Vector[(Instant, FiniteDuration)]] = Command("LATENCY", Command.NoKeys, Vector(History, Bytes.utf8(event)), Decode.vector(decodeLatencyHistory)) + private val decodeLatencyLatest: Frame => Either[DecodeError, LatencyEntry] = Decode.shape("latency latest [event, ts, latest, max]") { + case Frame.Array(Vector(Frame.BulkString(event), Frame.Integer(ts), Frame.Integer(latest), Frame.Integer(max))) => + Right(LatencyEntry(event.asUtf8String, Instant.ofEpochSecond(ts), latest.millis, max.millis)) + } + val latencyLatest: Command[Vector[LatencyEntry]] = Command("LATENCY", Command.NoKeys, Vector(Latest), Decode.vector(decodeLatencyLatest)) def latencyReset(events: String*): Command[Long] = @@ -293,7 +288,7 @@ private[sage] object Server { // --- cluster introspection (read-only; operator/mutation commands are deliberately not exposed) ---------------------------------------- - val clusterInfo: Command[String] = Command("CLUSTER", Command.NoKeys, Vector(ClInfo), Decode.text) + val clusterInfo: Command[String] = Command("CLUSTER", Command.NoKeys, Vector(Info), Decode.text) val clusterNodes: Command[String] = Command("CLUSTER", Command.NoKeys, Vector(ClNodes), Decode.text) val clusterMyId: Command[String] = Command("CLUSTER", Command.NoKeys, Vector(ClMyId), Decode.text) @@ -301,21 +296,10 @@ private[sage] object Server { Command("CLUSTER", Command.NoKeys, Vector(ClKeySlot, Bytes.utf8(key)), Decode.long) def clusterCountKeysInSlot(slot: Int): Command[Long] = - Command("CLUSTER", Command.NoKeys, Vector(ClCountKeys, Bytes.utf8(slot.toString)), Decode.long) + Command("CLUSTER", Command.NoKeys, Vector(ClCountKeys, Args.long(slot)), Decode.long) // --- decoders -------------------------------------------------------------------------------------------------------------------------- - private val decodeStringMap: Frame => Either[DecodeError, Map[String, String]] = - frame => - Decode.fieldMap(frame).flatMap { fields => - fields.foldLeft[Either[DecodeError, Map[String, String]]](Right(Map.empty)) { case (acc, (name, valueFrame)) => - for { - map <- acc - value <- Decode.text(valueFrame) - } yield map + (name -> value) - } - } - private def filterByArgs(filterBy: Option[CommandFilterBy]): Vector[Bytes] = filterBy match { case None => Vector.empty @@ -327,81 +311,47 @@ private[sage] object Server { private def decodePort(value: Long): Either[DecodeError, Int] = if (value >= 1L && value <= 65535L) Right(value.toInt) else Left(DecodeError("port in 1..65535", value.toString)) - private def decodeReplicas(frame: Frame): Either[DecodeError, Vector[ReplicaNode]] = - Decode.vector { - case Frame.Array(Vector(Frame.BulkString(host), Frame.BulkString(port), Frame.BulkString(offset))) => - for { - p <- port.asUtf8String.toLongOption.toRight(DecodeError("replica port", port.asUtf8String)).flatMap(decodePort) - o <- offset.asUtf8String.toLongOption.toRight(DecodeError("replica offset", offset.asUtf8String)) - } yield ReplicaNode(host.asUtf8String, p, o) - case other => Left(DecodeError("replica [host, port, offset]", Frame.describe(other))) - }(frame) - - // SLOWLOG GET and COMMANDLOG GET share this trailing [clientAddr, clientName] shape, absent on servers older than 4.0 - private def clientFields(tail: Vector[Frame]): (String, String) = { - val addr = tail.headOption.collect { case Frame.BulkString(b) => b.asUtf8String }.getOrElse("") - val name = tail.drop(1).headOption.collect { case Frame.BulkString(b) => b.asUtf8String }.getOrElse("") - (addr, name) - } - - private def decodeSlowLog(frame: Frame): Either[DecodeError, SlowLogEntry] = - frame match { - case Frame.Array(Frame.Integer(id) +: Frame.Integer(ts) +: Frame.Integer(micros) +: argsFrame +: tail) => - Decode.vector(Decode.utf8String)(argsFrame).map { command => - val (addr, name) = clientFields(tail) - SlowLogEntry(id, Instant.ofEpochSecond(ts), micros.micros, command, addr, name) - } - case other => Left(DecodeError("slowlog entry", Frame.describe(other))) + private val decodeReplicas: Frame => Either[DecodeError, Vector[ReplicaNode]] = Decode.vector(Decode.shape("replica [host, port, offset]") { + case Frame.Array(Vector(Frame.BulkString(host), Frame.BulkString(port), Frame.BulkString(offset))) => + for { + p <- Primitives.decodeLong("replica port in 1..65535", 1L, 65535L)(port) + o <- Primitives.decodeLong("replica offset", Long.MinValue, Long.MaxValue)(offset) + } yield ReplicaNode(host.asUtf8String, p.toInt, o) + }) + + // SLOWLOG GET and COMMANDLOG GET entries: [id, timestamp, metric, args, clientAddr, clientName]; the client fields are absent on servers + // older than 4.0 + private def logEntry[A](label: String)(build: (Long, Instant, Long, Vector[String], String, String) => A): Frame => Either[DecodeError, A] = + Decode.shape(label) { case Frame.Array(Frame.Integer(id) +: Frame.Integer(ts) +: Frame.Integer(metric) +: argsFrame +: tail) => + Decode.vector(Decode.utf8String)(argsFrame).map { command => + def client(i: Int) = tail.lift(i).collect { case Frame.BulkString(b) => b.asUtf8String }.getOrElse("") + build(id, Instant.ofEpochSecond(ts), metric, command, client(0), client(1)) + } } - private def decodeCommandLog(frame: Frame): Either[DecodeError, CommandLogEntry] = - frame match { - case Frame.Array(Frame.Integer(id) +: Frame.Integer(ts) +: Frame.Integer(metric) +: argsFrame +: tail) => - Decode.vector(Decode.utf8String)(argsFrame).map { command => - val (addr, name) = clientFields(tail) - CommandLogEntry(id, Instant.ofEpochSecond(ts), metric, command, addr, name) - } - case other => Left(DecodeError("commandlog entry", Frame.describe(other))) - } + private val decodeSlowLog: Frame => Either[DecodeError, SlowLogEntry] = + logEntry("slowlog entry")((id, ts, micros, command, addr, name) => SlowLogEntry(id, ts, micros.micros, command, addr, name)) - private def decodeLatencyLatest(frame: Frame): Either[DecodeError, LatencyEntry] = - frame match { - case Frame.Array(Vector(Frame.BulkString(event), Frame.Integer(ts), Frame.Integer(latest), Frame.Integer(max))) => - Right(LatencyEntry(event.asUtf8String, Instant.ofEpochSecond(ts), latest.millis, max.millis)) - case other => - Left(DecodeError("latency latest [event, ts, latest, max]", Frame.describe(other))) - } + private val decodeCommandLog: Frame => Either[DecodeError, CommandLogEntry] = logEntry("commandlog entry")(CommandLogEntry(_, _, _, _, _, _)) - private def decodeLatencyHistory(frame: Frame): Either[DecodeError, (Instant, FiniteDuration)] = - frame match { - case Frame.Array(Vector(Frame.Integer(ts), Frame.Integer(latency))) => Right((Instant.ofEpochSecond(ts), latency.millis)) - case other => Left(DecodeError("latency history [ts, latency]", Frame.describe(other))) - } + private val decodeLatencyHistory: Frame => Either[DecodeError, (Instant, FiniteDuration)] = Decode.shape("latency history [ts, latency]") { + case Frame.Array(Vector(Frame.Integer(ts), Frame.Integer(latency))) => Right((Instant.ofEpochSecond(ts), latency.millis)) + } private val decodeHistograms: Frame => Either[DecodeError, Map[String, CommandHistogram]] = { - case Frame.Map(entries) => - entries.foldLeft[Either[DecodeError, Map[String, CommandHistogram]]](Right(Map.empty)) { case (acc, (nameFrame, statsFrame)) => - for { - map <- acc - name <- Decode.utf8String(nameFrame) - hist <- decodeHistogram(statsFrame) - } yield map + (name -> hist) - } - case other => Left(DecodeError("latency histogram map", Frame.describe(other))) + val entry = Decode.pair(Decode.utf8String, decodeHistogram).tupled + Decode.shape("latency histogram map") { case Frame.Map(entries) => Decode.mapEntries(entries)(entry) } } private def decodeHistogram(frame: Frame): Either[DecodeError, CommandHistogram] = - frame match { - case Frame.Map(entries) => - val fields = entries.collect { case (Frame.BulkString(k), v) => k.asUtf8String -> v }.toMap - val calls = fields.get("calls").collect { case Frame.Integer(n) => n }.getOrElse(0L) - val buckets = fields.get("histogram_usec") match { - case Some(Frame.Map(bs)) => - bs.collect { case (Frame.Integer(bucket), Frame.Integer(count)) => bucket -> count }.toMap - case _ => Map.empty[Long, Long] - } - Right(CommandHistogram(calls, buckets)) - case other => Left(DecodeError("command histogram map", Frame.describe(other))) + Decode.fieldMap(frame).map { fields => + val calls = fields.get("calls").collect { case Frame.Integer(n) => n }.getOrElse(0L) + val buckets = fields.get("histogram_usec") match { + case Some(Frame.Map(bs)) => + bs.collect { case (Frame.Integer(bucket), Frame.Integer(count)) => bucket -> count }.toMap + case _ => Map.empty[Long, Long] + } + CommandHistogram(calls, buckets) } // a server may frame a flag list as a RESP3 Set or an Array @@ -411,40 +361,25 @@ private[sage] object Server { case other => Decode.vector(Decode.text)(other) } - private def decodeKeyAndFlags(frame: Frame): Either[DecodeError, (String, Set[String])] = - frame match { - case Frame.Array(Vector(Frame.BulkString(key), flagsFrame)) => - stringSeq(flagsFrame).map(flags => key.asUtf8String -> flags.toSet) - case other => Left(DecodeError("[key, [flags]]", Frame.describe(other))) - } + private val decodeKeyAndFlags: Frame => Either[DecodeError, (String, Set[String])] = Decode.shape("[key, [flags]]") { + case Frame.Array(Vector(Frame.BulkString(key), flagsFrame)) => stringSeq(flagsFrame).map(flags => key.asUtf8String -> flags.toSet) + } // COMMAND INFO yields one element per requested name; an unknown name is a null element, dropped here - private val decodeCommandInfos: Frame => Either[DecodeError, Vector[CommandInfo]] = { - case Frame.Array(elements) => - elements.foldLeft[Either[DecodeError, Vector[CommandInfo]]](Right(Vector.empty)) { (acc, element) => - acc.flatMap { infos => - element match { - case Frame.Null => Right(infos) - case other => decodeCommandInfo(other).map(infos :+ _) - } - } - } - case other => Left(DecodeError("COMMAND INFO array", Frame.describe(other))) + private val decodeCommandInfos: Frame => Either[DecodeError, Vector[CommandInfo]] = Decode.shape("COMMAND INFO array") { + case Frame.Array(elements) => Decode.each(elements)(Decode.nullable(decodeCommandInfo)).map(_.flatten) } - private def decodeCommandInfo(frame: Frame): Either[DecodeError, CommandInfo] = - frame match { - case Frame.Array( - Frame.BulkString(name) +: Frame.Integer(arity) +: flagsFrame +: Frame.Integer(firstKey) +: Frame.Integer(lastKey) +: Frame.Integer( - step - ) +: tail - ) => - for { - flags <- stringSeq(flagsFrame) - acl <- tail.headOption.fold[Either[DecodeError, Vector[String]]](Right(Vector.empty))(stringSeq) - } yield CommandInfo(name.asUtf8String, arity, flags.toSet, firstKey.toInt, lastKey.toInt, step.toInt, acl.toSet) - case other => - Left(DecodeError("command info entry", Frame.describe(other))) - } + private val decodeCommandInfo: Frame => Either[DecodeError, CommandInfo] = Decode.shape("command info entry") { + case Frame.Array( + Frame.BulkString(name) +: Frame.Integer(arity) +: flagsFrame +: Frame.Integer(firstKey) +: Frame.Integer(lastKey) +: Frame.Integer( + step + ) +: tail + ) => + for { + flags <- stringSeq(flagsFrame) + acl <- tail.headOption.fold[Either[DecodeError, Vector[String]]](Right(Vector.empty))(stringSeq) + } yield CommandInfo(name.asUtf8String, arity, flags.toSet, firstKey.toInt, lastKey.toInt, step.toInt, acl.toSet) + } } diff --git a/sage-core/src/main/scala/sage/commands/Sets.scala b/sage-core/src/main/scala/sage/commands/Sets.scala index 66254e91..e23853dc 100644 --- a/sage-core/src/main/scala/sage/commands/Sets.scala +++ b/sage-core/src/main/scala/sage/commands/Sets.scala @@ -1,9 +1,6 @@ package sage.commands -import sage.Bytes -import sage.SageException.DecodeError import sage.codec.{KeyCodec, ValueCodec} -import sage.protocol.Frame /** * Set members are values (a [[ValueCodec]]), like list elements. The set-returning reads decode the RESP3 Set frame into a `Set[V]`; @@ -11,13 +8,11 @@ import sage.protocol.Frame */ private[sage] object Sets { - private val Limit = Bytes.utf8("LIMIT") - def sAdd[K, V](key: K, first: V, rest: V*)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Long] = - Command("SADD", Command.FirstKey, keyCodec.encode(key) +: (first +: rest.toVector).map(valueCodec.encode), Decode.long) + Command("SADD", Command.FirstKey, Args.keyThen(key, first, rest)(valueCodec.encode), Decode.long) def sRem[K, V](key: K, first: V, rest: V*)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Long] = - Command("SREM", Command.FirstKey, keyCodec.encode(key) +: (first +: rest.toVector).map(valueCodec.encode), Decode.long) + Command("SREM", Command.FirstKey, Args.keyThen(key, first, rest)(valueCodec.encode), Decode.long) def sCard[K](key: K)(using keyCodec: KeyCodec[K]): Command[Long] = Command.read("SCARD", Command.FirstKey, Vector(keyCodec.encode(key)), Decode.long) @@ -26,12 +21,7 @@ private[sage] object Sets { Command.read("SISMEMBER", Command.FirstKey, Vector(keyCodec.encode(key), valueCodec.encode(member)), Decode.flag) def sMisMember[K, V](key: K, first: V, rest: V*)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Vector[Boolean]] = - Command.read( - "SMISMEMBER", - Command.FirstKey, - keyCodec.encode(key) +: (first +: rest.toVector).map(valueCodec.encode), - Decode.vector(Decode.flag) - ) + Command.read("SMISMEMBER", Command.FirstKey, Args.keyThen(key, first, rest)(valueCodec.encode), Decode.vector(Decode.flag)) def sMembers[K, V](key: K)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Set[V]] = Command.read("SMEMBERS", Command.FirstKey, Vector(keyCodec.encode(key)), Decode.set[V]) @@ -43,63 +33,39 @@ private[sage] object Sets { Command("SPOP", Command.FirstKey, Vector(keyCodec.encode(key)), Decode.optionalValue) def sPopCount[K, V](key: K, count: Long)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Set[V]] = - Command("SPOP", Command.FirstKey, Vector(keyCodec.encode(key), Bytes.utf8(count.toString)), Decode.set[V]) + Command("SPOP", Command.FirstKey, Vector(keyCodec.encode(key), Args.long(count)), Decode.set[V]) def sRandMember[K, V](key: K)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Option[V]] = Command.readUncacheable("SRANDMEMBER", Command.FirstKey, Vector(keyCodec.encode(key)), Decode.optionalValue) // a negative count may repeat members, so the reply is an ordered Array, not a Set def sRandMemberCount[K, V](key: K, count: Long)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Vector[V]] = - Command.readUncacheable("SRANDMEMBER", Command.FirstKey, Vector(keyCodec.encode(key), Bytes.utf8(count.toString)), Decode.vector(Decode.value[V])) + Command.readUncacheable("SRANDMEMBER", Command.FirstKey, Vector(keyCodec.encode(key), Args.long(count)), Decode.vector(Decode.value[V])) def sDiff[K, V](first: K, rest: K*)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Set[V]] = - setOp("SDIFF", first +: rest.toVector, Decode.set[V]) + KeyArgs.allKeys("SDIFF", first +: rest.toVector, Decode.set[V], readOnly = true) def sDiffStore[K](destination: K, first: K, rest: K*)(using keyCodec: KeyCodec[K]): Command[Long] = - storeOp("SDIFFSTORE", destination, first +: rest.toVector) + KeyArgs.allKeys("SDIFFSTORE", destination +: first +: rest.toVector, Decode.long) def sInter[K, V](first: K, rest: K*)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Set[V]] = - setOp("SINTER", first +: rest.toVector, Decode.set[V]) + KeyArgs.allKeys("SINTER", first +: rest.toVector, Decode.set[V], readOnly = true) def sInterStore[K](destination: K, first: K, rest: K*)(using keyCodec: KeyCodec[K]): Command[Long] = - storeOp("SINTERSTORE", destination, first +: rest.toVector) - - def sInterCard[K](first: K, rest: K*)(limit: Option[Long] = None)(using keyCodec: KeyCodec[K]): Command[Long] = { - val (keyIndices, prefix) = KeyArgs.numKeyed(first +: rest.toVector) - Command.read( - "SINTERCARD", - keyIndices, - args = prefix ++ limit.toVector.flatMap(n => Vector(Limit, Bytes.utf8(n.toString))), - Decode.long - ) - } + KeyArgs.allKeys("SINTERSTORE", destination +: first +: rest.toVector, Decode.long) + + def sInterCard[K](first: K, rest: K*)(limit: Option[Long] = None)(using keyCodec: KeyCodec[K]): Command[Long] = + KeyArgs.interCard("SINTERCARD", first +: rest.toVector, limit) def sUnion[K, V](first: K, rest: K*)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Set[V]] = - setOp("SUNION", first +: rest.toVector, Decode.set[V]) + KeyArgs.allKeys("SUNION", first +: rest.toVector, Decode.set[V], readOnly = true) def sUnionStore[K](destination: K, first: K, rest: K*)(using keyCodec: KeyCodec[K]): Command[Long] = - storeOp("SUNIONSTORE", destination, first +: rest.toVector) + KeyArgs.allKeys("SUNIONSTORE", destination +: first +: rest.toVector, Decode.long) def sScan[K, V](key: K, cursor: ScanCursor, pattern: Option[String] = None, count: Option[Long] = None)( using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V] ): Command[ScanPage[V]] = - Command.readCursor( - "SSCAN", - Command.FirstKey, - Vector(keyCodec.encode(key), ScanCursor.bytes(cursor)) ++ ScanArgs.options(pattern, count), - Decode.scanPage(Decode.vector(Decode.value[V])) - ) - - private def setOp[K, Out](name: String, keys: Vector[K], decode: Frame => Either[DecodeError, Out])( - using keyCodec: KeyCodec[K] - ): Command[Out] = { - val args = keys.map(keyCodec.encode) - Command.read(name, args.indices.toVector, args, decode) - } - - private def storeOp[K](name: String, destination: K, keys: Vector[K])(using keyCodec: KeyCodec[K]): Command[Long] = { - val args = (destination +: keys).map(keyCodec.encode) - Command(name, args.indices.toVector, args, Decode.long) - } + KeyArgs.keyScan("SSCAN", key, cursor, pattern, count)(Decode.vector(Decode.value[V])) } diff --git a/sage-core/src/main/scala/sage/commands/SortedSets.scala b/sage-core/src/main/scala/sage/commands/SortedSets.scala index 66f3eb3a..36be9368 100644 --- a/sage-core/src/main/scala/sage/commands/SortedSets.scala +++ b/sage-core/src/main/scala/sage/commands/SortedSets.scala @@ -3,6 +3,7 @@ package sage.commands import sage.Bytes import sage.SageException.DecodeError import sage.codec.{Doubles, KeyCodec, ValueCodec} +import sage.commands.Args.{Ch, Gt, Lt, Nx, Rev, Xx} import sage.protocol.Frame /** @@ -119,23 +120,15 @@ object ZRange { private[sage] object SortedSets { - private val Nx = Bytes.utf8("NX") - private val Xx = Bytes.utf8("XX") - private val Gt = Bytes.utf8("GT") - private val Lt = Bytes.utf8("LT") - private val Ch = Bytes.utf8("CH") private val Incr = Bytes.utf8("INCR") private val ByScore = Bytes.utf8("BYSCORE") private val ByLex = Bytes.utf8("BYLEX") - private val Rev = Bytes.utf8("REV") - private val LimitWord = Bytes.utf8("LIMIT") private val WithScores = Bytes.utf8("WITHSCORES") private val WithScore = Bytes.utf8("WITHSCORE") private val WeightsWord = Bytes.utf8("WEIGHTS") private val AggregateWord = Bytes.utf8("AGGREGATE") - private val MinWord = Bytes.utf8("MIN") - private val MaxWord = Bytes.utf8("MAX") - private val CountWord = Bytes.utf8("COUNT") + private val aggregateWord = Args.keywords(Aggregate.values) + private val minMaxArg = Args.keywords(MinMax.values) private val NegInfWord = Bytes.utf8("-inf") private val PosInfWord = Bytes.utf8("+inf") private val LexMin = Bytes.utf8("-") @@ -148,7 +141,7 @@ private[sage] object SortedSets { Command( "ZADD", Command.FirstKey, - (keyCodec.encode(key) +: conditionArgs(condition)) ++ (if (changed) Vector(Ch) else Vector.empty) ++ memberScoreArgs(first +: rest.toVector), + (keyCodec.encode(key) +: conditionArgs(condition)) ++ Args.flag(changed, Ch) ++ memberScoreArgs(first +: rest.toVector), Decode.long ) @@ -159,7 +152,7 @@ private[sage] object SortedSets { Command( "ZADD", Command.FirstKey, - (keyCodec.encode(key) +: conditionArgs(condition)) ++ Vector(Incr, scoreArg(score), valueCodec.encode(member)), + (keyCodec.encode(key) +: conditionArgs(condition)) ++ Vector(Incr, Args.double(score), valueCodec.encode(member)), Decode.optionalScore ) @@ -170,15 +163,10 @@ private[sage] object SortedSets { Command.read("ZSCORE", Command.FirstKey, Vector(keyCodec.encode(key), valueCodec.encode(member)), Decode.optionalScore) def zMScore[K, V](key: K, first: V, rest: V*)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Vector[Option[Double]]] = - Command.read( - "ZMSCORE", - Command.FirstKey, - keyCodec.encode(key) +: (first +: rest.toVector).map(valueCodec.encode), - Decode.vector(Decode.optionalScore) - ) + Command.read("ZMSCORE", Command.FirstKey, Args.keyThen(key, first, rest)(valueCodec.encode), Decode.vector(Decode.optionalScore)) def zIncrBy[K, V](key: K, member: V, increment: Double)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Double] = - Command("ZINCRBY", Command.FirstKey, Vector(keyCodec.encode(key), scoreArg(increment), valueCodec.encode(member)), Decode.score) + Command("ZINCRBY", Command.FirstKey, Vector(keyCodec.encode(key), Args.double(increment), valueCodec.encode(member)), Decode.double) def zRank[K, V](key: K, member: V)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Option[Long]] = Command.read("ZRANK", Command.FirstKey, Vector(keyCodec.encode(key), valueCodec.encode(member)), Decode.optionalLong) @@ -213,10 +201,10 @@ private[sage] object SortedSets { ) def zRem[K, V](key: K, first: V, rest: V*)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Long] = - Command("ZREM", Command.FirstKey, keyCodec.encode(key) +: (first +: rest.toVector).map(valueCodec.encode), Decode.long) + Command("ZREM", Command.FirstKey, Args.keyThen(key, first, rest)(valueCodec.encode), Decode.long) def zRemRangeByRank[K](key: K, start: Long, stop: Long)(using keyCodec: KeyCodec[K]): Command[Long] = - Command("ZREMRANGEBYRANK", Command.FirstKey, Vector(keyCodec.encode(key), Bytes.utf8(start.toString), Bytes.utf8(stop.toString)), Decode.long) + Command("ZREMRANGEBYRANK", Command.FirstKey, Vector(keyCodec.encode(key), Args.long(start), Args.long(stop)), Decode.long) def zRemRangeByScore[K](key: K, min: ScoreBoundary, max: ScoreBoundary)(using keyCodec: KeyCodec[K]): Command[Long] = Command("ZREMRANGEBYSCORE", Command.FirstKey, Vector(keyCodec.encode(key), scoreBoundaryArg(min), scoreBoundaryArg(max)), Decode.long) @@ -234,107 +222,87 @@ private[sage] object SortedSets { Command("ZPOPMAX", Command.FirstKey, Vector(keyCodec.encode(key)), Decode.optionalScoredMember[V]) def zPopMinCount[K, V](key: K, count: Long)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Vector[(V, Double)]] = - Command("ZPOPMIN", Command.FirstKey, Vector(keyCodec.encode(key), Bytes.utf8(count.toString)), Decode.scoredMembers[V]) + Command("ZPOPMIN", Command.FirstKey, Vector(keyCodec.encode(key), Args.long(count)), Decode.scoredMembers[V]) def zPopMaxCount[K, V](key: K, count: Long)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Vector[(V, Double)]] = - Command("ZPOPMAX", Command.FirstKey, Vector(keyCodec.encode(key), Bytes.utf8(count.toString)), Decode.scoredMembers[V]) + Command("ZPOPMAX", Command.FirstKey, Vector(keyCodec.encode(key), Args.long(count)), Decode.scoredMembers[V]) def zMpop[K, V](first: K, rest: K*)(minMax: MinMax, count: Option[Long] = None)( using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V] - ): Command[Option[(K, Vector[(V, Double)])]] = { - val (indices, prefix) = KeyArgs.numKeyed(first +: rest.toVector) - Command("ZMPOP", indices, (prefix :+ minMaxArg(minMax)) ++ countArg(count), mpopReply[K, V]) - } + ): Command[Option[(K, Vector[(V, Double)])]] = + KeyArgs.multiPop("ZMPOP", None, first +: rest.toVector, minMaxArg(minMax), count, Decode.scoredMembers[V], MpopLabel) def bzPopMin[K, V](first: K, rest: K*)( timeout: BlockTimeout )(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Option[(K, V, Double)]] = - blockingPop("BZPOPMIN", first, rest.toVector, timeout) + KeyArgs.blockingPop("BZPOPMIN", first +: rest.toVector, timeout, poppedMember[K, V]) def bzPopMax[K, V](first: K, rest: K*)( timeout: BlockTimeout )(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Option[(K, V, Double)]] = - blockingPop("BZPOPMAX", first, rest.toVector, timeout) + KeyArgs.blockingPop("BZPOPMAX", first +: rest.toVector, timeout, poppedMember[K, V]) def bzMpop[K, V](first: K, rest: K*)(minMax: MinMax, timeout: BlockTimeout, count: Option[Long] = None)( using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V] - ): Command[Option[(K, Vector[(V, Double)])]] = { - val keys = (first +: rest.toVector).map(keyCodec.encode) - Command( - "BZMPOP", - keyIndices = Vector.tabulate(keys.size)(_ + 2), - args = (BlockTimeout.wire(timeout) +: Bytes.utf8(keys.size.toString) +: keys :+ minMaxArg(minMax)) ++ countArg(count), - decode = mpopReply[K, V], - execution = Execution.Blocking - ) - } + ): Command[Option[(K, Vector[(V, Double)])]] = + KeyArgs.multiPop("BZMPOP", Some(timeout), first +: rest.toVector, minMaxArg(minMax), count, Decode.scoredMembers[V], MpopLabel) def zRandMember[K, V](key: K)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Option[V]] = Command.readUncacheable("ZRANDMEMBER", Command.FirstKey, Vector(keyCodec.encode(key)), Decode.optionalValue) def zRandMemberCount[K, V](key: K, count: Long)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Vector[V]] = - Command.readUncacheable("ZRANDMEMBER", Command.FirstKey, Vector(keyCodec.encode(key), Bytes.utf8(count.toString)), Decode.vector(Decode.value[V])) + Command.readUncacheable("ZRANDMEMBER", Command.FirstKey, Vector(keyCodec.encode(key), Args.long(count)), Decode.vector(Decode.value[V])) def zRandMemberWithScores[K, V](key: K, count: Long)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Vector[(V, Double)]] = Command.readUncacheable( "ZRANDMEMBER", Command.FirstKey, - Vector(keyCodec.encode(key), Bytes.utf8(count.toString), WithScores), + Vector(keyCodec.encode(key), Args.long(count), WithScores), Decode.scoredMembers[V] ) def zUnion[K, V](first: K, rest: K*)(weights: Option[Vector[Double]] = None, aggregate: Aggregate = Aggregate.Sum)( - using keyCodec: KeyCodec[K], - valueCodec: ValueCodec[V] - ): Command[Vector[V]] = { - val (indices, prefix) = KeyArgs.numKeyed(first +: rest.toVector) - Command.read("ZUNION", indices, prefix ++ weightsArgs(weights) ++ aggregateArgs(aggregate), Decode.vector(Decode.value[V])) - } + using KeyCodec[K], + ValueCodec[V] + ): Command[Vector[V]] = + combine("ZUNION", first +: rest.toVector, weights, aggregate, withScores = false, Decode.vector(Decode.value[V])) def zUnionWithScores[K, V](first: K, rest: K*)(weights: Option[Vector[Double]] = None, aggregate: Aggregate = Aggregate.Sum)( - using keyCodec: KeyCodec[K], - valueCodec: ValueCodec[V] - ): Command[Vector[(V, Double)]] = { - val (indices, prefix) = KeyArgs.numKeyed(first +: rest.toVector) - Command.read("ZUNION", indices, (prefix ++ weightsArgs(weights) ++ aggregateArgs(aggregate)) :+ WithScores, Decode.scoredMembers[V]) - } + using KeyCodec[K], + ValueCodec[V] + ): Command[Vector[(V, Double)]] = + combine("ZUNION", first +: rest.toVector, weights, aggregate, withScores = true, Decode.scoredMembers[V]) def zUnionStore[K](destination: K, first: K, rest: K*)(weights: Option[Vector[Double]] = None, aggregate: Aggregate = Aggregate.Sum)( using keyCodec: KeyCodec[K] ): Command[Long] = { - val (indices, prefix) = storeKeyed(destination, first +: rest.toVector) - Command("ZUNIONSTORE", indices, prefix ++ weightsArgs(weights) ++ aggregateArgs(aggregate), Decode.long) + val (indices, prefix) = KeyArgs.numKeyedAfter(keyCodec.encode(destination), first +: rest.toVector) + Command("ZUNIONSTORE", 0 +: indices, prefix ++ weightsArgs(weights) ++ aggregateArgs(aggregate), Decode.long) } def zInter[K, V](first: K, rest: K*)(weights: Option[Vector[Double]] = None, aggregate: Aggregate = Aggregate.Sum)( - using keyCodec: KeyCodec[K], - valueCodec: ValueCodec[V] - ): Command[Vector[V]] = { - val (indices, prefix) = KeyArgs.numKeyed(first +: rest.toVector) - Command.read("ZINTER", indices, prefix ++ weightsArgs(weights) ++ aggregateArgs(aggregate), Decode.vector(Decode.value[V])) - } + using KeyCodec[K], + ValueCodec[V] + ): Command[Vector[V]] = + combine("ZINTER", first +: rest.toVector, weights, aggregate, withScores = false, Decode.vector(Decode.value[V])) def zInterWithScores[K, V](first: K, rest: K*)(weights: Option[Vector[Double]] = None, aggregate: Aggregate = Aggregate.Sum)( - using keyCodec: KeyCodec[K], - valueCodec: ValueCodec[V] - ): Command[Vector[(V, Double)]] = { - val (indices, prefix) = KeyArgs.numKeyed(first +: rest.toVector) - Command.read("ZINTER", indices, (prefix ++ weightsArgs(weights) ++ aggregateArgs(aggregate)) :+ WithScores, Decode.scoredMembers[V]) - } + using KeyCodec[K], + ValueCodec[V] + ): Command[Vector[(V, Double)]] = + combine("ZINTER", first +: rest.toVector, weights, aggregate, withScores = true, Decode.scoredMembers[V]) def zInterStore[K](destination: K, first: K, rest: K*)(weights: Option[Vector[Double]] = None, aggregate: Aggregate = Aggregate.Sum)( using keyCodec: KeyCodec[K] ): Command[Long] = { - val (indices, prefix) = storeKeyed(destination, first +: rest.toVector) - Command("ZINTERSTORE", indices, prefix ++ weightsArgs(weights) ++ aggregateArgs(aggregate), Decode.long) + val (indices, prefix) = KeyArgs.numKeyedAfter(keyCodec.encode(destination), first +: rest.toVector) + Command("ZINTERSTORE", 0 +: indices, prefix ++ weightsArgs(weights) ++ aggregateArgs(aggregate), Decode.long) } - def zInterCard[K](first: K, rest: K*)(limit: Option[Long] = None)(using keyCodec: KeyCodec[K]): Command[Long] = { - val (indices, prefix) = KeyArgs.numKeyed(first +: rest.toVector) - Command.read("ZINTERCARD", indices, prefix ++ limit.toVector.flatMap(n => Vector(LimitWord, Bytes.utf8(n.toString))), Decode.long) - } + def zInterCard[K](first: K, rest: K*)(limit: Option[Long] = None)(using keyCodec: KeyCodec[K]): Command[Long] = + KeyArgs.interCard("ZINTERCARD", first +: rest.toVector, limit) def zDiff[K, V](first: K, rest: K*)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Vector[V]] = { val (indices, prefix) = KeyArgs.numKeyed(first +: rest.toVector) @@ -347,62 +315,26 @@ private[sage] object SortedSets { } def zDiffStore[K](destination: K, first: K, rest: K*)(using keyCodec: KeyCodec[K]): Command[Long] = { - val (indices, prefix) = storeKeyed(destination, first +: rest.toVector) - Command("ZDIFFSTORE", indices, prefix, Decode.long) + val (indices, prefix) = KeyArgs.numKeyedAfter(keyCodec.encode(destination), first +: rest.toVector) + Command("ZDIFFSTORE", 0 +: indices, prefix, Decode.long) } def zScan[K, V](key: K, cursor: ScanCursor, pattern: Option[String] = None, count: Option[Long] = None)( using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V] ): Command[ScanPage[(V, Double)]] = - Command.readCursor( - "ZSCAN", - Command.FirstKey, - Vector(keyCodec.encode(key), ScanCursor.bytes(cursor)) ++ ScanArgs.options(pattern, count), - Decode.scanPage(Decode.scoredMembersFlat[V]) - ) + KeyArgs.keyScan("ZSCAN", key, cursor, pattern, count)(Decode.scoredMembersFlat[V]) - private def blockingPop[K, V](name: String, first: K, rest: Vector[K], timeout: BlockTimeout)( - using keyCodec: KeyCodec[K], - valueCodec: ValueCodec[V] - ): Command[Option[(K, V, Double)]] = { - val keys = (first +: rest).map(keyCodec.encode) - Command( - name, - keyIndices = Vector.tabulate(keys.size)(identity), - args = keys :+ BlockTimeout.wire(timeout), - decode = { - case Frame.Null => Right(None) - case other => - Decode - .array3(Decode.key[K], Decode.value[V], Decode.score, "key, member and score or null") { (key, member, s) => - (key, member, s) - }(other) - .map(Some(_)) - }, - execution = Execution.Blocking - ) - } + private def poppedMember[K, V](using KeyCodec[K], ValueCodec[V]): Frame => Either[DecodeError, (K, V, Double)] = + Decode.array3(Decode.key[K], Decode.value[V], Decode.double, "key, member and score or null")((_, _, _)) - private def mpopReply[K, V](using KeyCodec[K], ValueCodec[V]): Frame => Either[DecodeError, Option[(K, Vector[(V, Double)])]] = { - case Frame.Null => Right(None) - case other => - Decode.array2(Decode.key[K], Decode.scoredMembers[V], "key and members or null")(_ -> _)(other).map(Some(_)) - } + private val MpopLabel = "key and members or null" - private val rankWithScore: Frame => Either[DecodeError, Option[(Long, Double)]] = { - case Frame.Null => Right(None) - case other => Decode.array2(Decode.long, Decode.score, "rank/score pair or null")(_ -> _)(other).map(Some(_)) - } - - private def storeKeyed[K](destination: K, keys: Vector[K])(using keyCodec: KeyCodec[K]): (Vector[Int], Vector[Bytes]) = { - val encoded = keys.map(keyCodec.encode) - val indices = 0 +: Vector.tabulate(encoded.size)(_ + 2) - (indices, keyCodec.encode(destination) +: Bytes.utf8(encoded.size.toString) +: encoded) - } + private val rankWithScore: Frame => Either[DecodeError, Option[(Long, Double)]] = + Decode.nullable(Decode.array2(Decode.long, Decode.double, "rank/score pair or null")(_ -> _)) private def memberScoreArgs[V](pairs: Vector[(V, Double)])(using valueCodec: ValueCodec[V]): Vector[Bytes] = - pairs.flatMap { case (member, score) => Vector(scoreArg(score), valueCodec.encode(member)) } + pairs.flatMap { case (member, score) => Vector(Args.double(score), valueCodec.encode(member)) } private def conditionArgs(condition: ZAddCondition): Vector[Bytes] = condition match { @@ -418,22 +350,18 @@ private[sage] object SortedSets { private def rangeArgs[V](range: ZRange[V])(using valueCodec: ValueCodec[V]): Vector[Bytes] = range match { case ZRange.ByRank(start, stop, rev) => - Vector(Bytes.utf8(start.toString), Bytes.utf8(stop.toString)) ++ (if (rev) Vector(Rev) else Vector.empty) - case ZRange.ByScore(min, max, limit, rev) => - val (a, b) = if (rev) (max, min) else (min, max) - (Vector(scoreBoundaryArg(a), scoreBoundaryArg(b), ByScore) ++ (if (rev) Vector(Rev) else Vector.empty)) ++ limitArgs(limit) - case ZRange.ByLex(min, max, limit, rev) => - val (a, b) = if (rev) (max, min) else (min, max) - (Vector(lexBoundaryArg[V](a), lexBoundaryArg[V](b), ByLex) ++ (if (rev) Vector(Rev) else Vector.empty)) ++ limitArgs(limit) + Vector(Args.long(start), Args.long(stop)) ++ Args.flag(rev, Rev) + case ZRange.ByScore(min, max, limit, rev) => bounded(scoreBoundaryArg(min), scoreBoundaryArg(max), ByScore, limit, rev) + case ZRange.ByLex(min, max, limit, rev) => bounded(lexBoundaryArg[V](min), lexBoundaryArg[V](max), ByLex, limit, rev) } - private def limitArgs(limit: Option[Limit]): Vector[Bytes] = - limit.toVector.flatMap(l => Vector(LimitWord, Bytes.utf8(l.offset.toString), Bytes.utf8(l.count.toString))) + private def bounded(min: Bytes, max: Bytes, by: Bytes, limit: Option[Limit], rev: Boolean): Vector[Bytes] = + (if (rev) Vector(max, min, by, Rev) else Vector(min, max, by)) ++ Args.limit(limit) private def scoreBoundaryArg(boundary: ScoreBoundary): Bytes = boundary match { - case ScoreBoundary.Inclusive(s) => Bytes.utf8(formatScore(s)) - case ScoreBoundary.Exclusive(s) => Bytes.utf8("(" + formatScore(s)) + case ScoreBoundary.Inclusive(s) => Args.double(s) + case ScoreBoundary.Exclusive(s) => Bytes.utf8("(" + Doubles.format(s)) case ScoreBoundary.NegInf => NegInfWord case ScoreBoundary.PosInf => PosInfWord } @@ -447,27 +375,22 @@ private[sage] object SortedSets { } private def aggregateArgs(aggregate: Aggregate): Vector[Bytes] = - aggregate match { - case Aggregate.Sum => Vector.empty - case Aggregate.Min => Vector(AggregateWord, MinWord) - case Aggregate.Max => Vector(AggregateWord, MaxWord) - } + if (aggregate == Aggregate.Sum) Vector.empty else Vector(AggregateWord, aggregateWord(aggregate)) + + private def combine[K: KeyCodec, Out]( + name: String, + keys: Vector[K], + weights: Option[Vector[Double]], + aggregate: Aggregate, + withScores: Boolean, + decode: Frame => Either[DecodeError, Out] + ): Command[Out] = { + val (indices, prefix) = KeyArgs.numKeyed(keys) + Command.read(name, indices, prefix ++ weightsArgs(weights) ++ aggregateArgs(aggregate) ++ Args.flag(withScores, WithScores), decode) + } private def weightsArgs(weights: Option[Vector[Double]]): Vector[Bytes] = - weights.toVector.flatMap(ws => WeightsWord +: ws.map(w => Bytes.utf8(formatScore(w)))) - - private def countArg(count: Option[Long]): Vector[Bytes] = - count.toVector.flatMap(c => Vector(CountWord, Bytes.utf8(c.toString))) - - private def minMaxArg(minMax: MinMax): Bytes = - minMax match { - case MinMax.Min => MinWord - case MinMax.Max => MaxWord - } - - private def scoreArg(score: Double): Bytes = Bytes.utf8(formatScore(score)) - - private def formatScore(value: Double): String = Doubles.format(value) + weights.toVector.flatMap(ws => WeightsWord +: ws.map(Args.double)) private def prefixed(prefix: Char, value: Bytes): Bytes = { val src = value.unsafeArray diff --git a/sage-core/src/main/scala/sage/commands/StreamInfo.scala b/sage-core/src/main/scala/sage/commands/StreamInfo.scala index 41254fe7..cf7896a9 100644 --- a/sage-core/src/main/scala/sage/commands/StreamInfo.scala +++ b/sage-core/src/main/scala/sage/commands/StreamInfo.scala @@ -1,12 +1,12 @@ package sage.commands import java.time.Instant -import java.util.concurrent.TimeUnit import scala.concurrent.duration.FiniteDuration import sage.SageException.DecodeError import sage.codec.{KeyCodec, ValueCodec} +import sage.commands.Args.Count import sage.protocol.Frame /** @@ -100,7 +100,7 @@ private[sage] object StreamInfo { Command.readUncacheable( "XINFO STREAM", Command.FirstKey, - Vector(keyCodec.encode(key), Full) ++ count.toVector.flatMap(n => Vector(CountWord, sage.Bytes.utf8(n.toString))), + Vector(keyCodec.encode(key), Full) ++ Args.optLong(Count, count), streamFullReply[F, V] ) @@ -111,167 +111,100 @@ private[sage] object StreamInfo { Command.readUncacheable("XINFO CONSUMERS", Command.FirstKey, Vector(keyCodec.encode(key), sage.Bytes.utf8(group)), Decode.vector(consumerReply)) private def streamReply[F, V](using KeyCodec[F], ValueCodec[V]): Frame => Either[DecodeError, StreamInfo[F, V]] = - frame => - Fields.of(frame).flatMap { f => - for { - length <- f.required("length", Decode.long) - rtKeys <- f.requiredOr("radix-tree-keys", Decode.long, 0L) - rtNodes <- f.requiredOr("radix-tree-nodes", Decode.long, 0L) - lastId <- f.required("last-generated-id", Streams.streamId) - maxDel <- f.optional("max-deleted-entry-id", Streams.streamId) - added <- f.optional("entries-added", Decode.long) - firstRec <- f.optional("recorded-first-entry-id", Streams.streamId) - groups <- f.requiredOr("groups", Decode.long, 0L) - firstE <- f.optional("first-entry", Streams.streamEntry[F, V]) - lastE <- f.optional("last-entry", Streams.streamEntry[F, V]) - } yield StreamInfo(length, rtKeys, rtNodes, lastId, maxDel, added, firstRec, groups, firstE, lastE) - } + Decode.fields { f => + for { + length <- f.required("length", Decode.long) + rtKeys <- f.requiredOr("radix-tree-keys", Decode.long, 0L) + rtNodes <- f.requiredOr("radix-tree-nodes", Decode.long, 0L) + lastId <- f.required("last-generated-id", Streams.streamId) + maxDel <- f.optional("max-deleted-entry-id", Streams.streamId) + added <- f.optional("entries-added", Decode.long) + firstRec <- f.optional("recorded-first-entry-id", Streams.streamId) + groups <- f.requiredOr("groups", Decode.long, 0L) + firstE <- f.optional("first-entry", Streams.streamEntry[F, V]) + lastE <- f.optional("last-entry", Streams.streamEntry[F, V]) + } yield StreamInfo(length, rtKeys, rtNodes, lastId, maxDel, added, firstRec, groups, firstE, lastE) + } private def streamFullReply[F, V](using KeyCodec[F], ValueCodec[V]): Frame => Either[DecodeError, StreamInfoFull[F, V]] = - frame => - Fields.of(frame).flatMap { f => - for { - length <- f.required("length", Decode.long) - rtKeys <- f.requiredOr("radix-tree-keys", Decode.long, 0L) - rtNodes <- f.requiredOr("radix-tree-nodes", Decode.long, 0L) - lastId <- f.required("last-generated-id", Streams.streamId) - maxDel <- f.optional("max-deleted-entry-id", Streams.streamId) - added <- f.optional("entries-added", Decode.long) - firstRec <- f.optional("recorded-first-entry-id", Streams.streamId) - entries <- f.optionalVector("entries", Streams.streamEntry[F, V]) - groups <- f.optionalVector("groups", fullGroupReply) - } yield StreamInfoFull(length, rtKeys, rtNodes, lastId, maxDel, added, firstRec, entries, groups) - } + Decode.fields { f => + for { + length <- f.required("length", Decode.long) + rtKeys <- f.requiredOr("radix-tree-keys", Decode.long, 0L) + rtNodes <- f.requiredOr("radix-tree-nodes", Decode.long, 0L) + lastId <- f.required("last-generated-id", Streams.streamId) + maxDel <- f.optional("max-deleted-entry-id", Streams.streamId) + added <- f.optional("entries-added", Decode.long) + firstRec <- f.optional("recorded-first-entry-id", Streams.streamId) + entries <- f.optionalVector("entries", Streams.streamEntry[F, V]) + groups <- f.optionalVector("groups", fullGroupReply) + } yield StreamInfoFull(length, rtKeys, rtNodes, lastId, maxDel, added, firstRec, entries, groups) + } private val groupReply: Frame => Either[DecodeError, GroupInfo] = - frame => - Fields.of(frame).flatMap { f => - for { - name <- f.required("name", Decode.utf8String) - consumers <- f.requiredOr("consumers", Decode.long, 0L) - pending <- f.requiredOr("pending", Decode.long, 0L) - lastId <- f.required("last-delivered-id", Streams.streamId) - read <- f.optional("entries-read", Decode.long) - lag <- f.optional("lag", Decode.long) - } yield GroupInfo(name, consumers, pending, lastId, read, lag) - } + Decode.fields { f => + for { + name <- f.required("name", Decode.utf8String) + consumers <- f.requiredOr("consumers", Decode.long, 0L) + pending <- f.requiredOr("pending", Decode.long, 0L) + lastId <- f.required("last-delivered-id", Streams.streamId) + read <- f.optional("entries-read", Decode.long) + lag <- f.optional("lag", Decode.long) + } yield GroupInfo(name, consumers, pending, lastId, read, lag) + } private val consumerReply: Frame => Either[DecodeError, ConsumerInfo] = - frame => - Fields.of(frame).flatMap { f => - for { - name <- f.required("name", Decode.utf8String) - pending <- f.requiredOr("pending", Decode.long, 0L) - idle <- f.required("idle", millisDuration) - inactive <- f.optional("inactive", millisDuration) - } yield ConsumerInfo(name, pending, idle, inactive) - } + Decode.fields { f => + for { + name <- f.required("name", Decode.utf8String) + pending <- f.requiredOr("pending", Decode.long, 0L) + idle <- f.required("idle", Decode.millisDuration) + inactive <- f.optional("inactive", Decode.millisDuration) + } yield ConsumerInfo(name, pending, idle, inactive) + } private val fullGroupReply: Frame => Either[DecodeError, FullGroupInfo] = - frame => - Fields.of(frame).flatMap { f => - for { - name <- f.required("name", Decode.utf8String) - lastId <- f.required("last-delivered-id", Streams.streamId) - pelCount <- f.optional("pel-count", Decode.long) - read <- f.optional("entries-read", Decode.long) - lag <- f.optional("lag", Decode.long) - pending <- f.optionalVector("pending", groupPendingReply) - consumers <- f.optionalVector("consumers", fullConsumerReply) - } yield FullGroupInfo(name, lastId, pelCount, read, lag, pending, consumers) - } + Decode.fields { f => + for { + name <- f.required("name", Decode.utf8String) + lastId <- f.required("last-delivered-id", Streams.streamId) + pelCount <- f.optional("pel-count", Decode.long) + read <- f.optional("entries-read", Decode.long) + lag <- f.optional("lag", Decode.long) + pending <- f.optionalVector("pending", groupPendingReply) + consumers <- f.optionalVector("consumers", fullConsumerReply) + } yield FullGroupInfo(name, lastId, pelCount, read, lag, pending, consumers) + } private val fullConsumerReply: Frame => Either[DecodeError, FullConsumerInfo] = - frame => - Fields.of(frame).flatMap { f => - for { - name <- f.required("name", Decode.utf8String) - seen <- f.optional("seen-time", millisInstant) - active <- f.optional("active-time", millisInstant) - pel <- f.optional("pel-count", Decode.long) - pending <- f.optionalVector("pending", consumerPendingReply) - } yield FullConsumerInfo(name, seen, active, pel, pending) - } - - // defined before the decoders that compose them: the array combinators force their decoder arguments when those vals initialize - private val millisDuration: Frame => Either[DecodeError, FiniteDuration] = - frame => Decode.long(frame).map(ms => FiniteDuration(ms, TimeUnit.MILLISECONDS)) - - private val millisInstant: Frame => Either[DecodeError, Instant] = - frame => Decode.long(frame).map(Instant.ofEpochMilli) + Decode.fields { f => + for { + name <- f.required("name", Decode.utf8String) + seen <- f.optional("seen-time", Decode.millisInstant) + active <- f.optional("active-time", Decode.millisInstant) + pel <- f.optional("pel-count", Decode.long) + pending <- f.optionalVector("pending", consumerPendingReply) + } yield FullConsumerInfo(name, seen, active, pel, pending) + } // A group-level FULL PEL row includes its owning consumer: [id, consumer, delivery-time-ms, delivery-count]. private val groupPendingReply: Frame => Either[DecodeError, FullPendingEntry] = - Decode.array4(Streams.streamId, Decode.utf8String, millisInstant, Decode.long, "group pending [id, consumer, delivery-time, delivery-count]") { - (id, consumer, delivered, count) => FullPendingEntry(id, Some(consumer), delivered, count) + Decode.array4( + Streams.streamId, + Decode.utf8String, + Decode.millisInstant, + Decode.long, + "group pending [id, consumer, delivery-time, delivery-count]" + ) { (id, consumer, delivered, count) => + FullPendingEntry(id, Some(consumer), delivered, count) } // a consumer-level FULL PEL row omits the (implied) consumer: [id, delivery-time-ms, delivery-count] private val consumerPendingReply: Frame => Either[DecodeError, FullPendingEntry] = - Decode.array3(Streams.streamId, millisInstant, Decode.long, "consumer pending [id, delivery-time, delivery-count]") { (id, delivered, count) => - FullPendingEntry(id, None, delivered, count) - } - - /** - * A lenient view over an introspection reply map: read fields by known name, ignore the rest. - */ - final private class Fields private (table: Map[String, Frame]) { - - def required[A](name: String, decode: Frame => Either[DecodeError, A]): Either[DecodeError, A] = - table.get(name).toRight(DecodeError(s"field '$name'", "absent")).flatMap(decode) - - // a field that is core to the reply but whose absence on some server we tolerate with a default rather than failing the whole decode - def requiredOr[A](name: String, decode: Frame => Either[DecodeError, A], fallback: A): Either[DecodeError, A] = - table.get(name) match { - case None | Some(Frame.Null) => Right(fallback) - case Some(frame) => decode(frame) - } - - def optional[A](name: String, decode: Frame => Either[DecodeError, A]): Either[DecodeError, Option[A]] = - table.get(name) match { - case None | Some(Frame.Null) => Right(None) - case Some(frame) => decode(frame).map(Some(_)) - } - - def optionalVector[A](name: String, element: Frame => Either[DecodeError, A]): Either[DecodeError, Vector[A]] = - table.get(name) match { - case None | Some(Frame.Null) => Right(Vector.empty) - case Some(frame) => Decode.vector(element)(frame) - } - } - - private object Fields { - - def of(frame: Frame): Either[DecodeError, Fields] = - frame match { - case Frame.Map(entries) => build(entries) - case Frame.Array(elements) if elements.length % 2 == 0 => - val pairs = Vector.newBuilder[(Frame, Frame)] - pairs.sizeHint(elements.length / 2) - var i = 0 - while (i < elements.length) { - pairs += (elements(i) -> elements(i + 1)) - i += 2 - } - build(pairs.result()) - case other => Left(DecodeError("introspection map", Frame.describe(other))) - } - - private def build(entries: Vector[(Frame, Frame)]): Either[DecodeError, Fields] = { - val builder = Map.newBuilder[String, Frame] - val it = entries.iterator - while (it.hasNext) { - val (keyFrame, valueFrame) = it.next() - keyFrame match { - case Frame.BulkString(bytes) => builder += bytes.asUtf8String -> valueFrame - case Frame.SimpleString(name) => builder += name -> valueFrame - case other => return Left(DecodeError("field name", Frame.describe(other))) - } - } - Right(new Fields(builder.result())) + Decode.array3(Streams.streamId, Decode.millisInstant, Decode.long, "consumer pending [id, delivery-time, delivery-count]") { + (id, delivered, count) => + FullPendingEntry(id, None, delivered, count) } - } - private val Full = sage.Bytes.utf8("FULL") - private val CountWord = sage.Bytes.utf8("COUNT") + private val Full = sage.Bytes.utf8("FULL") } diff --git a/sage-core/src/main/scala/sage/commands/Streams.scala b/sage-core/src/main/scala/sage/commands/Streams.scala index fdb958bc..5a62e873 100644 --- a/sage-core/src/main/scala/sage/commands/Streams.scala +++ b/sage-core/src/main/scala/sage/commands/Streams.scala @@ -6,7 +6,8 @@ import scala.concurrent.duration.FiniteDuration import sage.Bytes import sage.SageException.DecodeError -import sage.codec.{KeyCodec, ValueCodec} +import sage.codec.{KeyCodec, Primitives, ValueCodec} +import sage.commands.Args.{Count, LimitWord} import sage.protocol.Frame /** @@ -14,11 +15,12 @@ import sage.protocol.Frame * positions that accept special tokens (`*`, `-`/`+`, `$`, `>`, `(`) use separate sealed types that list the supported values. */ final case class StreamId(ms: Long, seq: Long) extends Ordered[StreamId] { - def compare(that: StreamId): Int = { + def compare(that: StreamId): Int = { val c = java.lang.Long.compareUnsigned(ms, that.ms) if (c != 0) c else java.lang.Long.compareUnsigned(seq, that.seq) } - private[commands] def wire: Bytes = Bytes.utf8(s"$ms-$seq") + private[commands] def text: String = s"${java.lang.Long.toUnsignedString(ms)}-${java.lang.Long.toUnsignedString(seq)}" + private[commands] def wire: Bytes = Bytes.utf8(text) } object StreamId { @@ -152,8 +154,6 @@ private[sage] object Streams { private val MinIdWord = Bytes.utf8("MINID") private val Eq = Bytes.utf8("=") private val Tilde = Bytes.utf8("~") - private val LimitWord = Bytes.utf8("LIMIT") - private val Count = Bytes.utf8("COUNT") private val Block = Bytes.utf8("BLOCK") private val StreamsWord = Bytes.utf8("STREAMS") private val Group = Bytes.utf8("GROUP") @@ -168,11 +168,8 @@ private[sage] object Streams { private val Force = Bytes.utf8("FORCE") private val JustId = Bytes.utf8("JUSTID") private val Ids = Bytes.utf8("IDS") - private val DelRef = Bytes.utf8("DELREF") - private val Acked = Bytes.utf8("ACKED") - private val Silent = Bytes.utf8("SILENT") - private val Fail = Bytes.utf8("FAIL") - private val Fatal = Bytes.utf8("FATAL") + private val policyWord = Args.keywords(StreamDeletionPolicy.values) + private val nackModeWire = Args.keywords(NackMode.values) private val IdmpDuration = Bytes.utf8("IDMP-DURATION") private val IdmpMaxSize = Bytes.utf8("IDMP-MAXSIZE") @@ -196,7 +193,7 @@ private[sage] object Streams { Command.read("XLEN", Command.FirstKey, Vector(keyCodec.encode(key)), Decode.long) def xDel[K](key: K)(first: StreamId, rest: StreamId*)(using keyCodec: KeyCodec[K]): Command[Long] = - Command("XDEL", Command.FirstKey, keyCodec.encode(key) +: (first +: rest.toVector).map(_.wire), Decode.long) + Command("XDEL", Command.FirstKey, Args.keyThen(key, first, rest)(_.wire), Decode.long) def xTrim[K](key: K, trim: Trimming, policy: StreamDeletionPolicy = StreamDeletionPolicy.KeepRef)(using keyCodec: KeyCodec[K]): Command[Long] = Command("XTRIM", Command.FirstKey, (keyCodec.encode(key) +: trimArgs(trim)) ++ policyArgs(policy), Decode.long) @@ -207,8 +204,8 @@ private[sage] object Streams { Command( "XSETID", Command.FirstKey, - (Vector(keyCodec.encode(key), groupStartWire(id)) ++ entriesAdded.toVector.flatMap(n => Vector(EntriesAdded, Bytes.utf8(n.toString)))) ++ - maxDeletedId.toVector.flatMap(d => Vector(MaxDeletedId, d.wire)), + (Vector(keyCodec.encode(key), groupStartWire(id)) ++ Args.optLong(EntriesAdded, entriesAdded)) ++ + Args.opt(MaxDeletedId, maxDeletedId)(_.wire), Decode.ok ) @@ -218,8 +215,8 @@ private[sage] object Streams { "XCFGSET", Command.FirstKey, keyCodec.encode(key) +: - (idmpDuration.toVector.flatMap(d => Vector(IdmpDuration, Bytes.utf8(d.toSeconds.toString))) ++ - idmpMaxSize.toVector.flatMap(n => Vector(IdmpMaxSize, Bytes.utf8(n.toString)))), + (Args.opt(IdmpDuration, idmpDuration)(d => Args.long(Math.ceilDiv(d.toNanos, 1000000000L))) ++ + Args.optLong(IdmpMaxSize, idmpMaxSize)), Decode.ok ) @@ -247,7 +244,7 @@ private[sage] object Streams { Command.read( name, Command.FirstKey, - Vector(keyCodec.encode(key), rangeWire(a), rangeWire(b)) ++ count.toVector.flatMap(n => Vector(Count, Bytes.utf8(n.toString))), + Vector(keyCodec.encode(key), rangeWire(a), rangeWire(b)) ++ Args.optLong(Count, count), Decode.vector(streamEntry[F, V]) ) @@ -258,7 +255,7 @@ private[sage] object Streams { fieldCodec: KeyCodec[F], valueCodec: ValueCodec[V] ): Command[Vector[(K, Vector[StreamEntry[F, V]])]] = { - val leading = countArg(count) ++ blockArg(block) + val leading = Args.optLong(Count, count) ++ blockArg(block) val (args, keys) = streamsArgs(first +: rest.toVector, leading, readWire) Command("XREAD", keys, args, readReply[K, F, V], execution = blockExecution(block), isReadOnly = true) } @@ -269,7 +266,7 @@ private[sage] object Streams { noAck: Boolean = false )(using keyCodec: KeyCodec[K], fieldCodec: KeyCodec[F], valueCodec: ValueCodec[V]): Command[Vector[(K, Vector[StreamEntry[F, V]])]] = { val leading = - Vector(Group, Bytes.utf8(group), Bytes.utf8(consumer)) ++ countArg(count) ++ blockArg(block) ++ (if (noAck) Vector(NoAck) else Vector.empty) + Vector(Group, Bytes.utf8(group), Bytes.utf8(consumer)) ++ Args.optLong(Count, count) ++ blockArg(block) ++ Args.flag(noAck, NoAck) val (args, keys) = streamsArgs(first +: rest.toVector, leading, groupReadWire) Command("XREADGROUP", keys, args, readReply[K, F, V], execution = blockExecution(block)) } @@ -285,8 +282,8 @@ private[sage] object Streams { Command( "XGROUP CREATE", Command.FirstKey, - (Vector(keyCodec.encode(key), Bytes.utf8(group), groupStartWire(id)) ++ (if (mkStream) Vector(MkStream) else Vector.empty)) ++ - entriesRead.toVector.flatMap(n => Vector(EntriesRead, Bytes.utf8(n.toString))), + (Vector(keyCodec.encode(key), Bytes.utf8(group), groupStartWire(id)) ++ Args.flag(mkStream, MkStream)) ++ + Args.optLong(EntriesRead, entriesRead), Decode.ok ) @@ -296,9 +293,7 @@ private[sage] object Streams { Command( "XGROUP SETID", Command.FirstKey, - Vector(keyCodec.encode(key), Bytes.utf8(group), groupStartWire(id)) ++ entriesRead.toVector.flatMap(n => - Vector(EntriesRead, Bytes.utf8(n.toString)) - ), + Vector(keyCodec.encode(key), Bytes.utf8(group), groupStartWire(id)) ++ Args.optLong(EntriesRead, entriesRead), Decode.ok ) @@ -321,7 +316,7 @@ private[sage] object Streams { Command( "XCLAIM", Command.FirstKey, - claimArgs(key, group, consumer, minIdle, first +: rest.toVector, idle, retryCount, force, justId = false), + claimArgs(key, group, consumer, minIdle, first +: rest.toVector, idle, retryCount, force), Decode.vector(streamEntry[F, V]) ) @@ -333,7 +328,7 @@ private[sage] object Streams { Command( "XCLAIM", Command.FirstKey, - claimArgs(key, group, consumer, minIdle, first +: rest.toVector, idle, retryCount, force, justId = true), + claimArgs(key, group, consumer, minIdle, first +: rest.toVector, idle, retryCount, force) :+ JustId, Decode.vector(streamId) ) @@ -349,7 +344,7 @@ private[sage] object Streams { fieldCodec: KeyCodec[F], valueCodec: ValueCodec[V] ): Command[XAutoClaimResult[F, V]] = - Command("XAUTOCLAIM", Command.FirstKey, autoClaimArgs(key, group, consumer, minIdle, start, count, justId = false), autoClaimReply[F, V]) + Command("XAUTOCLAIM", Command.FirstKey, autoClaimArgs(key, group, consumer, minIdle, start, count), autoClaimReply[F, V]) def xAutoClaimJustId[K]( key: K, @@ -361,7 +356,7 @@ private[sage] object Streams { )( using keyCodec: KeyCodec[K] ): Command[XAutoClaimJustIdResult] = - Command("XAUTOCLAIM", Command.FirstKey, autoClaimArgs(key, group, consumer, minIdle, start, count, justId = true), autoClaimJustIdReply) + Command("XAUTOCLAIM", Command.FirstKey, autoClaimArgs(key, group, consumer, minIdle, start, count) :+ JustId, autoClaimJustIdReply) // --- pending ------------------------------------------------------------ @@ -380,8 +375,8 @@ private[sage] object Streams { Command.readUncacheable( "XPENDING", Command.FirstKey, - (Vector(keyCodec.encode(key), Bytes.utf8(group)) ++ idle.toVector.flatMap(d => Vector(Idle, Bytes.utf8(TimeArgs.millis(d).toString)))) ++ - Vector(rangeWire(start), rangeWire(end), Bytes.utf8(count.toString)) ++ consumer.toVector.map(Bytes.utf8), + (Vector(keyCodec.encode(key), Bytes.utf8(group)) ++ Args.opt(Idle, idle)(d => Args.long(TimeArgs.millis(d)))) ++ + Vector(rangeWire(start), rangeWire(end), Args.long(count)) ++ consumer.toVector.map(Bytes.utf8), Decode.vector(pendingEntryElement) ) @@ -389,41 +384,30 @@ private[sage] object Streams { def xDelEx[K](key: K, policy: StreamDeletionPolicy = StreamDeletionPolicy.KeepRef)(first: StreamId, rest: StreamId*)( using keyCodec: KeyCodec[K] - ): Command[Vector[StreamEntryDeletion]] = { - val ids = first +: rest.toVector - Command( - "XDELEX", - Command.FirstKey, - (keyCodec.encode(key) +: policyArgs(policy)) ++ (Ids +: Bytes.utf8(ids.size.toString) +: ids.map(_.wire)), - Decode.vector(deletionElement) - ) - } + ): Command[Vector[StreamEntryDeletion]] = + Command("XDELEX", Command.FirstKey, (keyCodec.encode(key) +: policyArgs(policy)) ++ idsArgs(first, rest), Decode.vector(deletionElement)) def xAckDel[K](key: K, group: String, policy: StreamDeletionPolicy = StreamDeletionPolicy.KeepRef)(first: StreamId, rest: StreamId*)( using keyCodec: KeyCodec[K] - ): Command[Vector[StreamEntryDeletion]] = { - val ids = first +: rest.toVector + ): Command[Vector[StreamEntryDeletion]] = Command( "XACKDEL", Command.FirstKey, - (Vector(keyCodec.encode(key), Bytes.utf8(group)) ++ policyArgs(policy)) ++ (Ids +: Bytes.utf8(ids.size.toString) +: ids.map(_.wire)), + (Vector(keyCodec.encode(key), Bytes.utf8(group)) ++ policyArgs(policy)) ++ idsArgs(first, rest), Decode.vector(deletionElement) ) - } def xNack[K](key: K, group: String, mode: NackMode)(first: StreamId, rest: StreamId*)( retryCount: Option[Long] = None, force: Boolean = false - )(using keyCodec: KeyCodec[K]): Command[Long] = { - val ids = first +: rest.toVector + )(using keyCodec: KeyCodec[K]): Command[Long] = Command( "XNACK", Command.FirstKey, - (Vector(keyCodec.encode(key), Bytes.utf8(group), nackModeWire(mode), Ids, Bytes.utf8(ids.size.toString)) ++ ids.map(_.wire)) ++ - retryCount.toVector.flatMap(n => Vector(RetryCount, Bytes.utf8(n.toString))) ++ (if (force) Vector(Force) else Vector.empty), + (Vector(keyCodec.encode(key), Bytes.utf8(group), nackModeWire(mode)) ++ idsArgs(first, rest)) ++ + Args.optLong(RetryCount, retryCount) ++ Args.flag(force, Force), Decode.long ) - } // --- arg builders ------------------------------------------------------- @@ -432,10 +416,8 @@ private[sage] object Streams { fieldCodec: KeyCodec[F], valueCodec: ValueCodec[V] ): Vector[Bytes] = - (keyCodec.encode(key) +: (if (noMkStream) Vector(NoMkStream) else Vector.empty)) ++ - policyArgs(policy) ++ trim.toVector.flatMap(trimArgs) ++ (xAddIdWire(id) +: fields.flatMap { case (f, v) => - Vector(fieldCodec.encode(f), valueCodec.encode(v)) - }) + (keyCodec.encode(key) +: Args.flag(noMkStream, NoMkStream)) ++ + policyArgs(policy) ++ trim.toVector.flatMap(trimArgs) ++ (xAddIdWire(id) +: Args.pairs(fields)) private def claimArgs[K]( key: K, @@ -445,12 +427,11 @@ private[sage] object Streams { ids: Vector[StreamId], idle: Option[ClaimIdle], retryCount: Option[Long], - force: Boolean, - justId: Boolean + force: Boolean )(using keyCodec: KeyCodec[K]): Vector[Bytes] = - (Vector(keyCodec.encode(key), Bytes.utf8(group), Bytes.utf8(consumer), Bytes.utf8(TimeArgs.millis(minIdle).toString)) ++ ids.map(_.wire)) ++ - idleArgs(idle) ++ retryCount.toVector.flatMap(n => Vector(RetryCount, Bytes.utf8(n.toString))) ++ - (if (force) Vector(Force) else Vector.empty) ++ (if (justId) Vector(JustId) else Vector.empty) + (Vector(keyCodec.encode(key), Bytes.utf8(group), Bytes.utf8(consumer), Args.long(TimeArgs.millis(minIdle))) ++ ids.map(_.wire)) ++ + idleArgs(idle) ++ Args.optLong(RetryCount, retryCount) ++ + Args.flag(force, Force) private def autoClaimArgs[K]( key: K, @@ -458,44 +439,36 @@ private[sage] object Streams { consumer: String, minIdle: FiniteDuration, start: StreamId, - count: Option[Long], - justId: Boolean + count: Option[Long] )( using keyCodec: KeyCodec[K] ): Vector[Bytes] = - Vector(keyCodec.encode(key), Bytes.utf8(group), Bytes.utf8(consumer), Bytes.utf8(TimeArgs.millis(minIdle).toString), start.wire) ++ - count.toVector.flatMap(n => Vector(Count, Bytes.utf8(n.toString))) ++ (if (justId) Vector(JustId) else Vector.empty) + Vector(keyCodec.encode(key), Bytes.utf8(group), Bytes.utf8(consumer), Args.long(TimeArgs.millis(minIdle)), start.wire) ++ + Args.optLong(Count, count) + + private def idsArgs(first: StreamId, rest: Seq[StreamId]): Vector[Bytes] = + Ids +: Args.long(rest.length + 1) +: (first +: rest.toVector).map(_.wire) private def idleArgs(idle: Option[ClaimIdle]): Vector[Bytes] = idle.toVector.flatMap { - case ClaimIdle.Idle(duration) => Vector(Idle, Bytes.utf8(TimeArgs.millis(duration).toString)) - case ClaimIdle.At(timestamp) => Vector(Time, Bytes.utf8(TimeArgs.millis(timestamp).toString)) + case ClaimIdle.Idle(duration) => Vector(Idle, Args.long(TimeArgs.millis(duration))) + case ClaimIdle.At(timestamp) => Vector(Time, Args.long(TimeArgs.millis(timestamp))) } private def trimArgs(trim: Trimming): Vector[Bytes] = trim match { - case Trimming.Exact(threshold) => thresholdKeyword(threshold) +: Vector(Eq, thresholdValue(threshold)) - case Trimming.Approximate(threshold, limit) => - (thresholdKeyword(threshold) +: Vector(Tilde, thresholdValue(threshold))) ++ limit.toVector.flatMap(n => - Vector(LimitWord, Bytes.utf8(n.toString)) - ) + case Trimming.Exact(threshold) => thresholdArgs(threshold, Eq) + case Trimming.Approximate(threshold, limit) => thresholdArgs(threshold, Tilde) ++ Args.optLong(LimitWord, limit) } - private def thresholdKeyword(threshold: TrimThreshold): Bytes = threshold match { - case _: TrimThreshold.MaxLen => MaxLenWord - case _: TrimThreshold.MinId => MinIdWord - } - private def thresholdValue(threshold: TrimThreshold): Bytes = threshold match { - case TrimThreshold.MaxLen(c) => Bytes.utf8(c.toString) - case TrimThreshold.MinId(id) => id.wire - } + private def thresholdArgs(threshold: TrimThreshold, operator: Bytes): Vector[Bytes] = + threshold match { + case TrimThreshold.MaxLen(count) => Vector(MaxLenWord, operator, Args.long(count)) + case TrimThreshold.MinId(id) => Vector(MinIdWord, operator, id.wire) + } private def policyArgs(policy: StreamDeletionPolicy): Vector[Bytes] = - policy match { - case StreamDeletionPolicy.KeepRef => Vector.empty - case StreamDeletionPolicy.DelRef => Vector(DelRef) - case StreamDeletionPolicy.Acked => Vector(Acked) - } + if (policy == StreamDeletionPolicy.KeepRef) Vector.empty else Vector(policyWord(policy)) private def streamsArgs[K, I](keysAndIds: Vector[(K, I)], leading: Vector[Bytes], idWire: I => Bytes)( using keyCodec: KeyCodec[K] @@ -507,8 +480,7 @@ private[sage] object Streams { ((leading :+ StreamsWord) ++ keys ++ ids, keyIndices) } - private def countArg(count: Option[Long]): Vector[Bytes] = count.toVector.flatMap(n => Vector(Count, Bytes.utf8(n.toString))) - private def blockArg(block: Option[BlockTimeout]): Vector[Bytes] = block.toVector.flatMap(b => Vector(Block, BlockTimeout.millisWire(b))) + private def blockArg(block: Option[BlockTimeout]): Vector[Bytes] = Args.opt(Block, block)(BlockTimeout.millisWire) private def blockExecution(block: Option[BlockTimeout]): Execution = if (block.isDefined) Execution.Blocking else Execution.Ordinary // --- token wire forms --------------------------------------------------- @@ -516,7 +488,7 @@ private[sage] object Streams { private def xAddIdWire(id: XAddId): Bytes = id match { case XAddId.Auto => Star - case XAddId.AutoSeq(ms) => Bytes.utf8(s"$ms-*") + case XAddId.AutoSeq(ms) => Bytes.utf8(s"${java.lang.Long.toUnsignedString(ms)}-*") case XAddId.Explicit(sid) => sid.wire } @@ -525,7 +497,7 @@ private[sage] object Streams { case StreamRangeId.Min => Dash case StreamRangeId.Max => Plus case StreamRangeId.Inclusive(sid) => sid.wire - case StreamRangeId.Exclusive(sid) => Bytes.utf8(s"(${sid.ms}-${sid.seq}") + case StreamRangeId.Exclusive(sid) => Bytes.utf8("(" + sid.text) } private def readWire(id: ReadId): Bytes = @@ -547,44 +519,31 @@ private[sage] object Streams { case GroupStartId.At(sid) => sid.wire } - private def nackModeWire(mode: NackMode): Bytes = - mode match { - case NackMode.Silent => Silent - case NackMode.Fail => Fail - case NackMode.Fatal => Fatal - } - // --- decoders ----------------------------------------------------------- - private[commands] val streamId: Frame => Either[DecodeError, StreamId] = { + private[commands] val streamId: Frame => Either[DecodeError, StreamId] = Decode.shape("stream id") { case Frame.BulkString(raw) => parseId(raw.asUtf8String) case Frame.SimpleString(raw) => parseId(raw) - case other => Left(DecodeError("stream id", Frame.describe(other))) } + // both parts are unsigned 64-bit numbers, matching StreamId.compare private def parseId(text: String): Either[DecodeError, StreamId] = { val dash = text.indexOf('-') - if (dash < 0) text.toLongOption.map(ms => StreamId(ms, 0L)).toRight(DecodeError("stream id 'ms-seq'", s"'$text'")) - else - (text.substring(0, dash).toLongOption, text.substring(dash + 1).toLongOption) match { - case (Some(ms), Some(seq)) => Right(StreamId(ms, seq)) - case _ => Left(DecodeError("stream id 'ms-seq'", s"'$text'")) - } + val id = + if (dash < 0) unsignedLong(text).map(StreamId(_, 0L)) + else unsignedLong(text.substring(0, dash)).zip(unsignedLong(text.substring(dash + 1))).map(StreamId(_, _)) + id.toRight(DecodeError("stream id 'ms-seq'", s"'$text'")) } - private val optionalStreamId: Frame => Either[DecodeError, Option[StreamId]] = { - case Frame.Null => Right(None) - case other => streamId(other).map(Some(_)) - } + private def unsignedLong(text: String): Option[Long] = + try Some(java.lang.Long.parseUnsignedLong(text)) + catch { case _: NumberFormatException => None } + + private val optionalStreamId: Frame => Either[DecodeError, Option[StreamId]] = Decode.nullable(streamId) // an entry is `[id, [field, value, …]]`; a tombstone (claimed entry whose data was deleted) is `[id, nil]` - private[commands] def streamEntry[F, V](using KeyCodec[F], ValueCodec[V]): Frame => Either[DecodeError, StreamEntry[F, V]] = { - val fields: Frame => Either[DecodeError, Vector[(F, V)]] = { - case Frame.Null => Right(Vector.empty) - case other => Decode.flatPairs[F, V](other) - } - Decode.array2(streamId, fields, "stream entry [id, fields]")(StreamEntry(_, _)) - } + private[commands] def streamEntry[F, V](using KeyCodec[F], ValueCodec[V]): Frame => Either[DecodeError, StreamEntry[F, V]] = + Decode.array2(streamId, Decode.orEmpty(Decode.flatPairs[F, V]), "stream entry [id, fields]")(StreamEntry(_, _)) // XREAD/XREADGROUP reply: RESP3 map of stream-name -> entries (RESP2 array of [name, entries] pairs); null when nothing is ready private def readReply[K, F, V]( @@ -592,89 +551,64 @@ private[sage] object Streams { KeyCodec[F], ValueCodec[V] ): Frame => Either[DecodeError, Vector[(K, Vector[StreamEntry[F, V]])]] = { - val entries = Decode.vector(streamEntry[F, V]) - def pair(nameFrame: Frame, entriesFrame: Frame): Either[DecodeError, (K, Vector[StreamEntry[F, V]])] = - for { - name <- Decode.key[K](nameFrame) - decoded <- entries(entriesFrame) - } yield name -> decoded - frame => - frame match { - case Frame.Null => Right(Vector.empty) - case Frame.Map(rows) => Decode.each(rows) { case (n, e) => pair(n, e) } - case Frame.Array(rows) => - Decode.each(rows) { - case Frame.Array(Vector(n, e)) => pair(n, e) - case other => Left(DecodeError("stream [name, entries] pair", Frame.describe(other))) - } - case other => Left(DecodeError("stream read map or null", Frame.describe(other))) - } - } - - private def autoClaimReply[F, V](using KeyCodec[F], ValueCodec[V]): Frame => Either[DecodeError, XAutoClaimResult[F, V]] = { - case Frame.Array(Vector(cursorFrame, entriesFrame, deletedFrame)) => - for { - cursor <- streamId(cursorFrame) - entries <- Decode.vector(streamEntry[F, V])(entriesFrame) - deleted <- Decode.vector(streamId)(deletedFrame) - } yield XAutoClaimResult(cursor, entries, deleted) - // pre-7.0 omits the deleted-ids element - case Frame.Array(Vector(cursorFrame, entriesFrame)) => - for { - cursor <- streamId(cursorFrame) - entries <- Decode.vector(streamEntry[F, V])(entriesFrame) - } yield XAutoClaimResult(cursor, entries, Vector.empty) - case other => Left(DecodeError("xautoclaim [cursor, entries, deleted]", Frame.describe(other))) + val pair = Decode.pair(Decode.key[K], Decode.vector(streamEntry[F, V])) + Decode.shape("stream read map or null") { + case Frame.Null => Right(Vector.empty) + case Frame.Map(rows) => Decode.each(rows)(pair.tupled) + case Frame.Array(rows) => + Decode.each(rows) { + case Frame.Array(Vector(n, e)) => pair(n, e) + case other => Left(DecodeError("stream [name, entries] pair", Frame.describe(other))) + } + } } - private val autoClaimJustIdReply: Frame => Either[DecodeError, XAutoClaimJustIdResult] = { - case Frame.Array(Vector(cursorFrame, claimedFrame, deletedFrame)) => - for { - cursor <- streamId(cursorFrame) - claimed <- Decode.vector(streamId)(claimedFrame) - deleted <- Decode.vector(streamId)(deletedFrame) - } yield XAutoClaimJustIdResult(cursor, claimed, deleted) - case Frame.Array(Vector(cursorFrame, claimedFrame)) => - for { - cursor <- streamId(cursorFrame) - claimed <- Decode.vector(streamId)(claimedFrame) - } yield XAutoClaimJustIdResult(cursor, claimed, Vector.empty) - case other => Left(DecodeError("xautoclaim justid [cursor, ids, deleted]", Frame.describe(other))) + private def autoClaimReply[F, V](using KeyCodec[F], ValueCodec[V]): Frame => Either[DecodeError, XAutoClaimResult[F, V]] = + autoClaim(Decode.vector(streamEntry[F, V]), "xautoclaim [cursor, entries, deleted]")(XAutoClaimResult(_, _, _)) + + private val autoClaimJustIdReply: Frame => Either[DecodeError, XAutoClaimJustIdResult] = + autoClaim(Decode.vector(streamId), "xautoclaim justid [cursor, ids, deleted]")(XAutoClaimJustIdResult(_, _, _)) + + private def autoClaim[A, R](items: Frame => Either[DecodeError, Vector[A]], label: String)( + build: (StreamId, Vector[A], Vector[StreamId]) => R + ): Frame => Either[DecodeError, R] = { + val deletedIds = Decode.vector(streamId) + Decode.shape(label) { + // pre-7.0 omits the deleted-ids element + case Frame.Array(cursorFrame +: itemsFrame +: rest) if rest.length <= 1 => + for { + cursor <- streamId(cursorFrame) + decoded <- items(itemsFrame) + deleted <- rest.headOption.fold[Either[DecodeError, Vector[StreamId]]](Right(Vector.empty))(deletedIds) + } yield build(cursor, decoded, deleted) + } } - private val deletionElement: Frame => Either[DecodeError, StreamEntryDeletion] = { + private val deletionElement: Frame => Either[DecodeError, StreamEntryDeletion] = Decode.shape("deletion status -1/1/2") { case Frame.Integer(-1L) => Right(StreamEntryDeletion.NotFound) case Frame.Integer(1L) => Right(StreamEntryDeletion.Deleted) case Frame.Integer(2L) => Right(StreamEntryDeletion.Retained) - case other => Left(DecodeError("deletion status -1/1/2", Frame.describe(other))) } - // XPENDING summary: [total, min-id, max-id, [[consumer, count], …]]; an empty group replies [0, nil, nil, nil] - private val pendingSummaryReply: Frame => Either[DecodeError, PendingSummary] = { - val consumers: Frame => Either[DecodeError, Vector[(String, Long)]] = { - case Frame.Null => Right(Vector.empty) - case Frame.Array(rows) => Decode.each(rows)(consumerCount) - case other => Left(DecodeError("consumer counts array or null", Frame.describe(other))) - } - Decode.array4(Decode.long, optionalStreamId, optionalStreamId, consumers, "xpending summary")(PendingSummary(_, _, _, _)) + // XPENDING per-consumer counts come back as bulk-string integers + private val countText: Frame => Either[DecodeError, Long] = Decode.shape("integer") { + case Frame.Integer(value) => Right(value) + case Frame.BulkString(bytes) => Primitives.decodeLong("integer", Long.MinValue, Long.MaxValue)(bytes) } private val consumerCount: Frame => Either[DecodeError, (String, Long)] = Decode.array2(Decode.utf8String, countText, "[consumer, count] pair")(_ -> _) + // XPENDING summary: [total, min-id, max-id, [[consumer, count], …]]; an empty group replies [0, nil, nil, nil] + private val pendingSummaryReply: Frame => Either[DecodeError, PendingSummary] = { + val consumers = Decode.orEmpty(Decode.vector(consumerCount, "consumer counts array or null")) + Decode.array4(Decode.long, optionalStreamId, optionalStreamId, consumers, "xpending summary")(PendingSummary(_, _, _, _)) + } + // XPENDING extended row: [id, consumer, idle-ms, delivery-count] private val pendingEntryElement: Frame => Either[DecodeError, PendingEntry] = - Decode.array4(streamId, Decode.utf8String, Decode.long, Decode.long, "xpending entry [id, consumer, idle, count]") { - (id, consumer, idle, deliveries) => - PendingEntry(id, consumer, FiniteDuration(idle, java.util.concurrent.TimeUnit.MILLISECONDS), deliveries) - } - - // XPENDING per-consumer counts come back as bulk-string integers - private def countText(frame: Frame): Either[DecodeError, Long] = - frame match { - case Frame.Integer(value) => Right(value) - case Frame.BulkString(bytes) => bytes.asUtf8String.toLongOption.toRight(DecodeError("integer", s"bulk string '${bytes.asUtf8String}'")) - case other => Left(DecodeError("integer", Frame.describe(other))) - } + Decode.array4(streamId, Decode.utf8String, Decode.millisDuration, Decode.long, "xpending entry [id, consumer, idle, count]")( + PendingEntry(_, _, _, _) + ) } diff --git a/sage-core/src/main/scala/sage/commands/Strings.scala b/sage-core/src/main/scala/sage/commands/Strings.scala index dab8d768..db6b5a4a 100644 --- a/sage-core/src/main/scala/sage/commands/Strings.scala +++ b/sage-core/src/main/scala/sage/commands/Strings.scala @@ -6,7 +6,8 @@ import scala.concurrent.duration.FiniteDuration import sage.Bytes import sage.SageException.DecodeError -import sage.codec.{Doubles, KeyCodec, ValueCodec} +import sage.codec.{KeyCodec, ValueCodec} +import sage.commands.Args.{Get, Nx, Xx} import sage.protocol.Frame /** @@ -84,9 +85,6 @@ enum IncrExpiry { private[sage] object Strings { - private val Get = Bytes.utf8("GET") - private val Nx = Bytes.utf8("NX") - private val Xx = Bytes.utf8("XX") private val KeepTtl = Bytes.utf8("KEEPTTL") private val Persist = Bytes.utf8("PERSIST") private val Len = Bytes.utf8("LEN") @@ -111,7 +109,7 @@ private[sage] object Strings { Command("DECR", Command.FirstKey, Vector(keyCodec.encode(key)), Decode.long) def decrBy[K](key: K, decrement: Long)(using keyCodec: KeyCodec[K]): Command[Long] = - Command("DECRBY", Command.FirstKey, Vector(keyCodec.encode(key), Bytes.utf8(decrement.toString)), Decode.long) + Command("DECRBY", Command.FirstKey, Vector(keyCodec.encode(key), Args.long(decrement)), Decode.long) def get[K, V](key: K)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Option[V]] = Command.read("GET", Command.FirstKey, Vector(keyCodec.encode(key)), Decode.optionalValue) @@ -126,7 +124,7 @@ private[sage] object Strings { Command.read( "GETRANGE", Command.FirstKey, - Vector(keyCodec.encode(key), Bytes.utf8(start.toString), Bytes.utf8(end.toString)), + Vector(keyCodec.encode(key), Args.long(start), Args.long(end)), Decode.value ) @@ -134,21 +132,19 @@ private[sage] object Strings { Command("INCR", Command.FirstKey, Vector(keyCodec.encode(key)), Decode.long) def incrBy[K](key: K, increment: Long)(using keyCodec: KeyCodec[K]): Command[Long] = - Command("INCRBY", Command.FirstKey, Vector(keyCodec.encode(key), Bytes.utf8(increment.toString)), Decode.long) + Command("INCRBY", Command.FirstKey, Vector(keyCodec.encode(key), Args.long(increment)), Decode.long) def incrByFloat[K](key: K, increment: Double)(using keyCodec: KeyCodec[K]): Command[Double] = - Command("INCRBYFLOAT", Command.FirstKey, Vector(keyCodec.encode(key), Bytes.utf8(Doubles.format(increment))), Decode.double) + Command("INCRBYFLOAT", Command.FirstKey, Vector(keyCodec.encode(key), Args.double(increment)), Decode.double) - def mGet[K, V](first: K, rest: K*)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Vector[Option[V]]] = { - val keys = (first +: rest).iterator.map(keyCodec.encode).toVector - Command.read("MGET", keys.indices.toVector, keys, Decode.vector(Decode.optionalValue)) - } + def mGet[K, V](first: K, rest: K*)(using KeyCodec[K], ValueCodec[V]): Command[Vector[Option[V]]] = + KeyArgs.allKeys("MGET", first +: rest.toVector, Decode.vector(Decode.optionalValue), readOnly = true) def mSet[K, V](first: (K, V), rest: (K, V)*)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Unit] = - Command("MSET", msetKeyIndices(rest.size + 1), msetArgs(first +: rest.toVector), Decode.ok) + Command("MSET", msetKeyIndices(rest.size + 1), Args.pairs(first +: rest.toVector), Decode.ok) def mSetNx[K, V](first: (K, V), rest: (K, V)*)(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Boolean] = - Command("MSETNX", msetKeyIndices(rest.size + 1), msetArgs(first +: rest.toVector), Decode.flag) + Command("MSETNX", msetKeyIndices(rest.size + 1), Args.pairs(first +: rest.toVector), Decode.flag) /** * False when `condition` made the server skip the write. @@ -163,11 +159,7 @@ private[sage] object Strings { "SET", Command.FirstKey, Vector(keyCodec.encode(key), valueCodec.encode(value)) ++ conditionArgs(condition) ++ setExpiryArgs(expiry), - decode = { - case Frame.SimpleString("OK") => Right(true) - case Frame.Null => Right(false) - case other => Left(DecodeError("simple string 'OK' or null", Frame.describe(other))) - } + Decode.okOrNull ) /** @@ -190,7 +182,7 @@ private[sage] object Strings { Command( "SETRANGE", Command.FirstKey, - Vector(keyCodec.encode(key), Bytes.utf8(offset.toString), valueCodec.encode(value)), + Vector(keyCodec.encode(key), Args.long(offset), valueCodec.encode(value)), Decode.long ) @@ -210,8 +202,8 @@ private[sage] object Strings { "LCS", Vector(0, 1), Vector(keyCodec.encode(key1), keyCodec.encode(key2), Idx) ++ - minMatchLen.toVector.flatMap(n => Vector(MinMatchLen, Bytes.utf8(n.toString))) ++ - (if (withMatchLen) Vector(WithMatchLen) else Vector.empty), + Args.optLong(MinMatchLen, minMatchLen) ++ + Args.flag(withMatchLen, WithMatchLen), lcsMatches ) @@ -229,11 +221,11 @@ private[sage] object Strings { rest: (K, V)* )(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Command[Boolean] = { val pairs = first +: rest.toVector - val data = msetArgs(pairs) + val data = Args.pairs(pairs) Command( "MSETEX", Vector.tabulate(pairs.size)(i => 1 + i * 2), - (Bytes.utf8(pairs.size.toString) +: data) ++ conditionArgs(condition) ++ setExpiryArgs(expiry), + (Args.long(pairs.size) +: data) ++ conditionArgs(condition) ++ setExpiryArgs(expiry), Decode.flag ) } @@ -249,8 +241,8 @@ private[sage] object Strings { Command( "INCREX", Command.FirstKey, - Vector(keyCodec.encode(key), ByInt, Bytes.utf8(increment.toString)) ++ - incrExArgs(saturate, lowerBound.map(_.toString), upperBound.map(_.toString), expiry), + Vector(keyCodec.encode(key), ByInt, Args.long(increment)) ++ + incrExArgs(saturate, Args.optLong(LBound, lowerBound) ++ Args.optLong(UBound, upperBound), expiry), incrExResultLong ) @@ -265,22 +257,19 @@ private[sage] object Strings { Command( "INCREX", Command.FirstKey, - Vector(keyCodec.encode(key), ByFloat, Bytes.utf8(Doubles.format(increment))) ++ - incrExArgs(saturate, lowerBound.map(Doubles.format), upperBound.map(Doubles.format), expiry), + Vector(keyCodec.encode(key), ByFloat, Args.double(increment)) ++ + incrExArgs(saturate, Args.opt(LBound, lowerBound)(Args.double) ++ Args.opt(UBound, upperBound)(Args.double), expiry), incrExResultDouble ) - private def incrExArgs(saturate: Boolean, lowerBound: Option[String], upperBound: Option[String], expiry: IncrExpiry): Vector[Bytes] = - (if (saturate) Vector(Saturate) else Vector.empty) ++ - lowerBound.toVector.flatMap(x => Vector(LBound, Bytes.utf8(x))) ++ - upperBound.toVector.flatMap(x => Vector(UBound, Bytes.utf8(x))) ++ - incrExpiryArgs(expiry) + private def incrExArgs(saturate: Boolean, bounds: Vector[Bytes], expiry: IncrExpiry): Vector[Bytes] = + Args.flag(saturate, Saturate) ++ bounds ++ incrExpiryArgs(expiry) private def incrExpiryArgs(expiry: IncrExpiry): Vector[Bytes] = expiry match { case IncrExpiry.Keep => Vector.empty - case IncrExpiry.In(duration, onlyNoTtl) => TimeArgs.relative(duration) ++ (if (onlyNoTtl) Vector(Enx) else Vector.empty) - case IncrExpiry.At(timestamp, onlyNoTtl) => TimeArgs.absolute(timestamp) ++ (if (onlyNoTtl) Vector(Enx) else Vector.empty) + case IncrExpiry.In(duration, onlyNoTtl) => TimeArgs.relative(duration) ++ Args.flag(onlyNoTtl, Enx) + case IncrExpiry.At(timestamp, onlyNoTtl) => TimeArgs.absolute(timestamp) ++ Args.flag(onlyNoTtl, Enx) case IncrExpiry.Persist => Vector(Persist) } @@ -293,50 +282,29 @@ private[sage] object Strings { case DelexCondition.IfDigestNe(hash) => Vector(IfDne, Bytes.utf8(hash)) } - private val matchRange: Frame => Either[DecodeError, MatchRange] = { + private val matchRange: Frame => Either[DecodeError, MatchRange] = Decode.shape("match range [start, end]") { case Frame.Array(Vector(Frame.Integer(start), Frame.Integer(end))) => Right(MatchRange(start, end)) - case other => Left(DecodeError("match range [start, end]", Frame.describe(other))) } - private val lcsMatch: Frame => Either[DecodeError, LcsMatch] = { - case Frame.Array(Vector(a, b)) => - for { - ar <- matchRange(a) - br <- matchRange(b) - } yield LcsMatch(ar, br, None) - case Frame.Array(Vector(a, b, Frame.Integer(len))) => - for { - ar <- matchRange(a) - br <- matchRange(b) - } yield LcsMatch(ar, br, Some(len)) - case other => Left(DecodeError("lcs match", Frame.describe(other))) + private val lcsMatch: Frame => Either[DecodeError, LcsMatch] = Decode.shape("lcs match") { + case Frame.Array(Vector(a, b)) => matchRange(a).flatMap(ar => matchRange(b).map(LcsMatch(ar, _, None))) + case Frame.Array(Vector(a, b, Frame.Integer(len))) => matchRange(a).flatMap(ar => matchRange(b).map(LcsMatch(ar, _, Some(len)))) } - private val lcsMatches: Frame => Either[DecodeError, LcsMatches] = { - case Frame.Map(entries) => - val lookup = entries.collect { case (Frame.BulkString(name), value) => name.asUtf8String -> value }.toMap + private val lcsMatches: Frame => Either[DecodeError, LcsMatches] = + Decode.fields { f => for { - matchesFrame <- lookup.get("matches").toRight(DecodeError("lcs idx 'matches' field", "map without 'matches'")) - lenFrame <- lookup.get("len").toRight(DecodeError("lcs idx 'len' field", "map without 'len'")) - matches <- Decode.vector(lcsMatch)(matchesFrame) - len <- Decode.long(lenFrame) + matches <- f.required("matches", Decode.vector(lcsMatch)) + len <- f.required("len", Decode.long) } yield LcsMatches(matches, len) - case other => Left(DecodeError("lcs idx map", Frame.describe(other))) - } + } - private val incrExResultLong: Frame => Either[DecodeError, IncrExResult[Long]] = { + private val incrExResultLong: Frame => Either[DecodeError, IncrExResult[Long]] = Decode.shape("INCREX [value, applied] integers") { case Frame.Array(Vector(Frame.Integer(value), Frame.Integer(applied))) => Right(IncrExResult(value, applied)) - case other => Left(DecodeError("INCREX [value, applied] integers", Frame.describe(other))) } - private val incrExResultDouble: Frame => Either[DecodeError, IncrExResult[Double]] = { - case Frame.Array(Vector(valueFrame, appliedFrame)) => - for { - value <- Decode.score(valueFrame) - applied <- Decode.score(appliedFrame) - } yield IncrExResult(value, applied) - case other => Left(DecodeError("INCREX [value, applied] doubles", Frame.describe(other))) - } + private val incrExResultDouble: Frame => Either[DecodeError, IncrExResult[Double]] = + Decode.array2(Decode.double, Decode.double, "INCREX [value, applied] doubles")(IncrExResult(_, _)) private def conditionArgs(condition: SetCondition): Vector[Bytes] = condition match { @@ -361,8 +329,5 @@ private[sage] object Strings { case GetExpiry.Persist => Vector(Persist) } - private def msetArgs[K, V](pairs: Vector[(K, V)])(using keyCodec: KeyCodec[K], valueCodec: ValueCodec[V]): Vector[Bytes] = - pairs.flatMap { case (key, value) => Vector(keyCodec.encode(key), valueCodec.encode(value)) } - private def msetKeyIndices(pairs: Int): Vector[Int] = Vector.tabulate(pairs)(_ * 2) } diff --git a/sage-core/src/main/scala/sage/protocol/RespParser.scala b/sage-core/src/main/scala/sage/protocol/RespParser.scala index e89e037e..d5c38a99 100644 --- a/sage-core/src/main/scala/sage/protocol/RespParser.scala +++ b/sage-core/src/main/scala/sage/protocol/RespParser.scala @@ -23,10 +23,8 @@ final private[sage] class RespParser { private var writePos: Int = 0 private var failure: ProtocolError = null - // out-fields, avoiding a result-wrapper allocation per parsed value - private var cursor: Int = 0 - private var produced: Frame = null - private var failMessage: String = null + // out-field, avoiding a result-wrapper allocation per parsed value + private var produced: Frame = null private var numberOk: Boolean = false @@ -40,28 +38,15 @@ final private[sage] class RespParser { private[protocol] def unsafeBuffer: Array[Byte] = buf /** - * Returns every frame completed by `bytes`, in order. - */ - def feed(bytes: Bytes): Either[ProtocolError, Vector[Frame]] = { - val frames = Vector.newBuilder[Frame] - val array = bytes.unsafeArray - feed(array, 0, array.length)(frames += _) match { - case Some(error) => Left(error) - case None => Right(frames.result()) - } - } - - /** - * Parses every frame completed by `array(offset until offset + length)` and passes each one to `onFrame` in order. Unlike [[feed]], this - * overload does not create an input `Bytes` value or collect the frames in a `Vector`. It copies the slice into the parser's internal - * buffer and invokes the callback as frames are completed. + * Parses every frame completed by `array(offset until offset + length)` and passes each one to `onFrame` in order. The slice is copied + * into the parser's internal buffer. */ def feed(array: Array[Byte], offset: Int, length: Int)(onFrame: Frame => Unit): Option[ProtocolError] = if (failure != null) Some(failure) else if (!append(array, offset, length)) Some(poison("input exceeds the maximum buffer size")) else { parseLoop(onFrame) - if (failMessage != null) Some(poison(failMessage)) + if (failure != null) Some(failure) else { // partial aggregates remain on the stack, so an empty input buffer can be reset between reads. if (readPos == writePos) { @@ -122,52 +107,36 @@ final private[sage] class RespParser { } private def parseLoop(onFrame: Frame => Unit): Unit = { - var running = true - while (running) { + var status = Opened + while (status == Produced || status == Opened) { // close completed aggregates first: this also finalizes an empty aggregate (zero children) the instant it is opened - var closing = true - while (closing) - if (stackDepth > 0 && complete(stack(stackDepth - 1))) { - val top = stack(stackDepth - 1) - stack(stackDepth - 1) = null - stackDepth -= 1 - if (top.kind == Attr) closing = false // an attribute yields no value; the value it prefixes is produced next, for the same slot - else { - val value = build(top) - if (stackDepth == 0) { - onFrame(value) - closing = false - } else addChild(stack(stackDepth - 1), value) // re-check: the parent may now be complete too - } - } else closing = false - - val status = produceValue() - if (status == Produced) { - val value = produced - if (stackDepth == 0) onFrame(value) else addChild(stack(stackDepth - 1), value) - } else if (status == Opened) () // re-loop: the close pass finalizes it if empty, otherwise its children are produced next - else running = false + while (stackDepth > 0 && complete(stack(stackDepth - 1))) { + stackDepth -= 1 + val top = stack(stackDepth) + stack(stackDepth) = null + // an attribute yields no value; the value it prefixes is produced next, for the same slot + if (top.kind != Attr) deliver(build(top), onFrame) + } + status = produceValue() + if (status == Produced) deliver(produced, onFrame) } } + private def deliver(value: Frame, onFrame: Frame => Unit): Unit = + if (stackDepth == 0) onFrame(value) else addChild(stack(stackDepth - 1), value) + private def complete(agg: Agg): Boolean = agg.remaining == 0 && agg.pendingKey == null private def addChild(agg: Agg, value: Frame): Unit = agg.kind match { - case Map => + case Map | Attr => if (agg.pendingKey == null) agg.pendingKey = value else { agg.pairs += ((agg.pendingKey, value)) agg.pendingKey = null agg.remaining -= 1 } - case Attr => // discarded metadata: count off each pair without materializing it - if (agg.pendingKey == null) agg.pendingKey = value - else { - agg.pendingKey = null - agg.remaining -= 1 - } - case _ => + case _ => agg.elements += value agg.remaining -= 1 } @@ -182,29 +151,21 @@ final private[sage] class RespParser { } // Produces one value at `readPos`: Produced (`produced` set, `readPos` advanced), Opened (header pushed), Incomplete (`readPos` unmoved), - // or Invalid (`failMessage` set) + // or Invalid (the parser is poisoned) private def produceValue(): Int = if (readPos >= writePos) Incomplete else { val pos = readPos - (buf(pos).toChar: @switch) match { - case '+' => - val cr = findCrlf(pos + 1) - if (cr < 0) Incomplete else leaf(cr + 2, Frame.SimpleString(stringAt(pos + 1, cr))) - case '-' => - val cr = findCrlf(pos + 1) - if (cr < 0) Incomplete else leaf(cr + 2, Frame.SimpleError(stringAt(pos + 1, cr))) - case ':' => - val cr = findCrlf(pos + 1) - if (cr < 0) Incomplete - else { + val cr = findCrlf(pos + 1) + if (cr < 0) { if (FrameTypes.indexOf(buf(pos).toInt) >= 0) Incomplete else unknownType(buf(pos)) } + else + (buf(pos).toChar: @switch) match { + case '+' => leaf(cr + 2, Frame.SimpleString(stringAt(pos + 1, cr))) + case '-' => leaf(cr + 2, Frame.SimpleError(stringAt(pos + 1, cr))) + case ':' => val value = readLong(pos + 1, cr) if (!numberOk) fail(s"invalid integer: '${stringAt(pos + 1, cr)}'") else leaf(cr + 2, Frame.Integer(value)) - } - case ',' => - val cr = findCrlf(pos + 1) - if (cr < 0) Incomplete - else { + case ',' => val text = stringAt(pos + 1, cr) try { val value = text match { @@ -215,104 +176,69 @@ final private[sage] class RespParser { } leaf(cr + 2, Frame.Double(value)) } catch { case _: NumberFormatException => fail(s"invalid double: '$text'") } - } - case '#' => - val cr = findCrlf(pos + 1) - if (cr < 0) Incomplete - else if (cr != pos + 2) fail(s"invalid boolean: '${stringAt(pos + 1, cr)}'") - else - buf(pos + 1).toChar match { - case 't' => leaf(cr + 2, Frame.Bool(true)) - case 'f' => leaf(cr + 2, Frame.Bool(false)) - case _ => fail(s"invalid boolean: '${stringAt(pos + 1, cr)}'") - } - case '(' => - val cr = findCrlf(pos + 1) - if (cr < 0) Incomplete - else { + case '#' => + if (cr != pos + 2) fail(s"invalid boolean: '${stringAt(pos + 1, cr)}'") + else + buf(pos + 1).toChar match { + case 't' => leaf(cr + 2, Frame.Bool(true)) + case 'f' => leaf(cr + 2, Frame.Bool(false)) + case _ => fail(s"invalid boolean: '${stringAt(pos + 1, cr)}'") + } + case '(' => val text = stringAt(pos + 1, cr) try leaf(cr + 2, Frame.BigNumber(BigInt(text))) catch { case _: NumberFormatException => fail(s"invalid big number: '$text'") } - } - case '_' => - val cr = findCrlf(pos + 1) - if (cr < 0) Incomplete - else if (cr != pos + 1) fail(s"unexpected content in null frame: '${stringAt(pos + 1, cr)}'") - else leaf(cr + 2, Frame.Null) - case '$' => bulk(pos, allowNull = true, "invalid bulk string length", isError = false) - case '!' => bulk(pos, allowNull = false, "invalid bulk error length", isError = true) - case '=' => verbatim(pos) - case '*' => openElements(pos, Arr, allowNull = true) - case '~' => openElements(pos, Set, allowNull = false) - case '>' => openElements(pos, Push, allowNull = false) - case '%' => openPairs(pos, Map) - case '|' => openPairs(pos, Attr) - case other => - fail(f"unknown frame type byte 0x${other.toByte}%02x") - } + case '_' => + if (cr != pos + 1) fail(s"unexpected content in null frame: '${stringAt(pos + 1, cr)}'") + else leaf(cr + 2, Frame.Null) + case '$' => bulk(pos, cr, '$', -1, "invalid bulk string length") + case '!' => bulk(pos, cr, '!', 0, "invalid bulk error length") + case '=' => bulk(pos, cr, '=', 4, "invalid verbatim string length") + case '*' => open(pos, cr, Arr, -1, "invalid array length") + case '~' => open(pos, cr, Set, 0, "invalid set length") + case '>' => open(pos, cr, Push, 0, "invalid push length") + case '%' => open(pos, cr, Map, 0, "invalid map length") + case '|' => open(pos, cr, Attr, 0, "invalid attribute length") + case other => unknownType(other.toByte) + } } + private def unknownType(byte: Byte): Int = fail(f"unknown frame type byte 0x$byte%02x") + private def leaf(end: Int, frame: Frame): Int = { readPos = end produced = frame Produced } - private def bulk(pos: Int, allowNull: Boolean, lengthError: String, isError: Boolean): Int = { - val length = readLength(pos + 1, allowNull) - if (length == Incomplete) Incomplete - else if (length == Invalid) fail(s"$lengthError: '${headerText(pos + 1)}'") - else if (length == -1) leaf(cursor, Frame.Null) + private def bulk(pos: Int, cr: Int, kind: Char, minLength: Int, lengthError: String): Int = { + val length = readLength(pos + 1, cr) + if (length < minLength) fail(s"$lengthError: '${stringAt(pos + 1, cr)}'") + else if (length == -1) leaf(cr + 2, Frame.Null) else { - val start = cursor + val start = cr + 2 val end = payloadEnd(start, length) - if (end == Incomplete) Incomplete - else if (end == Invalid) Invalid // payloadEnd set failMessage - else { - val bytes = bytesAt(start, start + length) - leaf(end, if (isError) Frame.BulkError(bytes) else Frame.BulkString(bytes)) - } - } - } - - private def verbatim(pos: Int): Int = { - val length = readLength(pos + 1, allowNull = false) - if (length == Incomplete) Incomplete - else if (length < 4) fail(s"invalid verbatim string length: '${headerText(pos + 1)}'") // an Invalid sentinel is < 4 too - else { - val start = cursor - val end = payloadEnd(start, length) - if (end == Incomplete) Incomplete - else if (end == Invalid) Invalid + if (end < 0) end // Incomplete, or Invalid after payloadEnd poisoned the parser + else if (kind == '$') leaf(end, Frame.BulkString(bytesAt(start, start + length))) + else if (kind == '!') leaf(end, Frame.BulkError(bytesAt(start, start + length))) else if (buf(start + 3) != ':') fail("verbatim string missing ':' separator") else leaf(end, Frame.VerbatimString(stringAt(start, start + 3), bytesAt(start + 4, start + length))) } } - private def openElements(pos: Int, kind: Byte, allowNull: Boolean): Int = { - val count = readLength(pos + 1, allowNull) - if (count == Incomplete) Incomplete - else if (count == Invalid) fail(s"${lengthErrorFor(kind)}: '${headerText(pos + 1)}'") - else if (count == -1) leaf(cursor, Frame.Null) - else { - readPos = cursor - val agg = new Agg(kind, count) - agg.elements = Vector.newBuilder[Frame] - agg.elements.sizeHint(count) - push(agg) - } - } - - private def openPairs(pos: Int, kind: Byte): Int = { - val count = readLength(pos + 1, allowNull = false) - if (count == Incomplete) Incomplete - else if (count == Invalid) fail(s"${lengthErrorFor(kind)}: '${headerText(pos + 1)}'") + private def open(pos: Int, cr: Int, kind: Byte, minCount: Int, lengthError: String): Int = { + val count = readLength(pos + 1, cr) + if (count < minCount) fail(s"$lengthError: '${stringAt(pos + 1, cr)}'") + else if (count == -1) leaf(cr + 2, Frame.Null) else { - readPos = cursor + readPos = cr + 2 val agg = new Agg(kind, count) - if (kind == Map) { + if (kind == Map || kind == Attr) { agg.pairs = Vector.newBuilder[(Frame, Frame)] agg.pairs.sizeHint(count) + } else { + agg.elements = Vector.newBuilder[Frame] + agg.elements.sizeHint(count) } push(agg) } @@ -333,29 +259,18 @@ final private[sage] class RespParser { } // -1 is the RESP2 null marker; '+' is signed-integer syntax that the length grammar does not permit - private def readLength(pos: Int, allowNull: Boolean): Int = { - val cr = findCrlf(pos) - if (cr < 0) Incomplete - else if (buf(pos) == '+') Invalid + private def readLength(pos: Int, cr: Int): Int = + if (buf(pos) == '+') Invalid else { val value = readLong(pos, cr) - if (!numberOk || value > Int.MaxValue || value < -1 || (value == -1 && !allowNull)) Invalid - else { - cursor = cr + 2 - value.toInt - } + if (!numberOk || value > Int.MaxValue || value < -1) Invalid else value.toInt } - } // Long arithmetic: `start + length + 2` can overflow Int private def payloadEnd(start: Int, length: Int): Int = if ((writePos - start).toLong < length.toLong + 2) Incomplete - else if (buf(start + length) != '\r' || buf(start + length + 1) != '\n') { - failMessage = "missing CRLF after bulk payload" - Invalid - } else start + length + 2 - - private def headerText(pos: Int): String = stringAt(pos, findCrlf(pos)) + else if (buf(start + length) != '\r' || buf(start + length + 1) != '\n') fail("missing CRLF after bulk payload") + else start + length + 2 // index of the next CRLF's '\r', or -1 if the input ends first private def findCrlf(from: Int): Int = { @@ -404,31 +319,24 @@ final private[sage] class RespParser { Bytes.wrap(IArray.unsafeFromArray(java.util.Arrays.copyOfRange(buf, from, until))) private def fail(message: String): Int = { - failMessage = message + poison(message) Invalid } - // Incomplete/Invalid also serve as readLength/payloadEnd sentinels, so they must stay distinct from any valid Int position/length + // Incomplete/Invalid also serve as readLength/payloadEnd sentinels, so they must stay below -1 and every valid Int position/length final private val Incomplete: Int = Int.MinValue final private val Invalid: Int = Int.MinValue + 1 final private val Produced: Int = Int.MinValue + 2 final private val Opened: Int = Int.MinValue + 3 + private val FrameTypes = "+-:,#(_$!=*~>%|" + final private val Arr: Byte = 0 final private val Set: Byte = 1 final private val Push: Byte = 2 final private val Map: Byte = 3 final private val Attr: Byte = 4 - private def lengthErrorFor(kind: Byte): String = kind match { - case Arr => "invalid array length" - case Set => "invalid set length" - case Push => "invalid push length" - case Map => "invalid map length" - case Attr => "invalid attribute length" - case _ => "invalid aggregate length" - } - // largest unconsumed input the parser will buffer (the JVM's max array size) private inline def MaxBuffer: Long = Int.MaxValue - 8 diff --git a/sage-core/src/main/scala/sage/protocol/RespWriter.scala b/sage-core/src/main/scala/sage/protocol/RespWriter.scala index 3a3221d4..12c16d8d 100644 --- a/sage-core/src/main/scala/sage/protocol/RespWriter.scala +++ b/sage-core/src/main/scala/sage/protocol/RespWriter.scala @@ -1,53 +1,36 @@ package sage.protocol import sage.Bytes +import sage.codec.Primitives.{digitCount, writeDigits} /** * Encodes commands as RESP bytes. */ private[sage] object RespWriter { - // precomputes the buffer size to encode into one array and return it without copying. def writeCommand(name: String, args: Vector[Bytes]): Bytes = - if (name.indexOf(' ') < 0) { // encode a single-word command name directly without splitting it - val nameBytes = Bytes.utf8(name) - val count = 1L + args.length - val sink = new Sink(headerSize(count) + bulkSize(nameBytes.length) + argsSize(args)) - sink.writeByte('*') - sink.writeLong(count) - sink.writeCrlf() - writeBulk(nameBytes, sink) - writeArgs(args, sink) - sink.result() - } else { - val words = name.split(' ').filter(_.nonEmpty) - val wordBytes = words.map(Bytes.utf8) - val count = wordBytes.length.toLong + args.length - var bodySize = 0 - var s = 0 - while (s < wordBytes.length) { - bodySize += bulkSize(wordBytes(s).length) - s += 1 - } - val sink = new Sink(headerSize(count) + bodySize + argsSize(args)) - sink.writeByte('*') - sink.writeLong(count) - sink.writeCrlf() - var w = 0 - while (w < wordBytes.length) { - writeBulk(wordBytes(w), sink) - w += 1 + if (name.indexOf(' ') < 0) write(Bytes.utf8(name), args) + else + // a multi-word name such as "XGROUP CREATE" is sent as one bulk string per word + name.split(' ').iterator.filter(_.nonEmpty).map(Bytes.utf8).toVector match { + case first +: rest => write(first, rest ++ args) + case _ => write(Bytes.empty, args) } - writeArgs(args, sink) - sink.result() - } - private def writeArgs(args: Vector[Bytes], sink: Sink): Unit = { - var i = 0 + // precomputes the buffer size to encode into one array and return it without copying. + private def write(name: Bytes, args: Vector[Bytes]): Bytes = { + val count = 1L + args.length + val sink = new Sink(headerSize(count) + bulkSize(name.length) + argsSize(args)) + sink.writeByte('*') + sink.writeLong(count) + sink.writeCrlf() + writeBulk(name, sink) + var i = 0 while (i < args.length) { writeBulk(args(i), sink) i += 1 } + sink.result() } // keep this calculation aligned with writeBulk and the array header written above. @@ -65,17 +48,6 @@ private[sage] object RespWriter { total } - // shared with Sink.writeDigits to keep sizing and encoding consistent. - private def digitCount(value: Long): Int = { - var digits = 1 - var ceiling = 10L - while (digits < 19 && value >= ceiling) { - digits += 1 - ceiling *= 10 - } - digits - } - private def writeBulk(value: Bytes, sink: Sink): Unit = { sink.writeByte('$') sink.writeLong(value.length.toLong) @@ -84,64 +56,34 @@ private[sage] object RespWriter { sink.writeCrlf() } - /** - * An unsynchronized growable byte buffer (java.io.ByteArrayOutputStream locks on every call). - */ - final private class Sink(initialCapacity: Int) { + final private class Sink(size: Int) { - private var buf: Array[Byte] = new Array[Byte](initialCapacity) + private val buf: Array[Byte] = new Array[Byte](size) private var len: Int = 0 def writeByte(value: Int): Unit = { - ensure(1) buf(len) = value.toByte len += 1 } def writeCrlf(): Unit = { - ensure(2) buf(len) = '\r' buf(len + 1) = '\n' len += 2 } - def writeBytes(bytes: Bytes): Unit = - writeArray(bytes.unsafeArray) - - // only ever called with non-negative lengths and element counts - def writeLong(value: Long): Unit = writeDigits(value) - - // return the buffer directly when its size is exact. If it grew beyond the encoded length, copy only the used bytes. - def result(): Bytes = - if (len == buf.length) Bytes.wrap(IArray.unsafeFromArray(buf)) - else Bytes.wrap(IArray.unsafeFromArray(java.util.Arrays.copyOf(buf, len))) - - private def writeDigits(value: Long): Unit = { - val digits = digitCount(value) - ensure(digits) - var i = len + digits - 1 - var remaining = value - while (i >= len) { - buf(i) = ('0' + (remaining % 10).toInt).toByte - remaining /= 10 - i -= 1 - } - len += digits - } - - private def writeArray(array: Array[Byte]): Unit = { - ensure(array.length) + def writeBytes(bytes: Bytes): Unit = { + val array = bytes.unsafeArray System.arraycopy(array, 0, buf, len, array.length) len += array.length } - // use Long arithmetic to prevent capacity overflow for multi-gigabyte commands, then cap at the maximum array size. - private def ensure(extra: Int): Unit = - if (buf.length - len < extra) { - val needed = len.toLong + extra - var capacity = buf.length.toLong * 2 - while (capacity < needed) capacity *= 2 - buf = java.util.Arrays.copyOf(buf, math.min(capacity, Int.MaxValue - 8).toInt) - } + def writeLong(value: Long): Unit = { + val digits = digitCount(value) + writeDigits(buf, len, digits, value) + len += digits + } + + def result(): Bytes = Bytes.wrap(IArray.unsafeFromArray(buf)) } } diff --git a/sage-core/src/main/scala/sage/ratelimit/RateLimiter.scala b/sage-core/src/main/scala/sage/ratelimit/RateLimiter.scala index 5a4449d0..cf2aac08 100644 --- a/sage-core/src/main/scala/sage/ratelimit/RateLimiter.scala +++ b/sage-core/src/main/scala/sage/ratelimit/RateLimiter.scala @@ -4,7 +4,7 @@ import scala.concurrent.duration.* import sage.Bytes import sage.SageException.DecodeError -import sage.codec.KeyCodec +import sage.codec.{KeyCodec, Primitives} import sage.commands.* /** @@ -20,13 +20,13 @@ final case class RateLimiter[K](limit: RateLimit, namespace: String = RateLimite * Attempts to consume `cost` tokens for `subject`. Returns [[Decision.Allowed]] with the remaining balance, or [[Decision.Denied]] with the * time until enough tokens become available. Returns immediately. */ - def tryAcquire(subject: K, cost: Long = 1): Command[Decision] = eval(RateLimiter.Invocation.Eval, subject, cost, peek = false) + def tryAcquire(subject: K, cost: Long = 1): Command[Decision] = eval(cached = false, subject, cost, peek = false) /** * Checks `subject` without consuming a token. Elapsed time still refills the bucket. Returns [[Decision.Allowed]] when a token is available, * or [[Decision.Denied]] with the time until one becomes available. */ - def peek(subject: K): Command[Decision] = eval(RateLimiter.Invocation.Eval, subject, cost = 1, peek = true) + def peek(subject: K): Command[Decision] = eval(cached = false, subject, cost = 1, peek = true) /** * Clears `subject`'s bucket. Its next request starts with full capacity. @@ -34,18 +34,12 @@ final case class RateLimiter[K](limit: RateLimit, namespace: String = RateLimite def reset(subject: K): Command[Unit] = Command("DEL", Command.FirstKey, Vector(keyBytes(subject)), _ => Right(())) - private val capacityText = limit.capacity.toString - private val refillTokensText = limit.refillTokens.toString private val refillPeriodMicros = limit.refillPeriodMicros - private val refillPeriodText = refillPeriodMicros.toString - private val capacityArgument = Bytes.utf8(capacityText) - private val refillArgument = Bytes.utf8(refillTokensText) - private val periodArgument = Bytes.utf8(refillPeriodText) - private val policySignatureArgument = Bytes.utf8(s"$capacityText:$refillTokensText:$refillPeriodText") - private val keyPrefix = { - val ns = Bytes.utf8(namespace) - Bytes.concat(Vector(Bytes.utf8(s"${ns.length}:"), ns, Bytes.utf8(":"))) - } + private val capacityArgument = Primitives.encodeLong(limit.capacity) + private val refillArgument = Primitives.encodeLong(limit.refillTokens) + private val periodArgument = Primitives.encodeLong(refillPeriodMicros) + private val policySignatureArgument = Bytes.utf8(s"${limit.capacity}:${limit.refillTokens}:$refillPeriodMicros") + private val namespaced = SingleKeyScript.namespaced(namespace) private val policyProblem = if (limit.capacity <= 0) Some("capacity must be > 0") else if (limit.refillTokens <= 0) Some("refillTokens must be > 0") @@ -66,29 +60,19 @@ final case class RateLimiter[K](limit: RateLimit, namespace: String = RateLimite else None } - private[sage] def evalSha(subject: K, cost: Long, peek: Boolean = false): Command[Decision] = - eval(RateLimiter.Invocation.EvalSha, subject, cost, peek) - // test-only entry point that supplies the script clock explicitly. private[sage] def tryAcquireAt(subject: K, cost: Long, nowMicros: Long): Command[Decision] = - eval(RateLimiter.Invocation.Eval, subject, cost, peek = false, Some(nowMicros)) + eval(cached = false, subject, cost, peek = false, Some(nowMicros)) - // NOSCRIPT recovery for evalSha. Sending the body runs the check and caches the script on that node. - private[sage] def evalScript(subject: K, cost: Long, peek: Boolean): Command[Decision] = - eval(RateLimiter.Invocation.Eval, subject, cost, peek) + private def keyBytes(subject: K): Bytes = namespaced(keyCodec.encode(subject)) - // length framing distinguishes namespace `a` with subject `b:c` from namespace `a:b` with subject `c`. - private def keyBytes(subject: K): Bytes = Bytes.concat(Vector(keyPrefix, keyCodec.encode(subject))) - - private def eval(invocation: RateLimiter.Invocation, subject: K, cost: Long, peek: Boolean, now: Option[Long] = None): Command[Decision] = { - val costArgument = if (cost == 1L) RateLimiter.defaultCostArgument else Bytes.utf8(cost.toString) - val nowArgument = now match { - case None => Bytes.empty // an empty injected time uses the server's TIME value - case Some(value) => Bytes.utf8(value.toString) - } + // a cached call sends the digest; NOSCRIPT recovery sends the body, which runs the check and caches the script on that node. + private[sage] def eval(cached: Boolean, subject: K, cost: Long, peek: Boolean, now: Option[Long] = None): Command[Decision] = { + val costArgument = Primitives.encodeLong(cost) + val nowArgument = now.fold(Bytes.empty)(Primitives.encodeLong) // an empty injected time uses the server's TIME value val allArgs = Vector( - invocation.scriptReference, - RateLimiter.oneKeyArgument, + RateLimiter.compiled.reference(cached), + SingleKeyScript.NumKeys, keyBytes(subject), capacityArgument, refillArgument, @@ -96,9 +80,9 @@ final case class RateLimiter[K](limit: RateLimit, namespace: String = RateLimite policySignatureArgument, costArgument, nowArgument, - if (peek) RateLimiter.peekArgument else Bytes.empty + if (peek) Primitives.encodeLong(1L) else Bytes.empty ) - Command(invocation.verb, RateLimiter.scriptKeyIndices, allArgs, RateLimiter.decode, Execution.Ordinary) + Command(RateLimiter.compiled.verb(cached), SingleKeyScript.KeyIndices, allArgs, RateLimiter.decode) } } @@ -217,10 +201,7 @@ object RateLimiter { |return { allowed, tokens, timed_catchup, retry_wait, reset_wait } |""".stripMargin - private[sage] val sha: String = { - val digest = java.security.MessageDigest.getInstance("SHA-1").digest(script.getBytes(java.nio.charset.StandardCharsets.UTF_8)) - digest.iterator.map(b => f"${b & 0xff}%02x").mkString - } + private[sage] val compiled = SingleKeyScript(script) /** * The greatest retry/reset duration returned; a wait beyond it (severe clock rollback) saturates here, while the TTL still covers it. @@ -254,14 +235,4 @@ object RateLimiter { // Lua numbers are IEEE doubles, so integers (and capacity * refillPeriod) are held to 2^53 to stay exact private[sage] val maxExactInt: Long = 1L << 53 - - private val scriptKeyIndices: Vector[Int] = Vector(2) // 0 = script/sha, 1 = numkeys, 2 = the single key - private val oneKeyArgument: Bytes = Bytes.utf8("1") - private val defaultCostArgument: Bytes = Bytes.utf8("1") - private val peekArgument: Bytes = Bytes.utf8("1") - - private enum Invocation(val verb: String, val scriptReference: Bytes) { - case Eval extends Invocation("EVAL", Bytes.utf8(script)) - case EvalSha extends Invocation("EVALSHA", Bytes.utf8(sha)) - } } diff --git a/sage-core/src/test/scala/sage/cluster/RedirectSpec.scala b/sage-core/src/test/scala/sage/cluster/RedirectSpec.scala index 66fc9041..2de63b76 100644 --- a/sage-core/src/test/scala/sage/cluster/RedirectSpec.scala +++ b/sage-core/src/test/scala/sage/cluster/RedirectSpec.scala @@ -4,44 +4,39 @@ class RedirectSpec extends munit.FunSuite { test("MOVED parses into a permanent redirect") { assertEquals( - Redirect.parse("MOVED 3999 127.0.0.1:6381"), - Some(Redirect(RedirectKind.Moved, Slot.unsafe(3999), Node("127.0.0.1", 6381))) + Redirect.parse(RedirectKind.Moved, "3999 127.0.0.1:6381"), + Some(Redirect(RedirectKind.Moved, Node("127.0.0.1", 6381))) ) } test("ASK parses into a one-shot redirect") { assertEquals( - Redirect.parse("ASK 3999 127.0.0.1:6381"), - Some(Redirect(RedirectKind.Ask, Slot.unsafe(3999), Node("127.0.0.1", 6381))) + Redirect.parse(RedirectKind.Ask, "3999 127.0.0.1:6381"), + Some(Redirect(RedirectKind.Ask, Node("127.0.0.1", 6381))) ) } - test("an empty host is preserved for the runtime to resolve") { - assertEquals(Redirect.parse("MOVED 3999 :6381"), Some(Redirect(RedirectKind.Moved, Slot.unsafe(3999), Node("", 6381)))) + test("an empty host targets the node that sent the redirect") { + val from = Node("10.0.0.5", 7000) + assertEquals(Redirect.parse(RedirectKind.Moved, "3999 :6381").map(_.target(from)), Some(Node("10.0.0.5", 6381))) + assertEquals(Redirect.parse(RedirectKind.Moved, "3999 127.0.0.1:6381").map(_.target(from)), Some(Node("127.0.0.1", 6381))) } test("an IPv6 host keeps its colons, port taken after the last") { assertEquals( - Redirect.parse("MOVED 1 2001:db8::1:6379"), - Some(Redirect(RedirectKind.Moved, Slot.unsafe(1), Node("2001:db8::1", 6379))) + Redirect.parse(RedirectKind.Moved, "1 2001:db8::1:6379"), + Some(Redirect(RedirectKind.Moved, Node("2001:db8::1", 6379))) ) } - test("a non-redirect error is not a redirect") { - assertEquals(Redirect.parse("WRONGTYPE Operation against a key holding the wrong kind of value"), None) - assertEquals(Redirect.parse("ERR unknown command"), None) - } - test("malformed redirects parse to None") { - assertEquals(Redirect.parse("MOVED 3999"), None) - assertEquals(Redirect.parse("MOVED notaslot 127.0.0.1:6381"), None) - assertEquals(Redirect.parse("MOVED 3999 127.0.0.1:notaport"), None) - assertEquals(Redirect.parse("MOVED 99999 127.0.0.1:6381"), None) + assertEquals(Redirect.parse(RedirectKind.Moved, "3999"), None) + assertEquals(Redirect.parse(RedirectKind.Moved, "3999 127.0.0.1:notaport"), None) } test("an out-of-range port is rejected") { - assertEquals(Redirect.parse("MOVED 1 127.0.0.1:0"), None) - assertEquals(Redirect.parse("MOVED 1 127.0.0.1:-1"), None) - assertEquals(Redirect.parse("MOVED 1 127.0.0.1:70000"), None) + assertEquals(Redirect.parse(RedirectKind.Moved, "1 127.0.0.1:0"), None) + assertEquals(Redirect.parse(RedirectKind.Moved, "1 127.0.0.1:-1"), None) + assertEquals(Redirect.parse(RedirectKind.Moved, "1 127.0.0.1:70000"), None) } } diff --git a/sage-core/src/test/scala/sage/cluster/SplitPlanSpec.scala b/sage-core/src/test/scala/sage/cluster/SplitPlanSpec.scala index 8b19e1bf..6496c8fb 100644 --- a/sage-core/src/test/scala/sage/cluster/SplitPlanSpec.scala +++ b/sage-core/src/test/scala/sage/cluster/SplitPlanSpec.scala @@ -2,7 +2,7 @@ package sage.cluster import sage.Bytes import sage.cluster.TopologyFixtures.{keyed, keyless} -import sage.commands.{Command, Pipeline} +import sage.commands.Command class SplitPlanSpec extends munit.FunSuite { @@ -12,43 +12,44 @@ class SplitPlanSpec extends munit.FunSuite { private val sFoo = Slot.of(Bytes.utf8("foo")) private val sBar = Slot.of(Bytes.utf8("bar")) - private val twoNode = ClusterTopology.from( - Vector( - Shard(a, Vector.empty, Vector(SlotRange(sFoo, sFoo))), - Shard(b, Vector.empty, Vector(SlotRange(sBar, sBar))) - ) - ) + private val shardA = Shard(a, Vector.empty) + private val shardB = Shard(b, Vector.empty) + private val rangeA = SlotRange(sFoo, sFoo, a, Vector.empty) + private val twoNode = ClusterTopology.from(Vector(rangeA, SlotRange(sBar, sBar, b, Vector.empty))) test("commands group per node, positions kept in submission order") { - val plan = twoNode.split(Pipeline.sequence(Seq(keyed("foo"), keyed("bar"), keyed("foo")))) - assertEquals(plan.perNode, Vector(NodeGroup(a, Vector(0, 2)), NodeGroup(b, Vector(1)))) - assertEquals(plan.keyless, Vector.empty) - assertEquals(plan.rejected, Vector.empty) + val plan = twoNode.split(Vector(keyed("foo"), keyed("bar"), keyed("foo"))) + assertEquals(plan.perNode, Vector(NodeGroup(shardA, Vector(0, 2)), NodeGroup(shardB, Vector(1)))) } - test("keyless and rejected positions are partitioned out, every index placed once") { - val plan = twoNode.split(Pipeline.sequence(Seq(keyed("foo"), keyless, keyed("foo", "bar"), keyed("bar")))) - assertEquals(plan.perNode, Vector(NodeGroup(a, Vector(0)), NodeGroup(b, Vector(3)))) - assertEquals(plan.keyless, Vector(1)) - assertEquals(plan.rejected, Vector(2 -> Rejected.CrossSlot(Set(sFoo, sBar)))) + test("keyless positions join the first group in submission order, rejected positions are left out") { + val plan = twoNode.split(Vector(keyless, keyed("bar"), keyless, keyed("foo", "bar"), keyed("foo"))) + assertEquals(plan.perNode, Vector(NodeGroup(shardB, Vector(0, 1, 2)), NodeGroup(shardA, Vector(4)))) + assertEquals(plan.routes(3), Route.CrossSlot) + } + + test("keyless positions stay out of every group when no command has a key") { + val plan = twoNode.split(Vector(keyless, keyless)) + assertEquals(plan.perNode, Vector.empty) + assertEquals(plan.routes, Vector(Route.Keyless, Route.Keyless)) } test("an uncovered command is rejected as unowned, not dropped") { - val onlyA = ClusterTopology.from(Vector(Shard(a, Vector.empty, Vector(SlotRange(sFoo, sFoo))))) - val plan = onlyA.split(Pipeline.sequence(Seq(keyed("foo"), keyed("bar")))) - assertEquals(plan.perNode, Vector(NodeGroup(a, Vector(0)))) - assertEquals(plan.rejected, Vector(1 -> Rejected.Unowned(sBar))) + val onlyA = ClusterTopology.from(Vector(rangeA)) + val plan = onlyA.split(Vector(keyed("foo"), keyed("bar"))) + assertEquals(plan.perNode, Vector(NodeGroup(shardA, Vector(0)))) + assertEquals(plan.routes(1), Route.Unowned(sBar)) } test("a malformed command is rejected in place, not routed") { val malformed = Command("BAD", Vector(5), Vector(Bytes.utf8("k")), _ => Right(0L)) - val plan = twoNode.split(Pipeline.sequence(Seq(keyed("foo"), malformed))) - assertEquals(plan.perNode, Vector(NodeGroup(a, Vector(0)))) - assertEquals(plan.rejected, Vector(1 -> Rejected.Malformed)) + val plan = twoNode.split(Vector(keyed("foo"), malformed)) + assertEquals(plan.perNode, Vector(NodeGroup(shardA, Vector(0)))) + assertEquals(plan.routes(1), Route.Malformed) } test("an empty pipeline yields an empty plan") { - val plan = twoNode.split(Pipeline.sequence(Seq.empty[Command[Long]])) - assertEquals(plan, SplitPlan(Vector.empty, Vector.empty, Vector.empty)) + val plan = twoNode.split(Vector.empty) + assertEquals(plan, SplitPlan(Vector.empty, Vector.empty)) } } diff --git a/sage-core/src/test/scala/sage/cluster/TopologyFixtures.scala b/sage-core/src/test/scala/sage/cluster/TopologyFixtures.scala index 96fb77c4..71d066fb 100644 --- a/sage-core/src/test/scala/sage/cluster/TopologyFixtures.scala +++ b/sage-core/src/test/scala/sage/cluster/TopologyFixtures.scala @@ -13,6 +13,6 @@ object TopologyFixtures { val keyless: Command[Long] = Command("PING", Command.NoKeys, Vector.empty, _ => Right(0L)) - def covering(node: Node, from: Int, to: Int): Shard = - Shard(node, Vector.empty, Vector(SlotRange(Slot.unsafe(from), Slot.unsafe(to)))) + def covering(node: Node, from: Int, to: Int): SlotRange = + SlotRange(Slot.at(from).get, Slot.at(to).get, node, Vector.empty) } diff --git a/sage-core/src/test/scala/sage/cluster/TopologySpec.scala b/sage-core/src/test/scala/sage/cluster/TopologySpec.scala index f1b52df6..3725b8e3 100644 --- a/sage-core/src/test/scala/sage/cluster/TopologySpec.scala +++ b/sage-core/src/test/scala/sage/cluster/TopologySpec.scala @@ -9,35 +9,47 @@ class TopologySpec extends munit.FunSuite { private val a = Node("a", 6379) private val b = Node("b", 6379) - private val whole = ClusterTopology.from(Vector(covering(a, 0, Slot.Count - 1))) + private val wholeShard = Shard(a, Vector.empty) + private val whole = ClusterTopology.from(Vector(covering(a, 0, Slot.Count - 1))) test("nodeForSlot returns the owning master, None for an uncovered slot") { val partial = ClusterTopology.from(Vector(covering(a, 0, 100))) - assertEquals(partial.nodeForSlot(Slot.unsafe(50)), Some(a)) - assertEquals(partial.nodeForSlot(Slot.unsafe(200)), None) + assertEquals(partial.nodeForSlot(Slot.at(50).get), Some(a)) + assertEquals(partial.nodeForSlot(Slot.at(200).get), None) } test("overlapping ranges resolve last-listed-wins") { val topo = ClusterTopology.from(Vector(covering(a, 0, 10), covering(b, 5, 20))) - assertEquals(topo.nodeForSlot(Slot.unsafe(2)), Some(a)) - assertEquals(topo.nodeForSlot(Slot.unsafe(7)), Some(b)) - assertEquals(topo.nodeForSlot(Slot.unsafe(15)), Some(b)) + assertEquals(topo.nodeForSlot(Slot.at(2).get), Some(a)) + assertEquals(topo.nodeForSlot(Slot.at(7).get), Some(b)) + assertEquals(topo.nodeForSlot(Slot.at(15).get), Some(b)) } - test("shardForSlot exposes the owning shard's replicas, None for an uncovered slot") { + test("masters lists, in slot order, only the masters that route owns a slot for") { + val c = Node("c", 6379) + assertEquals(ClusterTopology.from(Vector(covering(a, 0, 10), covering(b, 11, 20))).masters, Vector(a, b)) + assertEquals(ClusterTopology.from(Vector(covering(a, 0, 100), covering(b, 0, 100))).masters, Vector(b)) + assertEquals(ClusterTopology.from(Vector(covering(c, 100, 50), covering(a, 0, 10))).masters, Vector(a)) + } + + test("a routed command carries the owning shard's replicas") { val r1 = Node("r1", 6379) - val shard = Shard(a, Vector(r1), Vector(SlotRange(Slot.unsafe(0), Slot.unsafe(100)))) - val topo = ClusterTopology.from(Vector(shard)) - assertEquals(topo.shardForSlot(Slot.unsafe(50)).map(_.replicas), Some(Vector(r1))) - assertEquals(topo.shardForSlot(Slot.unsafe(200)), None) + val range = SlotRange(Slot.at(0).get, Slot.at(Slot.Count - 1).get, a, Vector(r1)) + assertEquals(ClusterTopology.from(Vector(range)).route(keyed("foo")), Route.ToNode(Shard(a, Vector(r1)), Slot.of(Bytes.utf8("foo")))) + } + + test("ranges served by the same master form one shard listing the replicas of every range") { + val (r1, r2) = (Node("r1", 6379), Node("r2", 6379)) + val ranges = Vector(covering(a, 0, 10).copy(replicas = Vector(r1)), covering(a, 100, 110).copy(replicas = Vector(r2, r1))) + assertEquals(ClusterTopology.from(ranges).shards, Vector(Shard(a, Vector(r1, r2)))) } test("a single-slot command routes to its owner") { - assertEquals(whole.route(keyed("foo")), Route.ToNode(a, Slot.of(Bytes.utf8("foo")))) + assertEquals(whole.route(keyed("foo")), Route.ToNode(wholeShard, Slot.of(Bytes.utf8("foo")))) } test("a multi-key command sharing a slot routes to one node") { - assertEquals(whole.route(keyed("{tag}.a", "{tag}.b")), Route.ToNode(a, Slot.of(Bytes.utf8("tag")))) + assertEquals(whole.route(keyed("{tag}.a", "{tag}.b")), Route.ToNode(wholeShard, Slot.of(Bytes.utf8("tag")))) } test("a keyless command routes to any node") { @@ -48,7 +60,7 @@ class TopologySpec extends munit.FunSuite { val sFoo = Slot.of(Bytes.utf8("foo")) val sBar = Slot.of(Bytes.utf8("bar")) assertNotEquals(sFoo, sBar) - assertEquals(whole.route(keyed("foo", "bar")), Route.CrossSlot(Set(sFoo, sBar))) + assertEquals(whole.route(keyed("foo", "bar")), Route.CrossSlot) } test("a command on an uncovered slot is classified unowned") { diff --git a/sage-core/src/test/scala/sage/codec/CodecSpec.scala b/sage-core/src/test/scala/sage/codec/CodecSpec.scala index 0adc2011..8efe24cc 100644 --- a/sage-core/src/test/scala/sage/codec/CodecSpec.scala +++ b/sage-core/src/test/scala/sage/codec/CodecSpec.scala @@ -75,6 +75,19 @@ class CodecSpec extends munit.FunSuite { assertEquals(summon[ValueCodec[Long]].decode(Bytes.utf8("9223372036854775808")), Left(DecodeError("Long", "'9223372036854775808'"))) } + test("Int and Long decoding accepts ASCII decimal digits with an optional '-'") { + for (text <- List("+5", "\u0665", "\uff15", "", "-", " 5", "5 ")) { + assert(summon[KeyCodec[Int]].decode(Bytes.utf8(text)).isLeft, text) + assert(summon[ValueCodec[Long]].decode(Bytes.utf8(text)).isLeft, text) + } + for ((text, expected) <- List("05" -> 5, "-0" -> 0, "-007" -> -7, "-5" -> -5)) { + assertEquals(summon[ValueCodec[Int]].decode(Bytes.utf8(text)), Right(expected), text) + assertEquals(summon[KeyCodec[Long]].decode(Bytes.utf8(text)), Right(expected.toLong), text) + } + assertEquals(summon[ValueCodec[Int]].decode(Bytes.utf8("2147483648")), Left(DecodeError("Int", "'2147483648'"))) + assertEquals(summon[ValueCodec[Long]].decode(Bytes.utf8("-9223372036854775809")), Left(DecodeError("Long", "'-9223372036854775809'"))) + } + test("String decode rejects invalid UTF-8") { val invalid = Bytes.fromArray(Array(0xff.toByte, 0xfe.toByte)) summon[ValueCodec[String]].decode(invalid) match { diff --git a/sage-core/src/test/scala/sage/commands/ArraysSpec.scala b/sage-core/src/test/scala/sage/commands/ArraysSpec.scala index 9eaa7c97..0c03bcfa 100644 --- a/sage-core/src/test/scala/sage/commands/ArraysSpec.scala +++ b/sage-core/src/test/scala/sage/commands/ArraysSpec.scala @@ -6,55 +6,58 @@ import sage.protocol.Frames.{bulk, map} class ArraysSpec extends munit.FunSuite { test("ARGET decodes a value and a null") { - assertEquals(Reply.run(Arrays.arGet[String, String]("a", 1L), bulk("b")), Right(Some("b"))) - assertEquals(Reply.run(Arrays.arGet[String, String]("a", 1L), Frame.Null), Right(None)) + assertEquals(Reply.decode(Arrays.arGet[String, String]("a", 1L), bulk("b")).toEither, Right(Some("b"))) + assertEquals(Reply.decode(Arrays.arGet[String, String]("a", 1L), Frame.Null).toEither, Right(None)) } test("ARMGET and ARGETRANGE keep nils for empty slots") { val reply = Frame.Array(Vector(bulk("x"), Frame.Null, bulk("y"))) - assertEquals(Reply.run(Arrays.arMGet[String, String]("a", 0L, 1L, 2L), reply), Right(Vector(Some("x"), None, Some("y")))) - assertEquals(Reply.run(Arrays.arGetRange[String, String]("a", 0L, 2L), reply), Right(Vector(Some("x"), None, Some("y")))) + assertEquals(Reply.decode(Arrays.arMGet[String, String]("a", 0L, 1L, 2L), reply).toEither, Right(Vector(Some("x"), None, Some("y")))) + assertEquals(Reply.decode(Arrays.arGetRange[String, String]("a", 0L, 2L), reply).toEither, Right(Vector(Some("x"), None, Some("y")))) } test("ARLASTITEMS decodes the items in order") { val reply = Frame.Array(Vector(bulk("d"), bulk("e"))) - assertEquals(Reply.run(Arrays.arLastItems[String, String]("a", 2L), reply), Right(Vector("d", "e"))) + assertEquals(Reply.decode(Arrays.arLastItems[String, String]("a", 2L), reply).toEither, Right(Vector("d", "e"))) } test("ARNEXT decodes the next index and null when exhausted") { - assertEquals(Reply.run(Arrays.arNext("a"), Frame.Integer(2L)), Right(Some(2L))) - assertEquals(Reply.run(Arrays.arNext("a"), Frame.Null), Right(None)) + assertEquals(Reply.decode(Arrays.arNext("a"), Frame.Integer(2L)).toEither, Right(Some(2L))) + assertEquals(Reply.decode(Arrays.arNext("a"), Frame.Null).toEither, Right(None)) } test("ARSEEK decodes the cursor-set flag") { - assertEquals(Reply.run(Arrays.arSeek("a", 5L), Frame.Integer(1L)), Right(true)) - assertEquals(Reply.run(Arrays.arSeek("a", 5L), Frame.Integer(0L)), Right(false)) + assertEquals(Reply.decode(Arrays.arSeek("a", 5L), Frame.Integer(1L)).toEither, Right(true)) + assertEquals(Reply.decode(Arrays.arSeek("a", 5L), Frame.Integer(0L)).toEither, Right(false)) } test("ARSCAN and ARGREP WITHVALUES decode an array of [index, value] pairs") { val reply = Frame.Array( Vector(Frame.Array(Vector(Frame.Integer(0L), bulk("a"))), Frame.Array(Vector(Frame.Integer(5L), bulk("f")))) ) - assertEquals(Reply.run(Arrays.arScan[String, String]("a", 0L, 10L), reply), Right(Vector(0L -> "a", 5L -> "f"))) - assertEquals(Reply.run(Arrays.arGrepWithValues[String, String]("a", 0L, 10L)(ArMatch.Glob("*")), reply), Right(Vector(0L -> "a", 5L -> "f"))) + assertEquals(Reply.decode(Arrays.arScan[String, String]("a", 0L, 10L), reply).toEither, Right(Vector(0L -> "a", 5L -> "f"))) + assertEquals( + Reply.decode(Arrays.arGrepWithValues[String, String]("a", 0L, 10L)(ArMatch.Glob("*")), reply).toEither, + Right(Vector(0L -> "a", 5L -> "f")) + ) } test("ARGREP decodes matching indices") { val reply = Frame.Array(Vector(Frame.Integer(0L), Frame.Integer(2L))) - assertEquals(Reply.run(Arrays.arGrep("a", 0L, 10L)(ArMatch.Glob("ap*")), reply), Right(Vector(0L, 2L))) + assertEquals(Reply.decode(Arrays.arGrep("a", 0L, 10L)(ArMatch.Glob("ap*")), reply).toEither, Right(Vector(0L, 2L))) } test("AROP SUM/MIN/MAX decode a numeric bulk string or null") { - assertEquals(Reply.run(Arrays.arOpSum("a", 0L, 2L), bulk("60")), Right(Some(60.0))) - assertEquals(Reply.run(Arrays.arOpMin("a", 0L, 2L), bulk("10")), Right(Some(10.0))) - assertEquals(Reply.run(Arrays.arOpMax("a", 0L, 2L), Frame.Null), Right(None)) + assertEquals(Reply.decode(Arrays.arOpSum("a", 0L, 2L), bulk("60")).toEither, Right(Some(60.0))) + assertEquals(Reply.decode(Arrays.arOpMin("a", 0L, 2L), bulk("10")).toEither, Right(Some(10.0))) + assertEquals(Reply.decode(Arrays.arOpMax("a", 0L, 2L), Frame.Null).toEither, Right(None)) } test("AROP AND/OR/XOR decode an integer or null, MATCH/USED an integer") { - assertEquals(Reply.run(Arrays.arOpAnd("a", 0L, 2L), Frame.Integer(0L)), Right(Some(0L))) - assertEquals(Reply.run(Arrays.arOpXor("a", 0L, 2L), Frame.Null), Right(None)) - assertEquals(Reply.run(Arrays.arOpUsed("a", 0L, 2L), Frame.Integer(3L)), Right(3L)) - assertEquals(Reply.run(Arrays.arOpMatch("a", 0L, 2L, "v"), Frame.Integer(1L)), Right(1L)) + assertEquals(Reply.decode(Arrays.arOpAnd("a", 0L, 2L), Frame.Integer(0L)).toEither, Right(Some(0L))) + assertEquals(Reply.decode(Arrays.arOpXor("a", 0L, 2L), Frame.Null).toEither, Right(None)) + assertEquals(Reply.decode(Arrays.arOpUsed("a", 0L, 2L), Frame.Integer(3L)).toEither, Right(3L)) + assertEquals(Reply.decode(Arrays.arOpMatch("a", 0L, 2L, "v"), Frame.Integer(1L)).toEither, Right(1L)) } test("ARINFO decodes the core fields and leniently fills the structural ones") { @@ -66,7 +69,7 @@ class ArraysSpec extends munit.FunSuite { "slice-size" -> Frame.Integer(4096L) ) assertEquals( - Reply.run(Arrays.arInfo("a"), reply), + Reply.decode(Arrays.arInfo("a"), reply).toEither, Right(ArrayInfo(4L, 101L, 0L, slices = Some(1L), directorySize = None, superDirEntries = None, sliceSize = Some(4096L))) ) } @@ -79,7 +82,7 @@ class ArraysSpec extends munit.FunSuite { "sparse-slices" -> Frame.Integer(1L), "avg-sparse-size" -> Frame.Double(4.0) ) - Reply.run(Arrays.arInfoFull("a"), reply) match { + Reply.decode(Arrays.arInfoFull("a"), reply).toEither match { case Right(info) => assertEquals(info.count, 4L) assertEquals(info.sparseSlices, Some(1L)) @@ -90,6 +93,6 @@ class ArraysSpec extends munit.FunSuite { } test("ARINFO fails when a core field is absent") { - assert(Reply.run(Arrays.arInfo("a"), map("len" -> Frame.Integer(1L))).isLeft) + assert(Reply.decode(Arrays.arInfo("a"), map("len" -> Frame.Integer(1L))).toEither.isLeft) } } diff --git a/sage-core/src/test/scala/sage/commands/ClusterSpec.scala b/sage-core/src/test/scala/sage/commands/ClusterSpec.scala index 9f432465..53dfa620 100644 --- a/sage-core/src/test/scala/sage/commands/ClusterSpec.scala +++ b/sage-core/src/test/scala/sage/commands/ClusterSpec.scala @@ -1,7 +1,7 @@ package sage.commands import sage.SageException.DecodeError -import sage.cluster.{Node, Shard, Slot, SlotRange} +import sage.cluster.{Node, Slot, SlotRange} import sage.protocol.Frame import sage.protocol.Frames.bulk @@ -11,9 +11,11 @@ class ClusterSpec extends munit.FunSuite { private def node(host: String, port: Int, id: String): Frame = Frame.Array(Vector(bulk(host), int(port.toLong), bulk(id))) - private def run(frame: Frame): Either[DecodeError, Vector[Shard]] = Cluster.slots.decode(frame) + private val queried = Node("10.9.9.9", 7000) - test("decodes a range into a Shard with master and replicas") { + private def run(frame: Frame): Either[DecodeError, Vector[SlotRange]] = Cluster.slots(queried).decode(frame) + + test("decodes a range with its master and replicas") { val reply = Frame.Array( Vector( Frame.Array(Vector(int(0), int(5460), node("10.0.0.1", 6379, "m1"), node("10.0.0.2", 6379, "r1"))) @@ -21,11 +23,11 @@ class ClusterSpec extends munit.FunSuite { ) assertEquals( run(reply), - Right(Vector(Shard(Node("10.0.0.1", 6379), Vector(Node("10.0.0.2", 6379)), Vector(SlotRange(Slot.unsafe(0), Slot.unsafe(5460)))))) + Right(Vector(SlotRange(Slot.at(0).get, Slot.at(5460).get, Node("10.0.0.1", 6379), Vector(Node("10.0.0.2", 6379))))) ) } - test("merges multiple ranges owned by the same master into one Shard") { + test("keeps each range of the same master as listed") { val master = node("10.0.0.1", 6379, "m1") val reply = Frame.Array( Vector( @@ -37,11 +39,8 @@ class ClusterSpec extends munit.FunSuite { run(reply), Right( Vector( - Shard( - Node("10.0.0.1", 6379), - Vector.empty, - Vector(SlotRange(Slot.unsafe(0), Slot.unsafe(10)), SlotRange(Slot.unsafe(100), Slot.unsafe(110))) - ) + SlotRange(Slot.at(0).get, Slot.at(10).get, Node("10.0.0.1", 6379), Vector.empty), + SlotRange(Slot.at(100).get, Slot.at(110).get, Node("10.0.0.1", 6379), Vector.empty) ) ) ) @@ -66,10 +65,14 @@ class ClusterSpec extends munit.FunSuite { assert(run(reply).isLeft) } - test("a null endpoint decodes to the empty host the caller substitutes itself into") { - val master = Frame.Array(Vector(Frame.Null, int(6379), bulk("m1"))) - val reply = Frame.Array(Vector(Frame.Array(Vector(int(0), int(10), master)))) - assertEquals(run(reply).map(_.map(_.master)), Right(Vector(Node("", 6379)))) + test("a null or empty endpoint decodes to the host of the queried node") { + val master = Frame.Array(Vector(Frame.Null, int(6379), bulk("m1"))) + val replica = node("", 6380, "r1") + val reply = Frame.Array(Vector(Frame.Array(Vector(int(0), int(10), master, replica)))) + assertEquals( + run(reply).map(_.map(range => range.master +: range.replicas)), + Right(Vector(Vector(Node("10.9.9.9", 6379), Node("10.9.9.9", 6380)))) + ) } test("a `?` endpoint stays literal: it means an unknown node, not the queried one") { diff --git a/sage-core/src/test/scala/sage/commands/CommandSpec.scala b/sage-core/src/test/scala/sage/commands/CommandSpec.scala index 45cf99c7..ec757815 100644 --- a/sage-core/src/test/scala/sage/commands/CommandSpec.scala +++ b/sage-core/src/test/scala/sage/commands/CommandSpec.scala @@ -3,7 +3,7 @@ package sage.commands import sage.Bytes import sage.SageException.{DecodeError, ServerError} import sage.protocol.{Frame, RespParser} -import sage.protocol.Frames.bulk +import sage.protocol.Frames.{bulk, feed} class CommandSpec extends munit.FunSuite { @@ -37,6 +37,11 @@ class CommandSpec extends munit.FunSuite { test("multi-word command names encode one bulk string per word") { val command = Command[Unit]("CONFIG GET", Vector.empty, Vector(Bytes.utf8("maxmemory")), _ => Right(())) assertEquals(command.encode.asUtf8String, "*3\r\n$6\r\nCONFIG\r\n$3\r\nGET\r\n$9\r\nmaxmemory\r\n") + val spaced = Command[Unit](" CONFIG GET ", Vector.empty, Vector.empty, _ => Right(())) + assertEquals(spaced.encode.asUtf8String, "*2\r\n$6\r\nCONFIG\r\n$3\r\nGET\r\n") + // a blank name is sent as an empty bulk string so the first argument is never run as the command + val blank = Command[Unit](" ", Vector.empty, Vector(Bytes.utf8("GET")), _ => Right(())) + assertEquals(blank.encode.asUtf8String, "*2\r\n$0\r\n\r\n$3\r\nGET\r\n") } test("a command's encoded bytes parse back as an array of bulk strings") { @@ -54,12 +59,15 @@ class CommandSpec extends munit.FunSuite { } test("a top-level error frame becomes a ServerError for any command") { - assertEquals(Reply.run(Strings.get[String, String]("foo"), Frame.SimpleError("ERR oops")), Left(ServerError("ERR", "oops"))) - assertEquals(Reply.run(Strings.set("foo", "bar"), Frame.BulkError(Bytes.utf8("WRONGTYPE bad"))), Left(ServerError("WRONGTYPE", "bad"))) + assertEquals(Reply.decode(Strings.get[String, String]("foo"), Frame.SimpleError("ERR oops")).toEither, Left(ServerError("ERR", "oops"))) + assertEquals( + Reply.decode(Strings.set("foo", "bar"), Frame.BulkError(Bytes.utf8("WRONGTYPE bad"))).toEither, + Left(ServerError("WRONGTYPE", "bad")) + ) } test("an unexpected frame shape becomes a DecodeError naming expected and actual") { - Reply.run(Strings.mSet(("foo", "bar")), Frame.Integer(1)) match { + Reply.decode(Strings.mSet(("foo", "bar")), Frame.Integer(1)).toEither match { case Left(error: DecodeError) => assertEquals(error.expected, "simple string 'OK'") assertEquals(error.actual, "integer 1") @@ -69,7 +77,39 @@ class CommandSpec extends munit.FunSuite { test("map transforms the decoded result") { val exists = Strings.get[String, String]("foo").map(_.isDefined) - assertEquals(Reply.run(exists, bulk("bar")), Right(true)) - assertEquals(Reply.run(exists, Frame.Null), Right(false)) + assertEquals(Reply.decode(exists, bulk("bar")).toEither, Right(true)) + assertEquals(Reply.decode(exists, Frame.Null).toEither, Right(false)) + } + + test("enum keywords and wire names use root-locale case mapping, so a Turkish default locale still sends RIGHT and reads integer") { + val previous = java.util.Locale.getDefault + java.util.Locale.setDefault(java.util.Locale.forLanguageTag("tr-TR")) + try { + assert(Args.keywords(ListSide.values)(ListSide.Right).sameBytes(Bytes.utf8("RIGHT"))) + assertEquals(Decode.byLowerName(JsonType.Integer).get("integer"), Some(JsonType.Integer)) + } finally java.util.Locale.setDefault(previous) + } + + test("a First broadcast returns the first master's reply but fails when another master's reply does not decode") { + val flush = Scripting.scriptFlush() + val ok = Frame.SimpleString("OK") + assertEquals(flush.reduceReplies(ok, Vector(ok)), ok) + intercept[DecodeError](flush.reduceReplies(ok, Vector(ok, Frame.Integer(1L)))) + } + + test("a First broadcast whose decoder throws on a dropped reply fails with a DecodeError") { + val boom = new IllegalStateException("boom") + val command = Command[Unit]( + "SCRIPT", + Command.NoKeys, + Vector.empty, + { + case Frame.Integer(_) => throw boom + case _ => Right(()) + }, + allMasters = true + ) + val error = intercept[DecodeError](command.reduceReplies(Frame.SimpleString("OK"), Vector(Frame.Integer(1L)))) + assertEquals(error.getCause, boom) } } diff --git a/sage-core/src/test/scala/sage/commands/ConnectionSpec.scala b/sage-core/src/test/scala/sage/commands/ConnectionSpec.scala index f5688e74..95a664b2 100644 --- a/sage-core/src/test/scala/sage/commands/ConnectionSpec.scala +++ b/sage-core/src/test/scala/sage/commands/ConnectionSpec.scala @@ -7,11 +7,11 @@ import sage.protocol.Frames.{bulk, map} class ConnectionSpec extends munit.FunSuite { test("PING decodes PONG and an echoed message") { - assertEquals(Reply.run(Connection.ping(), Frame.SimpleString("PONG")), Right("PONG")) - assertEquals(Reply.run(Connection.ping(Some("hi")), bulk("hi")), Right("hi")) + assertEquals(Reply.decode(Connection.ping(), Frame.SimpleString("PONG")).toEither, Right("PONG")) + assertEquals(Reply.decode(Connection.ping(Some("hi")), bulk("hi")).toEither, Right("hi")) } - test("HELLO decodes the fields it needs and ignores unknown entries") { + test("HELLO accepts a proto 3 reply and ignores other entries") { val reply = map( "server" -> bulk("redis"), "version" -> bulk("7.4.0"), @@ -21,27 +21,27 @@ class ConnectionSpec extends munit.FunSuite { "role" -> bulk("master"), "modules" -> Frame.Array(Vector.empty) ) - assertEquals(Reply.run(Connection.hello(), reply), Right(HelloReply("redis", "7.4.0", 3, "master"))) + assertEquals(Reply.decode(Connection.hello(), reply).toEither, Right(())) } test("HELLO rejects a proto other than 3, including values beyond Int range") { def reply(proto: Long) = map("server" -> bulk("redis"), "version" -> bulk("7.4.0"), "proto" -> Frame.Integer(proto), "role" -> bulk("master")) - Reply.run(Connection.hello(), reply(2)) match { + Reply.decode(Connection.hello(), reply(2)).toEither match { case Left(error: DecodeError) => assertEquals(error.expected, "proto 3") assertEquals(error.actual, "proto 2") case other => fail(s"expected a DecodeError, got $other") } - Reply.run(Connection.hello(), reply(2147483648L)) match { + Reply.decode(Connection.hello(), reply(2147483648L)).toEither match { case Left(error: DecodeError) => assertEquals(error.actual, "proto 2147483648") case other => fail(s"expected a DecodeError, got $other") } } - test("HELLO reports a missing required field") { - Reply.run(Connection.hello(), map("server" -> bulk("redis"))) match { - case Left(error: DecodeError) => assertEquals(error.expected, "map entry 'version'") + test("HELLO reports a missing proto field") { + Reply.decode(Connection.hello(), map("server" -> bulk("redis"))).toEither match { + case Left(error: DecodeError) => assertEquals(error.expected, "field 'proto'") case other => fail(s"expected a DecodeError, got $other") } } diff --git a/sage-core/src/test/scala/sage/commands/FrameDecodeSpec.scala b/sage-core/src/test/scala/sage/commands/FrameDecodeSpec.scala index 27be427e..5f4fd14b 100644 --- a/sage-core/src/test/scala/sage/commands/FrameDecodeSpec.scala +++ b/sage-core/src/test/scala/sage/commands/FrameDecodeSpec.scala @@ -23,6 +23,8 @@ class FrameDecodeSpec extends munit.FunSuite { val array = Frame.Array(Vector(Frame.BulkString(Bytes.utf8("1")), Frame.BulkString(Bytes.utf8("2")))) assertEquals(array.asArray.map(_.length), Right(2)) assertEquals(array.asArrayOf[Int], Right(Vector(1, 2))) + assertEquals(Frame.Set(array.elements).asArrayOf[Int], Right(Vector(1, 2))) + assertEquals(Frame.Push(array.elements).asArrayOf[Int], Right(Vector(1, 2))) assertEquals(Frame.Integer(1).asArray, Left(DecodeError("array", "integer 1"))) } } diff --git a/sage-core/src/test/scala/sage/commands/FunctionsSpec.scala b/sage-core/src/test/scala/sage/commands/FunctionsSpec.scala index af562ffd..7b655376 100644 --- a/sage-core/src/test/scala/sage/commands/FunctionsSpec.scala +++ b/sage-core/src/test/scala/sage/commands/FunctionsSpec.scala @@ -9,7 +9,7 @@ import sage.protocol.Frames.{bulk, map} class FunctionsSpec extends munit.FunSuite { test("FCALL returns the raw frame and computes key indices") { - assertEquals(Reply.run(Functions.fCall("f", Seq("k"), Seq("a")), Frame.Integer(7L)), Right(Frame.Integer(7L))) + assertEquals(Reply.decode(Functions.fCall("f", Seq("k"), Seq("a")), Frame.Integer(7L)).toEither, Right(Frame.Integer(7L))) assertEquals(Functions.fCall("f", Seq("k1", "k2")).keyIndices, Vector(2, 3)) } @@ -44,7 +44,7 @@ class FunctionsSpec extends munit.FunSuite { ) ) assertEquals( - Reply.run(Functions.functionList(), reply), + Reply.decode(Functions.functionList(), reply).toEither, Right(Vector(LibraryInfo("mylib", "LUA", Vector(FunctionInfo("myfunc", None, Set("no-writes"))), None))) ) } @@ -53,7 +53,10 @@ class FunctionsSpec extends munit.FunSuite { val reply = Frame.Array( Vector(map("library_name" -> bulk("l"), "engine" -> bulk("LUA"), "functions" -> Frame.Array(Vector.empty), "library_code" -> bulk("#!lua"))) ) - assertEquals(Reply.run(Functions.functionList(withCode = true), reply), Right(Vector(LibraryInfo("l", "LUA", Vector.empty, Some("#!lua"))))) + assertEquals( + Reply.decode(Functions.functionList(withCode = true), reply).toEither, + Right(Vector(LibraryInfo("l", "LUA", Vector.empty, Some("#!lua")))) + ) } test("FUNCTION STATS decodes a running script and per-engine counts; null running_script is None") { @@ -62,15 +65,15 @@ class FunctionsSpec extends munit.FunSuite { "engines" -> map("LUA" -> map("libraries_count" -> Frame.Integer(2L), "functions_count" -> Frame.Integer(5L))) ) assertEquals( - Reply.run(Functions.functionStats, reply), + Reply.decode(Functions.functionStats, reply).toEither, Right(FunctionStats(Some(RunningScript("f", Vector("FCALL", "f"), 12.millis)), Map("LUA" -> EngineStats(2L, 5L)))) ) val idle = map("running_script" -> Frame.Null, "engines" -> Frame.Map(Vector.empty)) - assertEquals(Reply.run(Functions.functionStats, idle), Right(FunctionStats(None, Map.empty))) + assertEquals(Reply.decode(Functions.functionStats, idle).toEither, Right(FunctionStats(None, Map.empty))) } test("FUNCTION DUMP decodes the opaque payload as bytes") { val payload = Bytes.utf8("\u0000binary") - assertEquals(Reply.run(Functions.functionDump, Frame.BulkString(payload)).map(_.asUtf8String), Right(payload.asUtf8String)) + assertEquals(Reply.decode(Functions.functionDump, Frame.BulkString(payload)).toEither.map(_.asUtf8String), Right(payload.asUtf8String)) } } diff --git a/sage-core/src/test/scala/sage/commands/HashesSpec.scala b/sage-core/src/test/scala/sage/commands/HashesSpec.scala index 61e2686f..9431556e 100644 --- a/sage-core/src/test/scala/sage/commands/HashesSpec.scala +++ b/sage-core/src/test/scala/sage/commands/HashesSpec.scala @@ -7,31 +7,31 @@ class HashesSpec extends munit.FunSuite { test("HGETALL decodes a RESP3 map frame into a field-keyed map") { val reply = map("f1" -> bulk("v1"), "f2" -> bulk("v2")) - assertEquals(Reply.run(Hashes.hGetAll[String, String, String]("h"), reply), Right(Map("f1" -> "v1", "f2" -> "v2"))) - assertEquals(Reply.run(Hashes.hGetAll[String, String, String]("h"), Frame.Map(Vector.empty)), Right(Map.empty[String, String])) + assertEquals(Reply.decode(Hashes.hGetAll[String, String, String]("h"), reply).toEither, Right(Map("f1" -> "v1", "f2" -> "v2"))) + assertEquals(Reply.decode(Hashes.hGetAll[String, String, String]("h"), Frame.Map(Vector.empty)).toEither, Right(Map.empty[String, String])) } test("HMGET keeps missing fields as None positionally") { val reply = Frame.Array(Vector(bulk("v1"), Frame.Null, bulk("v3"))) - assertEquals(Reply.run(Hashes.hmGet[String, String, String]("h", "a", "b", "c"), reply), Right(Vector(Some("v1"), None, Some("v3")))) + assertEquals(Reply.decode(Hashes.hmGet[String, String, String]("h", "a", "b", "c"), reply).toEither, Right(Vector(Some("v1"), None, Some("v3")))) } test("HRANDFIELD decodes null as None and a single field as Some") { - assertEquals(Reply.run(Hashes.hRandField[String, String]("h"), Frame.Null), Right(None)) - assertEquals(Reply.run(Hashes.hRandField[String, String]("h"), bulk("f")), Right(Some("f"))) + assertEquals(Reply.decode(Hashes.hRandField[String, String]("h"), Frame.Null).toEither, Right(None)) + assertEquals(Reply.decode(Hashes.hRandField[String, String]("h"), bulk("f")).toEither, Right(Some("f"))) } test("HRANDFIELD WITHVALUES decodes the nested field/value pairs") { val reply = Frame.Array(Vector(Frame.Array(Vector(bulk("f1"), bulk("v1"))), Frame.Array(Vector(bulk("f2"), bulk("v2"))))) assertEquals( - Reply.run(Hashes.hRandFieldWithValues[String, String, String]("h", 2L), reply), + Reply.decode(Hashes.hRandFieldWithValues[String, String, String]("h", 2L), reply).toEither, Right(Vector("f1" -> "v1", "f2" -> "v2")) ) } test("HSCAN decodes the cursor and the flat field/value array into pairs") { val reply = Frame.Array(Vector(bulk("12"), Frame.Array(Vector(bulk("f1"), bulk("v1"), bulk("f2"), bulk("v2"))))) - Reply.run(Hashes.hScan[String, String, String]("h", ScanCursor.start), reply) match { + Reply.decode(Hashes.hScan[String, String, String]("h", ScanCursor.start), reply).toEither match { case Right(page) => assertEquals(page.items, Vector("f1" -> "v1", "f2" -> "v2")) assert(page.next.isDefined) @@ -41,12 +41,12 @@ class HashesSpec extends munit.FunSuite { test("HSCAN rejects an odd-length field/value array") { val reply = Frame.Array(Vector(bulk("0"), Frame.Array(Vector(bulk("f1"), bulk("v1"), bulk("f2"))))) - assert(Reply.run(Hashes.hScan[String, String, String]("h", ScanCursor.start), reply).isLeft) + assert(Reply.decode(Hashes.hScan[String, String, String]("h", ScanCursor.start), reply).toEither.isLeft) } test("HSCAN NOVALUES decodes a zero cursor as complete and the items as bare fields") { val reply = Frame.Array(Vector(bulk("0"), Frame.Array(Vector(bulk("f1"), bulk("f2"))))) - Reply.run(Hashes.hScanNoValues[String, String]("h", ScanCursor.start), reply) match { + Reply.decode(Hashes.hScanNoValues[String, String]("h", ScanCursor.start), reply).toEither match { case Right(page) => assertEquals(page.items, Vector("f1", "f2")) assertEquals(page.next, None) @@ -55,8 +55,8 @@ class HashesSpec extends munit.FunSuite { } test("HINCRBYFLOAT decodes the bulk-string double") { - assertEquals(Reply.run(Hashes.hIncrByFloat("h", "f", 1.5), bulk("10.75")), Right(10.75)) - assert(Reply.run(Hashes.hIncrByFloat("h", "f", 1.5), bulk("not-a-number")).isLeft) + assertEquals(Reply.decode(Hashes.hIncrByFloat("h", "f", 1.5), bulk("10.75")).toEither, Right(10.75)) + assert(Reply.decode(Hashes.hIncrByFloat("h", "f", 1.5), bulk("not-a-number")).toEither.isLeft) } test("every hash command routes on the hash key alone") { diff --git a/sage-core/src/test/scala/sage/commands/KeysSpec.scala b/sage-core/src/test/scala/sage/commands/KeysSpec.scala index 902b7c08..0dcf2fba 100644 --- a/sage-core/src/test/scala/sage/commands/KeysSpec.scala +++ b/sage-core/src/test/scala/sage/commands/KeysSpec.scala @@ -4,25 +4,27 @@ import java.time.Instant import scala.concurrent.duration.* +import sage.Bytes +import sage.SageException.DecodeError import sage.protocol.Frame import sage.protocol.Frames.bulk class KeysSpec extends munit.FunSuite { test("TTL and PTTL decode the sentinels and a remaining duration in their unit") { - assertEquals(Reply.run(Keys.ttl("k"), Frame.Integer(-2)), Right(Ttl.NoKey)) - assertEquals(Reply.run(Keys.ttl("k"), Frame.Integer(-1)), Right(Ttl.NoExpiry)) - assertEquals(Reply.run(Keys.ttl("k"), Frame.Integer(42)), Right(Ttl.Expires(42.seconds))) - assertEquals(Reply.run(Keys.pTtl("k"), Frame.Integer(42)), Right(Ttl.Expires(42.millis))) - assert(Reply.run(Keys.ttl("k"), Frame.Integer(-3)).isLeft) + assertEquals(Reply.decode(Keys.ttl("k"), Frame.Integer(-2)).toEither, Right(Ttl.NoKey)) + assertEquals(Reply.decode(Keys.ttl("k"), Frame.Integer(-1)).toEither, Right(Ttl.NoExpiry)) + assertEquals(Reply.decode(Keys.ttl("k"), Frame.Integer(42)).toEither, Right(Ttl.Expires(42.seconds))) + assertEquals(Reply.decode(Keys.pTtl("k"), Frame.Integer(42)).toEither, Right(Ttl.Expires(42.millis))) + assert(Reply.decode(Keys.ttl("k"), Frame.Integer(-3)).toEither.isLeft) } test("EXPIRETIME and PEXPIRETIME decode the sentinels and an absolute timestamp in their unit") { - assertEquals(Reply.run(Keys.expireTime("k"), Frame.Integer(-2)), Right(ExpiryTime.NoKey)) - assertEquals(Reply.run(Keys.expireTime("k"), Frame.Integer(-1)), Right(ExpiryTime.NoExpiry)) - assertEquals(Reply.run(Keys.expireTime("k"), Frame.Integer(2000000000L)), Right(ExpiryTime.At(Instant.ofEpochSecond(2000000000L)))) + assertEquals(Reply.decode(Keys.expireTime("k"), Frame.Integer(-2)).toEither, Right(ExpiryTime.NoKey)) + assertEquals(Reply.decode(Keys.expireTime("k"), Frame.Integer(-1)).toEither, Right(ExpiryTime.NoExpiry)) + assertEquals(Reply.decode(Keys.expireTime("k"), Frame.Integer(2000000000L)).toEither, Right(ExpiryTime.At(Instant.ofEpochSecond(2000000000L)))) assertEquals( - Reply.run(Keys.pExpireTime("k"), Frame.Integer(2000000000123L)), + Reply.decode(Keys.pExpireTime("k"), Frame.Integer(2000000000123L)).toEither, Right(ExpiryTime.At(Instant.ofEpochMilli(2000000000123L))) ) } @@ -37,10 +39,10 @@ class KeysSpec extends munit.FunSuite { "stream" -> RedisType.Stream ) expected.foreach { case (wire, tpe) => - assertEquals(Reply.run(Keys.typeOf("k"), Frame.SimpleString(wire)), Right(Some(tpe))) + assertEquals(Reply.decode(Keys.typeOf("k"), Frame.SimpleString(wire)).toEither, Right(Some(tpe))) } - assertEquals(Reply.run(Keys.typeOf("k"), Frame.SimpleString("none")), Right(None)) - assertEquals(Reply.run(Keys.typeOf("k"), Frame.SimpleString("ReJSON-RL")), Right(Some(RedisType.Other("ReJSON-RL")))) + assertEquals(Reply.decode(Keys.typeOf("k"), Frame.SimpleString("none")).toEither, Right(None)) + assertEquals(Reply.decode(Keys.typeOf("k"), Frame.SimpleString("ReJSON-RL")).toEither, Right(Some(RedisType.Other("ReJSON-RL")))) } test("SCAN decodes a mid-iteration page with a next cursor") { @@ -50,7 +52,7 @@ class KeysSpec extends munit.FunSuite { Frame.Array(Vector(bulk("a"), bulk("b"))) ) ) - Reply.run(Keys.scan[String](ScanCursor.start), reply) match { + Reply.decode(Keys.scan[String](ScanCursor.start), reply).toEither match { case Right(page) => assertEquals(page.items, Vector("a", "b")) assert(page.next.isDefined) @@ -60,7 +62,7 @@ class KeysSpec extends munit.FunSuite { test("SCAN decodes a zero cursor as iteration complete, even with keys in the page") { val reply = Frame.Array(Vector(bulk("0"), Frame.Array(Vector(bulk("last"))))) - Reply.run(Keys.scan[String](ScanCursor.start), reply) match { + Reply.decode(Keys.scan[String](ScanCursor.start), reply).toEither match { case Right(page) => assertEquals(page.items, Vector("last")) assertEquals(page.next, None) @@ -70,7 +72,7 @@ class KeysSpec extends munit.FunSuite { test("a returned cursor feeds the next SCAN call") { val reply = Frame.Array(Vector(bulk("17"), Frame.Array(Vector.empty))) - Reply.run(Keys.scan[String](ScanCursor.start), reply) match { + Reply.decode(Keys.scan[String](ScanCursor.start), reply).toEither match { case Right(ScanPage(_, Some(next))) => assertEquals(Keys.scan[String](next).args.head.asUtf8String, "17") case other => fail(s"expected a next cursor, got $other") @@ -78,13 +80,13 @@ class KeysSpec extends munit.FunSuite { } test("SCAN rejects a malformed reply shape") { - assert(Reply.run(Keys.scan[String](ScanCursor.start), Frame.Array(Vector(Frame.Integer(0)))).isLeft) - assert(Reply.run(Keys.scan[String](ScanCursor.start), Frame.Integer(0)).isLeft) + assert(Reply.decode(Keys.scan[String](ScanCursor.start), Frame.Array(Vector(Frame.Integer(0)))).toEither.isLeft) + assert(Reply.decode(Keys.scan[String](ScanCursor.start), Frame.Integer(0)).toEither.isLeft) } test("RANDOMKEY decodes null as None on an empty database") { - assertEquals(Reply.run(Keys.randomKey[String], Frame.Null), Right(None)) - assertEquals(Reply.run(Keys.randomKey[String], bulk("k")), Right(Some("k"))) + assertEquals(Reply.decode(Keys.randomKey[String], Frame.Null).toEither, Right(None)) + assertEquals(Reply.decode(Keys.randomKey[String], bulk("k")).toEither, Right(Some("k"))) } test("multi-key commands mark every position for the slot engine") { @@ -106,6 +108,12 @@ class KeysSpec extends munit.FunSuite { assertEquals(command.rawFrame.broadcast, BroadcastReduce.Concat) } + test("KEYS concatenates per-master key arrays and rejects a master reply that is not an array instead of reporting it as a key") { + val keys = Keys.keys[String]("*") + assertEquals(keys.reduceReplies(Frame.Array(Vector(bulk("a"))), Vector(Frame.Set(Vector(bulk("b"))))), Frame.Array(Vector(bulk("a"), bulk("b")))) + intercept[DecodeError](keys.reduceReplies(Frame.Array(Vector(bulk("a"))), Vector(bulk("b")))) + } + test("expire picks the wire command from the duration's precision") { assertEquals(Keys.expire("k", 90.seconds).name, "EXPIRE") assertEquals(Keys.expire("k", 90500.millis).name, "PEXPIRE") @@ -121,6 +129,12 @@ class KeysSpec extends munit.FunSuite { assertEquals(Keys.expireAt("k", Instant.ofEpochSecond(2000000000L, 1)).args(1).asUtf8String, "2000000000001") } + test("RESTORE expiries and the MIGRATE timeout never encode the wire value 0, which means no expiry or the default timeout") { + assertEquals(Keys.restore("k", Bytes.utf8("p"), RestoreExpiry.In(Duration.Zero)).args(1).asUtf8String, "1") + assertEquals(Keys.restore("k", Bytes.utf8("p"), RestoreExpiry.At(Instant.EPOCH)).args(1).asUtf8String, "1") + assertEquals(Keys.migrate("h", 6380, 0, 500.micros)("k").args(4).asUtf8String, "1") + } + test("an extreme instant saturates instead of throwing while building the command") { val command = Keys.expireAt("k", Instant.MAX) assertEquals(command.name, "PEXPIREAT") diff --git a/sage-core/src/test/scala/sage/commands/ListsSpec.scala b/sage-core/src/test/scala/sage/commands/ListsSpec.scala index ca6ff815..03f0ffe0 100644 --- a/sage-core/src/test/scala/sage/commands/ListsSpec.scala +++ b/sage-core/src/test/scala/sage/commands/ListsSpec.scala @@ -6,33 +6,33 @@ import sage.protocol.Frames.bulk class ListsSpec extends munit.FunSuite { test("LPOP with a count collapses a missing list's null to an empty vector") { - assertEquals(Reply.run(Lists.lPopCount[String, String]("l", 2L), Frame.Null), Right(Vector.empty[String])) + assertEquals(Reply.decode(Lists.lPopCount[String, String]("l", 2L), Frame.Null).toEither, Right(Vector.empty[String])) assertEquals( - Reply.run(Lists.lPopCount[String, String]("l", 2L), Frame.Array(Vector(bulk("a"), bulk("b")))), + Reply.decode(Lists.lPopCount[String, String]("l", 2L), Frame.Array(Vector(bulk("a"), bulk("b")))).toEither, Right(Vector("a", "b")) ) } test("LPOP without a count decodes null as None") { - assertEquals(Reply.run(Lists.lPop[String, String]("l"), Frame.Null), Right(None)) - assertEquals(Reply.run(Lists.lPop[String, String]("l"), bulk("a")), Right(Some("a"))) + assertEquals(Reply.decode(Lists.lPop[String, String]("l"), Frame.Null).toEither, Right(None)) + assertEquals(Reply.decode(Lists.lPop[String, String]("l"), bulk("a")).toEither, Right(Some("a"))) } test("LPOS decodes null as None, an integer as the index, and a count reply as a vector") { - assertEquals(Reply.run(Lists.lPos("l", "v"), Frame.Null), Right(None)) - assertEquals(Reply.run(Lists.lPos("l", "v"), Frame.Integer(3)), Right(Some(3L))) - assertEquals(Reply.run(Lists.lPosCount("l", "v", 0L), Frame.Array(Vector(Frame.Integer(1), Frame.Integer(3)))), Right(Vector(1L, 3L))) + assertEquals(Reply.decode(Lists.lPos("l", "v"), Frame.Null).toEither, Right(None)) + assertEquals(Reply.decode(Lists.lPos("l", "v"), Frame.Integer(3)).toEither, Right(Some(3L))) + assertEquals(Reply.decode(Lists.lPosCount("l", "v", 0L), Frame.Array(Vector(Frame.Integer(1), Frame.Integer(3)))).toEither, Right(Vector(1L, 3L))) } test("LMOVE decodes null as None when the source is empty") { - assertEquals(Reply.run(Lists.lMove[String, String]("s", "d", ListSide.Left, ListSide.Right), Frame.Null), Right(None)) - assertEquals(Reply.run(Lists.lMove[String, String]("s", "d", ListSide.Left, ListSide.Right), bulk("a")), Right(Some("a"))) + assertEquals(Reply.decode(Lists.lMove[String, String]("s", "d", ListSide.Left, ListSide.Right), Frame.Null).toEither, Right(None)) + assertEquals(Reply.decode(Lists.lMove[String, String]("s", "d", ListSide.Left, ListSide.Right), bulk("a")).toEither, Right(Some("a"))) } test("LMPOP decodes null as None and a key/values reply as the popped key and elements") { - assertEquals(Reply.run(Lists.lMpop[String, String]("a", "b")(ListSide.Left), Frame.Null), Right(None)) + assertEquals(Reply.decode(Lists.lMpop[String, String]("a", "b")(ListSide.Left), Frame.Null).toEither, Right(None)) val reply = Frame.Array(Vector(bulk("b"), Frame.Array(Vector(bulk("x"), bulk("y"))))) - assertEquals(Reply.run(Lists.lMpop[String, String]("a", "b")(ListSide.Left), reply), Right(Some(("b", Vector("x", "y"))))) + assertEquals(Reply.decode(Lists.lMpop[String, String]("a", "b")(ListSide.Left), reply).toEither, Right(Some(("b", Vector("x", "y"))))) } test("the list sides and insert positions encode to their wire words") { diff --git a/sage-core/src/test/scala/sage/commands/PipelineSpec.scala b/sage-core/src/test/scala/sage/commands/PipelineSpec.scala index aa6f53ae..1fd1daaf 100644 --- a/sage-core/src/test/scala/sage/commands/PipelineSpec.scala +++ b/sage-core/src/test/scala/sage/commands/PipelineSpec.scala @@ -10,32 +10,36 @@ class PipelineSpec extends munit.FunSuite { } test("tuple pipeline assembles the all-success tuple") { - val p = Pipeline.fromTuple((Connection.ping(), Strings.get[String, String]("k"))) - val out = p.toOut(Vector("PONG", Some("v"))) - assertEquals(out, ("PONG", Some("v"))) + val p = Pipeline.fromTuple((Connection.ping(), Strings.get[String, String]("k"))) + assertEquals(p.finish(Vector(Right("PONG"), Right(Some("v")))), Right(("PONG", Some("v")))) } test("tuple pipeline shapes per-position results, mixing success and failure") { - val p = Pipeline.fromTuple((Connection.ping(), Strings.incr[String]("n"))) - val results = p.toResults(Vector(Right("PONG"), Left(ServerError("WRONGTYPE", "")))) - assertEquals(results, (Right("PONG"), Left(ServerError("WRONGTYPE", "")))) + val p = Pipeline.fromTupleAttempt((Connection.ping(), Strings.incr[String]("n"))) + val results = p.finish(Vector(Right("PONG"), Left(ServerError("WRONGTYPE", "")))) + assertEquals(results, Right((Right("PONG"), Left(ServerError("WRONGTYPE", ""))))) } test("sequence preserves order and assembles a homogeneous vector") { val p = Pipeline.sequence(Vector("a", "b", "c").map(Strings.get[String, String])) assertEquals(p.commands.map(_.name), Vector("GET", "GET", "GET")) - assertEquals(p.toOut(Vector(Some("1"), None, Some("3"))), Vector(Some("1"), None, Some("3"))) + assertEquals(p.finish(Vector(Right(Some("1")), Right(None), Right(Some("3")))), Right(Vector(Some("1"), None, Some("3")))) } test("sequence shapes per-position results") { - val p = Pipeline.sequence(Vector(Strings.get[String, String]("k"), Strings.get[String, String]("j"))) - val results = p.toResults(Vector(Right(Some("v")), Left(DecodeError("bulk string", "integer")))) - assertEquals(results, Vector(Right(Some("v")), Left(DecodeError("bulk string", "integer")))) + val p = Pipeline.sequenceAttempt(Vector(Strings.get[String, String]("k"), Strings.get[String, String]("j"))) + val results = p.finish(Vector(Right(Some("v")), Left(DecodeError("bulk string", "integer")))) + assertEquals(results, Right(Vector(Right(Some("v")), Left(DecodeError("bulk string", "integer"))))) + } + + test("a strict pipeline fails with the first error in position order") { + val p = Pipeline.sequence(Vector("a", "b", "c").map(Strings.get[String, String])) + assertEquals(p.finish(Vector(Right(Some("1")), Left(ServerError("ERR", "x")), Left(ServerError("ERR", "y")))), Left(ServerError("ERR", "x"))) } test("an empty sequence carries no commands") { val p = Pipeline.sequence(Vector.empty[Command[Long]]) assertEquals(p.commands, Vector.empty) - assertEquals(p.toOut(Vector.empty), Vector.empty) + assertEquals(p.finish(Vector.empty), Right(Vector.empty)) } } diff --git a/sage-core/src/test/scala/sage/commands/PubsubSpec.scala b/sage-core/src/test/scala/sage/commands/PubsubSpec.scala index 4f9c5d92..4e879f50 100644 --- a/sage-core/src/test/scala/sage/commands/PubsubSpec.scala +++ b/sage-core/src/test/scala/sage/commands/PubsubSpec.scala @@ -1,5 +1,6 @@ package sage.commands +import sage.SageException.DecodeError import sage.protocol.Frame import sage.protocol.Frames.bulk @@ -40,66 +41,80 @@ class PubsubSpec extends munit.FunSuite with BroadcastFolds { val merge = fold(Pubsub.pubsubChannels()) val first = Frame.Array(Vector(bulk("news"), bulk("sport"))) val other = Frame.Array(Vector(bulk("news"), bulk("weather"))) - assertEquals(Reply.run(Pubsub.pubsubChannels(), merge(first, other)), Right(Vector("news", "sport", "weather"))) + assertEquals(Reply.decode(Pubsub.pubsubChannels(), merge(first, other)).toEither, Right(Vector("news", "sport", "weather"))) } test("CHANNELS merges an empty master in either position without losing the other's slice") { val merge = fold(Pubsub.pubsubChannels()) val occupied = Frame.Array(Vector(bulk("news"))) val bare = Frame.Array(Vector.empty) - assertEquals(Reply.run(Pubsub.pubsubChannels(), merge(occupied, bare)), Right(Vector("news"))) - assertEquals(Reply.run(Pubsub.pubsubChannels(), merge(bare, occupied)), Right(Vector("news"))) + assertEquals(Reply.decode(Pubsub.pubsubChannels(), merge(occupied, bare)).toEither, Right(Vector("news"))) + assertEquals(Reply.decode(Pubsub.pubsubChannels(), merge(bare, occupied)).toEither, Right(Vector("news"))) } - test("the CHANNELS merge passes a malformed reply through in either operand position, so it never hides behind a valid one") { + test("the CHANNELS merge fails on a malformed reply in either operand position, so it never hides behind a valid one") { val merge = fold(Pubsub.pubsubChannels()) val valid = Frame.Array(Vector(bulk("news"))) val bad = Frame.SimpleString("nonsense") - assertEquals(merge(valid, bad), bad) - assertEquals(merge(bad, valid), bad) - assert(Reply.run(Pubsub.pubsubChannels(), merge(valid, bad)).isLeft) + intercept[DecodeError](merge(valid, bad)) + intercept[DecodeError](merge(bad, valid)) } test("NUMSUB sums a channel's subscribers across masters instead of letting the decoder's Map keep only the last count") { val merge = fold(Pubsub.pubsubNumSub("news", "sport")) val first = Frame.Array(Vector(bulk("news"), Frame.Integer(1L), bulk("sport"), Frame.Integer(0L))) val other = Frame.Array(Vector(bulk("news"), Frame.Integer(2L), bulk("sport"), Frame.Integer(5L))) - assertEquals(Reply.run(Pubsub.pubsubNumSub("news", "sport"), merge(first, other)), Right(Map("news" -> 3L, "sport" -> 5L))) + assertEquals(Reply.decode(Pubsub.pubsubNumSub("news", "sport"), merge(first, other)).toEither, Right(Map("news" -> 3L, "sport" -> 5L))) } - test("the NUMSUB merge keeps each channel's first-seen position and admits a master that reports a channel the other does not") { - val merge = fold(Pubsub.pubsubNumSub("news", "sport")) - val first = Frame.Array(Vector(bulk("news"), Frame.Integer(1L))) - val other = Frame.Array(Vector(bulk("sport"), Frame.Integer(4L), bulk("news"), Frame.Integer(2L))) - val merged = merge(first, other) - assertEquals(Reply.run(Pubsub.pubsubNumSub("news", "sport"), merged), Right(Map("news" -> 3L, "sport" -> 4L))) - assertEquals(merged, Frame.Array(Vector(bulk("news"), Frame.Integer(3L), bulk("sport"), Frame.Integer(4L)))) + test("the NUMSUB merge admits a master that reports a channel the other does not") { + val merge = fold(Pubsub.pubsubNumSub("news", "sport")) + val first = Frame.Array(Vector(bulk("news"), Frame.Integer(1L))) + val other = Frame.Array(Vector(bulk("sport"), Frame.Integer(4L), bulk("news"), Frame.Integer(2L))) + assertEquals(Reply.decode(Pubsub.pubsubNumSub("news", "sport"), merge(first, other)).toEither, Right(Map("news" -> 3L, "sport" -> 4L))) + } + + test("NUMSUB with a repeated channel counts each master once, as a standalone server's reply does") { + val merge = fold(Pubsub.pubsubNumSub("news", "news")) + val reply = Frame.Array(Vector(bulk("news"), Frame.Integer(1L), bulk("news"), Frame.Integer(1L))) + assertEquals(Reply.decode(Pubsub.pubsubNumSub("news", "news"), reply).toEither, Right(Map("news" -> 1L))) + assertEquals(Reply.decode(Pubsub.pubsubNumSub("news", "news"), merge(reply, reply)).toEither, Right(Map("news" -> 2L))) } test("SHARDNUMSUB sums per channel too, so a shard channel keeps its owner's count when the other masters report zero") { val merge = fold(Pubsub.pubsubShardNumSub("orders")) val owner = Frame.Array(Vector(bulk("orders"), Frame.Integer(3L))) val bare = Frame.Array(Vector(bulk("orders"), Frame.Integer(0L))) - assertEquals(Reply.run(Pubsub.pubsubShardNumSub("orders"), merge(bare, owner)), Right(Map("orders" -> 3L))) + assertEquals(Reply.decode(Pubsub.pubsubShardNumSub("orders"), merge(bare, owner)).toEither, Right(Map("orders" -> 3L))) } - test("the NUMSUB merge passes an odd-length or mistyped reply through in either operand position") { + test("the NUMSUB merge fails on an odd-length or mistyped reply in either operand position") { val merge = fold(Pubsub.pubsubNumSub("news")) val valid = Frame.Array(Vector(bulk("news"), Frame.Integer(1L))) val odd = Frame.Array(Vector(bulk("news"))) val typed = Frame.Array(Vector(bulk("news"), bulk("1"))) - assertEquals(merge(valid, odd), odd) - assertEquals(merge(odd, valid), odd) - assertEquals(merge(valid, typed), typed) - assert(Reply.run(Pubsub.pubsubNumSub("news"), merge(valid, odd)).isLeft) + intercept[DecodeError](merge(valid, odd)) + intercept[DecodeError](merge(odd, valid)) + intercept[DecodeError](merge(valid, typed)) } - test("NUMPAT sums each master's pattern count, with checked overflow and malformed passthrough") { + test("NUMPAT sums each master's pattern count, with checked overflow and malformed replies rejected") { val merge = fold(Pubsub.pubsubNumPat) - assertEquals(Reply.run(Pubsub.pubsubNumPat, merge(Frame.Integer(2L), Frame.Integer(3L))), Right(5L)) + assertEquals(Reply.decode(Pubsub.pubsubNumPat, merge(Frame.Integer(2L), Frame.Integer(3L))).toEither, Right(5L)) val bad = Frame.SimpleString("nonsense") - assertEquals(merge(Frame.Integer(2L), bad), bad) - assertEquals(merge(bad, Frame.Integer(2L)), bad) + intercept[DecodeError](merge(Frame.Integer(2L), bad)) + intercept[DecodeError](merge(bad, Frame.Integer(2L))) intercept[ArithmeticException](merge(Frame.Integer(Long.MaxValue), Frame.Integer(1L))) } + + test("an invalidate push whose key list does not decode flushes the cache instead of being ignored") { + def decode(keys: Frame) = Invalidation.decode(Vector(bulk("invalidate"), keys)) + decode(Frame.Array(Vector(bulk("k")))) match { + case Some(Invalidation.Evict(keys)) => assertEquals(keys.map(_.asUtf8String), Vector("k")) + case other => fail(s"expected Evict, got $other") + } + assertEquals(decode(Frame.Null), Some(Invalidation.FlushAll)) + assertEquals(decode(Frame.Array(Vector(Frame.Integer(1)))), Some(Invalidation.FlushAll)) + assertEquals(decode(bulk("k")), Some(Invalidation.FlushAll)) + } } diff --git a/sage-core/src/test/scala/sage/commands/ScriptingSpec.scala b/sage-core/src/test/scala/sage/commands/ScriptingSpec.scala index 39094992..c29781b3 100644 --- a/sage-core/src/test/scala/sage/commands/ScriptingSpec.scala +++ b/sage-core/src/test/scala/sage/commands/ScriptingSpec.scala @@ -11,9 +11,9 @@ class ScriptingSpec extends munit.FunSuite with BroadcastFolds { private def existsFold: (Frame, Frame) => Frame = fold(Scripting.scriptExists("x")) test("EVAL returns the raw RESP3 frame untouched") { - assertEquals(Reply.run(Scripting.eval("return 1"), Frame.Integer(1L)), Right(Frame.Integer(1L))) + assertEquals(Reply.decode(Scripting.eval("return 1"), Frame.Integer(1L)).toEither, Right(Frame.Integer(1L))) val nested = Frame.Array(Vector(bulk("a"), Frame.Integer(2L))) - assertEquals(Reply.run(Scripting.eval("x", Seq("k")), nested), Right(nested)) + assertEquals(Reply.decode(Scripting.eval("x", Seq("k")), nested).toEither, Right(nested)) } test("EVAL computes numkeys and key indices from the key list") { @@ -35,11 +35,11 @@ class ScriptingSpec extends munit.FunSuite with BroadcastFolds { test("SCRIPT EXISTS decodes one flag per sha, in order") { val reply = Frame.Array(Vector(Frame.Integer(1L), Frame.Integer(0L), Frame.Integer(1L))) - assertEquals(Reply.run(Scripting.scriptExists("a", "b", "c"), reply), Right(Vector(true, false, true))) + assertEquals(Reply.decode(Scripting.scriptExists("a", "b", "c"), reply).toEither, Right(Vector(true, false, true))) } test("SCRIPT LOAD decodes the sha") { - assertEquals(Reply.run(Scripting.scriptLoad("return 1"), bulk("abc123")), Right("abc123")) + assertEquals(Reply.decode(Scripting.scriptLoad("return 1"), bulk("abc123")).toEither, Right("abc123")) } test("SCRIPT LOAD, FLUSH and EXISTS are All-Masters Commands; KILL and EVALSHA are not") { @@ -54,8 +54,9 @@ class ScriptingSpec extends munit.FunSuite with BroadcastFolds { assertEquals(existsFold(flags(1L, 1L), flags(1L, 0L)), flags(1L, 0L)) } - test("SCRIPT EXISTS surfaces a non-binary flag so the strict decode rejects it rather than coercing to false") { - assert(Reply.run(Scripting.scriptExists("a"), existsFold(flags(2L), flags(1L))).isLeft) + test("SCRIPT EXISTS rejects a non-binary flag rather than coercing it to false") { + intercept[DecodeError](existsFold(flags(2L), flags(1L))) + intercept[DecodeError](existsFold(flags(1L), flags(2L))) } test("SCRIPT EXISTS fails a length mismatch across masters rather than hiding the malformed reply") { diff --git a/sage-core/src/test/scala/sage/commands/ServerSpec.scala b/sage-core/src/test/scala/sage/commands/ServerSpec.scala index fa0e48b8..ae297e17 100644 --- a/sage-core/src/test/scala/sage/commands/ServerSpec.scala +++ b/sage-core/src/test/scala/sage/commands/ServerSpec.scala @@ -4,6 +4,7 @@ import java.time.Instant import scala.concurrent.duration.* +import sage.SageException.DecodeError import sage.protocol.Frame import sage.protocol.Frames.{bulk, map} @@ -11,36 +12,36 @@ class ServerSpec extends munit.FunSuite with BroadcastFolds { test("CONFIG GET decodes a RESP3 map and the RESP2 flat-array shape alike") { val resp3 = map("maxmemory" -> bulk("100mb"), "save" -> bulk("3600 1")) - assertEquals(Reply.run(Server.configGet("*"), resp3), Right(Map("maxmemory" -> "100mb", "save" -> "3600 1"))) + assertEquals(Reply.decode(Server.configGet("*"), resp3).toEither, Right(Map("maxmemory" -> "100mb", "save" -> "3600 1"))) val resp2 = Frame.Array(Vector(bulk("maxmemory"), bulk("100mb"))) - assertEquals(Reply.run(Server.configGet("*"), resp2), Right(Map("maxmemory" -> "100mb"))) + assertEquals(Reply.decode(Server.configGet("*"), resp2).toEither, Right(Map("maxmemory" -> "100mb"))) } test("TIME decodes seconds + microseconds into an Instant") { val reply = Frame.Array(Vector(bulk("1700000000"), bulk("123456"))) - assertEquals(Reply.run(Server.time, reply), Right(Instant.ofEpochSecond(1700000000L, 123456000L))) + assertEquals(Reply.decode(Server.time, reply).toEither, Right(Instant.ofEpochSecond(1700000000L, 123456000L))) } test("ROLE decodes master, replica, and sentinel forms") { val master = Frame.Array( Vector(bulk("master"), Frame.Integer(100L), Frame.Array(Vector(Frame.Array(Vector(bulk("127.0.0.1"), bulk("6380"), bulk("90")))))) ) - assertEquals(Reply.run(Server.role, master), Right(Role.Master(100L, Vector(ReplicaNode("127.0.0.1", 6380, 90L))))) + assertEquals(Reply.decode(Server.role, master).toEither, Right(Role.Master(100L, Vector(ReplicaNode("127.0.0.1", 6380, 90L))))) val replica = Frame.Array(Vector(bulk("slave"), bulk("127.0.0.1"), Frame.Integer(6379L), bulk("connected"), Frame.Integer(50L))) - assertEquals(Reply.run(Server.role, replica), Right(Role.Replica("127.0.0.1", 6379, "connected", 50L))) + assertEquals(Reply.decode(Server.role, replica).toEither, Right(Role.Replica("127.0.0.1", 6379, "connected", 50L))) val sentinel = Frame.Array(Vector(bulk("sentinel"), Frame.Array(Vector(bulk("master1"), bulk("master2"))))) - assertEquals(Reply.run(Server.role, sentinel), Right(Role.Sentinel(Vector("master1", "master2")))) + assertEquals(Reply.decode(Server.role, sentinel).toEither, Right(Role.Sentinel(Vector("master1", "master2")))) } test("ROLE rejects an out-of-range replica port instead of wrapping it") { val wrapping = Frame.Array(Vector(bulk("slave"), bulk("127.0.0.1"), Frame.Integer(Int.MaxValue.toLong + 1L), bulk("connected"), Frame.Integer(0L))) - assert(Reply.run(Server.role, wrapping).isLeft, "a port above Int.MaxValue must not wrap to a valid-looking port") + assert(Reply.decode(Server.role, wrapping).toEither.isLeft, "a port above Int.MaxValue must not wrap to a valid-looking port") val tooLarge = Frame.Array(Vector(bulk("master"), Frame.Integer(0L), Frame.Array(Vector(Frame.Array(Vector(bulk("127.0.0.1"), bulk("70000"), bulk("0"))))))) - assert(Reply.run(Server.role, tooLarge).isLeft, "a replica port outside 1..65535 must be a DecodeError") + assert(Reply.decode(Server.role, tooLarge).toEither.isLeft, "a replica port outside 1..65535 must be a DecodeError") } test("SLOWLOG GET decodes entries, defaulting client fields absent on old servers") { @@ -59,23 +60,26 @@ class ServerSpec extends munit.FunSuite with BroadcastFolds { ) ) assertEquals( - Reply.run(Server.slowLogGet(), withClient), + Reply.decode(Server.slowLogGet(), withClient).toEither, Right(Vector(SlowLogEntry(7L, Instant.ofEpochSecond(1700000000L), 150.micros, Vector("GET", "k"), "1.2.3.4:5", "app"))) ) val old = Frame.Array(Vector(Frame.Array(Vector(Frame.Integer(1L), Frame.Integer(10L), Frame.Integer(20L), Frame.Array(Vector(bulk("PING"))))))) - assertEquals(Reply.run(Server.slowLogGet(), old), Right(Vector(SlowLogEntry(1L, Instant.ofEpochSecond(10L), 20.micros, Vector("PING"), "", "")))) + assertEquals( + Reply.decode(Server.slowLogGet(), old).toEither, + Right(Vector(SlowLogEntry(1L, Instant.ofEpochSecond(10L), 20.micros, Vector("PING"), "", ""))) + ) } test("LATENCY LATEST decodes event rows") { val reply = Frame.Array(Vector(Frame.Array(Vector(bulk("command"), Frame.Integer(1700000000L), Frame.Integer(5L), Frame.Integer(20L))))) assertEquals( - Reply.run(Server.latencyLatest, reply), + Reply.decode(Server.latencyLatest, reply).toEither, Right(Vector(LatencyEntry("command", Instant.ofEpochSecond(1700000000L), 5.millis, 20.millis))) ) } test("WAITAOF decodes the [numlocal, numreplicas] pair; WAIT/WAITAOF ride the multiplexed connection and broadcast per master") { - assertEquals(Reply.run(Server.waitAof(1L, 0L, 1.second), Frame.Array(Vector(Frame.Integer(1L), Frame.Integer(2L)))), Right((1L, 2L))) + assertEquals(Reply.decode(Server.waitAof(1L, 0L, 1.second), Frame.Array(Vector(Frame.Integer(1L), Frame.Integer(2L)))).toEither, Right((1L, 2L))) assert(!Server.waitAof(1L, 0L, 1.second).isBlocking) assert(!Server.waitReplicas(1L, 1.second).isBlocking) assert(Server.waitAof(1L, 0L, 1.second).allMasters) @@ -109,7 +113,7 @@ class ServerSpec extends munit.FunSuite with BroadcastFolds { val waitFold = fold(Server.waitReplicas(1L, 1.second)) assertEquals(waitFold(Frame.Integer(2L), Frame.Integer(1L)), Frame.Integer(1L)) assertEquals(waitFold(Frame.Integer(0L), Frame.Integer(3L)), Frame.Integer(0L)) - assertEquals(Reply.run(Server.waitReplicas(1L, 1.second), Frame.Integer(1L)), Right(1L)) + assertEquals(Reply.decode(Server.waitReplicas(1L, 1.second), Frame.Integer(1L)).toEither, Right(1L)) val aofFold = fold(Server.waitAof(1L, 0L, 1.second)) assertEquals( @@ -118,22 +122,20 @@ class ServerSpec extends munit.FunSuite with BroadcastFolds { ) } - test("WAIT/WAITAOF folds pass a malformed reply through in either operand position, so it never hides behind a valid one") { + test("WAIT/WAITAOF folds fail on a malformed reply in either operand position, so it never hides behind a valid one") { val waitFold = fold(Server.waitReplicas(1L, 1.second)) val badInt = Frame.SimpleString("nonsense") - assertEquals(waitFold(Frame.Integer(2L), badInt), badInt) - assertEquals(waitFold(badInt, Frame.Integer(2L)), badInt) - assert(Reply.run(Server.waitReplicas(1L, 1.second), waitFold(Frame.Integer(2L), badInt)).isLeft) + intercept[DecodeError](waitFold(Frame.Integer(2L), badInt)) + intercept[DecodeError](waitFold(badInt, Frame.Integer(2L))) val aofFold = fold(Server.waitAof(1L, 0L, 1.second)) val validPair = Frame.Array(Vector(Frame.Integer(1L), Frame.Integer(2L))) val badPair = Frame.Array(Vector(Frame.Integer(1L))) - assertEquals(aofFold(validPair, badPair), badPair) - assertEquals(aofFold(badPair, validPair), badPair) - assert(Reply.run(Server.waitAof(1L, 0L, 1.second), aofFold(validPair, badPair)).isLeft) + intercept[DecodeError](aofFold(validPair, badPair)) + intercept[DecodeError](aofFold(badPair, validPair)) } - test("DBSIZE broadcasts per master and folds shard counts into the cluster total, with checked overflow and malformed passthrough") { + test("DBSIZE broadcasts per master and folds shard counts into the cluster total, with checked overflow and malformed replies rejected") { assert(Server.dbSize.allMasters) assert( Server.dbSize.requiresClusterWideTxResult, @@ -141,12 +143,11 @@ class ServerSpec extends munit.FunSuite with BroadcastFolds { ) val sum = fold(Server.dbSize) assertEquals(sum(Frame.Integer(10L), Frame.Integer(20L)), Frame.Integer(30L)) - assertEquals(Reply.run(Server.dbSize, Frame.Integer(30L)), Right(30L)) + assertEquals(Reply.decode(Server.dbSize, Frame.Integer(30L)).toEither, Right(30L)) intercept[ArithmeticException](sum(Frame.Integer(Long.MaxValue), Frame.Integer(1L))) val badInt = Frame.SimpleString("nonsense") - assertEquals(sum(Frame.Integer(10L), badInt), badInt) - assertEquals(sum(badInt, Frame.Integer(10L)), badInt) - assert(Reply.run(Server.dbSize, sum(Frame.Integer(10L), badInt)).isLeft) + intercept[DecodeError](sum(Frame.Integer(10L), badInt)) + intercept[DecodeError](sum(badInt, Frame.Integer(10L))) } test("MEMORY PURGE broadcasts per master, since purging one arbitrary master leaves every other one unpurged") { @@ -161,8 +162,8 @@ class ServerSpec extends munit.FunSuite with BroadcastFolds { } test("MEMORY USAGE decodes a present count and a missing key as None, and is keyed") { - assertEquals(Reply.run(Server.memoryUsage("k"), Frame.Integer(64L)), Right(Some(64L))) - assertEquals(Reply.run(Server.memoryUsage("k"), Frame.Null), Right(None)) + assertEquals(Reply.decode(Server.memoryUsage("k"), Frame.Integer(64L)).toEither, Right(Some(64L))) + assertEquals(Reply.decode(Server.memoryUsage("k"), Frame.Null).toEither, Right(None)) assertEquals(Server.memoryUsage("k").keyIndices, Vector(1)) } @@ -184,13 +185,13 @@ class ServerSpec extends munit.FunSuite with BroadcastFolds { ) ) assertEquals( - Reply.run(Server.commandInfo("get", "nope"), reply), + Reply.decode(Server.commandInfo("get", "nope"), reply).toEither, Right(Vector(CommandInfo("get", 2L, Set("readonly", "fast"), 1, 1, 1, Set("@read")))) ) } test("COMMAND GETKEYSANDFLAGS decodes key/flag pairs") { val reply = Frame.Array(Vector(Frame.Array(Vector(bulk("k"), Frame.Array(Vector(bulk("RW"), bulk("access"))))))) - assertEquals(Reply.run(Server.commandGetKeysAndFlags("SET", "k", "v"), reply), Right(Vector("k" -> Set("RW", "access")))) + assertEquals(Reply.decode(Server.commandGetKeysAndFlags("SET", "k", "v"), reply).toEither, Right(Vector("k" -> Set("RW", "access")))) } } diff --git a/sage-core/src/test/scala/sage/commands/SetsSpec.scala b/sage-core/src/test/scala/sage/commands/SetsSpec.scala index 1472d784..f1a03ef9 100644 --- a/sage-core/src/test/scala/sage/commands/SetsSpec.scala +++ b/sage-core/src/test/scala/sage/commands/SetsSpec.scala @@ -6,28 +6,28 @@ import sage.protocol.Frames.bulk class SetsSpec extends munit.FunSuite { test("SMEMBERS decodes a RESP3 set frame into a Set, empty included") { - assertEquals(Reply.run(Sets.sMembers[String, String]("s"), Frame.Set(Vector(bulk("a"), bulk("b")))), Right(Set("a", "b"))) - assertEquals(Reply.run(Sets.sMembers[String, String]("s"), Frame.Set(Vector.empty)), Right(Set.empty[String])) + assertEquals(Reply.decode(Sets.sMembers[String, String]("s"), Frame.Set(Vector(bulk("a"), bulk("b")))).toEither, Right(Set("a", "b"))) + assertEquals(Reply.decode(Sets.sMembers[String, String]("s"), Frame.Set(Vector.empty)).toEither, Right(Set.empty[String])) } test("a set reply rejects a plain array frame") { - assert(Reply.run(Sets.sMembers[String, String]("s"), Frame.Array(Vector(bulk("a")))).isLeft) + assert(Reply.decode(Sets.sMembers[String, String]("s"), Frame.Array(Vector(bulk("a")))).toEither.isLeft) } test("SPOP decodes null as None and a bulk string as the member; a count reads a set") { - assertEquals(Reply.run(Sets.sPop[String, String]("s"), Frame.Null), Right(None)) - assertEquals(Reply.run(Sets.sPop[String, String]("s"), bulk("a")), Right(Some("a"))) - assertEquals(Reply.run(Sets.sPopCount[String, String]("s", 2L), Frame.Set(Vector(bulk("a"), bulk("b")))), Right(Set("a", "b"))) + assertEquals(Reply.decode(Sets.sPop[String, String]("s"), Frame.Null).toEither, Right(None)) + assertEquals(Reply.decode(Sets.sPop[String, String]("s"), bulk("a")).toEither, Right(Some("a"))) + assertEquals(Reply.decode(Sets.sPopCount[String, String]("s", 2L), Frame.Set(Vector(bulk("a"), bulk("b")))).toEither, Right(Set("a", "b"))) } test("SMISMEMBER decodes the 0/1 array positionally") { val reply = Frame.Array(Vector(Frame.Integer(1), Frame.Integer(0), Frame.Integer(1))) - assertEquals(Reply.run(Sets.sMisMember("s", "a", "b", "c"), reply), Right(Vector(true, false, true))) + assertEquals(Reply.decode(Sets.sMisMember("s", "a", "b", "c"), reply).toEither, Right(Vector(true, false, true))) } test("SRANDMEMBER with a count keeps duplicates as an ordered vector") { val reply = Frame.Array(Vector(bulk("a"), bulk("a"), bulk("b"))) - assertEquals(Reply.run(Sets.sRandMemberCount[String, String]("s", -3L), reply), Right(Vector("a", "a", "b"))) + assertEquals(Reply.decode(Sets.sRandMemberCount[String, String]("s", -3L), reply).toEither, Right(Vector("a", "a", "b"))) } test("set commands route on their keys") { diff --git a/sage-core/src/test/scala/sage/commands/SortedSetsSpec.scala b/sage-core/src/test/scala/sage/commands/SortedSetsSpec.scala index 5339ff50..6d2828ff 100644 --- a/sage-core/src/test/scala/sage/commands/SortedSetsSpec.scala +++ b/sage-core/src/test/scala/sage/commands/SortedSetsSpec.scala @@ -8,54 +8,60 @@ class SortedSetsSpec extends munit.FunSuite { private def pair(member: String, score: Double): Frame = Frame.Array(Vector(bulk(member), Frame.Double(score))) test("ZSCORE decodes a RESP3 double, null as None, and rejects a bulk string") { - assertEquals(Reply.run(SortedSets.zScore[String, String]("z", "a"), Frame.Double(1.5)), Right(Some(1.5))) - assertEquals(Reply.run(SortedSets.zScore[String, String]("z", "a"), Frame.Null), Right(None)) - assert(Reply.run(SortedSets.zScore[String, String]("z", "a"), bulk("1.5")).isLeft) + assertEquals(Reply.decode(SortedSets.zScore[String, String]("z", "a"), Frame.Double(1.5)).toEither, Right(Some(1.5))) + assertEquals(Reply.decode(SortedSets.zScore[String, String]("z", "a"), Frame.Null).toEither, Right(None)) + assert(Reply.decode(SortedSets.zScore[String, String]("z", "a"), bulk("1.5")).toEither.isLeft) } test("ZADD INCR decodes the new score or None when the condition skipped the write") { - assertEquals(Reply.run(SortedSets.zAddIncr("z")("a", 1.0), Frame.Double(3.0)), Right(Some(3.0))) - assertEquals(Reply.run(SortedSets.zAddIncr("z", ZAddCondition.IfNotExists)("a", 1.0), Frame.Null), Right(None)) + assertEquals(Reply.decode(SortedSets.zAddIncr("z")("a", 1.0), Frame.Double(3.0)).toEither, Right(Some(3.0))) + assertEquals(Reply.decode(SortedSets.zAddIncr("z", ZAddCondition.IfNotExists)("a", 1.0), Frame.Null).toEither, Right(None)) } test("ZMSCORE decodes the score array, keeping missing members as None") { val reply = Frame.Array(Vector(Frame.Double(1.0), Frame.Null, Frame.Double(2.5))) - assertEquals(Reply.run(SortedSets.zMScore[String, String]("z", "a", "b", "c"), reply), Right(Vector(Some(1.0), None, Some(2.5)))) + assertEquals(Reply.decode(SortedSets.zMScore[String, String]("z", "a", "b", "c"), reply).toEither, Right(Vector(Some(1.0), None, Some(2.5)))) } test("ZRANGE WITHSCORES decodes the nested member/score pairs") { val reply = Frame.Array(Vector(pair("a", 1.0), pair("b", 2.0))) - assertEquals(Reply.run(SortedSets.zRangeWithScores[String, String]("z", ZRange.ByRank(0L, -1L)), reply), Right(Vector("a" -> 1.0, "b" -> 2.0))) + assertEquals( + Reply.decode(SortedSets.zRangeWithScores[String, String]("z", ZRange.ByRank(0L, -1L)), reply).toEither, + Right(Vector("a" -> 1.0, "b" -> 2.0)) + ) } test("ZPOPMIN decodes a flat member/score, an empty array as None, and a count as nested pairs") { - assertEquals(Reply.run(SortedSets.zPopMin[String, String]("z"), Frame.Array(Vector(bulk("a"), Frame.Double(1.0)))), Right(Some("a" -> 1.0))) - assertEquals(Reply.run(SortedSets.zPopMin[String, String]("z"), Frame.Array(Vector.empty)), Right(None)) + assertEquals( + Reply.decode(SortedSets.zPopMin[String, String]("z"), Frame.Array(Vector(bulk("a"), Frame.Double(1.0)))).toEither, + Right(Some("a" -> 1.0)) + ) + assertEquals(Reply.decode(SortedSets.zPopMin[String, String]("z"), Frame.Array(Vector.empty)).toEither, Right(None)) val counted = Frame.Array(Vector(pair("a", 1.0), pair("b", 2.0))) - assertEquals(Reply.run(SortedSets.zPopMinCount[String, String]("z", 2L), counted), Right(Vector("a" -> 1.0, "b" -> 2.0))) + assertEquals(Reply.decode(SortedSets.zPopMinCount[String, String]("z", 2L), counted).toEither, Right(Vector("a" -> 1.0, "b" -> 2.0))) } test("ZRANK WITHSCORE decodes the rank/score pair or None") { val reply = Frame.Array(Vector(Frame.Integer(2), Frame.Double(5.0))) - assertEquals(Reply.run(SortedSets.zRankWithScore[String, String]("z", "a"), reply), Right(Some((2L, 5.0)))) - assertEquals(Reply.run(SortedSets.zRankWithScore[String, String]("z", "a"), Frame.Null), Right(None)) + assertEquals(Reply.decode(SortedSets.zRankWithScore[String, String]("z", "a"), reply).toEither, Right(Some((2L, 5.0)))) + assertEquals(Reply.decode(SortedSets.zRankWithScore[String, String]("z", "a"), Frame.Null).toEither, Right(None)) } test("ZMPOP decodes null as None and a key with its scored members") { - assertEquals(Reply.run(SortedSets.zMpop[String, String]("a", "b")(MinMax.Min), Frame.Null), Right(None)) + assertEquals(Reply.decode(SortedSets.zMpop[String, String]("a", "b")(MinMax.Min), Frame.Null).toEither, Right(None)) val reply = Frame.Array(Vector(bulk("a"), Frame.Array(Vector(pair("x", 1.0))))) - assertEquals(Reply.run(SortedSets.zMpop[String, String]("a", "b")(MinMax.Min), reply), Right(Some(("a", Vector("x" -> 1.0))))) + assertEquals(Reply.decode(SortedSets.zMpop[String, String]("a", "b")(MinMax.Min), reply).toEither, Right(Some(("a", Vector("x" -> 1.0))))) } test("BZPOPMIN decodes null as None and the key, member, score triple") { - assertEquals(Reply.run(SortedSets.bzPopMin[String, String]("a")(BlockTimeout.Forever), Frame.Null), Right(None)) + assertEquals(Reply.decode(SortedSets.bzPopMin[String, String]("a")(BlockTimeout.Forever), Frame.Null).toEither, Right(None)) val reply = Frame.Array(Vector(bulk("a"), bulk("m"), Frame.Double(1.5))) - assertEquals(Reply.run(SortedSets.bzPopMin[String, String]("a")(BlockTimeout.Forever), reply), Right(Some(("a", "m", 1.5)))) + assertEquals(Reply.decode(SortedSets.bzPopMin[String, String]("a")(BlockTimeout.Forever), reply).toEither, Right(Some(("a", "m", 1.5)))) } test("ZSCAN decodes the flat member/score array with string scores, infinity included") { val reply = Frame.Array(Vector(bulk("0"), Frame.Array(Vector(bulk("a"), bulk("1.5"), bulk("b"), bulk("inf"))))) - Reply.run(SortedSets.zScan[String, String]("z", ScanCursor.start), reply) match { + Reply.decode(SortedSets.zScan[String, String]("z", ScanCursor.start), reply).toEither match { case Right(page) => assertEquals(page.items, Vector("a" -> 1.5, "b" -> Double.PositiveInfinity)) assertEquals(page.next, None) diff --git a/sage-core/src/test/scala/sage/commands/StreamsSpec.scala b/sage-core/src/test/scala/sage/commands/StreamsSpec.scala index 9aff088d..bab8f21b 100644 --- a/sage-core/src/test/scala/sage/commands/StreamsSpec.scala +++ b/sage-core/src/test/scala/sage/commands/StreamsSpec.scala @@ -12,15 +12,33 @@ class StreamsSpec extends munit.FunSuite { private def entry(id: String, fields: String*): Frame = Frame.Array(Vector(bulk(id), Frame.Array(fields.toVector.map(bulk)))) test("XADD decodes the generated id; NOMKSTREAM decodes null as None") { - assertEquals(Reply.run(Streams.xAdd("k")(("f", "v")), bulk("1526919030474-55")), Right(StreamId(1526919030474L, 55L))) - assertEquals(Reply.run(Streams.xAddNoMkStream("k")(("f", "v")), Frame.Null), Right(None)) - assertEquals(Reply.run(Streams.xAddNoMkStream("k")(("f", "v")), bulk("5-0")), Right(Some(StreamId(5L, 0L)))) + assertEquals(Reply.decode(Streams.xAdd("k")(("f", "v")), bulk("1526919030474-55")).toEither, Right(StreamId(1526919030474L, 55L))) + assertEquals(Reply.decode(Streams.xAddNoMkStream("k")(("f", "v")), Frame.Null).toEither, Right(None)) + assertEquals(Reply.decode(Streams.xAddNoMkStream("k")(("f", "v")), bulk("5-0")).toEither, Right(Some(StreamId(5L, 0L)))) + } + + test("stream ids are unsigned 64-bit numbers when decoded and encoded") { + val max = StreamId(-1L, -1L) + assertEquals(Reply.decode(Streams.xAdd("k")(("f", "v")), bulk("18446744073709551615-18446744073709551615")).toEither, Right(max)) + assertEquals(Streams.xDel("k")(max).args(1).asUtf8String, "18446744073709551615-18446744073709551615") + assertEquals( + Streams.xRange[String, String, String]("k", StreamRangeId.Exclusive(max)).args(1).asUtf8String, + "(18446744073709551615-18446744073709551615" + ) + assertEquals(Streams.xAdd("k", XAddId.AutoSeq(-1L))(("f", "v")).args.map(_.asUtf8String).contains("18446744073709551615-*"), true) + } + + test("XCFGSET rounds IDMP-DURATION up to whole seconds so ids are kept at least as long as requested") { + def wire(duration: FiniteDuration) = Streams.xCfgSet("s", idmpDuration = Some(duration)).args(2).asUtf8String + assertEquals(wire(FiniteDuration(500, TimeUnit.MILLISECONDS)), "1") + assertEquals(wire(FiniteDuration(1500, TimeUnit.MILLISECONDS)), "2") + assertEquals(wire(FiniteDuration(3, TimeUnit.SECONDS)), "3") } test("XRANGE decodes entries, preserving field order") { val reply = Frame.Array(Vector(entry("1-0", "a", "1", "b", "2"), entry("2-0", "c", "3"))) assertEquals( - Reply.run(Streams.xRange[String, String, String]("k"), reply), + Reply.decode(Streams.xRange[String, String, String]("k"), reply).toEither, Right(Vector(StreamEntry(StreamId(1L, 0L), Vector("a" -> "1", "b" -> "2")), StreamEntry(StreamId(2L, 0L), Vector("c" -> "3")))) ) } @@ -28,16 +46,16 @@ class StreamsSpec extends munit.FunSuite { test("XREAD decodes a RESP3 map of stream -> entries, and null/timeout as an empty vector") { val reply = map("s1" -> Frame.Array(Vector(entry("1-0", "f", "v")))) assertEquals( - Reply.run(Streams.xRead[String, String, String](("s1", ReadId.New))(), reply), + Reply.decode(Streams.xRead[String, String, String](("s1", ReadId.New))(), reply).toEither, Right(Vector("s1" -> Vector(StreamEntry(StreamId(1L, 0L), Vector("f" -> "v"))))) ) - assertEquals(Reply.run(Streams.xRead[String, String, String](("s1", ReadId.New))(), Frame.Null), Right(Vector.empty)) + assertEquals(Reply.decode(Streams.xRead[String, String, String](("s1", ReadId.New))(), Frame.Null).toEither, Right(Vector.empty)) } test("XREAD also accepts the RESP2 array-of-pairs shape") { val reply = Frame.Array(Vector(Frame.Array(Vector(bulk("s1"), Frame.Array(Vector(entry("1-0", "f", "v"))))))) assertEquals( - Reply.run(Streams.xRead[String, String, String](("s1", ReadId.New))(), reply), + Reply.decode(Streams.xRead[String, String, String](("s1", ReadId.New))(), reply).toEither, Right(Vector("s1" -> Vector(StreamEntry(StreamId(1L, 0L), Vector("f" -> "v"))))) ) } @@ -45,12 +63,12 @@ class StreamsSpec extends munit.FunSuite { test("XAUTOCLAIM decodes the cursor/entries/deleted triple and the pre-7.0 two-element form") { val three = Frame.Array(Vector(bulk("5-0"), Frame.Array(Vector(entry("1-0", "f", "v"))), Frame.Array(Vector(bulk("2-0"))))) assertEquals( - Reply.run(Streams.xAutoClaim[String, String, String]("k", "g", "c", FiniteDuration(1, TimeUnit.SECONDS)), three), + Reply.decode(Streams.xAutoClaim[String, String, String]("k", "g", "c", FiniteDuration(1, TimeUnit.SECONDS)), three).toEither, Right(XAutoClaimResult(StreamId(5L, 0L), Vector(StreamEntry(StreamId(1L, 0L), Vector("f" -> "v"))), Vector(StreamId(2L, 0L)))) ) val two = Frame.Array(Vector(bulk("0-0"), Frame.Array(Vector(entry("1-0", "f", "v"))))) assertEquals( - Reply.run(Streams.xAutoClaim[String, String, String]("k", "g", "c", FiniteDuration(1, TimeUnit.SECONDS)), two), + Reply.decode(Streams.xAutoClaim[String, String, String]("k", "g", "c", FiniteDuration(1, TimeUnit.SECONDS)), two).toEither, Right(XAutoClaimResult(StreamId(0L, 0L), Vector(StreamEntry(StreamId(1L, 0L), Vector("f" -> "v"))), Vector.empty)) ) } @@ -58,7 +76,7 @@ class StreamsSpec extends munit.FunSuite { test("XAUTOCLAIM tolerates a tombstone entry [id, nil] as an entry with no fields") { val reply = Frame.Array(Vector(bulk("0-0"), Frame.Array(Vector(Frame.Array(Vector(bulk("1-0"), Frame.Null)))), Frame.Array(Vector.empty))) assertEquals( - Reply.run(Streams.xAutoClaim[String, String, String]("k", "g", "c", FiniteDuration(1, TimeUnit.SECONDS)), reply), + Reply.decode(Streams.xAutoClaim[String, String, String]("k", "g", "c", FiniteDuration(1, TimeUnit.SECONDS)), reply).toEither, Right(XAutoClaimResult(StreamId(0L, 0L), Vector(StreamEntry(StreamId(1L, 0L), Vector.empty)), Vector.empty)) ) } @@ -73,17 +91,17 @@ class StreamsSpec extends munit.FunSuite { ) ) assertEquals( - Reply.run(Streams.xPending("k", "g"), populated), + Reply.decode(Streams.xPending("k", "g"), populated).toEither, Right(PendingSummary(2L, Some(StreamId(1L, 0L)), Some(StreamId(9L, 0L)), Vector("c1" -> 1L, "c2" -> 1L))) ) val empty = Frame.Array(Vector(Frame.Integer(0L), Frame.Null, Frame.Null, Frame.Null)) - assertEquals(Reply.run(Streams.xPending("k", "g"), empty), Right(PendingSummary(0L, None, None, Vector.empty))) + assertEquals(Reply.decode(Streams.xPending("k", "g"), empty).toEither, Right(PendingSummary(0L, None, None, Vector.empty))) } test("XPENDING extended decodes a row with its idle time and delivery count") { val reply = Frame.Array(Vector(Frame.Array(Vector(bulk("1-0"), bulk("c1"), Frame.Integer(5000L), Frame.Integer(3L))))) assertEquals( - Reply.run(Streams.xPendingExtended("k", "g"), reply), + Reply.decode(Streams.xPendingExtended("k", "g"), reply).toEither, Right(Vector(PendingEntry(StreamId(1L, 0L), "c1", FiniteDuration(5000L, TimeUnit.MILLISECONDS), 3L))) ) } @@ -91,7 +109,7 @@ class StreamsSpec extends munit.FunSuite { test("XDELEX decodes the per-id deletion status") { val reply = Frame.Array(Vector(Frame.Integer(1L), Frame.Integer(-1L), Frame.Integer(2L))) assertEquals( - Reply.run(Streams.xDelEx("k")(StreamId(1L, 0L), StreamId(2L, 0L), StreamId(3L, 0L)), reply), + Reply.decode(Streams.xDelEx("k")(StreamId(1L, 0L), StreamId(2L, 0L), StreamId(3L, 0L)), reply).toEither, Right(Vector(StreamEntryDeletion.Deleted, StreamEntryDeletion.NotFound, StreamEntryDeletion.Retained)) ) } @@ -109,7 +127,7 @@ class StreamsSpec extends munit.FunSuite { "first-entry" -> entry("1-0", "f", "v"), "last-entry" -> entry("5-0", "g", "w") ) - val info = Reply.run(StreamInfo.xInfoStream[String, String, String]("k"), withNew).toOption.get + val info = Reply.decode(StreamInfo.xInfoStream[String, String, String]("k"), withNew).toEither.toOption.get assertEquals(info.length, 2L) assertEquals(info.entriesAdded, Some(7L)) assertEquals(info.maxDeletedEntryId, Some(StreamId(3L, 0L))) @@ -124,7 +142,7 @@ class StreamsSpec extends munit.FunSuite { "first-entry" -> Frame.Null, "last-entry" -> Frame.Null ) - val old = Reply.run(StreamInfo.xInfoStream[String, String, String]("k"), legacy).toOption.get + val old = Reply.decode(StreamInfo.xInfoStream[String, String, String]("k"), legacy).toEither.toOption.get assertEquals(old.entriesAdded, None) assertEquals(old.maxDeletedEntryId, None) assertEquals(old.firstEntry, None) @@ -144,7 +162,7 @@ class StreamsSpec extends munit.FunSuite { ) ) assertEquals( - Reply.run(StreamInfo.xInfoGroups("k"), reply), + Reply.decode(StreamInfo.xInfoGroups("k"), reply).toEither, Right(Vector(GroupInfo("g1", 2L, 3L, StreamId(5L, 0L), Some(5L), None))) ) } @@ -168,7 +186,7 @@ class StreamsSpec extends munit.FunSuite { "entries" -> Frame.Array(Vector(entry("1-0", "f", "v"))), "groups" -> Frame.Array(Vector(group)) ) - val full = Reply.run(StreamInfo.xInfoStreamFull[String, String, String]("k"), reply).toOption.get + val full = Reply.decode(StreamInfo.xInfoStreamFull[String, String, String]("k"), reply).toEither.toOption.get assertEquals(full.entries.map(_.id), Vector(StreamId(1L, 0L))) val g = full.groups.head assertEquals(g.pending, Vector(FullPendingEntry(StreamId(1L, 0L), Some("c1"), java.time.Instant.ofEpochMilli(1000L), 2L))) diff --git a/sage-core/src/test/scala/sage/commands/StringsSpec.scala b/sage-core/src/test/scala/sage/commands/StringsSpec.scala index 325c668b..1f488570 100644 --- a/sage-core/src/test/scala/sage/commands/StringsSpec.scala +++ b/sage-core/src/test/scala/sage/commands/StringsSpec.scala @@ -8,17 +8,17 @@ import sage.protocol.Frames.bulk class StringsSpec extends munit.FunSuite { test("GET decodes a present value as Some and a missing key as None") { - assertEquals(Reply.run(Strings.get[String, String]("k"), bulk("v")), Right(Some("v"))) - assertEquals(Reply.run(Strings.get[String, String]("k"), Frame.Null), Right(None)) + assertEquals(Reply.decode(Strings.get[String, String]("k"), bulk("v")).toEither, Right(Some("v"))) + assertEquals(Reply.decode(Strings.get[String, String]("k"), Frame.Null).toEither, Right(None)) } test("SET decodes +OK as true and null as false") { - assertEquals(Reply.run(Strings.set("k", "v"), Frame.SimpleString("OK")), Right(true)) - assertEquals(Reply.run(Strings.set("k", "v", condition = SetCondition.IfNotExists), Frame.Null), Right(false)) + assertEquals(Reply.decode(Strings.set("k", "v"), Frame.SimpleString("OK")).toEither, Right(true)) + assertEquals(Reply.decode(Strings.set("k", "v", condition = SetCondition.IfNotExists), Frame.Null).toEither, Right(false)) } test("SET rejects an unexpected frame naming expected and actual") { - Reply.run(Strings.set("k", "v"), Frame.Integer(1)) match { + Reply.decode(Strings.set("k", "v"), Frame.Integer(1)).toEither match { case Left(error: DecodeError) => assertEquals(error.expected, "simple string 'OK' or null") assertEquals(error.actual, "integer 1") @@ -27,18 +27,18 @@ class StringsSpec extends munit.FunSuite { } test("setGet decodes the previous value and null when the key was absent") { - assertEquals(Reply.run(Strings.setGet("k", "v"), bulk("old")), Right(Some("old"))) - assertEquals(Reply.run(Strings.setGet[String, String]("k", "v"), Frame.Null), Right(None)) + assertEquals(Reply.decode(Strings.setGet("k", "v"), bulk("old")).toEither, Right(Some("old"))) + assertEquals(Reply.decode(Strings.setGet[String, String]("k", "v"), Frame.Null).toEither, Right(None)) } test("MGET decodes positionally with None for missing keys") { val reply = Frame.Array(Vector(bulk("1"), Frame.Null, bulk("3"))) - assertEquals(Reply.run(Strings.mGet[String, String]("a", "b", "c"), reply), Right(Vector(Some("1"), None, Some("3")))) + assertEquals(Reply.decode(Strings.mGet[String, String]("a", "b", "c"), reply).toEither, Right(Vector(Some("1"), None, Some("3")))) } test("MGET propagates an element decode failure") { val reply = Frame.Array(Vector(bulk("1"), Frame.Integer(2))) - assert(Reply.run(Strings.mGet[String, String]("a", "b"), reply).isLeft) + assert(Reply.decode(Strings.mGet[String, String]("a", "b"), reply).toEither.isLeft) } test("MGET and MSET mark key positions for the slot engine") { @@ -48,20 +48,20 @@ class StringsSpec extends munit.FunSuite { } test("INCRBYFLOAT decodes the float bulk string reply") { - assertEquals(Reply.run(Strings.incrByFloat("k", 0.1), bulk("3.0e3")), Right(3000.0)) - Reply.run(Strings.incrByFloat("k", 0.1), bulk("abc")) match { + assertEquals(Reply.decode(Strings.incrByFloat("k", 0.1), bulk("3.0e3")).toEither, Right(3000.0)) + Reply.decode(Strings.incrByFloat("k", 0.1), bulk("abc")).toEither match { case Left(error: DecodeError) => assertEquals(error.actual, "bulk string 'abc'") case other => fail(s"expected a DecodeError, got $other") } } test("GETRANGE decodes an empty bulk string for a missing key") { - assertEquals(Reply.run(Strings.getRange[String, String]("k", 0L, 4L), Frame.BulkString(Bytes.empty)), Right("")) + assertEquals(Reply.decode(Strings.getRange[String, String]("k", 0L, 4L), Frame.BulkString(Bytes.empty)).toEither, Right("")) } test("MSETNX decodes the flag and rejects other integers") { - assertEquals(Reply.run(Strings.mSetNx(("a", "1")), Frame.Integer(1)), Right(true)) - assertEquals(Reply.run(Strings.mSetNx(("a", "1")), Frame.Integer(0)), Right(false)) - assert(Reply.run(Strings.mSetNx(("a", "1")), Frame.Integer(2)).isLeft) + assertEquals(Reply.decode(Strings.mSetNx(("a", "1")), Frame.Integer(1)).toEither, Right(true)) + assertEquals(Reply.decode(Strings.mSetNx(("a", "1")), Frame.Integer(0)).toEither, Right(false)) + assert(Reply.decode(Strings.mSetNx(("a", "1")), Frame.Integer(2)).toEither.isLeft) } } diff --git a/sage-core/src/test/scala/sage/protocol/Frames.scala b/sage-core/src/test/scala/sage/protocol/Frames.scala index b7e2ed81..17fc84c7 100644 --- a/sage-core/src/test/scala/sage/protocol/Frames.scala +++ b/sage-core/src/test/scala/sage/protocol/Frames.scala @@ -1,6 +1,7 @@ package sage.protocol import sage.Bytes +import sage.SageException.ProtocolError /** * Frame constructors the specs share. @@ -9,5 +10,13 @@ object Frames { def bulk(value: String): Frame = Frame.BulkString(Bytes.utf8(value)) + extension (parser: RespParser) { + def feed(bytes: Bytes): Either[ProtocolError, Vector[Frame]] = { + val frames = Vector.newBuilder[Frame] + val array = bytes.unsafeArray + parser.feed(array, 0, array.length)(frames += _).toLeft(frames.result()) + } + } + def map(entries: (String, Frame)*): Frame = Frame.Map(entries.toVector.map { case (key, value) => bulk(key) -> value }) } diff --git a/sage-core/src/test/scala/sage/protocol/RespParserSpec.scala b/sage-core/src/test/scala/sage/protocol/RespParserSpec.scala index 9feb7a64..18609914 100644 --- a/sage-core/src/test/scala/sage/protocol/RespParserSpec.scala +++ b/sage-core/src/test/scala/sage/protocol/RespParserSpec.scala @@ -4,6 +4,7 @@ import scala.util.Random import sage.Bytes import sage.SageException.ProtocolError +import sage.protocol.Frames.feed class RespParserSpec extends munit.FunSuite { @@ -266,6 +267,10 @@ class RespParserSpec extends munit.FunSuite { assert(parseError("?\r\n").message.contains("unknown frame type")) } + test("rejects an unknown frame type byte before its line ends") { + assert(parseError("\u0015\u0003").message.contains("unknown frame type byte 0x15")) + } + test("rejects an invalid integer") { assert(parseError(":abc\r\n").message.contains("invalid integer")) } diff --git a/sage-core/src/test/scala/sage/ratelimit/RateLimiterSpec.scala b/sage-core/src/test/scala/sage/ratelimit/RateLimiterSpec.scala index f69a66ef..a30919f0 100644 --- a/sage-core/src/test/scala/sage/ratelimit/RateLimiterSpec.scala +++ b/sage-core/src/test/scala/sage/ratelimit/RateLimiterSpec.scala @@ -65,20 +65,20 @@ class RateLimiterSpec extends munit.FunSuite { assert(command.args.head.sameBytes(Bytes.utf8("9:ratelimit:u"))) } - test("evalSha builds an EVALSHA command carrying the script sha") { - val command = limiter.evalSha("u", 1) + test("a cached eval builds an EVALSHA command carrying the script sha") { + val command = limiter.eval(cached = true, "u", 1, peek = false) assertEquals(command.name, "EVALSHA") - assertEquals(command.args.head.asUtf8String, RateLimiter.sha) + assertEquals(command.args.head.asUtf8String, RateLimiter.compiled.sha) assertEquals(command.keyIndices, Vector(2)) } - test("evalScript is a keyed EVAL carrying the script body, and sha is a 40-char hex digest") { - val command = limiter.evalScript("u", 1, peek = false) + test("an uncached eval is a keyed EVAL carrying the script body, and sha is a 40-char hex digest") { + val command = limiter.eval(cached = false, "u", 1, peek = false) assertEquals(command.name, "EVAL") assertEquals(command.args.head.asUtf8String, RateLimiter.script) assertEquals(command.keyIndices, Vector(2)) - assertEquals(RateLimiter.sha.length, 40) - assert(RateLimiter.sha.forall(c => c.isDigit || ('a' to 'f').contains(c))) + assertEquals(RateLimiter.compiled.sha.length, 40) + assert(RateLimiter.compiled.sha.forall(c => c.isDigit || ('a' to 'f').contains(c))) } test("a Decision reports whether it admitted and how many tokens are left") { diff --git a/sage-opentelemetry/src/main/scala/sage/opentelemetry/OpenTelemetryCommandTracer.scala b/sage-opentelemetry/src/main/scala/sage/opentelemetry/OpenTelemetryCommandTracer.scala index 33ba38b7..6f4c1c98 100644 --- a/sage-opentelemetry/src/main/scala/sage/opentelemetry/OpenTelemetryCommandTracer.scala +++ b/sage-opentelemetry/src/main/scala/sage/opentelemetry/OpenTelemetryCommandTracer.scala @@ -1,7 +1,5 @@ package sage.opentelemetry -import java.util.concurrent.atomic.AtomicBoolean - import io.opentelemetry.api.{GlobalOpenTelemetry, OpenTelemetry} import io.opentelemetry.api.common.AttributeKey import io.opentelemetry.api.trace.{SpanKind, StatusCode, Tracer} @@ -61,7 +59,7 @@ object OpenTelemetryCommandTracer { * in-memory SDK. */ def apply(openTelemetry: OpenTelemetry, peerService: String = "redis"): CommandTracer = - new OpenTelemetryCommandTracer(openTelemetry.getTracer("sage"), peerService, () => Context.current()) + withContextProvider(openTelemetry, peerService, () => Context.current()) /** * Builds a tracer from the registered global `OpenTelemetry` instance. Use this method with an APM agent that registers itself globally, @@ -79,23 +77,20 @@ object OpenTelemetryCommandTracer { final private class Span(span: io.opentelemetry.api.trace.Span) extends CommandSpan { - // a fast failure can race with a late callback. End the span only for the first outcome. - private val ended = new AtomicBoolean(false) - def routedTo(node: Node): Unit = { span.setAttribute(ServerAddress, node.host) span.setAttribute(ServerPort, node.port.toLong): Unit } - def settled(outcome: Outcome): Unit = - if (ended.compareAndSet(false, true)) { - outcome match { - case Outcome.Succeeded => () - case Outcome.Failed(err) => - span.setStatus(StatusCode.ERROR, Option(err.getMessage).getOrElse(err.getClass.getName)) - span.recordException(err) - } - span.end() + // OpenTelemetry ignores calls on a span after it has ended, so a repeated call changes nothing. + def settled(outcome: Outcome): Unit = { + outcome match { + case Outcome.Succeeded => () + case Outcome.Failed(err) => + span.setStatus(StatusCode.ERROR, Option(err.getMessage).getOrElse(err.getClass.getName)) + span.recordException(err) } + span.end() + } } } diff --git a/sage-opentelemetry/src/test/scala/sage/opentelemetry/OpenTelemetryCommandTracerSpec.scala b/sage-opentelemetry/src/test/scala/sage/opentelemetry/OpenTelemetryCommandTracerSpec.scala index bbd8fd1a..7ef1adde 100644 --- a/sage-opentelemetry/src/test/scala/sage/opentelemetry/OpenTelemetryCommandTracerSpec.scala +++ b/sage-opentelemetry/src/test/scala/sage/opentelemetry/OpenTelemetryCommandTracerSpec.scala @@ -8,38 +8,40 @@ import io.opentelemetry.sdk.OpenTelemetrySdk import io.opentelemetry.sdk.testing.exporter.InMemorySpanExporter import io.opentelemetry.sdk.trace.`export`.SimpleSpanProcessor import io.opentelemetry.sdk.trace.SdkTracerProvider +import io.opentelemetry.sdk.trace.data.SpanData -import sage.Outcome +import sage.{CommandTracer, Outcome} import sage.cluster.Node import sage.commands.Command class OpenTelemetryCommandTracerSpec extends munit.FunSuite { // a fresh in-memory SDK + tracer per test, so finished spans never leak across tests - private def fixture: (OpenTelemetrySdk, InMemorySpanExporter) = { - val exporter = InMemorySpanExporter.create() - val sdk = OpenTelemetrySdk - .builder() - .setTracerProvider(SdkTracerProvider.builder().addSpanProcessor(SimpleSpanProcessor.create(exporter)).build()) - .build() - (sdk, exporter) + private class Traced(peerService: String = "redis") { + private val exporter = InMemorySpanExporter.create() + val sdk: OpenTelemetrySdk = + OpenTelemetrySdk.builder().setTracerProvider(SdkTracerProvider.builder().addSpanProcessor(SimpleSpanProcessor.create(exporter)).build()).build() + val tracer: CommandTracer = OpenTelemetryCommandTracer(sdk, peerService) + + def onlySpan: SpanData = + exporter.getFinishedSpanItems.asScala.toList match { + case List(span) => span + case spans => fail(s"expected one finished span, got $spans") + } } - private def get(span: io.opentelemetry.sdk.trace.data.SpanData, key: String): String = + private def get(span: SpanData, key: String): String = span.getAttributes.get(AttributeKey.stringKey(key)) private def command(name: String): Command[Unit] = Command(name, Command.NoKeys, Vector.empty, _ => Right(())) test("a successful command yields one CLIENT span named for the command, with the Datadog-parity attributes") { - val (sdk, exporter) = fixture - val tracer = OpenTelemetryCommandTracer(sdk) + val t = Traced() - tracer.onCommand(command("GET")).settled(Outcome.Succeeded) + t.tracer.onCommand(command("GET")).settled(Outcome.Succeeded) - val spans = exporter.getFinishedSpanItems.asScala.toList - assertEquals(spans.size, 1) - val span = spans.head + val span = t.onlySpan assertEquals(span.getName, "GET") assertEquals(span.getKind, SpanKind.CLIENT) assertEquals(get(span, "db.system"), "redis") @@ -50,83 +52,66 @@ class OpenTelemetryCommandTracerSpec extends munit.FunSuite { } test("peerService is configurable, collapsing a cluster's nodes into one dependency") { - val (sdk, exporter) = fixture - val tracer = OpenTelemetryCommandTracer(sdk, peerService = "my-redis") + val t = Traced("my-redis") - tracer.onCommand(command("SET")).settled(Outcome.Succeeded) + t.tracer.onCommand(command("SET")).settled(Outcome.Succeeded) - assertEquals(get(exporter.getFinishedSpanItems.asScala.head, "peer.service"), "my-redis") + assertEquals(get(t.onlySpan, "peer.service"), "my-redis") } test("routedTo sets the server address and port") { - val (sdk, exporter) = fixture - val tracer = OpenTelemetryCommandTracer(sdk) + val t = Traced() - val span = tracer.onCommand(command("GET")) + val span = t.tracer.onCommand(command("GET")) span.routedTo(Node("redis.internal", 6380)) span.settled(Outcome.Succeeded) - val data = exporter.getFinishedSpanItems.asScala.head + val data = t.onlySpan assertEquals(get(data, "server.address"), "redis.internal") assertEquals(data.getAttributes.get(AttributeKey.longKey("server.port")), java.lang.Long.valueOf(6380L)) } test("a failure records ERROR status and the exception") { - val (sdk, exporter) = fixture - val tracer = OpenTelemetryCommandTracer(sdk) + val t = Traced() - tracer.onCommand(command("GET")).settled(Outcome.Failed(new RuntimeException("boom"))) + t.tracer.onCommand(command("GET")).settled(Outcome.Failed(new RuntimeException("boom"))) - val data = exporter.getFinishedSpanItems.asScala.head + val data = t.onlySpan assertEquals(data.getStatus.getStatusCode, StatusCode.ERROR) assert(data.getEvents.asScala.exists(_.getName == "exception"), "expected a recorded exception event") } test("the span nests under the active context's span") { - val (sdk, exporter) = fixture - val tracer = OpenTelemetryCommandTracer(sdk) - val parent = sdk.getTracer("test").spanBuilder("request").startSpan() + val t = Traced() + val parent = t.sdk.getTracer("test").spanBuilder("request").startSpan() val scope = parent.makeCurrent() - try tracer.onCommand(command("GET")).settled(Outcome.Succeeded) + try t.tracer.onCommand(command("GET")).settled(Outcome.Succeeded) finally scope.close() - parent.end() - val spans = exporter.getFinishedSpanItems.asScala.toList - val redis = spans.find(_.getName == "GET").get - val server = spans.find(_.getName == "request").get - assertEquals(redis.getParentSpanContext.getSpanId, server.getSpanContext.getSpanId) - assertEquals(redis.getTraceId, server.getTraceId) + assertEquals(t.onlySpan.getParentSpanContext, parent.getSpanContext) } test("prepare captures the parent context up front, so a span started after the context is gone still nests under it") { - val (sdk, exporter) = fixture - val tracer = OpenTelemetryCommandTracer(sdk) - val parent = sdk.getTracer("test").spanBuilder("request").startSpan() + val t = Traced() + val parent = t.sdk.getTracer("test").spanBuilder("request").startSpan() // capture while the parent is current, then start the span after the scope is closed (mimicking a fetch on an offload worker) val scope = parent.makeCurrent() - val startSpan = tracer.prepare(command("GET")) + val startSpan = t.tracer.prepare(command("GET")) scope.close() startSpan().settled(Outcome.Succeeded) - parent.end() - val spans = exporter.getFinishedSpanItems.asScala.toList - val redis = spans.find(_.getName == "GET").get - assertEquals(redis.getParentSpanContext.getSpanId, parent.getSpanContext.getSpanId) - assertEquals(redis.getTraceId, parent.getSpanContext.getTraceId) + assertEquals(t.onlySpan.getParentSpanContext, parent.getSpanContext) } test("settling twice ends the span only once") { - val (sdk, exporter) = fixture - val tracer = OpenTelemetryCommandTracer(sdk) + val t = Traced() - val span = tracer.onCommand(command("GET")) + val span = t.tracer.onCommand(command("GET")) span.settled(Outcome.Succeeded) span.settled(Outcome.Failed(new RuntimeException("late"))) - val spans = exporter.getFinishedSpanItems.asScala.toList - assertEquals(spans.size, 1) - assertEquals(spans.head.getStatus.getStatusCode, StatusCode.UNSET) + assertEquals(t.onlySpan.getStatus.getStatusCode, StatusCode.UNSET) } }