[2.0.x] fix: Keep the interrupt flag in more thin-client paths (#9855)

**Problem**
Like #9848, some paths still swallow InterruptedException. The
NetworkTerminal request methods block on queue.take, so an
interrupted task thread gets a default value back and loses its
interrupt flag. removeChannel still uses put, and the client's boot
read thread ignores the interrupt meant to stop it.

**Solution**
Restore the flag in NetworkTerminal (in withThread only while the
terminal is open), use add in removeChannel, and stop the read
thread on interrupt. Add tests to NetworkChannelSpec.

Generated-by: Claude Opus 5.5

Co-authored-by: Claude Opus 5.5 <[email protected]>
This commit is contained in:
eugene yokota
2026-10-02 11:11:20 -04:00
committed by Eugene Yokota
co-authored by Claude Opus 5.5
parent 7bf30ff003
commit 7094f75c83
4 changed files with 70 additions and 55 deletions
@@ -475,7 +475,7 @@ class NetworkClient(
}
if (socket.isEmpty && readThreadAlive.get) {
try Thread.sleep(10)
catch { case _: InterruptedException => }
catch { case _: InterruptedException => readThreadAlive.set(false) }
}
}
} catch { case e: IOException => e.printStackTrace(System.err) }
@@ -172,14 +172,10 @@ private[sbt] final class CommandExchange {
channelBufferLock.synchronized {
Util.ignoreResult(channelBuffer -= c)
}
commandQueue.removeIf { e =>
e.source.map(_.channelName) == Some(c.name) && e.commandLine != Shutdown
}
currentExec.withFilter(_.source.map(_.channelName) == Some(c.name)).foreach { e =>
Util.ignoreResult(NetworkChannel.cancel(e.execId, e.execId.getOrElse("0"), force = false))
}
try commandQueue.put(Exec(s"${ContinuousCommands.stopWatch} ${c.name}", None))
catch { case _: InterruptedException => }
def isFromChannel(e: Exec): Boolean = e.source.exists(_.channelName == c.name)
commandQueue.removeIf { e => isFromChannel(e) && e.commandLine != Shutdown }
currentExec.foreach { e => if isFromChannel(e) then doCancel(e, force = false) }
Util.ignoreResult(commandQueue.add(Exec(s"${ContinuousCommands.stopWatch} ${c.name}", None)))
// Notify other servers to drop if idle when a real client disconnects
if (wasInitialized && !shuttingDown.get) notifyOtherServers()
}
@@ -496,15 +492,16 @@ private[sbt] final class CommandExchange {
commandQueue.add(exit)
()
}
private def cancel(e: Exec): Unit = {
if (e.commandLine.startsWith("console")) {
private def cancel(e: Exec): Unit =
if e.commandLine.startsWith("console") then
val terminal = Terminal.get
terminal.write(13, 13, 13, 4)
terminal.printStream.println("\nconsole session killed by remote sbt client")
} else {
Util.ignoreResult(NetworkChannel.cancel(e.execId, e.execId.getOrElse("0"), force = true))
}
}
else doCancel(e, force = true)
private def doCancel(e: Exec, force: Boolean): Unit =
Util.ignoreResult(NetworkChannel.cancel(e.execId, e.execId.getOrElse("0"), force = force))
/** Handle a dropIfIdle notification from another server. */
private[sbt] def handleDropIfIdle(): Unit = {
@@ -12,7 +12,7 @@ package server
import java.io.{ IOException, InputStream, OutputStream }
import java.net.{ Socket, SocketTimeoutException }
import java.util.concurrent.{ ConcurrentHashMap, LinkedBlockingQueue }
import java.util.concurrent.{ BlockingQueue, ConcurrentHashMap, LinkedBlockingQueue }
import java.util.concurrent.atomic.{ AtomicBoolean, AtomicReference }
import sbt.BasicCommandStrings.{ Shutdown, TerminateAction }
@@ -791,18 +791,26 @@ final class NetworkChannel(
()
} else throw new InterruptedException
}
private def withThread[R](f: => R, default: R): R = {
private def withThread[R](f: => R, default: R): R =
val t = Thread.currentThread
try {
try
blockedThreads.synchronized(blockedThreads.add(t))
f
} catch { case _: InterruptedException => default }
finally {
Util.ignoreResult(blockedThreads.synchronized(blockedThreads.remove(t)))
}
}
def getProperty[T](f: TerminalPropertiesResponse => T, default: T): Option[T] = {
if (closed.get || !isAttached) None
catch
case _: InterruptedException =>
if !closed.get then t.interrupt()
default
finally Util.ignoreResult(blockedThreads.synchronized(blockedThreads.remove(t)))
private def awaitResponse[A](queue: BlockingQueue[A]): Option[A] =
try Some(queue.take)
catch
case _: InterruptedException =>
Thread.currentThread.interrupt()
None
def getProperty[T](f: TerminalPropertiesResponse => T, default: T): Option[T] =
if closed.get || !isAttached then None
else
withThread(
{
@@ -811,7 +819,6 @@ final class NetworkChannel(
},
None
)
}
private def waitForPending(f: TerminalPropertiesResponse => Boolean): Boolean = {
if (closed.get || !isAttached) false
else
@@ -884,13 +891,12 @@ final class NetworkChannel(
override private[sbt] def getAttributes: Map[String, String] =
if (closed.get) Map.empty
else {
else
val queue = VirtualTerminal.sendTerminalAttributesQuery(
term.name,
jsonRpcRequest[TerminalAttributesQuery]
)
try {
val a = queue.take
awaitResponse(queue).fold(Map.empty[String, String]): a =>
Map(
"iflag" -> a.iflag,
"oflag" -> a.oflag,
@@ -898,10 +904,8 @@ final class NetworkChannel(
"lflag" -> a.lflag,
"cchars" -> a.cchars
)
} catch { case _: InterruptedException => Map.empty }
}
override private[sbt] def setAttributes(attributes: Map[String, String]): Unit =
if (!closed.get) {
if !closed.get then
val attrs = TerminalSetAttributesCommand(
iflag = attributes.getOrElse("iflag", ""),
oflag = attributes.getOrElse("oflag", ""),
@@ -914,48 +918,37 @@ final class NetworkChannel(
jsonRpcRequest[TerminalSetAttributesCommand],
attrs
)
try queue.take
catch { case _: InterruptedException => }
}
Util.ignoreResult(awaitResponse(queue))
override private[sbt] def getSizeImpl: (Int, Int) =
if (!closed.get) {
if !closed.get then
val queue =
VirtualTerminal.getTerminalSize(term.name, jsonRpcRequest[TerminalGetSizeQuery])
val res =
try queue.take
catch { case _: InterruptedException => TerminalGetSizeResponse(1, 1) }
val res = awaitResponse(queue).getOrElse(TerminalGetSizeResponse(1, 1))
(res.width, res.height)
} else (1, 1)
else (1, 1)
override def setSize(width: Int, height: Int): Unit =
if (!closed.get) {
if !closed.get then
val size = TerminalSetSizeCommand(width, height)
val queue =
VirtualTerminal.setTerminalSize(term.name, jsonRpcRequest[TerminalSetSizeCommand], size)
try queue.take
catch { case _: InterruptedException => }
}
private def setRawMode(toggle: Boolean): Unit = {
if (!closed.get || false) {
Util.ignoreResult(awaitResponse(queue))
private def setRawMode(toggle: Boolean): Unit =
if !closed.get || false then
val raw = TerminalSetRawModeCommand(toggle)
val queue = VirtualTerminal.setTerminalRawMode(
term.name,
jsonRpcRequest[TerminalSetRawModeCommand],
raw
)
try queue.take
catch { case _: InterruptedException => }
}
}
Util.ignoreResult(awaitResponse(queue))
override private[sbt] def enterRawMode(): Unit = setRawMode(true)
override private[sbt] def exitRawMode(): Unit = setRawMode(false)
override def setEchoEnabled(toggle: Boolean): Unit =
if (!closed.get) {
if !closed.get then
val echo = TerminalSetEchoCommand(toggle)
val queue =
VirtualTerminal.setTerminalEcho(term.name, jsonRpcRequest[TerminalSetEchoCommand], echo)
try queue.take
catch { case _: InterruptedException => () }
}
Util.ignoreResult(awaitResponse(queue))
override def flush(): Unit = doFlush()
override def toString: String = s"NetworkTerminal(${term.name})"
@@ -11,7 +11,8 @@ 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.internal.util.{ AttributeMap, ConsoleOut, GlobalLogging, MainAppender, Terminal, Util }
import sbt.internal.util.Terminal.TerminalImpl
import sbt.protocol.Serialization
import scala.jdk.CollectionConverters.*
import scala.util.Using
@@ -53,6 +54,30 @@ object NetworkChannelSpec extends BasicTestSuite:
val outcome = whileInterrupted(channel.publishBytes("bytes".getBytes, delimit = true))
assertSucceededAndStillInterrupted(outcome)
private val terminalRequests: Seq[(String, Terminal => Unit)] = Seq(
"reads the terminal width" -> (t => Util.ignoreResult(t.getWidth)),
"reads the terminal size" -> {
case t: TerminalImpl => Util.ignoreResult(t.getSizeImpl)
case t => sys.error(s"unexpected terminal: $t")
},
"reads the terminal attributes" -> (t => Util.ignoreResult(t.getAttributes)),
"sets the terminal attributes" -> (_.setAttributes(Map.empty)),
"sets the terminal size" -> (_.setSize(80, 24)),
"enters raw mode" -> (_.enterRawMode()),
"sets echo" -> (_.setEchoEnabled(false)),
)
terminalRequests.foreach: (label, request) =>
test(s"an interrupted thread that $label stays interrupted"):
withAttachedClient: channel =>
val outcome = whileInterrupted(request(channel.terminal))
assertSucceededAndStillInterrupted(outcome)
test("removing a channel from an interrupted thread keeps the interrupt flag"):
withAttachedClient: channel =>
val outcome = whileInterrupted(StandardMain.exchange.removeChannel(channel))
assertSucceededAndStillInterrupted(outcome)
private type Outcome = (thrown: Option[Exception], stillInterrupted: Boolean)
private given Using.Releasable[NetworkChannel] = _.shutdown(false)