mirror of
https://github.com/sbt/sbt.git
synced 2026-10-06 10:03:56 +02:00
@@ -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 }
|
||||
@@ -434,8 +434,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) {
|
||||
@@ -734,11 +733,11 @@ final class NetworkChannel(
|
||||
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))
|
||||
@@ -747,7 +746,7 @@ final class NetworkChannel(
|
||||
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]
|
||||
@@ -755,7 +754,7 @@ final class NetworkChannel(
|
||||
if (!list.isEmpty) 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))
|
||||
@@ -792,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(
|
||||
{
|
||||
@@ -812,7 +819,6 @@ final class NetworkChannel(
|
||||
},
|
||||
None
|
||||
)
|
||||
}
|
||||
private def waitForPending(f: TerminalPropertiesResponse => Boolean): Boolean = {
|
||||
if (closed.get || !isAttached) false
|
||||
else
|
||||
@@ -885,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,
|
||||
@@ -899,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", ""),
|
||||
@@ -915,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})"
|
||||
|
||||
@@ -8,7 +8,14 @@
|
||||
|
||||
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, Terminal, Util }
|
||||
import sbt.internal.util.Terminal.TerminalImpl
|
||||
import sbt.protocol.Serialization
|
||||
import scala.jdk.CollectionConverters.*
|
||||
import scala.util.Using
|
||||
import verify.BasicTestSuite
|
||||
|
||||
object NetworkChannelSpec extends BasicTestSuite:
|
||||
@@ -32,4 +39,132 @@ 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 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)
|
||||
|
||||
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