Skip to content
Merged
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
2 changes: 1 addition & 1 deletion http-core/src/main/resources/reference.conf
Original file line number Diff line number Diff line change
Expand Up @@ -359,7 +359,7 @@ pekko.http {
# exceeded while inflating a compressed message, the connection is closed
# with a WebSocket protocol error.
# Set to 0 to disable this limit.
max-allocation = 64k
max-allocation = 256k

permessage-deflate {
# Pekko HTTP uses the JDK Deflater/Inflater implementation for
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -93,7 +93,7 @@ private[http] object Handshake {
// - Origin header is optional and, if required, should be validated
// on higher levels (routing, application logic)
//
// TODO See #18709 Extension support is optional in WS and currently unsupported.
// WebSocket extension negotiation is optional. Currently only permessage-deflate is supported.
//
// these are not needed directly, we verify their presence and correctness only:
// - Upgrade
Expand Down Expand Up @@ -123,13 +123,16 @@ private[http] object Handshake {
case OptionVal.Some(p) => p.protocols
case _ => Nil
}
val clientRequestedExtensions = headers.collect {
val clientRequestedExtensions = headers.flatMap {
case extensions: `Sec-WebSocket-Extensions` => extensions.extensions
}.flatten
case _ => Nil
}
val perMessageDeflate =
PerMessageDeflate.negotiate(
clientRequestedExtensions,
settings.asInstanceOf[WebSocketSettingsImpl].compression)
settings match {
case impl: WebSocketSettingsImpl =>
PerMessageDeflate.negotiate(clientRequestedExtensions, impl.compression)
case _ => None
}

val header = new UpgradeToWebSocketLowLevel {
def requestedProtocols: Seq[String] = clientSupportedSubprotocols
Expand Down Expand Up @@ -203,16 +206,14 @@ private[http] object Handshake {
.join(messageHandler)
}

HttpResponse(
StatusCodes.SwitchingProtocols,
val extensionHeaders = perMessageDeflate.map(p => `Sec-WebSocket-Extensions`(Seq(p.responseExtension))).toList
val responseHeaders =
subprotocol.map(p => `Sec-WebSocket-Protocol`(Seq(p))).toList :::
List(
UpgradeHeader,
ConnectionUpgradeHeader,
`Sec-WebSocket-Accept`.forKey(key)) :::
perMessageDeflate.map(p => `Sec-WebSocket-Extensions`(Seq(p.responseExtension))).toList :::
List(
UpgradeToOtherProtocolResponseHeader(WebSocket.framing.join(frameHandler))))
List(UpgradeHeader, ConnectionUpgradeHeader, `Sec-WebSocket-Accept`.forKey(key)) :::
extensionHeaders :::
List(UpgradeToOtherProtocolResponseHeader(WebSocket.framing.join(frameHandler)))

HttpResponse(StatusCodes.SwitchingProtocols, responseHeaders)
}
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -35,7 +35,7 @@ import pekko.stream.stage.GraphStageLogic
import pekko.stream.stage.InHandler
import pekko.stream.stage.OutHandler
import pekko.stream.{ Attributes, FlowShape, Inlet, Outlet }
import pekko.util.ByteString
import pekko.util.{ ByteString, ByteStringBuilder }

import scala.collection.immutable
import scala.collection.immutable.ListMap
Expand All @@ -52,6 +52,16 @@ private[http] object PerMessageDeflate {
private val ServerNoContextTakeover = "server_no_context_takeover"
private val EmptyStoredBlock = ByteString(0x00, 0x00, 0xFF.toByte, 0xFF.toByte)

private[ws] trait CompressionFactory {
def newInflater(): Inflater
def newDeflater(compressionLevel: Int): Deflater
}

private object DefaultCompressionFactory extends CompressionFactory {
override def newInflater(): Inflater = new Inflater(true)
override def newDeflater(compressionLevel: Int): Deflater = new Deflater(compressionLevel, true)
}

final case class Negotiated(
responseExtension: WebSocketExtension,
serverNoContextTakeover: Boolean,
Expand All @@ -74,16 +84,28 @@ private[http] object PerMessageDeflate {
deflaterFlow)

private def inflaterFlow: Flow[FrameEventOrError, FrameEventOrError, NotUsed] =
Flow.fromGraph(new LifecycleMapConcatStage(
"PerMessageDeflate.inflater",
() => new InflaterFlow(clientNoContextTakeover, settings)))
createInflaterFlow(clientNoContextTakeover, settings, DefaultCompressionFactory)

private def deflaterFlow: Flow[FrameEvent, FrameEvent, NotUsed] =
Flow.fromGraph(new LifecycleMapConcatStage(
"PerMessageDeflate.deflater",
() => new DeflaterFlow(serverNoContextTakeover, settings)))
createDeflaterFlow(serverNoContextTakeover, settings, DefaultCompressionFactory)
}

private[ws] def createInflaterFlow(
noContextTakeover: Boolean,
settings: WebSocketCompressionSettingsImpl,
compressionFactory: CompressionFactory): Flow[FrameEventOrError, FrameEventOrError, NotUsed] =
Flow.fromGraph(new LifecycleMapConcatStage(
"PerMessageDeflate.inflater",
() => new InflaterFlow(noContextTakeover, settings, compressionFactory)))

private[ws] def createDeflaterFlow(
noContextTakeover: Boolean,
settings: WebSocketCompressionSettingsImpl,
compressionFactory: CompressionFactory): Flow[FrameEvent, FrameEvent, NotUsed] =
Flow.fromGraph(new LifecycleMapConcatStage(
"PerMessageDeflate.deflater",
() => new DeflaterFlow(noContextTakeover, settings, compressionFactory)))

def negotiate(
requested: immutable.Seq[WebSocketExtension],
settings: WebSocketCompressionSettingsImpl): Option[Negotiated] = {
Expand Down Expand Up @@ -129,20 +151,22 @@ private[http] object PerMessageDeflate {
}

private def validWindowBits(value: String): Boolean =
value.length <= 2 && value.forall(_.isDigit) && {
value.nonEmpty && value.length <= 2 && value.forall(_.isDigit) && {
val parsed = value.toInt
parsed >= 8 && parsed <= 15
}

private final class InflaterFlow(
noContextTakeover: Boolean,
settings: WebSocketCompressionSettingsImpl)
settings: WebSocketCompressionSettingsImpl,
compressionFactory: CompressionFactory)
extends LifecycleMapConcat[FrameEventOrError, FrameEventOrError] {
private var inflater = new Inflater(true)
private var inflater = compressionFactory.newInflater()
private var compressedFrame: Option[CompressedFrame] = None
private var compressedMessageInProgress = false
private var decompressedMessageBytes = 0L
private var bypassFrameInProgress = false
private val buffer = new Array[Byte](8192)

override def apply(event: FrameEventOrError): immutable.Iterable[FrameEventOrError] = event match {
case start @ FrameStart(header, data)
Expand Down Expand Up @@ -186,7 +210,7 @@ private[http] object PerMessageDeflate {
if (frame.appendTail) decompressedMessageBytes = 0L
if (frame.appendTail && noContextTakeover) {
inflater.end()
inflater = new Inflater(true)
inflater = compressionFactory.newInflater()
}
FrameStart(frame.header.copy(length = inflated.length), inflated) :: Nil
}
Expand All @@ -196,7 +220,6 @@ private[http] object PerMessageDeflate {
val input = if (appendTail) data ++ EmptyStoredBlock else data
inflater.setInput(input.toArrayUnsafe())
val output = new ByteArrayOutputStream(1024)
val buffer = new Array[Byte](1024)
var count = inflater.inflate(buffer)
while (count > 0) {
decompressedMessageBytes += count
Expand All @@ -218,12 +241,14 @@ private[http] object PerMessageDeflate {

private final class DeflaterFlow(
noContextTakeover: Boolean,
settings: WebSocketCompressionSettingsImpl)
settings: WebSocketCompressionSettingsImpl,
compressionFactory: CompressionFactory)
extends LifecycleMapConcat[FrameEvent, FrameEvent] {
private var deflater = new Deflater(settings.compressionLevel, true)
private var deflater = compressionFactory.newDeflater(settings.compressionLevel)
private var frame: Option[UncompressedFrame] = None
private var messageInProgress = false
private var bypassFrameInProgress = false
private val buffer = new Array[Byte](8192)

override def apply(event: FrameEvent): immutable.Iterable[FrameEvent] = event match {
case FrameStart(header, _)
Expand Down Expand Up @@ -269,15 +294,14 @@ private[http] object PerMessageDeflate {
val compressed = deflate(current.data, current.removeTail)
if (current.removeTail && noContextTakeover) {
deflater.end()
deflater = new Deflater(settings.compressionLevel, true)
deflater = compressionFactory.newDeflater(settings.compressionLevel)
}
FrameStart(current.header.copy(length = compressed.length), compressed) :: Nil
}

private def deflate(data: ByteString, removeTail: Boolean): ByteString = {
deflater.setInput(data.toArrayUnsafe())
val output = new ByteArrayOutputStream(1024)
val buffer = new Array[Byte](1024)
var count = deflater.deflate(buffer, 0, buffer.length, Deflater.SYNC_FLUSH)
while (count > 0) {
output.write(buffer, 0, count)
Expand Down Expand Up @@ -335,11 +359,40 @@ private[http] object PerMessageDeflate {
}
}

private final case class CompressedFrame(header: FrameHeader, data: ByteString, appendTail: Boolean) {
def append(next: ByteString): CompressedFrame = copy(data = data ++ next)
private final case class CompressedFrame(
header: FrameHeader,
fragments: Vector[ByteString],
length: Int,
appendTail: Boolean) {
def data: ByteString = compact(fragments, length)
def append(next: ByteString): CompressedFrame = copy(fragments = fragments :+ next, length = length + next.length)
}

private final case class UncompressedFrame(header: FrameHeader, data: ByteString, removeTail: Boolean) {
def append(next: ByteString): UncompressedFrame = copy(data = data ++ next)
private object CompressedFrame {
def apply(header: FrameHeader, data: ByteString, appendTail: Boolean): CompressedFrame =
CompressedFrame(header, Vector(data), data.length, appendTail)
}

private final case class UncompressedFrame(
header: FrameHeader,
fragments: Vector[ByteString],
length: Int,
removeTail: Boolean) {
def data: ByteString = compact(fragments, length)
def append(next: ByteString): UncompressedFrame = copy(fragments = fragments :+ next, length = length + next.length)
}

private object UncompressedFrame {
def apply(header: FrameHeader, data: ByteString, removeTail: Boolean): UncompressedFrame =
UncompressedFrame(header, Vector(data), data.length, removeTail)
}

private def compact(fragments: Vector[ByteString], length: Int): ByteString =
if (fragments.lengthCompare(1) == 0) fragments.head
else {
val builder = new ByteStringBuilder
builder.sizeHint(length)
fragments.foreach(builder.append)
builder.result()
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -101,7 +101,7 @@ private[pekko] object WebSocketCompressionSettingsImpl {
val Disabled: WebSocketCompressionSettingsImpl =
WebSocketCompressionSettingsImpl(
enabled = false,
maxAllocation = 0,
maxAllocation = 256 * 1024,
compressionLevel = 6,
preferredClientWindowSize = 15,
allowServerNoContext = false,
Expand Down
Loading