mirror of
https://github.com/sbt/sbt.git
synced 2026-10-06 10:03:56 +02:00
[2.x] fix: Keep the interrupt flag when writing to the thin client (#9848)
**Problem** In the default thin-client mode, writing to System.out or System.err from an interrupted thread throws InterruptedException and clears the interrupt flag. NetworkChannel enqueues output with `put` on unbounded LinkedBlockingQueues, and `put` acquires its lock with lockInterruptibly(). Logging libraries swallow the exception, so the interrupt is silently lost and code relying on it can hang. **Solution** Use `add` instead of `put` on the unbounded queues. They never block, so no other behavior changes. Add regression tests to NetworkChannelSpec. Fixes issue #9845 Generated-by: Claude Opus 5.5
This commit is contained in:
@@ -395,9 +395,7 @@ final class NetworkChannel(
|
||||
writeThread.setDaemon(true)
|
||||
|
||||
def publishBytes(event: Array[Byte], delimit: Boolean): Unit =
|
||||
try pendingWrites.put(event -> delimit)
|
||||
catch
|
||||
case _: InterruptedException =>
|
||||
Util.ignoreResult(pendingWrites.add(event -> delimit))
|
||||
|
||||
protected def onSettingQuery(execId: Option[String], req: SettingQuery) =
|
||||
if initialized then
|
||||
@@ -662,25 +660,25 @@ final class NetworkChannel(
|
||||
override def close(): Unit =
|
||||
forceFlush()
|
||||
override def write(b: Int): Unit = outputBuffer.synchronized {
|
||||
outputBuffer.put(b.toByte)
|
||||
Util.ignoreResult(outputBuffer.add(b.toByte))
|
||||
}
|
||||
override def flush(): Unit = flusher.flush()
|
||||
override def write(b: Array[Byte]): Unit = outputBuffer.synchronized {
|
||||
b.foreach(outputBuffer.put)
|
||||
b.foreach(outputBuffer.add)
|
||||
}
|
||||
override def write(b: Array[Byte], off: Int, len: Int): Unit =
|
||||
write(java.util.Arrays.copyOfRange(b, off, off + len))
|
||||
private lazy val errorStream: OutputStream = new OutputStream:
|
||||
private val buffer = new LinkedBlockingQueue[Byte]
|
||||
override def write(b: Int): Unit = buffer.synchronized {
|
||||
buffer.put(b.toByte)
|
||||
Util.ignoreResult(buffer.add(b.toByte))
|
||||
}
|
||||
override def flush(): Unit =
|
||||
val list = new java.util.ArrayList[Byte]
|
||||
buffer.synchronized(buffer.drainTo(list))
|
||||
if !list.isEmpty then jsonRpcNotify(Serialization.systemErr, list.asScala.toSeq)
|
||||
override def write(b: Array[Byte]): Unit = buffer.synchronized {
|
||||
b.foreach(buffer.put)
|
||||
b.foreach(buffer.add)
|
||||
}
|
||||
override def write(b: Array[Byte], off: Int, len: Int): Unit =
|
||||
write(java.util.Arrays.copyOfRange(b, off, off + len))
|
||||
|
||||
@@ -8,7 +8,13 @@
|
||||
|
||||
package sbt.internal.server
|
||||
|
||||
import java.io.{ File, OutputStream }
|
||||
import java.net.{ InetAddress, ServerSocket, Socket }
|
||||
import sbt.{ State, StandardMain }
|
||||
import sbt.internal.util.{ AttributeMap, ConsoleOut, GlobalLogging, MainAppender, Util }
|
||||
import sbt.protocol.Serialization
|
||||
import scala.jdk.CollectionConverters.*
|
||||
import scala.util.Using
|
||||
import verify.BasicTestSuite
|
||||
|
||||
object NetworkChannelSpec extends BasicTestSuite:
|
||||
@@ -32,4 +38,108 @@ object NetworkChannelSpec extends BasicTestSuite:
|
||||
kept.foreach(m => assert(!NetworkChannel.isCanceledOutput(m), s"must not drop: $m"))
|
||||
}
|
||||
|
||||
test("an interrupted thread can print to the client's STDOUT and stays interrupted"):
|
||||
withAttachedClient: channel =>
|
||||
val outcome = whileInterrupted(printLine(channel.terminal.outputStream, "out"))
|
||||
assertSucceededAndStillInterrupted(outcome)
|
||||
|
||||
test("an interrupted thread can print to the client's STDERR and stays interrupted"):
|
||||
withAttachedClient: channel =>
|
||||
val outcome = whileInterrupted(printLine(channel.terminal.errorStream, "err"))
|
||||
assertSucceededAndStillInterrupted(outcome)
|
||||
|
||||
test("an interrupted thread can publish bytes to the client and stays interrupted"):
|
||||
withAttachedClient: channel =>
|
||||
val outcome = whileInterrupted(channel.publishBytes("bytes".getBytes, delimit = true))
|
||||
assertSucceededAndStillInterrupted(outcome)
|
||||
|
||||
private type Outcome = (thrown: Option[Exception], stillInterrupted: Boolean)
|
||||
|
||||
private given Using.Releasable[NetworkChannel] = _.shutdown(false)
|
||||
|
||||
private def whileInterrupted(action: => Unit): Outcome =
|
||||
Thread.currentThread().interrupt()
|
||||
try
|
||||
val thrown =
|
||||
try
|
||||
action
|
||||
None
|
||||
catch case e: Exception => Some(e)
|
||||
(thrown = thrown, stillInterrupted = Thread.interrupted())
|
||||
finally Util.ignoreResult(Thread.interrupted())
|
||||
|
||||
private def assertSucceededAndStillInterrupted(outcome: Outcome): Unit =
|
||||
assert(outcome.thrown.isEmpty, s"expected no exception, but got ${outcome.thrown.orNull}")
|
||||
assert(outcome.stillInterrupted, "expected the interrupt flag to be kept, but it was cleared")
|
||||
|
||||
private def printLine(stream: OutputStream, text: String): Unit =
|
||||
stream.write(s"$text\n".getBytes)
|
||||
stream.flush()
|
||||
|
||||
private def withAttachedClient[A](test: NetworkChannel => A): A =
|
||||
withServerState:
|
||||
withChannelThreadsJoined:
|
||||
withLoopbackConnection: connection =>
|
||||
Using.resource(attachedChannel(connection))(test)
|
||||
|
||||
private def withChannelThreadsJoined[A](f: => A): A =
|
||||
val before = liveThreads
|
||||
try f
|
||||
finally (liveThreads -- before).filter(isChannelThread).foreach(_.join(5000))
|
||||
|
||||
private def liveThreads: Set[Thread] = Thread.getAllStackTraces.keySet.asScala.toSet
|
||||
|
||||
private val channelName = "interrupt-test"
|
||||
|
||||
private def isChannelThread(thread: Thread): Boolean =
|
||||
thread.getName.startsWith("sbt-networkchannel-") ||
|
||||
thread.getName.startsWith(s"sbt-$channelName-")
|
||||
|
||||
private def attachedChannel(connection: Socket): NetworkChannel =
|
||||
val channel = new NetworkChannel(
|
||||
name = channelName,
|
||||
connection = connection,
|
||||
auth = Set.empty,
|
||||
instance = null,
|
||||
handlers = Nil,
|
||||
mkUIThreadImpl = (_, _) => null,
|
||||
)
|
||||
val attachRequestId = "attached-id"
|
||||
channel.setInteractive(attachRequestId, value = false)
|
||||
channel
|
||||
|
||||
private def withLoopbackConnection[A](f: Socket => A): A =
|
||||
val loopback = InetAddress.getLoopbackAddress
|
||||
val anyFreePort = 0
|
||||
val backlog = 1
|
||||
Using.resource(new ServerSocket(anyFreePort, backlog, loopback)): server =>
|
||||
Using.resource(new Socket(loopback, server.getLocalPort)): _ =>
|
||||
Using.resource(server.accept())(f)
|
||||
|
||||
private def withServerState[A](f: => A): A =
|
||||
val previous = StandardMain.exchange.withState(Option(_))
|
||||
if previous.isEmpty then StandardMain.exchange.setState(minimalState)
|
||||
try f
|
||||
finally StandardMain.exchange.setState(previous.orNull)
|
||||
|
||||
private def minimalState: State =
|
||||
val logFile = File.createTempFile("network-channel-spec", ".log")
|
||||
logFile.deleteOnExit()
|
||||
State(
|
||||
configuration = null,
|
||||
definedCommands = Nil,
|
||||
exitHooks = Set.empty,
|
||||
onFailure = None,
|
||||
remainingCommands = Nil,
|
||||
history = State.newHistory,
|
||||
attributes = AttributeMap.empty,
|
||||
globalLogging = GlobalLogging.initial(
|
||||
MainAppender.globalDefault(ConsoleOut.globalProxy),
|
||||
logFile,
|
||||
ConsoleOut.globalProxy
|
||||
),
|
||||
currentCommand = None,
|
||||
next = State.Continue,
|
||||
)
|
||||
|
||||
end NetworkChannelSpec
|
||||
|
||||
Reference in New Issue
Block a user