@@ -2,8 +2,9 @@ package com.twitter.io
22
33import com .twitter .concurrent .AsyncStream
44import com .twitter .conversions .time ._
5- import com .twitter .util .{Await , Future , Promise }
6- import java .io .{ByteArrayInputStream , ByteArrayOutputStream , OutputStream }
5+ import com .twitter .util .{Await , Awaitable , Future , Promise }
6+ import java .io .{ByteArrayInputStream , ByteArrayOutputStream }
7+ import java .util .concurrent .atomic .AtomicBoolean
78import org .mockito .Mockito ._
89import org .scalatest .concurrent .{Eventually , IntegrationPatience }
910import org .scalatest .prop .GeneratorDrivenPropertyChecks
@@ -19,6 +20,8 @@ class ReaderTest
1920 private def arr (i : Int , j : Int ) = Array .range(i, j).map(_.toByte)
2021 private def buf (i : Int , j : Int ) = Buf .ByteArray .Owned (arr(i, j))
2122
23+ private def await [T ](t : Awaitable [T ]): T = Await .result(t, 5 .seconds)
24+
2225 private def toSeq (b : Option [Buf ]): Seq [Byte ] = b match {
2326 case None => fail(" Expected full buffer" )
2427 case Some (buf) =>
@@ -31,17 +34,17 @@ class ReaderTest
3134
3235 private def assertReadWhileReading (r : Reader [Buf ]): Unit = {
3336 val f = r.read(1 )
34- intercept[IllegalStateException ] { Await .result (r.read(1 )) }
37+ intercept[IllegalStateException ] { await (r.read(1 )) }
3538 assert(! f.isDefined)
3639 }
3740
3841 private def assertFailed (r : Reader [Buf ], p : Promise [Option [Buf ]]): Unit = {
3942 val f = r.read(1 )
4043 assert(! f.isDefined)
4144 p.setException(new Exception )
42- intercept[Exception ] { Await .result (f) }
43- intercept[Exception ] { Await .result (r.read(0 )) }
44- intercept[Exception ] { Await .result (r.read(1 )) }
45+ intercept[Exception ] { await (f) }
46+ intercept[Exception ] { await (r.read(0 )) }
47+ intercept[Exception ] { await (r.read(1 )) }
4548 }
4649
4750 test(" Reader.copy - source and destination equality" ) {
@@ -58,7 +61,7 @@ class ReaderTest
5861 }
5962 }
6063
61- Await .result (Future .join(f, g))
64+ await (Future .join(f, g))
6265
6366 val b = new ByteArrayOutputStream
6467 b.write(p)
@@ -76,7 +79,7 @@ class ReaderTest
7679 BufReader (Buf .Utf8 (s))
7780 }
7881 val buf = Reader .readAll(Reader .concat(AsyncStream .fromSeq(readers)))
79- Await .result (buf) should equal(Buf .Utf8 (ss.mkString))
82+ await (buf) should equal(Buf .Utf8 (ss.mkString))
8083 }
8184 }
8285
@@ -123,15 +126,15 @@ class ReaderTest
123126 }
124127 val combined = Reader .concat(head +:: tail)
125128 val buf = Reader .readAll(combined)
126- intercept[Exception ] { Await .result (buf) }
129+ intercept[Exception ] { await (buf) }
127130 assert(! p.isDefined)
128131 }
129132
130133 test(" Reader.fromStream closes resources on EOF read" ) {
131134 val in = spy(new ByteArrayInputStream (arr(0 , 10 )))
132135 val r = Reader .fromStream(in)
133136 val f = Reader .readAll(r)
134- assert(Await .result(f, 5 .seconds ) == buf(0 , 10 ))
137+ assert(await(f ) == buf(0 , 10 ))
135138 eventually {
136139 verify(in).close()
137140 }
@@ -145,4 +148,37 @@ class ReaderTest
145148 verify(in).close()
146149 }
147150 }
151+
152+ test(" Reader.fromAsyncStream completes when stream is empty" ) {
153+ val as = AsyncStream (buf(1 , 10 ))
154+ val r = Reader .fromAsyncStream(as)
155+ val f = Reader .readAll(r)
156+ assert(await(f) == buf(1 , 10 ))
157+ }
158+
159+ test(" Reader.fromAsyncStream fails on exceptional stream" ) {
160+ val as = AsyncStream .exception(new Exception ())
161+ val r = Reader .fromAsyncStream(as)
162+ val f = Reader .readAll(r)
163+ intercept[Exception ] { await(f) }
164+ }
165+
166+ test(" Reader.fromAsyncStream only evaluates tail when buffer is exhausted" ) {
167+ val tailEvaluated = new AtomicBoolean (false )
168+ def tail : AsyncStream [Buf ] = {
169+ tailEvaluated.set(true )
170+ AsyncStream .empty
171+ }
172+ val as = AsyncStream .mk(buf(0 , 10 ), tail)
173+ val r = Reader .fromAsyncStream(as)
174+
175+ // partially read the buffer
176+ await(r.read(9 ))
177+ assert(! tailEvaluated.get())
178+
179+ // read the rest of the buffer
180+ await(r.read(2 ))
181+ assert(tailEvaluated.get())
182+ }
183+
148184}
0 commit comments