From 67924298f29ddecb1d4cb126dd486a43cd4f3aff Mon Sep 17 00:00:00 2001 From: eugene yokota Date: Mon, 27 Jul 2026 19:52:30 -0400 Subject: [PATCH] [2.x] fix: Fixes sbtn stdin race condition (#9521) **Problem** There's a race condition between per-byte readSystemIn notification and one-byte-read thread lifecycle. Note: One-byte-read thread was introduced as a solution to the problem that switching the terminal between raw and canonical mode cannot happen if it's blocked by read. **Solution** This eliminates the thread lifecycle issue by keeping the thread alive throughout the lifecycle of sbtn itself. read is still called on demand by the server readSystemIn notification. --- build.sbt | 3 + .../sbt/internal/client/NetworkClient.scala | 75 ++++++++------- .../NetworkClientInputThreadRaceTest.scala | 94 +++++++++++++++++++ 3 files changed, 139 insertions(+), 33 deletions(-) create mode 100644 main-command/src/test/scala/sbt/internal/client/NetworkClientInputThreadRaceTest.scala diff --git a/build.sbt b/build.sbt index b8a99398e..77fa24bfb 100644 --- a/build.sbt +++ b/build.sbt @@ -638,6 +638,9 @@ lazy val commandProj = (project in file("main-command")) exclude[IncompatibleResultTypeProblem]("sbt.internal.client.NetworkClient.connection"), exclude[IncompatibleResultTypeProblem]("sbt.internal.client.NetworkClient.init"), exclude[DirectMissingMethodProblem]("sbt.internal.BootServerSocket.*"), + exclude[DirectMissingMethodProblem]( + "sbt.internal.client.NetworkClient#RawInputThread.stopped" + ), ), Compile / headerCreate / unmanagedSources := { val old = (Compile / headerCreate / unmanagedSources).value 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 37c308af9..37da42b03 100644 --- a/main-command/src/main/scala/sbt/internal/client/NetworkClient.scala +++ b/main-command/src/main/scala/sbt/internal/client/NetworkClient.scala @@ -16,7 +16,7 @@ import java.net.{ Socket, SocketException } import java.nio.file.Files import java.util.UUID import java.util.concurrent.atomic.{ AtomicBoolean, AtomicReference } -import java.util.concurrent.{ ConcurrentHashMap, LinkedBlockingQueue, TimeUnit } +import java.util.concurrent.{ ConcurrentHashMap, LinkedBlockingQueue, Semaphore, TimeUnit } import sbt.BasicCommandStrings.{ DashDashDetachStdio, DashDashServer, Shutdown, TerminateAction } import sbt.internal.langserver.{ LogMessageParams, MessageType, PublishDiagnosticsParams } @@ -167,16 +167,14 @@ class NetworkClient( private val stdinBytes = new LinkedBlockingQueue[Integer] private val inLock = new Object - private val inputThread = new AtomicReference[RawInputThread] + // A single persistent reader for the life of the client. + private val inputThread = new RawInputThread private val exitClean = new AtomicBoolean(true) private val inClientSideRun = new AtomicBoolean(false) private val sbtProcess = new AtomicReference[Process](null) private class ConnectionRefusedException(t: Throwable) extends Throwable(t) private class ServerFailedException extends Exception - private def startInputThread(): Unit = inputThread.get match { - case null => inputThread.set(new RawInputThread) - case _ => - } + private[client] def startInputThread(): Unit = inputThread.request() private lazy val log: Logger = new Logger { def trace(t: => Throwable): Unit = () def success(message: => String): Unit = () @@ -302,16 +300,12 @@ class NetworkClient( console.appendLog(Level.Info, s"${if (log) "sbt server " else ""}disconnected") } stdinBytes.offer(-1) - Option(inputThread.get).foreach(_.close()) + inputThread.close() Option(interactiveThread.get).foreach(_.interrupt) } case `readSystemIn` => startInputThread() - case `cancelReadSystemIn` => - inputThread.get match { - case null => - case t => t.close() - } - case _ => self.onNotification(msg) + case `cancelReadSystemIn` => inputThread.cancel() + case _ => self.onNotification(msg) } } override protected def onRequest(msg: JsonRpcRequestMessage): Unit = self.onRequest(msg) @@ -570,7 +564,7 @@ class NetworkClient( } // Clean up stderr temp file on successful startup serverStderrFile.foreach(_.delete()) - if (attached.get && !stdinBytes.isEmpty) Option(inputThread.get).foreach(_.drain()) + if (attached.get && !stdinBytes.isEmpty) inputThread.drain() } /** Called on the response for a returning message. */ @@ -607,7 +601,7 @@ class NetworkClient( case msg if attachUUID.get == msg.id => attachUUID.set(null) attached.set(true) - Option(inputThread.get).foreach(_.drain()) + inputThread.drain() () } def completeExec(execId: String, exitCode: Int) = { @@ -1123,28 +1117,42 @@ class NetworkClient( try sendExecCommand("exit") finally c.close() } - Option(inputThread.get).foreach(_.interrupt()) + inputThread.close() } catch { case t: Throwable => t.printStackTrace(); throw t } - private class RawInputThread extends Thread("sbt-read-input-thread") with AutoCloseable { + /** + * Reads stdin on behalf of the server, which asks for it one byte at a time via + * `readSystemIn`/`cancelReadSystemIn` notifications. The design here answers two problems: + * + * - (2020, #5828/#5863/#5856) Switching the terminal between raw and canonical mode can't + * happen while a read is blocked on it. So a read must exist only for as long as the + * server has actually asked for a byte, never sitting on the terminal unrequested. + * - (2026, #9507) Satisfying that by spawning a thread per byte that exits once forwarded + * races the next request against that exit: a `readSystemIn` arriving mid-exit is silently + * dropped, and since nothing else will ever ask for that byte again, the session stops + * accepting input. + */ + private class RawInputThread extends Thread("sbt-read-input-thread") with AutoCloseable: setDaemon(true) + private val stopped = AtomicBoolean(false) + private val readGate = Semaphore(0) start() - val stopped = new AtomicBoolean(false) - override final def run(): Unit = { - def read(): Unit = { - val b = inputStream.read - inLock.synchronized(stdinBytes.offer(b)) - if (attached.get()) drain() - } - try read() - catch { case _: InterruptedException | NonFatal(_) => stopped.set(true) } - finally { - inputThread.set(null) - } - } + override final def run(): Unit = + while !stopped.get do + try + readGate.acquire() + if !stopped.get then + val b = inputStream.read + inLock.synchronized(stdinBytes.offer(b)) + if attached.get() then drain() + if b == -1 then stopped.set(true) + catch case _: InterruptedException | NonFatal(_) => () + + def request(): Unit = readGate.release() + def cancel(): Unit = interrupt() def drain(): Unit = inLock.synchronized { while (!stdinBytes.isEmpty) { val byte = stdinBytes.poll() @@ -1152,10 +1160,11 @@ class NetworkClient( } } - override def close(): Unit = { + override def close(): Unit = + stopped.set(true) + readGate.release() RawInputThread.this.interrupt() - } - } + end RawInputThread } object NetworkClient { diff --git a/main-command/src/test/scala/sbt/internal/client/NetworkClientInputThreadRaceTest.scala b/main-command/src/test/scala/sbt/internal/client/NetworkClientInputThreadRaceTest.scala new file mode 100644 index 000000000..6f7a18e14 --- /dev/null +++ b/main-command/src/test/scala/sbt/internal/client/NetworkClientInputThreadRaceTest.scala @@ -0,0 +1,94 @@ +/* + * sbt + * Copyright 2023, Scala center + * Copyright 2011 - 2022, Lightbend, Inc. + * Copyright 2008 - 2010, Mark Harrah + * Licensed under Apache License 2.0 (see LICENSE) + */ + +package sbt.internal.client + +import java.io.{ ByteArrayOutputStream, InputStream, PrintStream } +import java.util.concurrent.atomic.AtomicInteger +import java.util.concurrent.{ ExecutorService, Executors, TimeUnit } +import sbt.util.Level +import scala.util.Using +import verify.BasicTestSuite + +/** + * Regression test for #9507: a `readSystemIn` request arriving while the previous reader thread + * was silently dropped mid-teardown, and nothing else would ever ask for + * that byte again. This resulted in the session stopping to accepting input. + */ +object NetworkClientInputThreadRaceTest extends BasicTestSuite: + + final val chainLength = 40 + final val perSessionTimeoutMillis = 2000L + + val dummyConsole = new ConsoleInterface: + def appendLog(level: Level.Value, message: => String): Unit = () + def success(msg: String): Unit = () + + val nullPrintStream = new PrintStream(new ByteArrayOutputStream()) + + def newClient(in: InputStream): NetworkClient = + new NetworkClient( + NetworkClient.parseArgs(Array("compile")), + dummyConsole, + in, + nullPrintStream, + nullPrintStream, + false, + ) + + test("startInputThread should not drop a readSystemIn request under contention"): + val numSessions = 50 + withFakeLoad: + val wedged = (1 to numSessions).map(_ => session).sum + assert(wedged == 0, s"$wedged/$numSessions sessions permanently wedged (expected 0)") + + /** + * reads succeed instantly (as if from an already-filled paste buffer) + */ + class ChainedInputStream extends InputStream: + val reads = new AtomicInteger(0) + @volatile var onRead: Int => Unit = _ => () + override def read(): Int = + val n = reads.incrementAndGet() + onRead(n) + 'a'.toInt + + def session: Int = + Using.resource(new ChainedInputStream): in => + Using.resource(newClient(in)): client => + val dispatcher: ExecutorService = Executors.newSingleThreadExecutor() + def requestNext(): Unit = + dispatcher.submit((() => client.startInputThread()): Runnable): Unit + in.onRead = n => if n < chainLength then requestNext() + try + requestNext() + val deadline = System.nanoTime() + TimeUnit.MILLISECONDS.toNanos(perSessionTimeoutMillis) + while in.reads.get() < chainLength && System.nanoTime() < deadline do Thread.sleep(1) + // the session is wedged + if in.reads.get() < chainLength then 1 + else 0 + finally + dispatcher.shutdownNow() + client.close() + + /* + * Tests f under a busy spin to emulate contention/jitter. + */ + def withFakeLoad[A1](f: => A1): A1 = + val fakeLoad = (1 to Runtime.getRuntime.availableProcessors).toList.map: _ => + val busyThread = new Thread(() => + var x = 0L + while !Thread.currentThread.isInterrupted do x += System.nanoTime() + ) + busyThread.setDaemon(true) + busyThread.start() + busyThread + try + f + finally fakeLoad.foreach(_.interrupt()) +end NetworkClientInputThreadRaceTest