diff --git a/main-actions/src/main/scala/sbt/internal/WorkerExchange.scala b/main-actions/src/main/scala/sbt/internal/WorkerExchange.scala index fe22ca67c..eae02d416 100644 --- a/main-actions/src/main/scala/sbt/internal/WorkerExchange.scala +++ b/main-actions/src/main/scala/sbt/internal/WorkerExchange.scala @@ -24,8 +24,9 @@ import scala.sys.process.{ BasicIO, Process, ProcessIO } import scala.collection.mutable import scala.collection.concurrent.TrieMap import scala.collection.mutable.ListBuffer -import scala.concurrent.{ Await, Promise } +import scala.concurrent.{ Await, Future, Promise } import scala.concurrent.duration.* +import scala.util.Try import scala.util.control.NonFatal object WorkerExchange: @@ -83,11 +84,13 @@ object WorkerExchange: IO.classLocationPath(classOf[Gson]).toFile, ) val inputRef = Promise[OutputStream]() + val responsesRead = Promise[Unit]() def runAccepter(out: OutputStream, in: InputStream): Unit = inputRef.success(out) val scanner = Scanner(in, "UTF-8") while scanner.hasNextLine() do notifyListeners(scanner.nextLine()) - val (connArgs, closer): (Seq[String], Option[AutoCloseable]) = connectionType match + responsesRead.trySuccess(()) + val (connArgs, closer) = connectionType match case WorkerConnection.Tcp => val serverSocket = Retry(ServerSocket(0, 1, loopback)) val accepter = Thread(() => { @@ -130,20 +133,24 @@ object WorkerExchange: val onStdoutLine: String => Unit = connectionType match case WorkerConnection.Stdio => notifyListeners case _ => (line) => scala.Console.out.println(line) + def readStdout(stdout: InputStream): Unit = + BasicIO.processFully(onStdoutLine)(stdout) + if connectionType == WorkerConnection.Stdio then responsesRead.trySuccess(()) val processIo = ProcessIO( in = (input) => (connectionType match case WorkerConnection.Stdio => inputRef.success(input) case _ => () ), - out = BasicIO.processFully(onStdoutLine), + out = readStdout, err = BasicIO.processFully((line) => scala.Console.err.println(line)), ) val forkWithIo = fo.withOutputStrategy(OutputStrategy.CustomInputOutput(processIo)) val p = Fork.java.fork(forkWithIo, options) val forkTimeout = fo.connectionTimeout.getOrElse(30.seconds) val input = Await.result(inputRef.future, forkTimeout) - WorkerProxy(input, p, options, closer) + WorkerProxy(input, p, options, closer, responsesRead.future) + end startWorker /** Generates a fresh path suitable for binding a `WorkerConnection.Ipc` socket. */ def newIpcSocketPath(): NioPath = @@ -178,6 +185,7 @@ class WorkerProxy( val process: Process, val options: Seq[String], closer: Option[AutoCloseable], + responsesRead: Future[Unit], ) extends AutoCloseable: lazy val inputStream = PrintStream(input) def close(): Unit = @@ -194,11 +202,21 @@ class WorkerProxy( val watch = Thread(() => { while process.isAlive() do Thread.sleep(100) + Try(Await.ready(responsesRead, WorkerProxy.responseDrainTimeout)) WorkerExchange.listeners.foreach(_.notifyExit(process)) }) watch.start() end WorkerProxy +object WorkerProxy: + /** + * How long to keep reading a worker's responses after it exits, before reporting the exit. + * The reader reaches end of stream as soon as it has consumed everything the worker wrote; + * the bound only matters for a worker that died before connecting. + */ + private val responseDrainTimeout = 10.seconds +end WorkerProxy + abstract class WorkerResponseListener extends Function1[String, Unit]: def notifyExit(p: Process): Unit diff --git a/sbt-app/src/sbt-test/tests/fork-slow-listener/build.sbt b/sbt-app/src/sbt-test/tests/fork-slow-listener/build.sbt new file mode 100644 index 000000000..9f889b515 --- /dev/null +++ b/sbt-app/src/sbt-test/tests/fork-slow-listener/build.sbt @@ -0,0 +1,11 @@ +scalaVersion := "3.9.0" +Test / fork := true +libraryDependencies += "org.scalameta" %% "munit" % "1.0.4" % Test + +Test / testListeners += new TestsListener: + def doInit(): Unit = () + def startGroup(name: String): Unit = () + def testEvent(event: TestEvent): Unit = Thread.sleep(50) + def endGroup(name: String, t: Throwable): Unit = () + def endGroup(name: String, result: TestResult): Unit = () + def doComplete(finalResult: TestResult): Unit = () diff --git a/sbt-app/src/sbt-test/tests/fork-slow-listener/src/test/scala/ManyTests.scala b/sbt-app/src/sbt-test/tests/fork-slow-listener/src/test/scala/ManyTests.scala new file mode 100644 index 000000000..9f1ae0f39 --- /dev/null +++ b/sbt-app/src/sbt-test/tests/fork-slow-listener/src/test/scala/ManyTests.scala @@ -0,0 +1,2 @@ +class ManyTests extends munit.FunSuite: + (1 to 40).foreach(i => test(s"test $i")(assert(i < 40))) diff --git a/sbt-app/src/sbt-test/tests/fork-slow-listener/test b/sbt-app/src/sbt-test/tests/fork-slow-listener/test new file mode 100644 index 000000000..5a9f22365 --- /dev/null +++ b/sbt-app/src/sbt-test/tests/fork-slow-listener/test @@ -0,0 +1 @@ +-> test