diff --git a/main/src/main/scala/sbt/internal/server/NetworkChannel.scala b/main/src/main/scala/sbt/internal/server/NetworkChannel.scala index ff1773f2f..734204258 100644 --- a/main/src/main/scala/sbt/internal/server/NetworkChannel.scala +++ b/main/src/main/scala/sbt/internal/server/NetworkChannel.scala @@ -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)) diff --git a/main/src/test/scala/sbt/internal/server/NetworkChannelSpec.scala b/main/src/test/scala/sbt/internal/server/NetworkChannelSpec.scala index 3a5db2825..36e8675f8 100644 --- a/main/src/test/scala/sbt/internal/server/NetworkChannelSpec.scala +++ b/main/src/test/scala/sbt/internal/server/NetworkChannelSpec.scala @@ -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