From 982629c3e80440c12115e5a8225933c90d30693d Mon Sep 17 00:00:00 2001 From: Rico Date: Mon, 26 Jan 2026 13:33:47 +0100 Subject: [PATCH 1/3] Guard Jetty header writes against _fields NPE --- .../scala/io/udash/rest/RestServlet.scala | 25 +++++++++++++------ 1 file changed, 17 insertions(+), 8 deletions(-) diff --git a/rest/.jvm/src/main/scala/io/udash/rest/RestServlet.scala b/rest/.jvm/src/main/scala/io/udash/rest/RestServlet.scala index c569a3532..2399da199 100644 --- a/rest/.jvm/src/main/scala/io/udash/rest/RestServlet.scala +++ b/rest/.jvm/src/main/scala/io/udash/rest/RestServlet.scala @@ -123,16 +123,17 @@ class RestServlet( } private def setResponseHeaders(response: HttpServletResponse, code: Int, headers: IMapping[PlainValue]): Unit = { - response.setStatus(code) + ignoreJettyFieldsNpe(response.setStatus(code)) headers.entries.foreach { - case (name, PlainValue(value)) => response.addHeader(name, value) + case (name, PlainValue(value)) => + ignoreJettyFieldsNpe(response.addHeader(name, value)) } } private def writeNonEmptyBody(response: HttpServletResponse, body: HttpBody.NonEmpty): Unit = { val bytes = body.bytes - response.setContentType(body.contentType) - response.setContentLength(bytes.length) + ignoreJettyFieldsNpe(response.setContentType(body.contentType)) + ignoreJettyFieldsNpe(response.setContentLength(bytes.length)) response.getOutputStream.write(bytes) } @@ -148,14 +149,14 @@ class RestServlet( case single: StreamedBody.Single => Task.eval(writeNonEmptyBody(response, single.body)) case binary: StreamedBody.RawBinary => - response.setContentType(binary.contentType) + ignoreJettyFieldsNpe(response.setContentType(binary.contentType)) binary.content .foreachL { chunk => response.getOutputStream.write(chunk) response.getOutputStream.flush() } case jsonList: StreamedBody.JsonList => - response.setContentType(jsonList.contentType) + ignoreJettyFieldsNpe(response.setContentType(jsonList.contentType)) jsonList.elements .bufferTumbling(jsonList.customBatchSize.getOrElse(defaultStreamingBatchSize)) .switchIfEmpty(Observable(Seq.empty)) @@ -217,13 +218,21 @@ class RestServlet( } private def writeFailure(response: HttpServletResponse, message: Opt[String]): Unit = { - response.setStatus(500) + ignoreJettyFieldsNpe(response.setStatus(500)) message.foreach { msg => - response.setContentType(s"text/plain;charset=utf-8") + ignoreJettyFieldsNpe(response.setContentType(s"text/plain;charset=utf-8")) response.getWriter.write(msg) } } + private def ignoreJettyFieldsNpe(op: => Unit): Unit = + try { + op + } catch { + case e: NullPointerException if Option(e.getMessage).exists(_.contains("_fields")) => + () + } + private def readParameters(request: HttpServletRequest): RestParameters = { // can't use request.getPathInfo because it decodes the URL before we can split it val pathPrefix = request.getContextPath.orEmpty + request.getServletPath.orEmpty From 17172476a3ff6225d5cd211aae12ac7878df3d60 Mon Sep 17 00:00:00 2001 From: Rico Date: Tue, 3 Feb 2026 13:39:16 +0100 Subject: [PATCH 2/3] Guard response writes against Jetty _fields NPE Wrap HttpServletResponse/streams to ignore Jetty _fields NPE and keep response handling intact. --- .../scala/io/udash/rest/RestServlet.scala | 74 +++++++++++++++---- 1 file changed, 58 insertions(+), 16 deletions(-) diff --git a/rest/.jvm/src/main/scala/io/udash/rest/RestServlet.scala b/rest/.jvm/src/main/scala/io/udash/rest/RestServlet.scala index 2399da199..a5e4314cd 100644 --- a/rest/.jvm/src/main/scala/io/udash/rest/RestServlet.scala +++ b/rest/.jvm/src/main/scala/io/udash/rest/RestServlet.scala @@ -14,8 +14,8 @@ import org.slf4j.{Logger, LoggerFactory} import java.io.{ByteArrayOutputStream, EOFException} import java.util.concurrent.atomic.AtomicBoolean -import javax.servlet.http.{HttpServlet, HttpServletRequest, HttpServletResponse} -import javax.servlet.{AsyncEvent, AsyncListener} +import javax.servlet.http.{HttpServlet, HttpServletRequest, HttpServletResponse, HttpServletResponseWrapper} +import javax.servlet.{AsyncEvent, AsyncListener, ServletOutputStream, WriteListener} import scala.annotation.tailrec import scala.concurrent.duration.* @@ -94,20 +94,21 @@ class RestServlet( // readRequest must execute in Jetty thread but we want exceptions to be handled uniformly, hence the Try val udashRequest = Try(readRequest(request)) + val safeResponse = new SafeHttpServletResponse(response) val cancelable = (for { restRequest <- Task.fromTry(udashRequest) restResponse <- handleRequest(restRequest) - _ <- Task(setResponseHeaders(response, restResponse.code, restResponse.headers)) - _ <- writeResponseBody(response, restResponse) + _ <- Task(setResponseHeaders(safeResponse, restResponse.code, restResponse.headers)) + _ <- writeResponseBody(safeResponse, restResponse) } yield ()).executeAsync.runAsync { case Right(_) => asyncContext.complete() case Left(e: HttpErrorException) => - completeWith(writeResponse(response, e.toResponse)) + completeWith(writeResponse(safeResponse, e.toResponse)) case Left(e) => logger.error("Failed to handle REST request", e) - completeWith(writeFailure(response, e.getMessage.opt)) + completeWith(writeFailure(safeResponse, e.getMessage.opt)) } asyncContext.setTimeout(handleTimeout.toMillis) @@ -115,25 +116,29 @@ class RestServlet( def onComplete(event: AsyncEvent): Unit = () def onTimeout(event: AsyncEvent): Unit = { cancelable.cancel() - completeWith(writeFailure(response, s"server operation timed out after $handleTimeout".opt)) + completeWith(writeFailure(safeResponse, s"server operation timed out after $handleTimeout".opt)) } def onError(event: AsyncEvent): Unit = () def onStartAsync(event: AsyncEvent): Unit = () }) } - private def setResponseHeaders(response: HttpServletResponse, code: Int, headers: IMapping[PlainValue]): Unit = { - ignoreJettyFieldsNpe(response.setStatus(code)) + private def setResponseHeaders( + response: HttpServletResponse, + code: Int, + headers: IMapping[PlainValue], + ): Unit = { + response.setStatus(code) headers.entries.foreach { case (name, PlainValue(value)) => - ignoreJettyFieldsNpe(response.addHeader(name, value)) + response.addHeader(name, value) } } private def writeNonEmptyBody(response: HttpServletResponse, body: HttpBody.NonEmpty): Unit = { val bytes = body.bytes - ignoreJettyFieldsNpe(response.setContentType(body.contentType)) - ignoreJettyFieldsNpe(response.setContentLength(bytes.length)) + response.setContentType(body.contentType) + response.setContentLength(bytes.length) response.getOutputStream.write(bytes) } @@ -149,14 +154,14 @@ class RestServlet( case single: StreamedBody.Single => Task.eval(writeNonEmptyBody(response, single.body)) case binary: StreamedBody.RawBinary => - ignoreJettyFieldsNpe(response.setContentType(binary.contentType)) + response.setContentType(binary.contentType) binary.content .foreachL { chunk => response.getOutputStream.write(chunk) response.getOutputStream.flush() } case jsonList: StreamedBody.JsonList => - ignoreJettyFieldsNpe(response.setContentType(jsonList.contentType)) + response.setContentType(jsonList.contentType) jsonList.elements .bufferTumbling(jsonList.customBatchSize.getOrElse(defaultStreamingBatchSize)) .switchIfEmpty(Observable(Seq.empty)) @@ -218,9 +223,9 @@ class RestServlet( } private def writeFailure(response: HttpServletResponse, message: Opt[String]): Unit = { - ignoreJettyFieldsNpe(response.setStatus(500)) + response.setStatus(500) message.foreach { msg => - ignoreJettyFieldsNpe(response.setContentType(s"text/plain;charset=utf-8")) + response.setContentType(s"text/plain;charset=utf-8") response.getWriter.write(msg) } } @@ -233,6 +238,43 @@ class RestServlet( () } + private final class SafeServletOutputStream(delegate: ServletOutputStream) extends ServletOutputStream { + override def isReady: Boolean = delegate.isReady + override def setWriteListener(listener: WriteListener): Unit = delegate.setWriteListener(listener) + override def write(b: Int): Unit = ignoreJettyFieldsNpe(delegate.write(b)) + override def write(b: Array[Byte]): Unit = ignoreJettyFieldsNpe(delegate.write(b)) + override def write(b: Array[Byte], off: Int, len: Int): Unit = + ignoreJettyFieldsNpe(delegate.write(b, off, len)) + override def flush(): Unit = ignoreJettyFieldsNpe(delegate.flush()) + override def close(): Unit = ignoreJettyFieldsNpe(delegate.close()) + } + + private final class SafePrintWriter(delegate: java.io.PrintWriter) extends java.io.PrintWriter(delegate) { + override def write(buf: Array[Char], off: Int, len: Int): Unit = ignoreJettyFieldsNpe(super.write(buf, off, len)) + override def write(buf: Array[Char]): Unit = ignoreJettyFieldsNpe(super.write(buf)) + override def write(s: String, off: Int, len: Int): Unit = ignoreJettyFieldsNpe(super.write(s, off, len)) + override def write(s: String): Unit = ignoreJettyFieldsNpe(super.write(s)) + override def write(c: Int): Unit = ignoreJettyFieldsNpe(super.write(c)) + override def flush(): Unit = ignoreJettyFieldsNpe(super.flush()) + override def close(): Unit = ignoreJettyFieldsNpe(super.close()) + } + + private final class SafeHttpServletResponse(response: HttpServletResponse) + extends HttpServletResponseWrapper(response) { + override def setStatus(sc: Int): Unit = ignoreJettyFieldsNpe(super.setStatus(sc)) + override def setHeader(name: String, value: String): Unit = ignoreJettyFieldsNpe(super.setHeader(name, value)) + override def addHeader(name: String, value: String): Unit = ignoreJettyFieldsNpe(super.addHeader(name, value)) + override def setDateHeader(name: String, date: Long): Unit = ignoreJettyFieldsNpe(super.setDateHeader(name, date)) + override def addDateHeader(name: String, date: Long): Unit = ignoreJettyFieldsNpe(super.addDateHeader(name, date)) + override def setIntHeader(name: String, value: Int): Unit = ignoreJettyFieldsNpe(super.setIntHeader(name, value)) + override def addIntHeader(name: String, value: Int): Unit = ignoreJettyFieldsNpe(super.addIntHeader(name, value)) + override def setContentType(`type`: String): Unit = ignoreJettyFieldsNpe(super.setContentType(`type`)) + override def setContentLength(len: Int): Unit = ignoreJettyFieldsNpe(super.setContentLength(len)) + override def setContentLengthLong(len: Long): Unit = ignoreJettyFieldsNpe(super.setContentLengthLong(len)) + override def getOutputStream: ServletOutputStream = new SafeServletOutputStream(super.getOutputStream) + override def getWriter: java.io.PrintWriter = new SafePrintWriter(super.getWriter) + } + private def readParameters(request: HttpServletRequest): RestParameters = { // can't use request.getPathInfo because it decodes the URL before we can split it val pathPrefix = request.getContextPath.orEmpty + request.getServletPath.orEmpty From 268c2ec3ce294a1d312d9de0e27b4b32de118a64 Mon Sep 17 00:00:00 2001 From: Rico Date: Wed, 4 Feb 2026 10:10:23 +0100 Subject: [PATCH 3/3] Handle response lifecycle without masking NPE Stop writes after completion/commit, map NPE to 500, and log full stacktrace for diagnosis. --- .../scala/io/udash/rest/RestServlet.scala | 110 +++++++++++------- 1 file changed, 69 insertions(+), 41 deletions(-) diff --git a/rest/.jvm/src/main/scala/io/udash/rest/RestServlet.scala b/rest/.jvm/src/main/scala/io/udash/rest/RestServlet.scala index a5e4314cd..f6ef98afe 100644 --- a/rest/.jvm/src/main/scala/io/udash/rest/RestServlet.scala +++ b/rest/.jvm/src/main/scala/io/udash/rest/RestServlet.scala @@ -82,19 +82,28 @@ class RestServlet( override def service(request: HttpServletRequest, response: HttpServletResponse): Unit = { val asyncContext = request.startAsync() val completed = new AtomicBoolean(false) + val responseActive = new AtomicBoolean(true) + val requestInfo = { + val query = Option(request.getQueryString).map(q => s"?$q").getOrElse("") + s"${request.getMethod} ${request.getRequestURI}$query" + } // Need to protect asyncContext from being completed twice because after a timeout the // servlet may recycle the same context instance between subsequent requests (not cool) // https://stackoverflow.com/a/27744537 def completeWith(code: => Unit): Unit = if (!completed.getAndSet(true)) { - code - asyncContext.complete() + try { + code + } finally { + responseActive.set(false) + asyncContext.complete() + } } // readRequest must execute in Jetty thread but we want exceptions to be handled uniformly, hence the Try val udashRequest = Try(readRequest(request)) - val safeResponse = new SafeHttpServletResponse(response) + val safeResponse = new SafeHttpServletResponse(response, responseActive, requestInfo) val cancelable = (for { restRequest <- Task.fromTry(udashRequest) @@ -103,9 +112,12 @@ class RestServlet( _ <- writeResponseBody(safeResponse, restResponse) } yield ()).executeAsync.runAsync { case Right(_) => - asyncContext.complete() + completeWith(()) case Left(e: HttpErrorException) => completeWith(writeResponse(safeResponse, e.toResponse)) + case Left(e: NullPointerException) => + logger.error("Failed to handle REST request (NPE)", e) + completeWith(writeFailure(safeResponse, "Internal server error".opt)) case Left(e) => logger.error("Failed to handle REST request", e) completeWith(writeFailure(safeResponse, e.getMessage.opt)) @@ -230,47 +242,63 @@ class RestServlet( } } - private def ignoreJettyFieldsNpe(op: => Unit): Unit = - try { - op - } catch { - case e: NullPointerException if Option(e.getMessage).exists(_.contains("_fields")) => - () + private final class SafeHttpServletResponse( + response: HttpServletResponse, + responseActive: AtomicBoolean, + requestInfo: String, + ) + extends HttpServletResponseWrapper(response) { + private val suppressionLogged = new AtomicBoolean(false) + + private def logSuppressed(opName: String, reason: String, exception: Throwable = null): Unit = + if (suppressionLogged.compareAndSet(false, true)) { + val message = s"Suppressed response write ($opName): $reason for $requestInfo" + if (exception == null) { + logger.warn(message) + } else { + logger.warn(message, exception) + } + } + + private def guard(opName: String)(op: => Unit): Unit = { + if (!responseActive.get() || super.isCommitted) { + logSuppressed(opName, "response already completed or committed") + } else { + op + } } - private final class SafeServletOutputStream(delegate: ServletOutputStream) extends ServletOutputStream { - override def isReady: Boolean = delegate.isReady - override def setWriteListener(listener: WriteListener): Unit = delegate.setWriteListener(listener) - override def write(b: Int): Unit = ignoreJettyFieldsNpe(delegate.write(b)) - override def write(b: Array[Byte]): Unit = ignoreJettyFieldsNpe(delegate.write(b)) - override def write(b: Array[Byte], off: Int, len: Int): Unit = - ignoreJettyFieldsNpe(delegate.write(b, off, len)) - override def flush(): Unit = ignoreJettyFieldsNpe(delegate.flush()) - override def close(): Unit = ignoreJettyFieldsNpe(delegate.close()) - } + private final class SafeServletOutputStream(delegate: ServletOutputStream) extends ServletOutputStream { + override def isReady: Boolean = delegate.isReady + override def setWriteListener(listener: WriteListener): Unit = delegate.setWriteListener(listener) + override def write(b: Int): Unit = guard("write")(delegate.write(b)) + override def write(b: Array[Byte]): Unit = guard("write")(delegate.write(b)) + override def write(b: Array[Byte], off: Int, len: Int): Unit = + guard("write")(delegate.write(b, off, len)) + override def flush(): Unit = guard("flush")(delegate.flush()) + override def close(): Unit = guard("close")(delegate.close()) + } - private final class SafePrintWriter(delegate: java.io.PrintWriter) extends java.io.PrintWriter(delegate) { - override def write(buf: Array[Char], off: Int, len: Int): Unit = ignoreJettyFieldsNpe(super.write(buf, off, len)) - override def write(buf: Array[Char]): Unit = ignoreJettyFieldsNpe(super.write(buf)) - override def write(s: String, off: Int, len: Int): Unit = ignoreJettyFieldsNpe(super.write(s, off, len)) - override def write(s: String): Unit = ignoreJettyFieldsNpe(super.write(s)) - override def write(c: Int): Unit = ignoreJettyFieldsNpe(super.write(c)) - override def flush(): Unit = ignoreJettyFieldsNpe(super.flush()) - override def close(): Unit = ignoreJettyFieldsNpe(super.close()) - } + private final class SafePrintWriter(delegate: java.io.PrintWriter) extends java.io.PrintWriter(delegate) { + override def write(buf: Array[Char], off: Int, len: Int): Unit = guard("write")(super.write(buf, off, len)) + override def write(buf: Array[Char]): Unit = guard("write")(super.write(buf)) + override def write(s: String, off: Int, len: Int): Unit = guard("write")(super.write(s, off, len)) + override def write(s: String): Unit = guard("write")(super.write(s)) + override def write(c: Int): Unit = guard("write")(super.write(c)) + override def flush(): Unit = guard("flush")(super.flush()) + override def close(): Unit = guard("close")(super.close()) + } - private final class SafeHttpServletResponse(response: HttpServletResponse) - extends HttpServletResponseWrapper(response) { - override def setStatus(sc: Int): Unit = ignoreJettyFieldsNpe(super.setStatus(sc)) - override def setHeader(name: String, value: String): Unit = ignoreJettyFieldsNpe(super.setHeader(name, value)) - override def addHeader(name: String, value: String): Unit = ignoreJettyFieldsNpe(super.addHeader(name, value)) - override def setDateHeader(name: String, date: Long): Unit = ignoreJettyFieldsNpe(super.setDateHeader(name, date)) - override def addDateHeader(name: String, date: Long): Unit = ignoreJettyFieldsNpe(super.addDateHeader(name, date)) - override def setIntHeader(name: String, value: Int): Unit = ignoreJettyFieldsNpe(super.setIntHeader(name, value)) - override def addIntHeader(name: String, value: Int): Unit = ignoreJettyFieldsNpe(super.addIntHeader(name, value)) - override def setContentType(`type`: String): Unit = ignoreJettyFieldsNpe(super.setContentType(`type`)) - override def setContentLength(len: Int): Unit = ignoreJettyFieldsNpe(super.setContentLength(len)) - override def setContentLengthLong(len: Long): Unit = ignoreJettyFieldsNpe(super.setContentLengthLong(len)) + override def setStatus(sc: Int): Unit = guard("setStatus")(super.setStatus(sc)) + override def setHeader(name: String, value: String): Unit = guard("setHeader")(super.setHeader(name, value)) + override def addHeader(name: String, value: String): Unit = guard("addHeader")(super.addHeader(name, value)) + override def setDateHeader(name: String, date: Long): Unit = guard("setDateHeader")(super.setDateHeader(name, date)) + override def addDateHeader(name: String, date: Long): Unit = guard("addDateHeader")(super.addDateHeader(name, date)) + override def setIntHeader(name: String, value: Int): Unit = guard("setIntHeader")(super.setIntHeader(name, value)) + override def addIntHeader(name: String, value: Int): Unit = guard("addIntHeader")(super.addIntHeader(name, value)) + override def setContentType(`type`: String): Unit = guard("setContentType")(super.setContentType(`type`)) + override def setContentLength(len: Int): Unit = guard("setContentLength")(super.setContentLength(len)) + override def setContentLengthLong(len: Long): Unit = guard("setContentLengthLong")(super.setContentLengthLong(len)) override def getOutputStream: ServletOutputStream = new SafeServletOutputStream(super.getOutputStream) override def getWriter: java.io.PrintWriter = new SafePrintWriter(super.getWriter) }