mirror of
https://github.com/sbt/sbt.git
synced 2026-10-06 10:03:56 +02:00
[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:
committed by
Eugene Yokota
co-authored by
Claude Opus 5.5
parent
7bf30ff003
commit
7094f75c83
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user