He-Pin commented on code in PR #1114:
URL: https://github.com/apache/pekko-http/pull/1114#discussion_r3505159616


##########
http-core/src/main/scala/org/apache/pekko/http/impl/engine/ws/PerMessageDeflate.scala:
##########
@@ -0,0 +1,345 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ *    https://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.pekko.http.impl.engine.ws
+
+import java.io.ByteArrayOutputStream
+import java.util.Random
+import java.util.zip.Deflater
+import java.util.zip.Inflater
+import java.util.zip.DataFormatException
+
+import org.apache.pekko
+import pekko.NotUsed
+import pekko.annotation.InternalApi
+import pekko.http.impl.settings.WebSocketCompressionSettingsImpl
+import pekko.http.scaladsl.model.headers.WebSocketExtension
+import pekko.stream.scaladsl.BidiFlow
+import pekko.stream.scaladsl.Flow
+import pekko.stream.stage.GraphStage
+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 scala.collection.immutable
+import scala.collection.immutable.ListMap
+
+/**
+ * INTERNAL API
+ */
+@InternalApi
+private[http] object PerMessageDeflate {
+  private val ExtensionName = "permessage-deflate"
+  private val ClientMaxWindowBits = "client_max_window_bits"
+  private val ServerMaxWindowBits = "server_max_window_bits"
+  private val ClientNoContextTakeover = "client_no_context_takeover"
+  private val ServerNoContextTakeover = "server_no_context_takeover"
+  private val EmptyStoredBlock = ByteString(0x00, 0x00, 0xFF.toByte, 
0xFF.toByte)
+
+  final case class Negotiated(
+      responseExtension: WebSocketExtension,
+      serverNoContextTakeover: Boolean,
+      clientNoContextTakeover: Boolean,
+      settings: WebSocketCompressionSettingsImpl) {
+    def bidiFlow: BidiFlow[FrameEventOrError, FrameEventOrError, FrameEvent, 
FrameEvent, NotUsed] =
+      BidiFlow.fromFlows(inflaterFlow, deflaterFlow)
+
+    def frameEventBidiFlow(
+        maskRandom: () => Random): BidiFlow[FrameEvent, FrameEvent, 
FrameEvent, FrameEvent, NotUsed] =
+      BidiFlow.fromFlows(
+        Flow[FrameEvent]
+          .via(Masking.unmaskIf(condition = true))
+          .via(inflaterFlow)
+          .map {
+            case frame: FrameEvent => frame
+            case FrameError(ex)    => throw ex
+          }
+          .via(Masking.maskIf(condition = true, maskRandom)),
+        deflaterFlow)
+
+    private def inflaterFlow: Flow[FrameEventOrError, FrameEventOrError, 
NotUsed] =
+      Flow.fromGraph(new LifecycleMapConcatStage(
+        "PerMessageDeflate.inflater",
+        () => new InflaterFlow(clientNoContextTakeover, settings)))
+
+    private def deflaterFlow: Flow[FrameEvent, FrameEvent, NotUsed] =
+      Flow.fromGraph(new LifecycleMapConcatStage(
+        "PerMessageDeflate.deflater",
+        () => new DeflaterFlow(serverNoContextTakeover, settings)))
+  }
+
+  def negotiate(
+      requested: immutable.Seq[WebSocketExtension],
+      settings: WebSocketCompressionSettingsImpl): Option[Negotiated] = {
+    if (!settings.enabled) None
+    else {
+      requested.collectFirst(Function.unlift { extension =>
+        if (extension.name.equalsIgnoreCase(ExtensionName)) 
negotiate(extension, settings) else None
+      })
+    }
+  }
+
+  private def negotiate(
+      extension: WebSocketExtension,
+      settings: WebSocketCompressionSettingsImpl): Option[Negotiated] = {
+    var responseParams = ListMap.empty[String, String]
+    var clientNoContext = false
+    var serverNoContext = false
+    var accepted = true
+
+    extension.params.foreach {
+      case (ClientMaxWindowBits, value) =>
+        if (value.isEmpty) responseParams += ClientMaxWindowBits -> 
settings.preferredClientWindowSize.toString
+        else if (validWindowBits(value)) responseParams += ClientMaxWindowBits 
-> value
+        else accepted = false
+      case (ServerMaxWindowBits, value) =>
+        if (value == "15") responseParams += ServerMaxWindowBits -> value
+        else accepted = false
+      case (ClientNoContextTakeover, "") =>
+        clientNoContext = settings.preferredClientNoContext
+        if (clientNoContext) responseParams += ClientNoContextTakeover -> ""
+      case (ServerNoContextTakeover, "") =>
+        if (settings.allowServerNoContext) {
+          serverNoContext = true
+          responseParams += ServerNoContextTakeover -> ""
+        } else accepted = false
+      case _ =>
+        accepted = false
+    }
+
+    if (accepted) {
+      Some(Negotiated(WebSocketExtension(ExtensionName, responseParams), 
serverNoContext, clientNoContext, settings))
+    } else None
+  }
+
+  private def validWindowBits(value: String): Boolean =
+    value.length <= 2 && value.forall(_.isDigit) && {
+      val parsed = value.toInt
+      parsed >= 8 && parsed <= 15
+    }
+
+  private final class InflaterFlow(
+      noContextTakeover: Boolean,
+      settings: WebSocketCompressionSettingsImpl)
+      extends LifecycleMapConcat[FrameEventOrError, FrameEventOrError] {
+    private var inflater = new Inflater(true)
+    private var compressedFrame: Option[CompressedFrame] = None
+    private var compressedMessageInProgress = false
+    private var decompressedMessageBytes = 0L
+    private var bypassFrameInProgress = false
+
+    override def apply(event: FrameEventOrError): 
immutable.Iterable[FrameEventOrError] = event match {
+      case start @ FrameStart(header, data)
+          if header.rsv1 &&
+          (header.opcode == Protocol.Opcode.Text ||
+          header.opcode == Protocol.Opcode.Binary) =>
+        if (compressedMessageInProgress || compressedFrame.isDefined)
+          throw new ProtocolException("Unexpected data frame while fragmented 
message is open")
+        if (header.rsv2 || header.rsv3) throw new 
ProtocolException("Unexpected reserved bit for compressed message")
+        compressedMessageInProgress = !header.fin
+        compressedFrame = Some(CompressedFrame(header.copy(rsv1 = false, 
length = 0), data, appendTail = header.fin))
+        if (start.lastPart) finishFrame() else Nil
+      case start @ FrameStart(header, _) if bypassFrameInProgress =>
+        throw new ProtocolException(s"Unexpected frame ${header.opcode} while 
frame data is open")
+      case start @ FrameStart(header, _)
+          if (compressedFrame.isDefined || compressedMessageInProgress) && 
header.opcode.isControl =>
+        bypassFrameInProgress = !start.lastPart
+        start :: Nil
+      case start @ FrameStart(header, data)
+          if compressedMessageInProgress && header.opcode == 
Protocol.Opcode.Continuation =>
+        if (header.rsv1 || header.rsv2 || header.rsv3)
+          throw new ProtocolException("Unexpected reserved bit for 
continuation frame")
+        compressedMessageInProgress = !header.fin
+        compressedFrame = Some(CompressedFrame(header.copy(length = 0), data, 
appendTail = header.fin))
+        if (start.lastPart) finishFrame() else Nil
+      case start @ FrameStart(header, _) if compressedFrame.isDefined || 
compressedMessageInProgress =>
+        throw new ProtocolException(s"Unexpected frame ${header.opcode} while 
fragmented message is open")
+      case data: FrameData if bypassFrameInProgress =>
+        bypassFrameInProgress = !data.lastPart
+        data :: Nil
+      case data: FrameData if compressedFrame.isDefined =>
+        compressedFrame = compressedFrame.map(_.append(data.data))
+        if (data.lastPart) finishFrame() else Nil
+      case other => other :: Nil
+    }
+
+    private def finishFrame(): immutable.Iterable[FrameEventOrError] = {
+      val frame = compressedFrame.get
+      compressedFrame = None
+      val inflated = inflate(frame.data, frame.appendTail)
+      if (frame.appendTail) decompressedMessageBytes = 0L
+      if (frame.appendTail && noContextTakeover) {
+        inflater.end()
+        inflater = new Inflater(true)
+      }
+      FrameStart(frame.header.copy(length = inflated.length), inflated) :: Nil
+    }
+
+    private def inflate(data: ByteString, appendTail: Boolean): ByteString = {
+      try {
+        val input = if (appendTail) data ++ EmptyStoredBlock else data
+        inflater.setInput(input.toArray)
+        val output = new ByteArrayOutputStream()
+        val buffer = new Array[Byte](1024)

Review Comment:
   The 1024-byte buffer is allocated fresh on every `inflate()` call. Since 
this method runs per frame within a single connection, it might be worth making 
this a class-level field to reduce GC pressure. Also, 1024 is pretty small — 
something like 8192 would cut down loop iterations for larger messages without 
meaningfully increasing memory usage.



##########
http-core/src/main/scala/org/apache/pekko/http/impl/engine/ws/Handshake.scala:
##########
@@ -122,26 +123,47 @@ private[http] object Handshake {
             case OptionVal.Some(p) => p.protocols
             case _                 => Nil
           }
+          val clientRequestedExtensions = headers.collect {
+            case extensions: `Sec-WebSocket-Extensions` => 
extensions.extensions
+          }.flatten
+          val perMessageDeflate =
+            PerMessageDeflate.negotiate(
+              clientRequestedExtensions,
+              settings.asInstanceOf[WebSocketSettingsImpl].compression)

Review Comment:
   This unchecked cast will throw a `ClassCastException` if someone provides a 
custom `WebSocketSettings` implementation that isn't `WebSocketSettingsImpl`. 
Might be worth a pattern match with a graceful fallback (e.g. skip compression 
negotiation for unknown settings types), or at least documenting that 
`WebSocketSettings` must be a `WebSocketSettingsImpl`.



##########
http-core/src/main/scala/org/apache/pekko/http/impl/engine/ws/PerMessageDeflate.scala:
##########
@@ -0,0 +1,345 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ *    https://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.pekko.http.impl.engine.ws
+
+import java.io.ByteArrayOutputStream
+import java.util.Random
+import java.util.zip.Deflater
+import java.util.zip.Inflater
+import java.util.zip.DataFormatException
+
+import org.apache.pekko
+import pekko.NotUsed
+import pekko.annotation.InternalApi
+import pekko.http.impl.settings.WebSocketCompressionSettingsImpl
+import pekko.http.scaladsl.model.headers.WebSocketExtension
+import pekko.stream.scaladsl.BidiFlow
+import pekko.stream.scaladsl.Flow
+import pekko.stream.stage.GraphStage
+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 scala.collection.immutable
+import scala.collection.immutable.ListMap
+
+/**
+ * INTERNAL API
+ */
+@InternalApi
+private[http] object PerMessageDeflate {
+  private val ExtensionName = "permessage-deflate"
+  private val ClientMaxWindowBits = "client_max_window_bits"
+  private val ServerMaxWindowBits = "server_max_window_bits"
+  private val ClientNoContextTakeover = "client_no_context_takeover"
+  private val ServerNoContextTakeover = "server_no_context_takeover"
+  private val EmptyStoredBlock = ByteString(0x00, 0x00, 0xFF.toByte, 
0xFF.toByte)
+
+  final case class Negotiated(
+      responseExtension: WebSocketExtension,
+      serverNoContextTakeover: Boolean,
+      clientNoContextTakeover: Boolean,
+      settings: WebSocketCompressionSettingsImpl) {
+    def bidiFlow: BidiFlow[FrameEventOrError, FrameEventOrError, FrameEvent, 
FrameEvent, NotUsed] =
+      BidiFlow.fromFlows(inflaterFlow, deflaterFlow)
+
+    def frameEventBidiFlow(
+        maskRandom: () => Random): BidiFlow[FrameEvent, FrameEvent, 
FrameEvent, FrameEvent, NotUsed] =
+      BidiFlow.fromFlows(
+        Flow[FrameEvent]
+          .via(Masking.unmaskIf(condition = true))
+          .via(inflaterFlow)
+          .map {
+            case frame: FrameEvent => frame
+            case FrameError(ex)    => throw ex
+          }
+          .via(Masking.maskIf(condition = true, maskRandom)),
+        deflaterFlow)
+
+    private def inflaterFlow: Flow[FrameEventOrError, FrameEventOrError, 
NotUsed] =
+      Flow.fromGraph(new LifecycleMapConcatStage(
+        "PerMessageDeflate.inflater",
+        () => new InflaterFlow(clientNoContextTakeover, settings)))
+
+    private def deflaterFlow: Flow[FrameEvent, FrameEvent, NotUsed] =
+      Flow.fromGraph(new LifecycleMapConcatStage(
+        "PerMessageDeflate.deflater",
+        () => new DeflaterFlow(serverNoContextTakeover, settings)))
+  }
+
+  def negotiate(
+      requested: immutable.Seq[WebSocketExtension],
+      settings: WebSocketCompressionSettingsImpl): Option[Negotiated] = {
+    if (!settings.enabled) None
+    else {
+      requested.collectFirst(Function.unlift { extension =>
+        if (extension.name.equalsIgnoreCase(ExtensionName)) 
negotiate(extension, settings) else None
+      })
+    }
+  }
+
+  private def negotiate(
+      extension: WebSocketExtension,
+      settings: WebSocketCompressionSettingsImpl): Option[Negotiated] = {
+    var responseParams = ListMap.empty[String, String]
+    var clientNoContext = false
+    var serverNoContext = false
+    var accepted = true
+
+    extension.params.foreach {
+      case (ClientMaxWindowBits, value) =>
+        if (value.isEmpty) responseParams += ClientMaxWindowBits -> 
settings.preferredClientWindowSize.toString
+        else if (validWindowBits(value)) responseParams += ClientMaxWindowBits 
-> value
+        else accepted = false
+      case (ServerMaxWindowBits, value) =>
+        if (value == "15") responseParams += ServerMaxWindowBits -> value
+        else accepted = false
+      case (ClientNoContextTakeover, "") =>
+        clientNoContext = settings.preferredClientNoContext
+        if (clientNoContext) responseParams += ClientNoContextTakeover -> ""
+      case (ServerNoContextTakeover, "") =>
+        if (settings.allowServerNoContext) {
+          serverNoContext = true
+          responseParams += ServerNoContextTakeover -> ""
+        } else accepted = false
+      case _ =>
+        accepted = false
+    }
+
+    if (accepted) {
+      Some(Negotiated(WebSocketExtension(ExtensionName, responseParams), 
serverNoContext, clientNoContext, settings))
+    } else None
+  }
+
+  private def validWindowBits(value: String): Boolean =
+    value.length <= 2 && value.forall(_.isDigit) && {
+      val parsed = value.toInt
+      parsed >= 8 && parsed <= 15
+    }
+
+  private final class InflaterFlow(
+      noContextTakeover: Boolean,
+      settings: WebSocketCompressionSettingsImpl)
+      extends LifecycleMapConcat[FrameEventOrError, FrameEventOrError] {
+    private var inflater = new Inflater(true)
+    private var compressedFrame: Option[CompressedFrame] = None
+    private var compressedMessageInProgress = false
+    private var decompressedMessageBytes = 0L
+    private var bypassFrameInProgress = false
+
+    override def apply(event: FrameEventOrError): 
immutable.Iterable[FrameEventOrError] = event match {
+      case start @ FrameStart(header, data)
+          if header.rsv1 &&
+          (header.opcode == Protocol.Opcode.Text ||
+          header.opcode == Protocol.Opcode.Binary) =>
+        if (compressedMessageInProgress || compressedFrame.isDefined)
+          throw new ProtocolException("Unexpected data frame while fragmented 
message is open")
+        if (header.rsv2 || header.rsv3) throw new 
ProtocolException("Unexpected reserved bit for compressed message")
+        compressedMessageInProgress = !header.fin
+        compressedFrame = Some(CompressedFrame(header.copy(rsv1 = false, 
length = 0), data, appendTail = header.fin))
+        if (start.lastPart) finishFrame() else Nil
+      case start @ FrameStart(header, _) if bypassFrameInProgress =>
+        throw new ProtocolException(s"Unexpected frame ${header.opcode} while 
frame data is open")
+      case start @ FrameStart(header, _)
+          if (compressedFrame.isDefined || compressedMessageInProgress) && 
header.opcode.isControl =>
+        bypassFrameInProgress = !start.lastPart
+        start :: Nil
+      case start @ FrameStart(header, data)
+          if compressedMessageInProgress && header.opcode == 
Protocol.Opcode.Continuation =>
+        if (header.rsv1 || header.rsv2 || header.rsv3)
+          throw new ProtocolException("Unexpected reserved bit for 
continuation frame")
+        compressedMessageInProgress = !header.fin
+        compressedFrame = Some(CompressedFrame(header.copy(length = 0), data, 
appendTail = header.fin))
+        if (start.lastPart) finishFrame() else Nil
+      case start @ FrameStart(header, _) if compressedFrame.isDefined || 
compressedMessageInProgress =>
+        throw new ProtocolException(s"Unexpected frame ${header.opcode} while 
fragmented message is open")
+      case data: FrameData if bypassFrameInProgress =>
+        bypassFrameInProgress = !data.lastPart
+        data :: Nil
+      case data: FrameData if compressedFrame.isDefined =>
+        compressedFrame = compressedFrame.map(_.append(data.data))
+        if (data.lastPart) finishFrame() else Nil
+      case other => other :: Nil
+    }
+
+    private def finishFrame(): immutable.Iterable[FrameEventOrError] = {
+      val frame = compressedFrame.get
+      compressedFrame = None
+      val inflated = inflate(frame.data, frame.appendTail)
+      if (frame.appendTail) decompressedMessageBytes = 0L
+      if (frame.appendTail && noContextTakeover) {
+        inflater.end()
+        inflater = new Inflater(true)
+      }
+      FrameStart(frame.header.copy(length = inflated.length), inflated) :: Nil
+    }
+
+    private def inflate(data: ByteString, appendTail: Boolean): ByteString = {
+      try {
+        val input = if (appendTail) data ++ EmptyStoredBlock else data
+        inflater.setInput(input.toArray)
+        val output = new ByteArrayOutputStream()
+        val buffer = new Array[Byte](1024)
+        var count = inflater.inflate(buffer)
+        while (count > 0) {
+          decompressedMessageBytes += count
+          if (settings.maxAllocation > 0 && decompressedMessageBytes > 
settings.maxAllocation)
+            throw new ProtocolException("WebSocket decompressed message 
exceeds configured maximum allocation")
+          output.write(buffer, 0, count)
+          count = inflater.inflate(buffer)
+        }
+        ByteString.fromArray(output.toByteArray)
+      } catch {
+        case ex: DataFormatException =>
+          throw new ProtocolException(s"Invalid WebSocket compressed message: 
${ex.getMessage}")
+      }
+    }
+
+    override def close(): Unit =
+      inflater.end()
+  }
+
+  private final class DeflaterFlow(
+      noContextTakeover: Boolean,
+      settings: WebSocketCompressionSettingsImpl)
+      extends LifecycleMapConcat[FrameEvent, FrameEvent] {
+    private var deflater = new Deflater(settings.compressionLevel, true)
+    private var frame: Option[UncompressedFrame] = None
+    private var messageInProgress = false
+    private var bypassFrameInProgress = false
+
+    override def apply(event: FrameEvent): immutable.Iterable[FrameEvent] = 
event match {
+      case FrameStart(header, _)
+          if (header.opcode == Protocol.Opcode.Text ||
+          header.opcode == Protocol.Opcode.Binary) &&
+          (header.rsv1 || header.rsv2 || header.rsv3) =>
+        throw new ProtocolException("Unexpected reserved bit for outbound 
WebSocket message")
+      case FrameStart(header, _)
+          if header.opcode == Protocol.Opcode.Continuation &&
+          (header.rsv1 || header.rsv2 || header.rsv3) =>
+        throw new ProtocolException("Unexpected reserved bit for outbound 
WebSocket continuation frame")
+      case start @ FrameStart(header, data)
+          if header.opcode == Protocol.Opcode.Text ||
+          header.opcode == Protocol.Opcode.Binary =>
+        if (messageInProgress || frame.isDefined)
+          throw new ProtocolException("Unexpected data frame while fragmented 
message is open")
+        messageInProgress = !header.fin
+        frame = Some(UncompressedFrame(header.copy(length = 0, rsv1 = true), 
data, removeTail = header.fin))
+        if (start.lastPart) finishFrame() else Nil
+      case start @ FrameStart(header, _) if bypassFrameInProgress =>
+        throw new ProtocolException(s"Unexpected frame ${header.opcode} while 
frame data is open")
+      case start @ FrameStart(header, _) if (frame.isDefined || 
messageInProgress) && header.opcode.isControl =>
+        bypassFrameInProgress = !start.lastPart
+        start :: Nil
+      case start @ FrameStart(header, data) if messageInProgress && 
header.opcode == Protocol.Opcode.Continuation =>
+        messageInProgress = !header.fin
+        frame = Some(UncompressedFrame(header.copy(length = 0), data, 
removeTail = header.fin))
+        if (start.lastPart) finishFrame() else Nil
+      case start @ FrameStart(header, _) if frame.isDefined || 
messageInProgress =>
+        throw new ProtocolException(s"Unexpected frame ${header.opcode} while 
fragmented message is open")
+      case data: FrameData if bypassFrameInProgress =>
+        bypassFrameInProgress = !data.lastPart
+        data :: Nil
+      case data: FrameData if frame.isDefined =>
+        frame = frame.map(_.append(data.data))
+        if (data.lastPart) finishFrame() else Nil
+      case other => other :: Nil
+    }
+
+    private def finishFrame(): immutable.Iterable[FrameEvent] = {
+      val current = frame.get
+      frame = None
+      val compressed = deflate(current.data, current.removeTail)
+      if (current.removeTail && noContextTakeover) {
+        deflater.end()
+        deflater = new Deflater(settings.compressionLevel, true)
+      }
+      FrameStart(current.header.copy(length = compressed.length), compressed) 
:: Nil
+    }
+
+    private def deflate(data: ByteString, removeTail: Boolean): ByteString = {
+      deflater.setInput(data.toArray)
+      val output = new ByteArrayOutputStream()
+      val buffer = new Array[Byte](1024)

Review Comment:
   Same as above — this buffer could be a field of `DeflaterFlow` and sized to 
8192 or so. The deflater loop will spin many more times than necessary on 
larger payloads with a 1KB buffer.



##########
http-core/src/main/scala/org/apache/pekko/http/impl/engine/ws/PerMessageDeflate.scala:
##########
@@ -0,0 +1,345 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ *    https://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.pekko.http.impl.engine.ws
+
+import java.io.ByteArrayOutputStream
+import java.util.Random
+import java.util.zip.Deflater
+import java.util.zip.Inflater
+import java.util.zip.DataFormatException
+
+import org.apache.pekko
+import pekko.NotUsed
+import pekko.annotation.InternalApi
+import pekko.http.impl.settings.WebSocketCompressionSettingsImpl
+import pekko.http.scaladsl.model.headers.WebSocketExtension
+import pekko.stream.scaladsl.BidiFlow
+import pekko.stream.scaladsl.Flow
+import pekko.stream.stage.GraphStage
+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 scala.collection.immutable
+import scala.collection.immutable.ListMap
+
+/**
+ * INTERNAL API
+ */
+@InternalApi
+private[http] object PerMessageDeflate {
+  private val ExtensionName = "permessage-deflate"
+  private val ClientMaxWindowBits = "client_max_window_bits"
+  private val ServerMaxWindowBits = "server_max_window_bits"
+  private val ClientNoContextTakeover = "client_no_context_takeover"
+  private val ServerNoContextTakeover = "server_no_context_takeover"
+  private val EmptyStoredBlock = ByteString(0x00, 0x00, 0xFF.toByte, 
0xFF.toByte)
+
+  final case class Negotiated(
+      responseExtension: WebSocketExtension,
+      serverNoContextTakeover: Boolean,
+      clientNoContextTakeover: Boolean,
+      settings: WebSocketCompressionSettingsImpl) {
+    def bidiFlow: BidiFlow[FrameEventOrError, FrameEventOrError, FrameEvent, 
FrameEvent, NotUsed] =
+      BidiFlow.fromFlows(inflaterFlow, deflaterFlow)
+
+    def frameEventBidiFlow(
+        maskRandom: () => Random): BidiFlow[FrameEvent, FrameEvent, 
FrameEvent, FrameEvent, NotUsed] =
+      BidiFlow.fromFlows(
+        Flow[FrameEvent]
+          .via(Masking.unmaskIf(condition = true))
+          .via(inflaterFlow)
+          .map {
+            case frame: FrameEvent => frame
+            case FrameError(ex)    => throw ex
+          }
+          .via(Masking.maskIf(condition = true, maskRandom)),
+        deflaterFlow)
+
+    private def inflaterFlow: Flow[FrameEventOrError, FrameEventOrError, 
NotUsed] =
+      Flow.fromGraph(new LifecycleMapConcatStage(
+        "PerMessageDeflate.inflater",
+        () => new InflaterFlow(clientNoContextTakeover, settings)))
+
+    private def deflaterFlow: Flow[FrameEvent, FrameEvent, NotUsed] =
+      Flow.fromGraph(new LifecycleMapConcatStage(
+        "PerMessageDeflate.deflater",
+        () => new DeflaterFlow(serverNoContextTakeover, settings)))
+  }
+
+  def negotiate(
+      requested: immutable.Seq[WebSocketExtension],
+      settings: WebSocketCompressionSettingsImpl): Option[Negotiated] = {
+    if (!settings.enabled) None
+    else {
+      requested.collectFirst(Function.unlift { extension =>
+        if (extension.name.equalsIgnoreCase(ExtensionName)) 
negotiate(extension, settings) else None
+      })
+    }
+  }
+
+  private def negotiate(
+      extension: WebSocketExtension,
+      settings: WebSocketCompressionSettingsImpl): Option[Negotiated] = {
+    var responseParams = ListMap.empty[String, String]
+    var clientNoContext = false
+    var serverNoContext = false
+    var accepted = true
+
+    extension.params.foreach {
+      case (ClientMaxWindowBits, value) =>
+        if (value.isEmpty) responseParams += ClientMaxWindowBits -> 
settings.preferredClientWindowSize.toString
+        else if (validWindowBits(value)) responseParams += ClientMaxWindowBits 
-> value
+        else accepted = false
+      case (ServerMaxWindowBits, value) =>
+        if (value == "15") responseParams += ServerMaxWindowBits -> value
+        else accepted = false
+      case (ClientNoContextTakeover, "") =>
+        clientNoContext = settings.preferredClientNoContext
+        if (clientNoContext) responseParams += ClientNoContextTakeover -> ""
+      case (ServerNoContextTakeover, "") =>
+        if (settings.allowServerNoContext) {
+          serverNoContext = true
+          responseParams += ServerNoContextTakeover -> ""
+        } else accepted = false
+      case _ =>
+        accepted = false
+    }
+
+    if (accepted) {
+      Some(Negotiated(WebSocketExtension(ExtensionName, responseParams), 
serverNoContext, clientNoContext, settings))
+    } else None
+  }
+
+  private def validWindowBits(value: String): Boolean =
+    value.length <= 2 && value.forall(_.isDigit) && {
+      val parsed = value.toInt
+      parsed >= 8 && parsed <= 15
+    }
+
+  private final class InflaterFlow(
+      noContextTakeover: Boolean,
+      settings: WebSocketCompressionSettingsImpl)
+      extends LifecycleMapConcat[FrameEventOrError, FrameEventOrError] {
+    private var inflater = new Inflater(true)
+    private var compressedFrame: Option[CompressedFrame] = None
+    private var compressedMessageInProgress = false
+    private var decompressedMessageBytes = 0L
+    private var bypassFrameInProgress = false
+
+    override def apply(event: FrameEventOrError): 
immutable.Iterable[FrameEventOrError] = event match {
+      case start @ FrameStart(header, data)
+          if header.rsv1 &&
+          (header.opcode == Protocol.Opcode.Text ||
+          header.opcode == Protocol.Opcode.Binary) =>
+        if (compressedMessageInProgress || compressedFrame.isDefined)
+          throw new ProtocolException("Unexpected data frame while fragmented 
message is open")
+        if (header.rsv2 || header.rsv3) throw new 
ProtocolException("Unexpected reserved bit for compressed message")
+        compressedMessageInProgress = !header.fin
+        compressedFrame = Some(CompressedFrame(header.copy(rsv1 = false, 
length = 0), data, appendTail = header.fin))
+        if (start.lastPart) finishFrame() else Nil
+      case start @ FrameStart(header, _) if bypassFrameInProgress =>
+        throw new ProtocolException(s"Unexpected frame ${header.opcode} while 
frame data is open")
+      case start @ FrameStart(header, _)
+          if (compressedFrame.isDefined || compressedMessageInProgress) && 
header.opcode.isControl =>
+        bypassFrameInProgress = !start.lastPart
+        start :: Nil
+      case start @ FrameStart(header, data)
+          if compressedMessageInProgress && header.opcode == 
Protocol.Opcode.Continuation =>
+        if (header.rsv1 || header.rsv2 || header.rsv3)
+          throw new ProtocolException("Unexpected reserved bit for 
continuation frame")
+        compressedMessageInProgress = !header.fin
+        compressedFrame = Some(CompressedFrame(header.copy(length = 0), data, 
appendTail = header.fin))
+        if (start.lastPart) finishFrame() else Nil
+      case start @ FrameStart(header, _) if compressedFrame.isDefined || 
compressedMessageInProgress =>
+        throw new ProtocolException(s"Unexpected frame ${header.opcode} while 
fragmented message is open")
+      case data: FrameData if bypassFrameInProgress =>
+        bypassFrameInProgress = !data.lastPart
+        data :: Nil
+      case data: FrameData if compressedFrame.isDefined =>
+        compressedFrame = compressedFrame.map(_.append(data.data))
+        if (data.lastPart) finishFrame() else Nil
+      case other => other :: Nil
+    }
+
+    private def finishFrame(): immutable.Iterable[FrameEventOrError] = {
+      val frame = compressedFrame.get
+      compressedFrame = None
+      val inflated = inflate(frame.data, frame.appendTail)
+      if (frame.appendTail) decompressedMessageBytes = 0L
+      if (frame.appendTail && noContextTakeover) {
+        inflater.end()
+        inflater = new Inflater(true)
+      }
+      FrameStart(frame.header.copy(length = inflated.length), inflated) :: Nil
+    }
+
+    private def inflate(data: ByteString, appendTail: Boolean): ByteString = {
+      try {
+        val input = if (appendTail) data ++ EmptyStoredBlock else data
+        inflater.setInput(input.toArray)
+        val output = new ByteArrayOutputStream()
+        val buffer = new Array[Byte](1024)
+        var count = inflater.inflate(buffer)
+        while (count > 0) {
+          decompressedMessageBytes += count
+          if (settings.maxAllocation > 0 && decompressedMessageBytes > 
settings.maxAllocation)
+            throw new ProtocolException("WebSocket decompressed message 
exceeds configured maximum allocation")
+          output.write(buffer, 0, count)
+          count = inflater.inflate(buffer)
+        }
+        ByteString.fromArray(output.toByteArray)
+      } catch {
+        case ex: DataFormatException =>
+          throw new ProtocolException(s"Invalid WebSocket compressed message: 
${ex.getMessage}")
+      }
+    }
+
+    override def close(): Unit =
+      inflater.end()
+  }
+
+  private final class DeflaterFlow(
+      noContextTakeover: Boolean,
+      settings: WebSocketCompressionSettingsImpl)
+      extends LifecycleMapConcat[FrameEvent, FrameEvent] {
+    private var deflater = new Deflater(settings.compressionLevel, true)
+    private var frame: Option[UncompressedFrame] = None
+    private var messageInProgress = false
+    private var bypassFrameInProgress = false
+
+    override def apply(event: FrameEvent): immutable.Iterable[FrameEvent] = 
event match {
+      case FrameStart(header, _)
+          if (header.opcode == Protocol.Opcode.Text ||
+          header.opcode == Protocol.Opcode.Binary) &&
+          (header.rsv1 || header.rsv2 || header.rsv3) =>
+        throw new ProtocolException("Unexpected reserved bit for outbound 
WebSocket message")
+      case FrameStart(header, _)
+          if header.opcode == Protocol.Opcode.Continuation &&
+          (header.rsv1 || header.rsv2 || header.rsv3) =>
+        throw new ProtocolException("Unexpected reserved bit for outbound 
WebSocket continuation frame")
+      case start @ FrameStart(header, data)
+          if header.opcode == Protocol.Opcode.Text ||
+          header.opcode == Protocol.Opcode.Binary =>
+        if (messageInProgress || frame.isDefined)
+          throw new ProtocolException("Unexpected data frame while fragmented 
message is open")
+        messageInProgress = !header.fin
+        frame = Some(UncompressedFrame(header.copy(length = 0, rsv1 = true), 
data, removeTail = header.fin))
+        if (start.lastPart) finishFrame() else Nil
+      case start @ FrameStart(header, _) if bypassFrameInProgress =>
+        throw new ProtocolException(s"Unexpected frame ${header.opcode} while 
frame data is open")
+      case start @ FrameStart(header, _) if (frame.isDefined || 
messageInProgress) && header.opcode.isControl =>
+        bypassFrameInProgress = !start.lastPart
+        start :: Nil
+      case start @ FrameStart(header, data) if messageInProgress && 
header.opcode == Protocol.Opcode.Continuation =>
+        messageInProgress = !header.fin
+        frame = Some(UncompressedFrame(header.copy(length = 0), data, 
removeTail = header.fin))
+        if (start.lastPart) finishFrame() else Nil
+      case start @ FrameStart(header, _) if frame.isDefined || 
messageInProgress =>
+        throw new ProtocolException(s"Unexpected frame ${header.opcode} while 
fragmented message is open")
+      case data: FrameData if bypassFrameInProgress =>
+        bypassFrameInProgress = !data.lastPart
+        data :: Nil
+      case data: FrameData if frame.isDefined =>
+        frame = frame.map(_.append(data.data))
+        if (data.lastPart) finishFrame() else Nil
+      case other => other :: Nil
+    }
+
+    private def finishFrame(): immutable.Iterable[FrameEvent] = {
+      val current = frame.get
+      frame = None
+      val compressed = deflate(current.data, current.removeTail)
+      if (current.removeTail && noContextTakeover) {
+        deflater.end()
+        deflater = new Deflater(settings.compressionLevel, true)
+      }
+      FrameStart(current.header.copy(length = compressed.length), compressed) 
:: Nil
+    }
+
+    private def deflate(data: ByteString, removeTail: Boolean): ByteString = {
+      deflater.setInput(data.toArray)
+      val output = new ByteArrayOutputStream()
+      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)
+        count = deflater.deflate(buffer, 0, buffer.length, Deflater.SYNC_FLUSH)
+      }
+      val bytes = ByteString.fromArray(output.toByteArray)
+      if (removeTail && bytes.endsWith(EmptyStoredBlock)) 
bytes.dropRight(EmptyStoredBlock.length) else bytes
+    }
+
+    override def close(): Unit =
+      deflater.end()
+  }
+
+  private trait LifecycleMapConcat[-In, +Out] extends (In => 
immutable.Iterable[Out]) {
+    def close(): Unit
+  }
+
+  private final class LifecycleMapConcatStage[In, Out](
+      name: String,
+      create: () => LifecycleMapConcat[In, Out])
+      extends GraphStage[FlowShape[In, Out]] {
+    private val in = Inlet[In](s"$name.in")
+    private val out = Outlet[Out](s"$name.out")
+    override val shape: FlowShape[In, Out] = FlowShape(in, out)
+
+    override def createLogic(inheritedAttributes: Attributes): GraphStageLogic 
=
+      new GraphStageLogic(shape) with InHandler with OutHandler {
+        private val handler = create()
+        private var pending = Iterator.empty[Out]
+        private var upstreamFinished = false
+
+        override def onPush(): Unit = {
+          pending = handler(grab(in)).iterator
+          pushOrPull()
+        }
+
+        override def onPull(): Unit =
+          pushOrPull()
+
+        override def onUpstreamFinish(): Unit = {
+          upstreamFinished = true
+          if (!pending.hasNext) completeStage()
+        }
+
+        override def postStop(): Unit =
+          handler.close()
+
+        private def pushOrPull(): Unit =
+          if (pending.hasNext) push(out, pending.next())
+          else if (upstreamFinished) completeStage()
+          else if (!hasBeenPulled(in)) pull(in)
+
+        setHandler(in, this)
+        setHandler(out, this)
+      }
+  }
+
+  private final case class CompressedFrame(header: FrameHeader, data: 
ByteString, appendTail: Boolean) {
+    def append(next: ByteString): CompressedFrame = copy(data = data ++ next)
+  }

Review Comment:
   '`data ++ next` creates a new ByteString copy on every fragment. For heavily 
fragmented messages this becomes O(n²) in the total payload size. Consider 
accumulating fragments in a `Vector[ByteString]` (or a `ByteStringBuilder`) and 
concatenating once in `finishFrame()`. Same applies to 
`UncompressedFrame.append` below.'



##########
http-core/src/main/resources/reference.conf:
##########
@@ -348,6 +348,44 @@ pekko.http {
 
       # Enable verbose debug logging for all ingoing and outgoing frames
       log-frames = false
+
+      compression {
+        # Whether the server should support WebSocket compression using the 
RFC 7692
+        # permessage-deflate extension. Compression is negotiated during the
+        # WebSocket handshake and is only used when the client requests it.
+        enabled = true
+
+        # Maximum size of a decompressed WebSocket message. If this value is
+        # 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

Review Comment:
   64 KB feels quite conservative for a default. Real-world WebSocket apps 
(collaborative editing, gaming, large JSON payloads) routinely exceed this. 
Tomcat has no limit by default, Jetty uses 128 KB for text message size. 
Something like 256 KB or 1 MB might be a more practical default while still 
protecting against decompression bombs. Users who need tighter limits can 
always lower it.



##########
http-core/src/main/scala/org/apache/pekko/http/impl/engine/ws/PerMessageDeflate.scala:
##########
@@ -0,0 +1,345 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ *    https://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.pekko.http.impl.engine.ws
+
+import java.io.ByteArrayOutputStream
+import java.util.Random
+import java.util.zip.Deflater
+import java.util.zip.Inflater
+import java.util.zip.DataFormatException
+
+import org.apache.pekko
+import pekko.NotUsed
+import pekko.annotation.InternalApi
+import pekko.http.impl.settings.WebSocketCompressionSettingsImpl
+import pekko.http.scaladsl.model.headers.WebSocketExtension
+import pekko.stream.scaladsl.BidiFlow
+import pekko.stream.scaladsl.Flow
+import pekko.stream.stage.GraphStage
+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 scala.collection.immutable
+import scala.collection.immutable.ListMap
+
+/**
+ * INTERNAL API
+ */
+@InternalApi
+private[http] object PerMessageDeflate {
+  private val ExtensionName = "permessage-deflate"
+  private val ClientMaxWindowBits = "client_max_window_bits"
+  private val ServerMaxWindowBits = "server_max_window_bits"
+  private val ClientNoContextTakeover = "client_no_context_takeover"
+  private val ServerNoContextTakeover = "server_no_context_takeover"
+  private val EmptyStoredBlock = ByteString(0x00, 0x00, 0xFF.toByte, 
0xFF.toByte)
+
+  final case class Negotiated(
+      responseExtension: WebSocketExtension,
+      serverNoContextTakeover: Boolean,
+      clientNoContextTakeover: Boolean,
+      settings: WebSocketCompressionSettingsImpl) {
+    def bidiFlow: BidiFlow[FrameEventOrError, FrameEventOrError, FrameEvent, 
FrameEvent, NotUsed] =
+      BidiFlow.fromFlows(inflaterFlow, deflaterFlow)
+
+    def frameEventBidiFlow(
+        maskRandom: () => Random): BidiFlow[FrameEvent, FrameEvent, 
FrameEvent, FrameEvent, NotUsed] =
+      BidiFlow.fromFlows(
+        Flow[FrameEvent]
+          .via(Masking.unmaskIf(condition = true))
+          .via(inflaterFlow)
+          .map {
+            case frame: FrameEvent => frame
+            case FrameError(ex)    => throw ex
+          }
+          .via(Masking.maskIf(condition = true, maskRandom)),
+        deflaterFlow)
+
+    private def inflaterFlow: Flow[FrameEventOrError, FrameEventOrError, 
NotUsed] =
+      Flow.fromGraph(new LifecycleMapConcatStage(
+        "PerMessageDeflate.inflater",
+        () => new InflaterFlow(clientNoContextTakeover, settings)))
+
+    private def deflaterFlow: Flow[FrameEvent, FrameEvent, NotUsed] =
+      Flow.fromGraph(new LifecycleMapConcatStage(
+        "PerMessageDeflate.deflater",
+        () => new DeflaterFlow(serverNoContextTakeover, settings)))
+  }
+
+  def negotiate(
+      requested: immutable.Seq[WebSocketExtension],
+      settings: WebSocketCompressionSettingsImpl): Option[Negotiated] = {
+    if (!settings.enabled) None
+    else {
+      requested.collectFirst(Function.unlift { extension =>
+        if (extension.name.equalsIgnoreCase(ExtensionName)) 
negotiate(extension, settings) else None
+      })
+    }
+  }
+
+  private def negotiate(
+      extension: WebSocketExtension,
+      settings: WebSocketCompressionSettingsImpl): Option[Negotiated] = {
+    var responseParams = ListMap.empty[String, String]
+    var clientNoContext = false
+    var serverNoContext = false
+    var accepted = true
+
+    extension.params.foreach {
+      case (ClientMaxWindowBits, value) =>
+        if (value.isEmpty) responseParams += ClientMaxWindowBits -> 
settings.preferredClientWindowSize.toString
+        else if (validWindowBits(value)) responseParams += ClientMaxWindowBits 
-> value
+        else accepted = false
+      case (ServerMaxWindowBits, value) =>
+        if (value == "15") responseParams += ServerMaxWindowBits -> value
+        else accepted = false
+      case (ClientNoContextTakeover, "") =>
+        clientNoContext = settings.preferredClientNoContext
+        if (clientNoContext) responseParams += ClientNoContextTakeover -> ""
+      case (ServerNoContextTakeover, "") =>
+        if (settings.allowServerNoContext) {
+          serverNoContext = true
+          responseParams += ServerNoContextTakeover -> ""
+        } else accepted = false
+      case _ =>
+        accepted = false
+    }
+
+    if (accepted) {
+      Some(Negotiated(WebSocketExtension(ExtensionName, responseParams), 
serverNoContext, clientNoContext, settings))
+    } else None
+  }
+
+  private def validWindowBits(value: String): Boolean =
+    value.length <= 2 && value.forall(_.isDigit) && {

Review Comment:
   Nit: if this ever gets called with an empty string (it won't today because 
the empty case is handled before reaching here), `value.toInt` would throw 
`NumberFormatException`. A quick `value.nonEmpty &&` guard at the start would 
make this more defensive for future callers.



##########
http-core/src/main/scala/org/apache/pekko/http/impl/engine/ws/PerMessageDeflate.scala:
##########
@@ -0,0 +1,345 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ *    https://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.pekko.http.impl.engine.ws
+
+import java.io.ByteArrayOutputStream
+import java.util.Random
+import java.util.zip.Deflater
+import java.util.zip.Inflater
+import java.util.zip.DataFormatException
+
+import org.apache.pekko
+import pekko.NotUsed
+import pekko.annotation.InternalApi
+import pekko.http.impl.settings.WebSocketCompressionSettingsImpl
+import pekko.http.scaladsl.model.headers.WebSocketExtension
+import pekko.stream.scaladsl.BidiFlow
+import pekko.stream.scaladsl.Flow
+import pekko.stream.stage.GraphStage
+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 scala.collection.immutable
+import scala.collection.immutable.ListMap
+
+/**
+ * INTERNAL API
+ */
+@InternalApi
+private[http] object PerMessageDeflate {
+  private val ExtensionName = "permessage-deflate"
+  private val ClientMaxWindowBits = "client_max_window_bits"
+  private val ServerMaxWindowBits = "server_max_window_bits"
+  private val ClientNoContextTakeover = "client_no_context_takeover"
+  private val ServerNoContextTakeover = "server_no_context_takeover"
+  private val EmptyStoredBlock = ByteString(0x00, 0x00, 0xFF.toByte, 
0xFF.toByte)
+
+  final case class Negotiated(
+      responseExtension: WebSocketExtension,
+      serverNoContextTakeover: Boolean,
+      clientNoContextTakeover: Boolean,
+      settings: WebSocketCompressionSettingsImpl) {
+    def bidiFlow: BidiFlow[FrameEventOrError, FrameEventOrError, FrameEvent, 
FrameEvent, NotUsed] =
+      BidiFlow.fromFlows(inflaterFlow, deflaterFlow)
+
+    def frameEventBidiFlow(
+        maskRandom: () => Random): BidiFlow[FrameEvent, FrameEvent, 
FrameEvent, FrameEvent, NotUsed] =
+      BidiFlow.fromFlows(
+        Flow[FrameEvent]
+          .via(Masking.unmaskIf(condition = true))
+          .via(inflaterFlow)
+          .map {
+            case frame: FrameEvent => frame
+            case FrameError(ex)    => throw ex
+          }
+          .via(Masking.maskIf(condition = true, maskRandom)),
+        deflaterFlow)
+
+    private def inflaterFlow: Flow[FrameEventOrError, FrameEventOrError, 
NotUsed] =
+      Flow.fromGraph(new LifecycleMapConcatStage(
+        "PerMessageDeflate.inflater",
+        () => new InflaterFlow(clientNoContextTakeover, settings)))
+
+    private def deflaterFlow: Flow[FrameEvent, FrameEvent, NotUsed] =
+      Flow.fromGraph(new LifecycleMapConcatStage(
+        "PerMessageDeflate.deflater",
+        () => new DeflaterFlow(serverNoContextTakeover, settings)))
+  }
+
+  def negotiate(
+      requested: immutable.Seq[WebSocketExtension],
+      settings: WebSocketCompressionSettingsImpl): Option[Negotiated] = {
+    if (!settings.enabled) None
+    else {
+      requested.collectFirst(Function.unlift { extension =>
+        if (extension.name.equalsIgnoreCase(ExtensionName)) 
negotiate(extension, settings) else None
+      })
+    }
+  }
+
+  private def negotiate(
+      extension: WebSocketExtension,
+      settings: WebSocketCompressionSettingsImpl): Option[Negotiated] = {
+    var responseParams = ListMap.empty[String, String]
+    var clientNoContext = false
+    var serverNoContext = false
+    var accepted = true
+
+    extension.params.foreach {
+      case (ClientMaxWindowBits, value) =>
+        if (value.isEmpty) responseParams += ClientMaxWindowBits -> 
settings.preferredClientWindowSize.toString
+        else if (validWindowBits(value)) responseParams += ClientMaxWindowBits 
-> value
+        else accepted = false
+      case (ServerMaxWindowBits, value) =>
+        if (value == "15") responseParams += ServerMaxWindowBits -> value
+        else accepted = false
+      case (ClientNoContextTakeover, "") =>
+        clientNoContext = settings.preferredClientNoContext
+        if (clientNoContext) responseParams += ClientNoContextTakeover -> ""
+      case (ServerNoContextTakeover, "") =>
+        if (settings.allowServerNoContext) {
+          serverNoContext = true
+          responseParams += ServerNoContextTakeover -> ""
+        } else accepted = false
+      case _ =>
+        accepted = false
+    }
+
+    if (accepted) {
+      Some(Negotiated(WebSocketExtension(ExtensionName, responseParams), 
serverNoContext, clientNoContext, settings))
+    } else None
+  }
+
+  private def validWindowBits(value: String): Boolean =
+    value.length <= 2 && value.forall(_.isDigit) && {
+      val parsed = value.toInt
+      parsed >= 8 && parsed <= 15
+    }
+
+  private final class InflaterFlow(
+      noContextTakeover: Boolean,
+      settings: WebSocketCompressionSettingsImpl)
+      extends LifecycleMapConcat[FrameEventOrError, FrameEventOrError] {
+    private var inflater = new Inflater(true)
+    private var compressedFrame: Option[CompressedFrame] = None
+    private var compressedMessageInProgress = false
+    private var decompressedMessageBytes = 0L
+    private var bypassFrameInProgress = false
+
+    override def apply(event: FrameEventOrError): 
immutable.Iterable[FrameEventOrError] = event match {
+      case start @ FrameStart(header, data)
+          if header.rsv1 &&
+          (header.opcode == Protocol.Opcode.Text ||
+          header.opcode == Protocol.Opcode.Binary) =>
+        if (compressedMessageInProgress || compressedFrame.isDefined)
+          throw new ProtocolException("Unexpected data frame while fragmented 
message is open")
+        if (header.rsv2 || header.rsv3) throw new 
ProtocolException("Unexpected reserved bit for compressed message")
+        compressedMessageInProgress = !header.fin
+        compressedFrame = Some(CompressedFrame(header.copy(rsv1 = false, 
length = 0), data, appendTail = header.fin))
+        if (start.lastPart) finishFrame() else Nil
+      case start @ FrameStart(header, _) if bypassFrameInProgress =>
+        throw new ProtocolException(s"Unexpected frame ${header.opcode} while 
frame data is open")
+      case start @ FrameStart(header, _)
+          if (compressedFrame.isDefined || compressedMessageInProgress) && 
header.opcode.isControl =>
+        bypassFrameInProgress = !start.lastPart
+        start :: Nil
+      case start @ FrameStart(header, data)
+          if compressedMessageInProgress && header.opcode == 
Protocol.Opcode.Continuation =>
+        if (header.rsv1 || header.rsv2 || header.rsv3)
+          throw new ProtocolException("Unexpected reserved bit for 
continuation frame")
+        compressedMessageInProgress = !header.fin
+        compressedFrame = Some(CompressedFrame(header.copy(length = 0), data, 
appendTail = header.fin))
+        if (start.lastPart) finishFrame() else Nil
+      case start @ FrameStart(header, _) if compressedFrame.isDefined || 
compressedMessageInProgress =>
+        throw new ProtocolException(s"Unexpected frame ${header.opcode} while 
fragmented message is open")
+      case data: FrameData if bypassFrameInProgress =>
+        bypassFrameInProgress = !data.lastPart
+        data :: Nil
+      case data: FrameData if compressedFrame.isDefined =>
+        compressedFrame = compressedFrame.map(_.append(data.data))
+        if (data.lastPart) finishFrame() else Nil
+      case other => other :: Nil
+    }
+
+    private def finishFrame(): immutable.Iterable[FrameEventOrError] = {
+      val frame = compressedFrame.get
+      compressedFrame = None
+      val inflated = inflate(frame.data, frame.appendTail)
+      if (frame.appendTail) decompressedMessageBytes = 0L
+      if (frame.appendTail && noContextTakeover) {
+        inflater.end()
+        inflater = new Inflater(true)
+      }
+      FrameStart(frame.header.copy(length = inflated.length), inflated) :: Nil
+    }
+
+    private def inflate(data: ByteString, appendTail: Boolean): ByteString = {
+      try {
+        val input = if (appendTail) data ++ EmptyStoredBlock else data
+        inflater.setInput(input.toArray)
+        val output = new ByteArrayOutputStream()
+        val buffer = new Array[Byte](1024)
+        var count = inflater.inflate(buffer)
+        while (count > 0) {
+          decompressedMessageBytes += count
+          if (settings.maxAllocation > 0 && decompressedMessageBytes > 
settings.maxAllocation)
+            throw new ProtocolException("WebSocket decompressed message 
exceeds configured maximum allocation")
+          output.write(buffer, 0, count)
+          count = inflater.inflate(buffer)
+        }
+        ByteString.fromArray(output.toByteArray)
+      } catch {
+        case ex: DataFormatException =>
+          throw new ProtocolException(s"Invalid WebSocket compressed message: 
${ex.getMessage}")
+      }
+    }
+
+    override def close(): Unit =
+      inflater.end()
+  }
+
+  private final class DeflaterFlow(
+      noContextTakeover: Boolean,
+      settings: WebSocketCompressionSettingsImpl)
+      extends LifecycleMapConcat[FrameEvent, FrameEvent] {
+    private var deflater = new Deflater(settings.compressionLevel, true)
+    private var frame: Option[UncompressedFrame] = None
+    private var messageInProgress = false
+    private var bypassFrameInProgress = false
+
+    override def apply(event: FrameEvent): immutable.Iterable[FrameEvent] = 
event match {
+      case FrameStart(header, _)
+          if (header.opcode == Protocol.Opcode.Text ||
+          header.opcode == Protocol.Opcode.Binary) &&
+          (header.rsv1 || header.rsv2 || header.rsv3) =>
+        throw new ProtocolException("Unexpected reserved bit for outbound 
WebSocket message")
+      case FrameStart(header, _)
+          if header.opcode == Protocol.Opcode.Continuation &&
+          (header.rsv1 || header.rsv2 || header.rsv3) =>
+        throw new ProtocolException("Unexpected reserved bit for outbound 
WebSocket continuation frame")
+      case start @ FrameStart(header, data)
+          if header.opcode == Protocol.Opcode.Text ||
+          header.opcode == Protocol.Opcode.Binary =>
+        if (messageInProgress || frame.isDefined)
+          throw new ProtocolException("Unexpected data frame while fragmented 
message is open")
+        messageInProgress = !header.fin
+        frame = Some(UncompressedFrame(header.copy(length = 0, rsv1 = true), 
data, removeTail = header.fin))
+        if (start.lastPart) finishFrame() else Nil
+      case start @ FrameStart(header, _) if bypassFrameInProgress =>
+        throw new ProtocolException(s"Unexpected frame ${header.opcode} while 
frame data is open")
+      case start @ FrameStart(header, _) if (frame.isDefined || 
messageInProgress) && header.opcode.isControl =>
+        bypassFrameInProgress = !start.lastPart
+        start :: Nil
+      case start @ FrameStart(header, data) if messageInProgress && 
header.opcode == Protocol.Opcode.Continuation =>
+        messageInProgress = !header.fin
+        frame = Some(UncompressedFrame(header.copy(length = 0), data, 
removeTail = header.fin))
+        if (start.lastPart) finishFrame() else Nil
+      case start @ FrameStart(header, _) if frame.isDefined || 
messageInProgress =>
+        throw new ProtocolException(s"Unexpected frame ${header.opcode} while 
fragmented message is open")
+      case data: FrameData if bypassFrameInProgress =>
+        bypassFrameInProgress = !data.lastPart
+        data :: Nil
+      case data: FrameData if frame.isDefined =>
+        frame = frame.map(_.append(data.data))
+        if (data.lastPart) finishFrame() else Nil
+      case other => other :: Nil
+    }
+
+    private def finishFrame(): immutable.Iterable[FrameEvent] = {
+      val current = frame.get
+      frame = None
+      val compressed = deflate(current.data, current.removeTail)
+      if (current.removeTail && noContextTakeover) {
+        deflater.end()
+        deflater = new Deflater(settings.compressionLevel, true)
+      }
+      FrameStart(current.header.copy(length = compressed.length), compressed) 
:: Nil
+    }
+
+    private def deflate(data: ByteString, removeTail: Boolean): ByteString = {
+      deflater.setInput(data.toArray)
+      val output = new ByteArrayOutputStream()
+      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)
+        count = deflater.deflate(buffer, 0, buffer.length, Deflater.SYNC_FLUSH)
+      }
+      val bytes = ByteString.fromArray(output.toByteArray)
+      if (removeTail && bytes.endsWith(EmptyStoredBlock)) 
bytes.dropRight(EmptyStoredBlock.length) else bytes
+    }
+
+    override def close(): Unit =
+      deflater.end()
+  }
+
+  private trait LifecycleMapConcat[-In, +Out] extends (In => 
immutable.Iterable[Out]) {
+    def close(): Unit
+  }
+
+  private final class LifecycleMapConcatStage[In, Out](
+      name: String,
+      create: () => LifecycleMapConcat[In, Out])
+      extends GraphStage[FlowShape[In, Out]] {
+    private val in = Inlet[In](s"$name.in")
+    private val out = Outlet[Out](s"$name.out")
+    override val shape: FlowShape[In, Out] = FlowShape(in, out)
+
+    override def createLogic(inheritedAttributes: Attributes): GraphStageLogic 
=
+      new GraphStageLogic(shape) with InHandler with OutHandler {
+        private val handler = create()
+        private var pending = Iterator.empty[Out]
+        private var upstreamFinished = false
+
+        override def onPush(): Unit = {
+          pending = handler(grab(in)).iterator
+          pushOrPull()
+        }
+
+        override def onPull(): Unit =
+          pushOrPull()
+
+        override def onUpstreamFinish(): Unit = {

Review Comment:
   If upstream finishes while a fragmented compressed message is still in 
progress (`compressedMessageInProgress == true` or 
`compressedFrame.isDefined`), the pending data is silently dropped here. Not 
sure if that's the right behavior — maybe emit a protocol error or at least log 
a warning so it's not a silent data loss?



-- 
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.

To unsubscribe, e-mail: [email protected]

For queries about this service, please contact Infrastructure at:
[email protected]


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

Reply via email to