This is an automated email from the ASF dual-hosted git repository.

pjfanning pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/pekko-http.git


The following commit(s) were added to refs/heads/main by this push:
     new 6f5670fa9 Improve WebSocket compression internals and resource cleanup 
(#1173)
6f5670fa9 is described below

commit 6f5670fa9ab9f1f22879f8f9145612a267cc28ca
Author: Matthias Kurz <[email protected]>
AuthorDate: Thu Jul 23 23:14:45 2026 +0200

    Improve WebSocket compression internals and resource cleanup (#1173)
    
    * 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.
    
    * 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.
    
    * 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 +-
 .../pekko/http/impl/engine/ws/Handshake.scala      |  31 ++--
 .../http/impl/engine/ws/PerMessageDeflate.scala    |  93 ++++++++---
 .../http/impl/settings/WebSocketSettingsImpl.scala |   2 +-
 .../http/impl/engine/ws/WebSocketServerSpec.scala  | 171 +++++++++++++++++++++
 5 files changed, 262 insertions(+), 37 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/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..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
@@ -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
@@ -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] = {
@@ -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)
@@ -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
     }
@@ -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
@@ -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, _)
@@ -269,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
     }
@@ -277,7 +302,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 +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()
+    }
 }
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..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 = 0,
+      maxAllocation = 256 * 1024,
       compressionLevel = 6,
       preferredClientWindowSize = 15,
       allowServerNoContext = false,
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()
+    }
+  }
 }


---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]

Reply via email to