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]
