From eb59d1fa9c827b8ccd9790938e33490105cc3523 Mon Sep 17 00:00:00 2001 From: Matthias Kurz Date: Wed, 1 Jul 2026 13:00:18 +0200 Subject: [PATCH 1/3] Address WebSocket compression review feedback #61 * Reuse 8 KiB buffers in the inflater and deflater flows. * Accumulate split frame data without repeatedly copying earlier chunks. * Handle extension parameters and custom WebSocket settings defensively. * Simplify extension header handling and align disabled compression defaults. --- .../pekko/http/impl/engine/ws/Handshake.scala | 31 ++++++------- .../impl/engine/ws/PerMessageDeflate.scala | 45 +++++++++++++++---- .../impl/settings/WebSocketSettingsImpl.scala | 2 +- 3 files changed, 54 insertions(+), 24 deletions(-) diff --git a/http-core/src/main/scala/org/apache/pekko/http/impl/engine/ws/Handshake.scala b/http-core/src/main/scala/org/apache/pekko/http/impl/engine/ws/Handshake.scala index e98a900c5..b9b293c9e 100644 --- a/http-core/src/main/scala/org/apache/pekko/http/impl/engine/ws/Handshake.scala +++ b/http-core/src/main/scala/org/apache/pekko/http/impl/engine/ws/Handshake.scala @@ -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 @@ -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 @@ -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) } } diff --git a/http-core/src/main/scala/org/apache/pekko/http/impl/engine/ws/PerMessageDeflate.scala b/http-core/src/main/scala/org/apache/pekko/http/impl/engine/ws/PerMessageDeflate.scala index 96d2c282d..cf909e8a6 100644 --- a/http-core/src/main/scala/org/apache/pekko/http/impl/engine/ws/PerMessageDeflate.scala +++ b/http-core/src/main/scala/org/apache/pekko/http/impl/engine/ws/PerMessageDeflate.scala @@ -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 @@ -129,7 +129,7 @@ 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 } @@ -143,6 +143,7 @@ private[http] object PerMessageDeflate { 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) @@ -196,7 +197,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 @@ -224,6 +224,7 @@ private[http] object PerMessageDeflate { 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, _) @@ -277,7 +278,6 @@ private[http] object PerMessageDeflate { 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) @@ -335,11 +335,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() + } } diff --git a/http-core/src/main/scala/org/apache/pekko/http/impl/settings/WebSocketSettingsImpl.scala b/http-core/src/main/scala/org/apache/pekko/http/impl/settings/WebSocketSettingsImpl.scala index 831ca8c97..15b6dcaf1 100644 --- a/http-core/src/main/scala/org/apache/pekko/http/impl/settings/WebSocketSettingsImpl.scala +++ b/http-core/src/main/scala/org/apache/pekko/http/impl/settings/WebSocketSettingsImpl.scala @@ -101,7 +101,7 @@ private[pekko] object WebSocketCompressionSettingsImpl { val Disabled: WebSocketCompressionSettingsImpl = WebSocketCompressionSettingsImpl( enabled = false, - maxAllocation = 0, + maxAllocation = 64 * 1024, compressionLevel = 6, preferredClientWindowSize = 15, allowServerNoContext = false, From 964f62152ee84c9be338553cd6689c2f718469cd Mon Sep 17 00:00:00 2001 From: Matthias Kurz Date: Tue, 7 Jul 2026 13:07:06 +0200 Subject: [PATCH 2/3] Verify WebSocket compression resource cleanup #61 * Inject inflater and deflater creation per materialized flow instead of using mutable global test hooks. * Verify cleanup after completion, failure, cancellation, protocol errors, and incomplete frame or message state. * Ensure every created compression resource is released exactly once. --- .../impl/engine/ws/PerMessageDeflate.scala | 48 +++-- .../impl/engine/ws/WebSocketServerSpec.scala | 171 ++++++++++++++++++ 2 files changed, 207 insertions(+), 12 deletions(-) diff --git a/http-core/src/main/scala/org/apache/pekko/http/impl/engine/ws/PerMessageDeflate.scala b/http-core/src/main/scala/org/apache/pekko/http/impl/engine/ws/PerMessageDeflate.scala index cf909e8a6..7bf669aa9 100644 --- a/http-core/src/main/scala/org/apache/pekko/http/impl/engine/ws/PerMessageDeflate.scala +++ b/http-core/src/main/scala/org/apache/pekko/http/impl/engine/ws/PerMessageDeflate.scala @@ -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, @@ -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] = { @@ -136,9 +158,10 @@ private[http] object PerMessageDeflate { 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 @@ -187,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 } @@ -218,9 +241,10 @@ 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 @@ -270,7 +294,7 @@ 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 } diff --git a/http-core/src/test/scala/org/apache/pekko/http/impl/engine/ws/WebSocketServerSpec.scala b/http-core/src/test/scala/org/apache/pekko/http/impl/engine/ws/WebSocketServerSpec.scala index 0146cc948..bffc62f9a 100644 --- a/http-core/src/test/scala/org/apache/pekko/http/impl/engine/ws/WebSocketServerSpec.scala +++ b/http-core/src/test/scala/org/apache/pekko/http/impl/engine/ws/WebSocketServerSpec.scala @@ -14,6 +14,9 @@ package org.apache.pekko.http.impl.engine.ws import java.io.ByteArrayOutputStream +import java.util.concurrent.CountDownLatch +import java.util.concurrent.TimeUnit +import java.util.concurrent.atomic.AtomicInteger import java.util.zip.Deflater import java.util.zip.Inflater @@ -24,8 +27,10 @@ import pekko.http.scaladsl.model.AttributeKeys.webSocketUpgrade import pekko.stream.Materializer import pekko.stream.scaladsl.{ Flow, Keep, Sink, Source } import pekko.stream.testkit.Utils +import pekko.stream.testkit.scaladsl.TestSink import pekko.util.ByteString import pekko.http.impl.engine.server.HttpServerTestSetupBase +import pekko.http.impl.settings.WebSocketCompressionSettingsImpl import pekko.http.impl.settings.WebSocketSettingsImpl import pekko.http.impl.util.PekkoSpecWithMaterializer @@ -601,6 +606,118 @@ class WebSocketServerSpec extends PekkoSpecWithMaterializer("pekko.http.server.w } } + "release compression resources after normal completion" in Utils.assertAllStagesStopped { + val tracking = new TrackingCompression + + Source.empty[FrameEventOrError].via(inflaterFlow(tracking)).runWith(Sink.ignore).futureValue + Source.empty[FrameEvent].via(deflaterFlow(tracking)).runWith(Sink.ignore).futureValue + + tracking.awaitAllEnded() + } + + "release compression resources after upstream failure" in Utils.assertAllStagesStopped { + val tracking = new TrackingCompression + val failure = new RuntimeException("test failure") + + Source.failed[FrameEventOrError](failure) + .via(inflaterFlow(tracking)) + .runWith(Sink.ignore) + .failed + .futureValue shouldEqual failure + Source.failed[FrameEvent](failure) + .via(deflaterFlow(tracking)) + .runWith(Sink.ignore) + .failed + .futureValue shouldEqual failure + + tracking.awaitAllEnded() + } + + "release compression resources after downstream cancellation" in Utils.assertAllStagesStopped { + val tracking = new TrackingCompression + val inflaterProbe = + Source.maybe[FrameEventOrError].via(inflaterFlow(tracking)).runWith(TestSink[FrameEventOrError]()) + val deflaterProbe = Source.maybe[FrameEvent].via(deflaterFlow(tracking)).runWith(TestSink[FrameEvent]()) + + inflaterProbe.cancel() + deflaterProbe.cancel() + + tracking.awaitAllEnded() + } + + "release compression resources after protocol errors" in Utils.assertAllStagesStopped { + val tracking = new TrackingCompression + val invalidInbound = + FrameEvent.fullFrame( + Protocol.Opcode.Text, + None, + ByteString(0xFF, 0xFF, 0xFF), + fin = true, + rsv1 = true) + val invalidOutbound = + FrameEvent.fullFrame(Protocol.Opcode.Text, None, ByteString("reserved"), fin = true, rsv1 = true) + + Source.single[FrameEventOrError](invalidInbound) + .via(inflaterFlow(tracking)) + .runWith(Sink.ignore) + .failed + .futureValue shouldBe a[ProtocolException] + Source.single[FrameEvent](invalidOutbound) + .via(deflaterFlow(tracking)) + .runWith(Sink.ignore) + .failed + .futureValue shouldBe a[ProtocolException] + + tracking.awaitAllEnded() + } + + "release compression resources with incomplete compression state" in Utils.assertAllStagesStopped { + val tracking = new TrackingCompression + val payload = ByteString("unfinished compressed message") + val (firstCompressedFragment, _) = deflatePerMessageFrames(payload, splitAt = 12) + val incompleteInboundMessage = + FrameEvent.fullFrame( + Protocol.Opcode.Text, + None, + firstCompressedFragment, + fin = false, + rsv1 = true) + val incompleteInboundFrame = + FrameStart( + FrameHeader( + Protocol.Opcode.Text, + None, + length = firstCompressedFragment.length + 1, + fin = true, + rsv1 = true), + firstCompressedFragment) + val incompleteOutboundMessage = + FrameEvent.fullFrame(Protocol.Opcode.Text, None, payload, fin = false) + val incompleteOutboundFrame = + FrameStart( + FrameHeader(Protocol.Opcode.Text, None, length = payload.length + 1, fin = true), + payload) + + Source(List[FrameEventOrError](incompleteInboundMessage)) + .via(inflaterFlow(tracking)) + .runWith(Sink.ignore) + .futureValue + Source(List[FrameEventOrError](incompleteInboundFrame)) + .via(inflaterFlow(tracking)) + .runWith(Sink.ignore) + .futureValue + Source(List[FrameEvent](incompleteOutboundMessage)) + .via(deflaterFlow(tracking)) + .runWith(Sink.ignore) + .futureValue + Source(List[FrameEvent](incompleteOutboundFrame)) + .via(deflaterFlow(tracking)) + .runWith(Sink.ignore) + .futureValue + + tracking.awaitAllEnded() + } + "fail invalid compressed messages with a protocol error" in Utils.assertAllStagesStopped { new TestSetup { sendWebSocketRequest("Sec-WebSocket-Extensions: permessage-deflate\r\n") @@ -1404,4 +1521,58 @@ class WebSocketServerSpec extends PekkoSpecWithMaterializer("pekko.http.server.w inflater.end() } } + + private val compressionSettings = WebSocketCompressionSettingsImpl.Disabled.copy(enabled = true) + + private def inflaterFlow(compressionFactory: PerMessageDeflate.CompressionFactory) = + PerMessageDeflate.createInflaterFlow( + noContextTakeover = false, + compressionSettings, + compressionFactory) + + private def deflaterFlow(compressionFactory: PerMessageDeflate.CompressionFactory) = + PerMessageDeflate.createDeflaterFlow( + noContextTakeover = false, + compressionSettings, + compressionFactory) + + private final class TrackingCompression extends PerMessageDeflate.CompressionFactory { + private val inflaterCreated = new AtomicInteger + private val deflaterCreated = new AtomicInteger + private val inflaterEnded = new AtomicInteger + private val deflaterEnded = new AtomicInteger + private val inflaterEndLatch = new CountDownLatch(1) + private val deflaterEndLatch = new CountDownLatch(1) + + override def newInflater(): Inflater = { + inflaterCreated.incrementAndGet() + new Inflater(true) { + override def end(): Unit = { + inflaterEnded.incrementAndGet() + inflaterEndLatch.countDown() + super.end() + } + } + } + + override def newDeflater(level: Int): Deflater = { + deflaterCreated.incrementAndGet() + new Deflater(level, true) { + override def end(): Unit = { + deflaterEnded.incrementAndGet() + deflaterEndLatch.countDown() + super.end() + } + } + } + + def awaitAllEnded(): Unit = { + inflaterCreated.get() should be > 0 + deflaterCreated.get() should be > 0 + inflaterEndLatch.await(3.seconds.toMillis, TimeUnit.MILLISECONDS) shouldEqual true + deflaterEndLatch.await(3.seconds.toMillis, TimeUnit.MILLISECONDS) shouldEqual true + inflaterEnded.get() shouldEqual inflaterCreated.get() + deflaterEnded.get() shouldEqual deflaterCreated.get() + } + } } From 164b9f4b047d27b216d96cc72f7a602341700500 Mon Sep 17 00:00:00 2001 From: Matthias Kurz Date: Thu, 23 Jul 2026 17:54:19 +0200 Subject: [PATCH 3/3] Increase WebSocket decompression limit #61 Use the 256 KiB default agreed during review of the initial WebSocket compression implementation. Keep the disabled settings baseline aligned with reference.conf. --- http-core/src/main/resources/reference.conf | 2 +- .../apache/pekko/http/impl/settings/WebSocketSettingsImpl.scala | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/http-core/src/main/resources/reference.conf b/http-core/src/main/resources/reference.conf index 1784bfd93..b47d6c48e 100644 --- a/http-core/src/main/resources/reference.conf +++ b/http-core/src/main/resources/reference.conf @@ -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 diff --git a/http-core/src/main/scala/org/apache/pekko/http/impl/settings/WebSocketSettingsImpl.scala b/http-core/src/main/scala/org/apache/pekko/http/impl/settings/WebSocketSettingsImpl.scala index 15b6dcaf1..5288e4e33 100644 --- a/http-core/src/main/scala/org/apache/pekko/http/impl/settings/WebSocketSettingsImpl.scala +++ b/http-core/src/main/scala/org/apache/pekko/http/impl/settings/WebSocketSettingsImpl.scala @@ -101,7 +101,7 @@ private[pekko] object WebSocketCompressionSettingsImpl { val Disabled: WebSocketCompressionSettingsImpl = WebSocketCompressionSettingsImpl( enabled = false, - maxAllocation = 64 * 1024, + maxAllocation = 256 * 1024, compressionLevel = 6, preferredClientWindowSize = 15, allowServerNoContext = false,