Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 12 additions & 4 deletions io/js/src/main/scala/fs2/io/ioplatform.scala
Original file line number Diff line number Diff line change
Expand Up @@ -192,17 +192,25 @@ private[fs2] trait ioplatform {
def writeWritable[F[_]](
writable: F[Writable],
endAfterUse: Boolean = true
)(implicit F: Async[F]): Pipe[F, Byte, Nothing] =
)(implicit F: Async[F]): Pipe[F, Byte, Nothing] = stream =>
writeWritableIncremental(writable, endAfterUse)(F)(stream).drain

/** Writes all bytes to the specified `Writable`.
*/
def writeWritableIncremental[F[_]](
writable: F[Writable],
endAfterUse: Boolean = true
)(implicit F: Async[F]): Pipe[F, Byte, Int] =
in =>
Stream
.eval(writable)
.flatMap { writable =>
val writes = in.chunks.foreach { chunk =>
F.async[Unit] { cb =>
val writes = in.chunks.evalMap { chunk =>
F.async[Int] { cb =>
F.delay {
writable.write(
chunk.toUint8Array,
e => cb(e.filterNot(_ == null).toLeft(()).leftMap(js.JavaScriptException))
e => cb(e.filterNot(_ == null).toLeft(chunk.size).leftMap(js.JavaScriptException))
)
Some(F.delay(writable.destroy()))
}
Expand Down
7 changes: 7 additions & 0 deletions io/js/src/main/scala/fs2/io/net/SocketPlatform.scala
Original file line number Diff line number Diff line change
Expand Up @@ -112,8 +112,15 @@ private[net] trait SocketCompanionPlatform {
override def write(bytes: Chunk[Byte]): F[Unit] =
Stream.chunk(bytes).through(writes).compile.drain

override def writeIncremental(bytes: fs2.Chunk[Byte]): fs2.Stream[F, Int] =
Stream.chunk(bytes).through(writesIncremental)

override def writes: Pipe[F, Byte, Nothing] =
writeWritable(F.pure(sock), endAfterUse = false)

override def writesIncremental: Pipe[F, Byte, Int] =
writeWritableIncremental(F.pure(sock), endAfterUse = false)

}

}
16 changes: 16 additions & 0 deletions io/jvm-native/src/main/scala/fs2/io/net/SocketPlatform.scala
Original file line number Diff line number Diff line change
Expand Up @@ -148,6 +148,22 @@ private[net] trait SocketCompanionPlatform {
}
}

override def writeIncremental(bytes: Chunk[Byte]): Stream[F, Int] = {
def go(buff: ByteBuffer): Pull[F, Int, Unit] =
Pull
.eval(F.async[Int] { cb =>
ch.write(buff, null, new IntCompletionHandler(cb))
F.delay(Some(endOfOutput.voidError))
})
.flatMap { written =>
Pull.output1(written) >>
go(buff).whenA(written >= 0 && buff.remaining() > 0)
}

Stream.resource(writeMutex.lock) >>
Stream.eval(F.delay(bytes.toByteBuffer)).flatMap(go(_).stream)
}

override def localAddress: F[SocketAddress[IpAddress]] =
asyncInstance.pure(address.asIpUnsafe)

Expand Down
10 changes: 10 additions & 0 deletions io/jvm/src/main/scala/fs2/io/net/AsyncUnixSocketsProvider.scala
Original file line number Diff line number Diff line change
Expand Up @@ -117,6 +117,16 @@ private[net] object AsyncUnixSocketsProvider {
}
}

override def writeIncremental(bytes: Chunk[Byte]): Stream[F, Int] = {
def go(buff: ByteBuffer): Pull[F, Int, Unit] =
Pull
.eval(evalOnVirtualThreadIfAvailable(F.blocking(ch.write(buff))).cancelable(close))
.flatMap(written => Pull.output1(written) >> go(buff).whenA(buff.remaining() > 0))

Stream.resource(writeMutex.lock) *>
Stream.eval(F.delay(bytes.toByteBuffer)).flatMap(go(_).stream)
}

private def raiseIpAddressError[A]: F[A] =
F.raiseError(new UnsupportedOperationException("Unix sockets do not use IP addressing"))

Expand Down
14 changes: 14 additions & 0 deletions io/jvm/src/main/scala/fs2/io/net/SelectingSocket.scala
Original file line number Diff line number Diff line change
Expand Up @@ -77,6 +77,20 @@ private final class SelectingSocket[F[_]: LiftIO] private (
}
}

override def writeIncremental(bytes: Chunk[Byte]): Stream[F, Int] = {
def go(buf: ByteBuffer): Pull[F, Int, Unit] =
Pull.eval(F.delay(ch.write(buf))).flatMap { written =>
if (buf.remaining() > 0)
Pull.eval(selector.select(ch, OP_WRITE).to).void >>
Pull.output1(written) >>
go(buf)
else Pull.output1(written)
}

Stream.resource(writeMutex.lock) >>
Stream.eval(F.delay(bytes.toByteBuffer)).flatMap(go(_).stream)
}

def isOpen: F[Boolean] = F.delay(ch.isOpen)

def endOfOutput: F[Unit] =
Expand Down
11 changes: 5 additions & 6 deletions io/jvm/src/main/scala/fs2/io/net/tls/TLSContextPlatform.scala
Original file line number Diff line number Diff line change
Expand Up @@ -34,7 +34,6 @@ import javax.net.ssl.{
TrustManagerFactory,
X509ExtendedTrustManager
}
import cats.Applicative
import cats.effect.kernel.{Async, Resource}
import cats.syntax.all._
import com.comcast.ip4s.{IpAddress, SocketAddress}
Expand Down Expand Up @@ -159,8 +158,8 @@ private[tls] trait TLSContextCompanionPlatform { self: TLSContext.type =>
.eval(
engine(
new TLSEngine.Binding[F] {
def write(data: Chunk[Byte]): F[Unit] =
socket.write(data)
def write(data: Chunk[Byte]): Stream[F, Int] =
socket.writeIncremental(data)
def read(maxBytes: Int): F[Option[Chunk[Byte]]] =
socket.read(maxBytes)
},
Expand Down Expand Up @@ -194,9 +193,9 @@ private[tls] trait TLSContextCompanionPlatform { self: TLSContext.type =>
.eval(
engine(
new TLSEngine.Binding[F] {
def write(data: Chunk[Byte]): F[Unit] =
if (data.isEmpty) Applicative[F].unit
else socket.write(data, remoteAddress)
def write(data: Chunk[Byte]): Stream[F, Int] =
if (data.isEmpty) Stream.empty
else Stream.eval(socket.write(data, remoteAddress).as(data.size))
def read(maxBytes: Int): F[Option[Chunk[Byte]]] =
socket.read.map(p => Some(p.bytes))
},
Expand Down
105 changes: 63 additions & 42 deletions io/jvm/src/main/scala/fs2/io/net/tls/TLSEngine.scala
Original file line number Diff line number Diff line change
Expand Up @@ -42,12 +42,13 @@ private[tls] trait TLSEngine[F[_]] {
def stopWrap: F[Unit]
def stopUnwrap: F[Unit]
def write(data: Chunk[Byte]): F[Unit]
def writeIncremental(data: Chunk[Byte]): Stream[F, Int]
def read(maxBytes: Int): F[Option[Chunk[Byte]]]
}

private[tls] object TLSEngine {
trait Binding[F[_]] {
def write(data: Chunk[Byte]): F[Unit]
def write(data: Chunk[Byte]): Stream[F, Int]
def read(maxBytes: Int): F[Option[Chunk[Byte]]]
}

Expand Down Expand Up @@ -85,41 +86,54 @@ private[tls] object TLSEngine {
def stopUnwrap = Sync[F].delay(engine.closeInbound()).attempt.void

def write(data: Chunk[Byte]): F[Unit] =
writeMutex.lock.surround(write0(data))
writeIncremental(data).compile.drain

private def write0(data: Chunk[Byte]): F[Unit] =
wrapBuffer.input(data) >> wrap
def writeIncremental(data: Chunk[Byte]): Stream[F, Int] =
Stream.resource(writeMutex.lock) >> write0(data)

private def write0(data: Chunk[Byte]): Stream[F, Int] =
Stream.eval(wrapBuffer.input(data)) >> wrap.stream

/** Performs a wrap operation on the underlying engine. */
private def wrap: F[Unit] =
wrapBuffer
.perform(engine.wrap(_, _))
.flatTap(result => log(s"wrap result: $result"))
private def wrap: Pull[F, Int, Unit] =
Pull
.eval(
wrapBuffer
.perform(engine.wrap(_, _))
.flatTap(result => log(s"wrap result: $result"))
)
.flatMap { result =>
result.getStatus match {
case SSLEngineResult.Status.OK =>
doWrite >> {
doWrite.pull.echo >> {
result.getHandshakeStatus match {
case SSLEngineResult.HandshakeStatus.NOT_HANDSHAKING =>
wrapBuffer.inputRemains
Pull
.eval(wrapBuffer.inputRemains)
.flatMap(x => wrap.whenA(x > 0 && result.bytesConsumed > 0))
case _ =>
handshakeMutex.lock
.surround(stepHandshake(result, true)) >> wrap
// TODO: is this the right way to use and close the resource?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

I'm not quite sure how resources are scoped in a Pull. Is this right?

Pull.scope(
Stream
.resource(handshakeMutex.lock)
.flatMap(_ => stepHandshake(result, true).stream)
.pull
.echo
) >> wrap
}
}
case SSLEngineResult.Status.BUFFER_UNDERFLOW =>
doWrite
doWrite.pull.echo
case SSLEngineResult.Status.BUFFER_OVERFLOW =>
wrapBuffer.expandOutput >> wrap
Pull.eval(wrapBuffer.expandOutput) >> wrap
case SSLEngineResult.Status.CLOSED =>
stopWrap
Pull.eval(stopWrap)
}
}

private def doWrite: F[Unit] =
wrapBuffer.output(Int.MaxValue).flatMap { out =>
if (out.isEmpty) Applicative[F].unit
private def doWrite: Stream[F, Int] =
Stream.eval(wrapBuffer.output(Int.MaxValue)).flatMap { out =>
if (out.isEmpty) Stream.empty
else binding.write(out)
}

Expand Down Expand Up @@ -169,7 +183,7 @@ private[tls] object TLSEngine {
unwrap(maxBytes)
case _ =>
handshakeMutex.lock
.surround(stepHandshake(result, false)) >> unwrap(
.surround(stepHandshake(result, false).stream.compile.drain) >> unwrap(
maxBytes
)
}
Expand All @@ -194,68 +208,75 @@ private[tls] object TLSEngine {
private def stepHandshake(
result: SSLEngineResult,
lastOperationWrap: Boolean
): F[Unit] =
): Pull[F, Int, Unit] =
result.getHandshakeStatus match {
case SSLEngineResult.HandshakeStatus.NOT_HANDSHAKING =>
Applicative[F].unit
Pull.done
case SSLEngineResult.HandshakeStatus.FINISHED =>
unwrapBuffer.inputRemains.flatMap { remaining =>
Pull.eval(unwrapBuffer.inputRemains).flatMap { remaining =>
if (remaining > 0) unwrapHandshake
else Applicative[F].unit
else Pull.done
}
case SSLEngineResult.HandshakeStatus.NEED_TASK =>
sslEngineTaskRunner.runDelegatedTasks >> (if (lastOperationWrap) wrapHandshake
else unwrapHandshake)
Pull.eval(sslEngineTaskRunner.runDelegatedTasks) >>
(if (lastOperationWrap) wrapHandshake
else unwrapHandshake)
case SSLEngineResult.HandshakeStatus.NEED_WRAP =>
wrapHandshake
case SSLEngineResult.HandshakeStatus.NEED_UNWRAP =>
unwrapBuffer.inputRemains.flatMap { remaining =>
Pull.eval(unwrapBuffer.inputRemains).flatMap { remaining =>
if (remaining > 0 && result.getStatus != SSLEngineResult.Status.BUFFER_UNDERFLOW)
unwrapHandshake
else
binding.read(engine.getSession.getPacketBufferSize).flatMap {
case Some(c) => unwrapBuffer.input(c) >> unwrapHandshake
case None => stopUnwrap
Pull.eval(binding.read(engine.getSession.getPacketBufferSize)).flatMap {
case Some(c) => Pull.eval(unwrapBuffer.input(c)) >> unwrapHandshake
case None => Pull.eval(stopUnwrap)
}
}
case SSLEngineResult.HandshakeStatus.NEED_UNWRAP_AGAIN =>
unwrapHandshake
}

/** Performs a wrap operation as part of handshaking. */
private def wrapHandshake: F[Unit] =
wrapBuffer
.perform(engine.wrap(_, _))
.flatTap(result => log(s"wrapHandshake result: $result"))
private def wrapHandshake: Pull[F, Int, Unit] =
Pull
.eval(
wrapBuffer
.perform(engine.wrap(_, _))
.flatTap(result => log(s"wrapHandshake result: $result"))
)
.flatMap { result =>
result.getStatus match {
case SSLEngineResult.Status.OK | SSLEngineResult.Status.BUFFER_UNDERFLOW =>
doWrite >> stepHandshake(
doWrite.pull.echo >> stepHandshake(
result,
true
)
case SSLEngineResult.Status.BUFFER_OVERFLOW =>
wrapBuffer.expandOutput >> wrapHandshake
Pull.eval(wrapBuffer.expandOutput) >> wrapHandshake
case SSLEngineResult.Status.CLOSED =>
stopWrap >> stopUnwrap
Pull.eval(stopWrap >> stopUnwrap)
}
}

/** Performs an unwrap operation as part of handshaking. */
private def unwrapHandshake: F[Unit] =
unwrapBuffer
.perform(engine.unwrap(_, _))
.flatTap(result => log(s"unwrapHandshake result: $result"))
private def unwrapHandshake: Pull[F, Int, Unit] =
Pull
.eval(
unwrapBuffer
.perform(engine.unwrap(_, _))
.flatTap(result => log(s"unwrapHandshake result: $result"))
)
.flatMap { result =>
result.getStatus match {
case SSLEngineResult.Status.OK =>
stepHandshake(result, false)
case SSLEngineResult.Status.BUFFER_UNDERFLOW =>
stepHandshake(result, false)
case SSLEngineResult.Status.BUFFER_OVERFLOW =>
unwrapBuffer.expandOutput >> unwrapHandshake
Pull.eval(unwrapBuffer.expandOutput) >> unwrapHandshake
case SSLEngineResult.Status.CLOSED =>
stopWrap >> stopUnwrap
Pull.eval(stopWrap >> stopUnwrap)
}
}
}
Expand Down
3 changes: 3 additions & 0 deletions io/jvm/src/main/scala/fs2/io/net/tls/TLSSocketPlatform.scala
Original file line number Diff line number Diff line change
Expand Up @@ -56,6 +56,9 @@ private[tls] trait TLSSocketCompanionPlatform { self: TLSSocket.type =>
def write(bytes: Chunk[Byte]): F[Unit] =
engine.write(bytes)

override def writeIncremental(bytes: Chunk[Byte]): Stream[F, Int] =
engine.writeIncremental(bytes)

private def read0(maxBytes: Int): F[Option[Chunk[Byte]]] =
engine.read(maxBytes)

Expand Down
12 changes: 12 additions & 0 deletions io/jvm/src/test/scala/fs2/io/net/tls/TLSSocketSuite.scala
Original file line number Diff line number Diff line change
Expand Up @@ -251,6 +251,18 @@ class TLSSocketSuite extends TLSSuite {
raw.write(b) >> IO(totalWritten += b.size) >> IO(totalWritten >= limit)
.ifM(endOfOutput, IO.unit)
}

override def writeIncremental(bytes: Chunk[Byte]): Stream[IO, Int] =
if (totalWritten >= limit) Stream.eval(endOfOutput).drain
else {
val b = bytes.take(limit - totalWritten)
raw.writeIncremental(b) ++ Stream
.eval(
IO(totalWritten += b.size) >> IO(totalWritten >= limit).ifM(endOfOutput, IO.unit)
)
.drain
}

}

// Setup an HTTPS echo server & a client that starts a TLS handshake but only sends the first few bytes and then signals no more output
Expand Down
Loading
Loading