Skip to content
12 changes: 12 additions & 0 deletions benchmark/src/main/scala/fs2/benchmark/GroupWithinBenchmark.scala
Original file line number Diff line number Diff line change
Expand Up @@ -59,4 +59,16 @@ class GroupWithinBenchmark {
.compile
.drain
.unsafeRunSync()

@Benchmark
def groupWithinChunkedUpstream(): Unit =
Stream
.range(0, rangeLength)
.chunkN(bufferSize / 4 + 1)
.unchunks
.covary[IO]
.groupWithin(bufferSize, bufferWindow)
.compile
.drain
.unsafeRunSync()
}
24 changes: 24 additions & 0 deletions core/shared/src/main/scala/fs2/Chunk.scala
Original file line number Diff line number Diff line change
Expand Up @@ -207,6 +207,30 @@ abstract class Chunk[+O] extends Serializable with ChunkPlatform[O] with ChunkRu
}
}

/** Splits this chunk into groups of `n` elements, the last of which may be smaller.
*
* Like `List#grouped`, an empty chunk has no groups, and `n` must be positive.
*/
def grouped(n: Int): Chunk[Chunk[O]] = {
require(n > 0, s"n must be positive, but got ${n.toString}")

if (isEmpty) Chunk.empty
else {
val numGroups = (size - 1) / n + 1 // ceil divide
Comment thread
mmienko marked this conversation as resolved.
val groups = new Array[Chunk[O]](numGroups)
var i = 0
var chunk = this
while (i < numGroups) {
val (group, remaining) = chunk.splitAt(n)
groups(i) = group
chunk = remaining
i += 1
}

Chunk.array(groups)
}
}

/** Gets the first element of this chunk. */
def head: Option[O] = if (isEmpty) None else Some(apply(0))

Expand Down
163 changes: 67 additions & 96 deletions core/shared/src/main/scala/fs2/Stream.scala
Original file line number Diff line number Diff line change
Expand Up @@ -1487,13 +1487,16 @@ final class Stream[+F[_], +O] private[fs2] (private[fs2] val underlying: Pull[F,
go(None, this).stream
}

/** Splits this stream into a stream of chunks of elements, such that
/** Splits this stream into a stream of chunks of elements, such that:
*
* 1. each chunk in the output has at most `outputSize` elements, and
*
* 2. the concatenation of those chunks, which is obtained by calling
* `unchunks`, yields the same element sequence as this stream.
*
* As `this` stream emits input elements, the result stream them in a
* waiting buffer, until it has enough elements to emit next chunk.
* As `this` stream emits input elements, the result stream accumulates
* them in a waiting buffer, until it has enough elements to emit the
* next chunk. Accumulation acts on chunks of input elements for efficiency.
*
* To avoid holding input elements for too long, this method takes a
* `timeout`. This timeout is reset after each output chunk is emitted.
Expand All @@ -1520,110 +1523,78 @@ final class Stream[+F[_], +O] private[fs2] (private[fs2] val underlying: Pull[F,
def groupWithin[F2[x] >: F[x]](
chunkSize: Int,
timeout: FiniteDuration
)(implicit F: Temporal[F2]): Stream[F2, Chunk[O]] = {

case class JunctionBuffer[T](
data: Vector[T],
endOfSupply: Option[Either[Throwable, Unit]],
endOfDemand: Option[Either[Throwable, Unit]]
) {
def splitAt(n: Int): (JunctionBuffer[T], JunctionBuffer[T]) =
if (this.data.size >= n) {
val (head, tail) = this.data.splitAt(n.toInt)
(this.copy(tail), this.copy(head))
} else {
(this.copy(Vector.empty), this)
}
}

val outputLong = chunkSize.toLong
fs2.Stream.force {
for {
demand <- Semaphore[F2](outputLong)
supply <- Semaphore[F2](0L)
buffer <- Ref[F2].of(
JunctionBuffer[O](Vector.empty[O], endOfSupply = None, endOfDemand = None)
)
} yield {
/* - Buffer: stores items from input to be sent on next output chunk
* - Demand Semaphore: to avoid adding too many items to buffer
* - Supply: counts filled positions for next output chunk */
def enqueue(t: O): F2[Boolean] =
for {
_ <- demand.acquire
buf <- buffer.modify(buf => (buf.copy(buf.data :+ t), buf))
_ <- supply.release
} yield buf.endOfDemand.isEmpty

val dequeueNextOutput: F2[Option[Vector[O]]] = {
// Trigger: waits until the supply buffer is full (with acquireN)
val waitSupply = supply.acquireN(outputLong).guaranteeCase {
case Outcome.Succeeded(_) => supply.releaseN(outputLong)
case _ => F.unit
}
)(implicit F: Temporal[F2]): Stream[F2, Chunk[O]] =
Stream.force {
require(chunkSize > 0, s"chunkSize must be > 0, but got ${chunkSize.toString}")

val onTimeout: F2[Long] =
for {
_ <- supply.acquire // waits until there is at least one element in buffer
m <- supply.available
k = m.min(outputLong - 1)
b <- supply.tryAcquireN(k)
} yield if (b) k + 1 else 1

// in JS cancellation doesn't always seem to run, so race conditions should restore state on their own
for {
acq <- F.race(F.sleep(timeout), waitSupply).flatMap {
case Left(_) => onTimeout
case Right(_) => supply.acquireN(outputLong).as(outputLong)
}
buf <- buffer.modify(_.splitAt(acq.toInt))
_ <- demand.releaseN(buf.data.size.toLong)
res <- buf.endOfSupply match {
case Some(Left(error)) => F.raiseError(error)
case Some(Right(_)) if buf.data.isEmpty => F.pure(None)
case _ => F.pure(Some(buf.data))
}
} yield res
}
val StopConsumer = none[Chunk[O]].pure[F2]
val Skip = Chunk.empty[O].some.pure[F2]

def endSupply(result: Either[Throwable, Unit]): F2[Unit] =
buffer.update(_.copy(endOfSupply = Some(result))) *> supply.releaseN(
// enough supply for 2 iterations of the race loop in case of upstream
// interruption: so that downstream can terminate immediately
outputLong * 2
)
final case class Buffer[A](chunk: Chunk[A], done: Option[ExitCase]) {
def size: Int = chunk.size
def isEmpty: Boolean = chunk.isEmpty
def nonEmpty: Boolean = chunk.nonEmpty
def isFull: Boolean = size >= chunkSize
def isDone: Boolean = done.isDefined
}

def endDemand(result: Either[Throwable, Unit]): F2[Unit] =
buffer.update(_.copy(endOfDemand = Some(result))) *> demand.releaseN(Int.MaxValue)
ConditionedRef.of[F2, Buffer[O]](Buffer(chunk = Chunk.empty[O], done = none)).map { buffer =>
Comment thread
mmienko marked this conversation as resolved.
val producer = chunks
.evalMap { chunk =>
buffer
.updateAndGet(b => b.copy(chunk = b.chunk ++ chunk))
.flatMap(b => F.whenA(b.isFull)(buffer.waitUntil(!_.isFull)))
}
.onFinalizeCase(exitCase => buffer.update(_.copy(done = exitCase.some)))
.compile
.drain

def toEnding(ec: ExitCase): Either[Throwable, Unit] = ec match {
case ExitCase.Succeeded => Right(())
case ExitCase.Errored(e) => Left(e)
case ExitCase.Canceled => Right(())
def take(b: Buffer[O], n: Int): (Buffer[O], F2[Option[Chunk[O]]]) = {
val (taken, remaining) = b.chunk.splitAt(n)
b.copy(chunk = remaining) -> taken.some.pure[F2]
}

val enqueueAsync = F.start {
this
.evalMap(enqueue)
.forall(identity)
.onFinalizeCase(ec => endSupply(toEnding(ec)))
.compile
.drain
// None means upstream is done, and the buffer is drained.
def takeOrExit(timedOut: Boolean): F2[Option[Chunk[O]]] =
buffer.modify { b =>
if (b.nonEmpty && b.isDone) take(b, n = b.size)
else if (b.isFull) take(b, n = b.size - b.size % chunkSize)
else if (b.nonEmpty && timedOut) take(b, n = b.size)
else
b -> (b.done match {
case None /* not full & no-timeout */ => Skip
case Some(ExitCase.Errored(e)) => F.raiseError[Option[Chunk[O]]](e)
case Some(_) /* empty */ => StopConsumer
})
}.flatten

val onTimeout = takeOrExit(timedOut = true).flatMap {
case Some(batch) if batch.isEmpty =>
buffer.waitUntil(b => b.nonEmpty || b.isDone) >> takeOrExit(timedOut = true)
case result => result.pure[F2]
}

val outputStream: Stream[F2, Chunk[O]] =
Stream
.eval(dequeueNextOutput)
.repeat
.collectWhile { case Some(data) => Chunk.from(data) }
val nextBatch: F2[Option[Chunk[O]]] =
// Potentially skip starting timer fiber if buffer is full
takeOrExit(timedOut = false).flatMap {
case Some(batch) if batch.isEmpty =>
F.race(F.sleep(timeout), buffer.waitUntil(b => b.isFull || b.isDone))
.flatMap {
case Left(_ /* timeout */ ) => onTimeout
case Right(_ /* full batch or done */ ) => takeOrExit(timedOut = false)
}
case result => result.pure[F2]
}

Stream
.bracketCase(enqueueAsync) { case (upstream, exitCase) =>
endDemand(toEnding(exitCase)) *> upstream.cancel
} >> outputStream
def emitBatches: Pull[F2, Chunk[O], Unit] =
Pull.eval(nextBatch).flatMap {
case Some(batch) => Pull.output(batch.grouped(chunkSize)) >> emitBatches
case None => Pull.done
}

Stream.bracket(producer.start)(_.cancel) >> emitBatches.stream
}
}
}

/** If `this` terminates with `Stream.raiseError(e)`, invoke `h(e)`.
*
Expand Down
110 changes: 110 additions & 0 deletions core/shared/src/main/scala/fs2/concurrent/ConditionedRef.scala
Original file line number Diff line number Diff line change
@@ -0,0 +1,110 @@
/*
* Copyright (c) 2013 Functional Streams for Scala
*
* Permission is hereby granted, free of charge, to any person obtaining a copy of
* this software and associated documentation files (the "Software"), to deal in
* the Software without restriction, including without limitation the rights to
* use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies of
* the Software, and to permit persons to whom the Software is furnished to do so,
* subject to the following conditions:
*
* The above copyright notice and this permission notice shall be included in all
* copies or substantial portions of the Software.
*
* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
* IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS
* FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR
* COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER
* IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN
* CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
*/

package fs2
package concurrent

import cats.effect._
import cats.effect.implicits._
import cats.syntax.all._

/** A `Ref` whose value can be waited on.
*
* Waiters register a predicate on the value. Every `modify` evaluates the registered predicates
* against the new value and wakes exactly the waiters whose predicate is satisfied. Unlike [[SignallingRef]],
* waiters are not woken by unrelated updates.
*/
private[fs2] sealed trait ConditionedRef[F[_], A] {

def get: F[A]

/** Atomically updates the value and wakes every waiter whose predicate condition holds for the new value. */
def modify[B](f: A => (A, B)): F[B]

def update(f: A => A): F[Unit]

def updateAndGet(f: A => A): F[A] =
modify { a =>
val newA = f(a)
(newA, newA)
}

def set(a: A): F[Unit] = update(_ => a)

/** Completes if the predicate, `p` holds for the current value, or a later `modify` sets a value that satisfies `p`.
*
* `p` may no longer hold by the time this completes: act on the value through `modify`.
*/
def waitUntil(p: A => Boolean): F[Unit]
}

private[fs2] object ConditionedRef {

def of[F[_], A](initial: A)(implicit F: Concurrent[F]): F[ConditionedRef[F, A]] =
F.ref(State[F, A](value = initial, waiters = Nil)).map(new Impl(_))

private final class Waiter[F[_], A](val accepts: A => Boolean, val wake: Deferred[F, Unit]) {
def wakeUp: F[Boolean] = wake.complete(())
}

private final case class State[F[_], A](value: A, waiters: List[Waiter[F, A]]) {
def register(waiter: Waiter[F, A]): State[F, A] = copy(waiters = waiter :: waiters)
def deregister(waiter: Waiter[F, A]): State[F, A] =
copy(waiters = waiters.filterNot(_ eq waiter))
}

private final class Impl[F[_], A](state: Ref[F, State[F, A]])(implicit F: Concurrent[F])
extends ConditionedRef[F, A] {

def get: F[A] = state.get.map(_.value)

def modify[B](f: A => (A, B)): F[B] =
state.flatModify { s => // uncancellable to avoid losing wake-up signals
val (value, result) = f(s.value)
// Scan first instead of calling partition to avoid extra allocations which have a small impact on perf.
// This optimizes for the case where writes infrequently trigger wake-ups.
if (!s.waiters.exists(_.accepts(value)))

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Have you tried partitioning first to avoid two traversals of the list?

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We're dealing with a single consumer within groupWithin and the predicate is cheap vs. allocating new lists each time in partition where in practice the first if-clause should execute more often.

Arguably, this could be premature optimization (and I can run quick bechmarks to confirm) or impl should reflect spsc. Open to suggestions, but I'm starting to lean towards a scoped impl of this class for spsc rather than a more generic mpmc. Probably separate impls, ConditiedRef.mpmc, ConditiedRef.spsc (only this would be needed for PR), etc. will avoid confusion for future maintenance too.

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

So there's a small hit when buffer size is small, but overall not very noticable. I'm fine with changing it to call partition first and have a single traversal.

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Ok, ran some more benchmarks and there's a drop in perf in the scenarios where there is no chunking (single elements), so I would prefer to keep it for my use case.

State(value = value, waiters = s.waiters) -> result.pure[F]
else {
val (toWake, waiting) = s.waiters.partition(_.accepts(value))
State(value = value, waiters = waiting) -> toWake.traverse_(_.wakeUp).as(result)
}
}

def update(f: A => A): F[Unit] = modify(a => (f(a), ()))

def waitUntil(p: A => Boolean): F[Unit] =
get.flatMap { value =>
if (p(value)) F.unit
Comment thread
mmienko marked this conversation as resolved.
else
F.deferred[Unit].flatMap { wake =>
val waiter = new Waiter(accepts = p, wake = wake)
state.flatModifyFull { (poll, s) =>
if (p(s.value)) s -> F.unit
else
s.register(waiter) -> poll(wake.get).onCancel {
state.update(_.deregister(waiter))
}
}
}
}
}
}
23 changes: 23 additions & 0 deletions core/shared/src/test/scala/fs2/ChunkSuite.scala
Original file line number Diff line number Diff line change
Expand Up @@ -131,6 +131,24 @@ class ChunkSuite extends Fs2Suite {
}
}

test("Chunk.grouped of an empty chunk has no groups") {
assertEquals(Chunk.empty[Int].grouped(3), Chunk.empty)
}

test("Chunk.grouped rejects a non-positive group size") {
List(0, -1).foreach { n =>
intercept[IllegalArgumentException](Chunk(1, 2, 3).grouped(n))
intercept[IllegalArgumentException](Chunk.empty[Int].grouped(n))
intercept[IllegalArgumentException](List(1, 2, 3).grouped(n))
}
}

test("Chunk.grouped returns a chunk that fits in one group as is") {
forAll { (c: Chunk[Int]) =>
if (c.nonEmpty) assert(c.grouped(c.size).head.get eq c)
}
}

class OddStringExtractor {
val callCounter: AtomicInteger = new AtomicInteger(0)

Expand Down Expand Up @@ -167,6 +185,11 @@ class ChunkSuite extends Fs2Suite {
property("isEmpty") {
forAll((c: Chunk[A]) => assertEquals(c.isEmpty, c.toList.isEmpty))
}
property("grouped") {
forAll(genChunk, Gen.choose(1, 50)) { (c: Chunk[A], n: Int) =>
assertEquals(c.grouped(n).map(_.toList).toList, c.toList.grouped(n).toList)
}
}
property("toArray") {
forAll { (c: Chunk[A]) =>
assertEquals(c.toArray.toVector, c.toVector)
Expand Down
Loading
Loading