This is an automated email from the ASF dual-hosted git repository.
pjfanning pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/pekko-http.git
The following commit(s) were added to refs/heads/main by this push:
new 7591e2e53 Ensure deflater/inflater end() is called in
DeflateCompressor and GzipCompressor (#1142)
7591e2e53 is described below
commit 7591e2e53f5cd997388096961a4734c74ced05f6
Author: PJ Fanning <[email protected]>
AuthorDate: Fri Jul 10 13:37:50 2026 +0100
Ensure deflater/inflater end() is called in DeflateCompressor and
GzipCompressor (#1142)
* Ensure deflater/inflater end() is called in DeflateCompressor and
GzipCompressor
- DeflateCompressor: add deflaterEnded flag + idempotent endDeflater()
method;
call endDeflater() in finishWithBuffer() instead of deflater.end()
directly
- GzipDecompressor: add postStop() calling inflater.end() + add
createInflater()
factory method for testability
- StreamUtils.byteStringTransformer: add optional cleanup callback, call it
in
postStop() when the stage is stopped before onUpstreamFinish()
- Encoder.singleUseEncoderFlow: pass cleanup callback that calls
endDeflater()
on DeflateCompressor instances to handle stream cancellation/failure
- DeflateSpec: add tests for deflater cleanup on normal finish and
cancellation
- GzipSpec: add tests for inflater cleanup on normal finish and
cancellation,
and deflater cleanup on normal finish and cancellation
* compile issues
* Add explicit cleanup() to Compressor/DeflateCompressor; GzipCompressor
inherits it
* revert GzipCompressor changes
* Update GzipSpec.scala
* Update GzipSpec.scala
* Update GzipSpec.scala
* review comments
* Refactor onUpstreamFinish method for clarity
---------
Co-authored-by: copilot-swe-agent[bot]
<[email protected]>
---
.../apache/pekko/http/impl/util/StreamUtils.scala | 11 ++-
.../pekko/http/scaladsl/coding/DeflateSpec.scala | 44 ++++++++++
.../pekko/http/scaladsl/coding/GzipSpec.scala | 96 +++++++++++++++++-----
.../pekko/http/scaladsl/coding/Deflate.scala | 2 +-
.../http/scaladsl/coding/DeflateCompressor.scala | 11 ++-
.../pekko/http/scaladsl/coding/Encoder.scala | 5 +-
.../apache/pekko/http/scaladsl/coding/Gzip.scala | 2 +-
7 files changed, 144 insertions(+), 27 deletions(-)
diff --git
a/http-core/src/main/scala/org/apache/pekko/http/impl/util/StreamUtils.scala
b/http-core/src/main/scala/org/apache/pekko/http/impl/util/StreamUtils.scala
index 44f35f3c5..5f43eca17 100644
--- a/http-core/src/main/scala/org/apache/pekko/http/impl/util/StreamUtils.scala
+++ b/http-core/src/main/scala/org/apache/pekko/http/impl/util/StreamUtils.scala
@@ -43,12 +43,18 @@ private[http] object StreamUtils {
* Creates a transformer that will call `f` for each incoming ByteString and
output its result. After the complete
* input has been read it will call `finish` once to determine the final
ByteString to post to the output.
* Empty ByteStrings are discarded.
+ * If the stage is stopped before the input is fully consumed (e.g. on
downstream cancellation or upstream failure),
+ * `cleanup` is called to release any resources held by the transformer.
*/
def byteStringTransformer(
- f: ByteString => ByteString, finish: () => ByteString):
GraphStage[FlowShape[ByteString, ByteString]] =
+ f: ByteString => ByteString,
+ finish: () => ByteString,
+ cleanup: () => Unit = () => ()): GraphStage[FlowShape[ByteString,
ByteString]] =
new SimpleLinearGraphStage[ByteString] {
override def createLogic(inheritedAttributes: Attributes):
GraphStageLogic =
new GraphStageLogic(shape) with InHandler with OutHandler {
+ private var finished = false
+
override def onPush(): Unit = {
val data = f(grab(in))
if (data.nonEmpty) push(out, data)
@@ -59,10 +65,13 @@ private[http] object StreamUtils {
override def onUpstreamFinish(): Unit = {
val data = finish()
+ finished = true
if (data.nonEmpty) emit(out, data)
completeStage()
}
+ override def postStop(): Unit = if (!finished) cleanup()
+
setHandlers(in, out, this)
}
}
diff --git
a/http-tests/src/test/scala/org/apache/pekko/http/scaladsl/coding/DeflateSpec.scala
b/http-tests/src/test/scala/org/apache/pekko/http/scaladsl/coding/DeflateSpec.scala
index b654b360e..4e99c03bc 100644
---
a/http-tests/src/test/scala/org/apache/pekko/http/scaladsl/coding/DeflateSpec.scala
+++
b/http-tests/src/test/scala/org/apache/pekko/http/scaladsl/coding/DeflateSpec.scala
@@ -92,11 +92,41 @@ class DeflateSpec extends CoderSpec {
decodeWith(inflater, streamEncode(smallTextBytes).dropRight(5))
inflater.endCalls.get() shouldEqual 1
}
+ "release the deflater when encoding completes" in {
+ val tracking = new TrackingDeflater
+ Source.single(smallTextBytes)
+ .via(encoderWith(tracking).encoderFlow)
+ .runWith(Sink.ignore)
+ .awaitResult(3.seconds.dilated)
+ tracking.awaitEnd(3.seconds.dilated)
+ tracking.endCalls.get() shouldEqual 1
+ }
+ "release the deflater when encoding is cancelled early" in {
+ val tracking = new TrackingDeflater
+ Source.single(largeTextBytes)
+ .via(encoderWith(tracking).encoderFlow)
+ .take(1)
+ .runWith(Sink.ignore)
+ .awaitResult(3.seconds.dilated)
+ // postStop() (which calls end()) is dispatched to the stage actor after
the
+ // Sink.ignore future completes, so we must wait for end() itself rather
than
+ // for the stream future to avoid a race.
+ tracking.awaitEnd(3.seconds.dilated)
+ tracking.endCalls.get() shouldEqual 1
+ }
}
private def decodeWith(inflater: TrackingInflater, bytes: ByteString):
ByteString =
decoderWith(inflater).decode(bytes)(SystemMaterializer(system).materializer).awaitResult(3.seconds.dilated)
+ @nowarn("msg=deprecated")
+ private def encoderWith(tracking: TrackingDeflater): Deflate =
+ new Deflate(Encoder.DefaultFilter) {
+ override private[http] def newCompressor: DeflateCompressor = new
DeflateCompressor() {
+ override protected lazy val deflater: java.util.zip.Deflater = tracking
+ }
+ }
+
@nowarn("msg=deprecated")
private def decoderWith(inflater: TrackingInflater): StreamDecoder =
new StreamDecoder {
@@ -123,6 +153,20 @@ class DeflateSpec extends CoderSpec {
}
}
+ private class TrackingDeflater extends
java.util.zip.Deflater(Deflater.DEFAULT_COMPRESSION, false) {
+ val endCalls = new AtomicInteger
+ private val endLatch = new CountDownLatch(1)
+
+ def awaitEnd(atMost: FiniteDuration): Unit =
+ endLatch.await(atMost.toMillis, TimeUnit.MILLISECONDS)
+
+ override def end(): Unit = {
+ endCalls.incrementAndGet()
+ endLatch.countDown()
+ super.end()
+ }
+ }
+
private def encodeMessage(request: HttpRequest, compressionLevel: Int,
noWrap: Boolean): HttpRequest = {
@nowarn("msg=deprecated .* is internal API")
val deflaterWithoutWrapping = new Deflate(Encoder.DefaultFilter) {
diff --git
a/http-tests/src/test/scala/org/apache/pekko/http/scaladsl/coding/GzipSpec.scala
b/http-tests/src/test/scala/org/apache/pekko/http/scaladsl/coding/GzipSpec.scala
index c7f749201..a085413b5 100644
---
a/http-tests/src/test/scala/org/apache/pekko/http/scaladsl/coding/GzipSpec.scala
+++
b/http-tests/src/test/scala/org/apache/pekko/http/scaladsl/coding/GzipSpec.scala
@@ -15,17 +15,19 @@ package org.apache.pekko.http.scaladsl.coding
import java.io.{ InputStream, OutputStream }
import java.nio.charset.StandardCharsets
+import java.util.concurrent.{ CountDownLatch, TimeUnit }
import java.util.concurrent.atomic.AtomicInteger
-import java.util.zip.{ GZIPInputStream, GZIPOutputStream, ZipException }
+import java.util.zip.{ GZIPInputStream, GZIPOutputStream, Inflater,
ZipException }
+import scala.annotation.nowarn
+import scala.concurrent.ExecutionContext.Implicits.global
import scala.concurrent.duration._
import org.apache.pekko
import pekko.http.impl.util._
-import pekko.http.scaladsl.model.headers.HttpEncodings
import pekko.stream.SystemMaterializer
import pekko.stream.scaladsl.{ Sink, Source }
-import pekko.testkit.TestDuration
+import pekko.testkit._
import pekko.util.ByteString
class GzipSpec extends CoderSpec {
@@ -56,23 +58,29 @@ class GzipSpec extends CoderSpec {
val ex = the[RuntimeException] thrownBy
ourDecode(brokenCompress("abcdefghijkl"))
ex.ultimateCause.getMessage should equal("Truncated GZIP stream")
}
+ "throw early if header is corrupt" in {
+ val cause = (the[RuntimeException] thrownBy ourDecode(ByteString(0, 1,
2, 3, 4))).ultimateCause
+ cause should ((be(a[ZipException]) and have).message("Not in GZIP
format"))
+ }
"release the inflater when decoding completes" in {
- val inflater = new TrackingInflater
-
- decodeWith(inflater, streamEncode(smallTextBytes)) should
readAs(smallText)
- inflater.endCalls.get() shouldEqual 1
+ val tracking = new TrackingInflater
+ decodeWith(tracking, streamEncode(smallTextBytes)) should
readAs(smallText)
+ tracking.endCalls.get() shouldEqual 1
}
"release the inflater when decoding is cancelled early" in {
- val inflater = new TrackingInflater
- val compressed = streamEncode(largeTextBytes)
+ val tracking = new TrackingInflater
- Source.single(compressed)
- .via(decoderWith(inflater).withMaxBytesPerChunk(1).decoderFlow)
+ Source.single(streamEncode(largeTextBytes))
+ .via(decoderWith(tracking).withMaxBytesPerChunk(1).decoderFlow)
.take(1)
.runWith(Sink.ignore)
.awaitResult(3.seconds.dilated)
- inflater.endCalls.get() shouldEqual 1
+ // postStop() (which calls end()) is dispatched to the stage actor after
the
+ // Sink.ignore future completes, so we must wait for end() itself rather
than
+ // for the stream future to avoid a race.
+ tracking.awaitEnd(3.seconds.dilated)
+ tracking.endCalls.get() shouldEqual 1
}
"release the inflater when decoding fails on truncation" in {
val inflater = new TrackingInflater
@@ -81,31 +89,75 @@ class GzipSpec extends CoderSpec {
ex.ultimateCause.getMessage should equal("Truncated GZIP stream")
inflater.endCalls.get() shouldEqual 1
}
- "throw early if header is corrupt" in {
- val cause = (the[RuntimeException] thrownBy ourDecode(ByteString(0, 1,
2, 3, 4))).ultimateCause
- cause should ((be(a[ZipException]) and have).message("Not in GZIP
format"))
+ "release the deflater when encoding completes" in {
+ val tracking = new TrackingDeflater
+ Source.single(smallTextBytes)
+ .via(encoderWith(tracking).encoderFlow)
+ .runWith(Sink.ignore)
+ .awaitResult(3.seconds.dilated)
+ tracking.awaitEnd(3.seconds.dilated)
+ tracking.endCalls.get() shouldEqual 1
+ }
+ "release the deflater when encoding is cancelled early" in {
+ val tracking = new TrackingDeflater
+ Source.single(largeTextBytes)
+ .via(encoderWith(tracking).encoderFlow)
+ .take(1)
+ .runWith(Sink.ignore)
+ .awaitResult(3.seconds.dilated)
+ // postStop() (which calls end()) is dispatched to the stage actor after
the
+ // Sink.ignore future completes, so we must wait for end() itself rather
than
+ // for the stream future to avoid a race.
+ tracking.awaitEnd(3.seconds.dilated)
+ tracking.endCalls.get() shouldEqual 1
}
}
- private def decodeWith(inflater: TrackingInflater, bytes: ByteString):
ByteString =
-
decoderWith(inflater).decode(bytes)(SystemMaterializer(system).materializer).awaitResult(3.seconds.dilated)
-
- private def decoderWith(inflater: TrackingInflater): StreamDecoder =
- new StreamDecoder {
- override val encoding = HttpEncodings.gzip
+ private def decodeWith(tracking: TrackingInflater, bytes: ByteString):
ByteString =
+
decoderWith(tracking).decode(bytes)(SystemMaterializer(system).materializer).awaitResult(3.seconds.dilated)
+ @nowarn("msg=deprecated")
+ private def decoderWith(tracking: TrackingInflater): Gzip =
+ new Gzip(Encoder.DefaultFilter) {
override def newDecompressorStage(maxBytesPerChunk: Int) =
() =>
new GzipDecompressor(maxBytesPerChunk) {
- override protected[coding] def createInflater() = inflater
+ override protected[coding] def createInflater(): Inflater =
tracking
}
}
+ @nowarn("msg=deprecated")
+ private def encoderWith(tracking: TrackingDeflater): Gzip =
+ new Gzip(Encoder.DefaultFilter) {
+ override private[http] def newCompressor: GzipCompressor = new
GzipCompressor() {
+ override protected lazy val deflater: java.util.zip.Deflater = tracking
+ }
+ }
+
private class TrackingInflater extends java.util.zip.Inflater(true) {
val endCalls = new AtomicInteger
+ private val endLatch = new CountDownLatch(1)
+
+ def awaitEnd(atMost: FiniteDuration): Unit =
+ endLatch.await(atMost.toMillis, TimeUnit.MILLISECONDS)
+
+ override def end(): Unit = {
+ endCalls.incrementAndGet()
+ endLatch.countDown()
+ super.end()
+ }
+ }
+
+ private class TrackingDeflater extends
java.util.zip.Deflater(java.util.zip.Deflater.DEFAULT_COMPRESSION, true) {
+ val endCalls = new AtomicInteger
+ private val endLatch = new CountDownLatch(1)
+
+ def awaitEnd(atMost: FiniteDuration): Unit =
+ endLatch.await(atMost.toMillis, TimeUnit.MILLISECONDS)
override def end(): Unit = {
endCalls.incrementAndGet()
+ endLatch.countDown()
super.end()
}
}
diff --git
a/http/src/main/scala/org/apache/pekko/http/scaladsl/coding/Deflate.scala
b/http/src/main/scala/org/apache/pekko/http/scaladsl/coding/Deflate.scala
index 5c3c39d5b..6263b5f12 100644
--- a/http/src/main/scala/org/apache/pekko/http/scaladsl/coding/Deflate.scala
+++ b/http/src/main/scala/org/apache/pekko/http/scaladsl/coding/Deflate.scala
@@ -27,7 +27,7 @@ class Deflate private[http] (compressionLevel: Int, val
messageFilter: HttpMessa
}
val encoding = HttpEncodings.deflate
- private[http] def newCompressor = new DeflateCompressor(compressionLevel)
+ private[http] def newCompressor: DeflateCompressor = new
DeflateCompressor(compressionLevel)
def newDecompressorStage(maxBytesPerChunk: Int) = () => new
DeflateDecompressor(maxBytesPerChunk)
@InternalApi // used by javadsl.coding.Coder
diff --git
a/http/src/main/scala/org/apache/pekko/http/scaladsl/coding/DeflateCompressor.scala
b/http/src/main/scala/org/apache/pekko/http/scaladsl/coding/DeflateCompressor.scala
index 87a9d1bdc..bd84ebc3d 100644
---
a/http/src/main/scala/org/apache/pekko/http/scaladsl/coding/DeflateCompressor.scala
+++
b/http/src/main/scala/org/apache/pekko/http/scaladsl/coding/DeflateCompressor.scala
@@ -33,6 +33,7 @@ private[coding] class DeflateCompressor private[coding]
(compressionLevel: Int)
def this() = this(DeflateCompressor.DefaultCompressionLevel)
protected lazy val deflater = new Deflater(compressionLevel, false)
+ private var deflaterEnded = false
override final def compressAndFlush(input: ByteString): ByteString = {
val buffer = newTempBuffer(input.size)
@@ -60,10 +61,18 @@ private[coding] class DeflateCompressor private[coding]
(compressionLevel: Int)
protected def finishWithBuffer(buffer: Array[Byte]): ByteString = {
deflater.finish()
val res = drainDeflater(deflater, buffer)
- deflater.end()
+ endDeflater()
res
}
+ private[coding] def endDeflater(): Unit =
+ if (!deflaterEnded) {
+ deflaterEnded = true
+ deflater.end()
+ }
+
+ private[coding] override def cleanup(): Unit = endDeflater()
+
private def newTempBuffer(size: Int = 65536): Array[Byte] = {
// The default size is somewhat arbitrary, we'd like to guess a better
value but Deflater/zlib
// is buffering in an unpredictable manner.
diff --git
a/http/src/main/scala/org/apache/pekko/http/scaladsl/coding/Encoder.scala
b/http/src/main/scala/org/apache/pekko/http/scaladsl/coding/Encoder.scala
index 7255fd9c0..b6c61d0a5 100644
--- a/http/src/main/scala/org/apache/pekko/http/scaladsl/coding/Encoder.scala
+++ b/http/src/main/scala/org/apache/pekko/http/scaladsl/coding/Encoder.scala
@@ -56,7 +56,7 @@ trait Encoder {
def encodeChunk(bytes: ByteString): ByteString =
compressor.compressAndFlush(bytes)
def finish(): ByteString = compressor.finish()
- StreamUtils.byteStringTransformer(encodeChunk, () => finish())
+ StreamUtils.byteStringTransformer(encodeChunk, () => finish(), () =>
compressor.cleanup())
}
}
@@ -100,4 +100,7 @@ abstract class Compressor {
/** Combines `compress` + `finish` */
def compressAndFinish(input: ByteString): ByteString
+
+ /** Release any native resources held by this compressor. Idempotent. */
+ private[coding] def cleanup(): Unit = ()
}
diff --git
a/http/src/main/scala/org/apache/pekko/http/scaladsl/coding/Gzip.scala
b/http/src/main/scala/org/apache/pekko/http/scaladsl/coding/Gzip.scala
index e7e2926e4..86750cfca 100644
--- a/http/src/main/scala/org/apache/pekko/http/scaladsl/coding/Gzip.scala
+++ b/http/src/main/scala/org/apache/pekko/http/scaladsl/coding/Gzip.scala
@@ -27,7 +27,7 @@ class Gzip private[http] (compressionLevel: Int, val
messageFilter: HttpMessage
}
val encoding = HttpEncodings.gzip
- private[http] def newCompressor = new GzipCompressor(compressionLevel)
+ private[http] def newCompressor: GzipCompressor = new
GzipCompressor(compressionLevel)
def newDecompressorStage(maxBytesPerChunk: Int) = () => new
GzipDecompressor(maxBytesPerChunk)
@InternalApi // used by javadsl.coding.Coder
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]