Sitelet https://github.com/twitter/util/commit/2316aa5dcf68780d125124d9eb7cf62ba585844a
Skip to content

Commit 2316aa5

Browse files
David Rusekjenkins
authored andcommitted
util-core: Introduce Reader.framed
Problem Network protocols often frame data based on a size field preceding the message so that a reader knows how much data to read before decoding the message. For instance, this is the case with the mux protocol. Solution Add a Reader[Buf] implementation that will frame Bufs based on a given pattern. JIRA Issues: CSL-6778 Differential Revision: https://phabricator.twitter.biz/D212396
1 parent 845620b commit 2316aa5

3 files changed

Lines changed: 114 additions & 9 deletions

File tree

‎CHANGELOG.rst‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,9 @@ New Features
1919
* util-core: Introducing `Reader.chunked` that chunks the output of a given reader.
2020
``PHAB_ID=D206676``
2121

22+
* util-core: Added Reader#framed for consuming data framed by a user supplied function.
23+
``PHAB_ID=D212396``
24+
2225
* util-security: Add `NullSslSession` related objects for use with non-existent
2326
`SSLSession`s. ``PHAB_ID=D201421``
2427

‎util-core/src/main/scala/com/twitter/io/Reader.scala‎

Lines changed: 43 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -57,6 +57,7 @@ object Reader {
5757
}
5858

5959
def read(n: Int): Future[Option[Buf]] = synchronized {
60+
// flatMap to `this` to prevent allocating
6061
if (state.isEmpty) r.read(Int.MaxValue).flatMap(this)
6162
else {
6263
val result = state.slice(0, chunkSize)
@@ -68,6 +69,40 @@ object Reader {
6869
def discard(): Unit = r.discard()
6970
}
7071

72+
// see Reader.framed
73+
private final class Framed(r: Reader[Buf], framer: Buf => Seq[Buf])
74+
extends Reader[Buf] with (Option[Buf] => Future[Option[Buf]]) {
75+
76+
private[this] var frames: Seq[Buf] = Seq.empty
77+
78+
// we only enter here when `frames` is empty.
79+
def apply(in: Option[Buf]): Future[Option[Buf]] = synchronized {
80+
in match {
81+
case Some(data) =>
82+
frames = framer(data)
83+
read(Int.MaxValue)
84+
case None =>
85+
Future.None
86+
}
87+
}
88+
89+
def read(n: Int): Future[Option[Buf]] = synchronized {
90+
frames match {
91+
case nextFrame :: rst =>
92+
frames = rst
93+
Future.value(Some(nextFrame))
94+
case _ =>
95+
// flatMap to `this` to prevent allocating
96+
r.read(Int.MaxValue).flatMap(this)
97+
}
98+
}
99+
100+
def discard(): Unit = synchronized {
101+
frames = Seq.empty
102+
r.discard()
103+
}
104+
}
105+
71106
/**
72107
* Read the entire bytestream presented by `r`.
73108
*/
@@ -258,4 +293,12 @@ object Reader {
258293
* }}}
259294
*/
260295
def copy(r: Reader[Buf], w: Writer[Buf]): Future[Unit] = copy(r, w, Writer.BufferSize)
296+
297+
/**
298+
* Wraps a [[ Reader[Buf] ]] and emits frames as decided by `framer`.
299+
*
300+
* @note The returned `Reader` may not be thread safe depending on the behavior
301+
* of the framer.
302+
*/
303+
def framed(r: Reader[Buf], framer: Buf => Seq[Buf]): Reader[Buf] = new Framed(r, framer)
261304
}

‎util-core/src/test/scala/com/twitter/io/ReaderTest.scala‎

Lines changed: 68 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -2,14 +2,16 @@ package com.twitter.io
22

33
import com.twitter.concurrent.AsyncStream
44
import com.twitter.conversions.time._
5+
import com.twitter.conversions.storage._
56
import com.twitter.util.{Await, Awaitable, Future, Promise}
67
import java.io.{ByteArrayInputStream, ByteArrayOutputStream}
78
import java.util.concurrent.atomic.AtomicBoolean
89
import org.mockito.Mockito._
9-
import org.scalacheck.Gen
10+
import org.scalacheck.{Arbitrary, Gen}
1011
import org.scalatest.concurrent.{Eventually, IntegrationPatience}
1112
import org.scalatest.prop.GeneratorDrivenPropertyChecks
1213
import org.scalatest.{FunSuite, Matchers}
14+
import scala.annotation.tailrec
1315

1416
class ReaderTest
1517
extends FunSuite
@@ -81,16 +83,16 @@ class ReaderTest
8183
} yield (s, i)
8284

8385
forAll(stringAndChunk) { case (s, i) =>
84-
val r = Reader.chunked(Reader.fromBuf(Buf.Utf8(s)), i)
86+
val r = Reader.chunked(Reader.fromBuf(Buf.Utf8(s)), i)
8587

86-
def readLoop(): Unit = await(r.read(Int.MaxValue)) match {
87-
case Some(b) =>
88-
assert(b.length <= i)
89-
readLoop()
90-
case None => ()
91-
}
88+
def readLoop(): Unit = await(r.read(Int.MaxValue)) match {
89+
case Some(b) =>
90+
assert(b.length <= i)
91+
readLoop()
92+
case None => ()
93+
}
9294

93-
readLoop()
95+
readLoop()
9496
}
9597
}
9698

@@ -202,4 +204,61 @@ class ReaderTest
202204
assert(tailEvaluated.get())
203205
}
204206

207+
test("Reader.framed reads framed data") {
208+
val getByteArrays: Gen[Seq[Buf]] = Gen.listOf(
209+
for {
210+
// limit arrays to a few kilobytes, otherwise we may generate a very large amount of data
211+
numBytes <- Gen.choose(0.bytes.inBytes, 2.kilobytes.inBytes)
212+
bytes <- Gen.containerOfN[Array, Byte](numBytes.toInt, Arbitrary.arbitrary[Byte])
213+
} yield Buf.ByteArray.Owned(bytes)
214+
)
215+
216+
forAll(getByteArrays) { buffers: Seq[Buf] =>
217+
val buffersWithLength = buffers.map(buf => Buf.U32BE(buf.length).concat(buf))
218+
219+
val r = Reader.framed(BufReader(Buf(buffersWithLength)), new ReaderTest.U32BEFramer())
220+
221+
// read all of the frames
222+
buffers.foreach { buf =>
223+
assert(await(r.read(Int.MaxValue)).contains(buf))
224+
}
225+
226+
// make sure the reader signals EOF
227+
assert(await(r.read(Int.MaxValue)).isEmpty)
228+
}
229+
}
230+
231+
test("Reader.framed reads empty frames") {
232+
val r = Reader.framed(BufReader(Buf.U32BE(0)), new ReaderTest.U32BEFramer())
233+
assert(await(r.read(Int.MaxValue)).contains(Buf.Empty))
234+
assert(await(r.read(Int.MaxValue)).isEmpty)
235+
}
236+
237+
}
238+
239+
object ReaderTest {
240+
241+
/**
242+
* Used to test Reader.framed, extract fields in terms of
243+
* frames, signified by a 32-bit BE value preceding
244+
* each frame.
245+
*/
246+
private class U32BEFramer() extends (Buf => Seq[Buf]) {
247+
var state: Buf = Buf.Empty
248+
249+
@tailrec
250+
private def loop(acc: Seq[Buf], buf: Buf): Seq[Buf] = {
251+
buf match {
252+
case Buf.U32BE(l, d) if d.length >= l =>
253+
loop(acc :+ d.slice(0, l), d.slice(l, d.length))
254+
case _ =>
255+
state = buf
256+
acc
257+
}
258+
}
259+
260+
def apply(buf: Buf): Seq[Buf] = synchronized {
261+
loop(Seq.empty, state concat buf)
262+
}
263+
}
205264
}

0 commit comments

Comments
 (0)