diff --git a/main-command/src/main/scala/sbt/internal/client/NetworkClient.scala b/main-command/src/main/scala/sbt/internal/client/NetworkClient.scala index 03cf6b165..a98abc5ee 100644 --- a/main-command/src/main/scala/sbt/internal/client/NetworkClient.scala +++ b/main-command/src/main/scala/sbt/internal/client/NetworkClient.scala @@ -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) } diff --git a/main/src/main/scala/sbt/internal/CommandExchange.scala b/main/src/main/scala/sbt/internal/CommandExchange.scala index 75ded6960..889c358ca 100644 --- a/main/src/main/scala/sbt/internal/CommandExchange.scala +++ b/main/src/main/scala/sbt/internal/CommandExchange.scala @@ -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 = { diff --git a/main/src/main/scala/sbt/internal/server/NetworkChannel.scala b/main/src/main/scala/sbt/internal/server/NetworkChannel.scala index 93a3ee651..05809670b 100644 --- a/main/src/main/scala/sbt/internal/server/NetworkChannel.scala +++ b/main/src/main/scala/sbt/internal/server/NetworkChannel.scala @@ -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})" diff --git a/main/src/test/scala/sbt/internal/server/NetworkChannelSpec.scala b/main/src/test/scala/sbt/internal/server/NetworkChannelSpec.scala index 36e8675f8..83a043243 100644 --- a/main/src/test/scala/sbt/internal/server/NetworkChannelSpec.scala +++ b/main/src/test/scala/sbt/internal/server/NetworkChannelSpec.scala @@ -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)