[2.x] feat: Persistent worker for test (#9678)

This commit is contained in:
eugene yokota
2026-09-25 03:30:09 -04:00
committed by GitHub
parent f9ca5f5ee6
commit c8a0007d6e
13 changed files with 315 additions and 24 deletions
+23 -13
View File
@@ -50,6 +50,8 @@ private[sbt] object ForkTests:
log: Logger,
parallelism: Option[Int],
virtualClasspath: Boolean,
persistentWorker: Boolean,
maxPoolSize: Int,
tags: (Tag, Int)*
): Task[TestOutput] =
import std.TaskExtra.*
@@ -75,11 +77,14 @@ private[sbt] object ForkTests:
parallel = config.parallel,
parallelism = parallelism,
virtualClasspath = virtualClasspath,
persistentWorker = persistentWorker && virtualClasspath,
maxPoolSize = maxPoolSize,
).tagw(config.tags*)
.tagw(tags*)
.dependsOn(all(opts.setup)*)
.flatMap: results =>
all(opts.cleanup).join.map(_ => results)
end if
end apply
private def mainTestTask(
@@ -92,6 +97,8 @@ private[sbt] object ForkTests:
parallel: Boolean,
parallelism: Option[Int],
virtualClasspath: Boolean,
persistentWorker: Boolean,
maxPoolSize: Int,
): Task[TestOutput] =
std.TaskExtra.task {
val testListeners = opts.testListeners.flatMap:
@@ -145,19 +152,22 @@ private[sbt] object ForkTests:
)
testListeners.foreach(_.doInit())
val result =
val w = WorkerExchange.startWorker(fork, if virtualClasspath then Nil else cpFiles)
val wl = React(randomId, log, opts.testListeners, resultsAcc, w.process)
try
WorkerExchange.registerListener(wl)
val paramJson = g.toJson(param, param.getClass)
val json = jsonRpcRequest(randomId, "test", paramJson)
w.println(json)
if wl.blockForResponse() != 0 then
throw MessageOnlyException("Forked test harness failed")
testOutputResult
finally
w.close()
WorkerExchange.unregisterListener(wl)
WorkerExchange.withWorker(
fork,
if virtualClasspath then Nil else cpFiles,
persistentWorker,
maxPoolSize,
): w =>
val wl = React(randomId, log, opts.testListeners, resultsAcc, w.process)
try
WorkerExchange.registerListener(wl)
val paramJson = g.toJson(param, param.getClass)
val json = jsonRpcRequest(randomId, "test", paramJson)
w.println(json)
if wl.blockForResponse() != 0 then
throw MessageOnlyException("Forked test harness failed")
testOutputResult
finally WorkerExchange.unregisterListener(wl)
testListeners.foreach(_.doComplete(result.overall))
result
} // end task
@@ -15,20 +15,26 @@ import java.net.{ InetAddress, ServerSocket, StandardProtocolFamily, UnixDomainS
import java.nio.channels.{ ServerSocketChannel, SocketChannel }
import java.nio.file.{ Files, Path as NioPath }
import java.util.Scanner
import java.util.concurrent.ConcurrentLinkedQueue
import sbt.io.IO
import sbt.internal.io.Retry
import sbt.internal.worker1.*
import sbt.protocol.DuplexChannels
import sbt.testing.Framework
import scala.sys.process.{ BasicIO, Process, ProcessIO }
import scala.collection.mutable
import scala.collection.concurrent.TrieMap
import scala.collection.mutable
import scala.collection.mutable.ListBuffer
import scala.concurrent.{ Await, Promise }
import scala.concurrent.duration.*
import scala.util.control.NonFatal
object WorkerExchange:
/**
* What makes two test runs interchangeable enough to reuse the same worker JVM
*/
private case class WorkerKey(fo: ForkOptions)
val listeners: mutable.ListBuffer[WorkerResponseListener] = ListBuffer.empty
private val loopback = InetAddress.getByName(null)
private val jdkIpcSupportCache = TrieMap.empty[Option[File], Boolean]
@@ -69,6 +75,12 @@ object WorkerExchange:
jdkIpcSupportCache.getOrElseUpdate(javaHome, doDetect)
end supportsUnixDomainSockets
// Idle checked-in workers, oldest first; outlives any single test task execution.
private val idleWorkers = new ConcurrentLinkedQueue[(WorkerKey, WorkerProxy)]()
// Best-effort cleanup: forked children outlive this JVM otherwise.
Runtime.getRuntime.addShutdownHook(new Thread(() => closeIdleWorkers()))
/**
* Start a worker process.
*/
@@ -76,6 +88,14 @@ object WorkerExchange:
fo: ForkOptions,
extraCp: Seq[File],
connectionType: WorkerConnection,
): WorkerProxy =
startWorker(fo, extraCp, connectionType, persistent = false)
def startWorker(
fo: ForkOptions,
extraCp: Seq[File],
connectionType: WorkerConnection,
persistent: Boolean,
): WorkerProxy =
// put extraCp first so we can shadow the WorkerMain class
val fullCp = extraCp ++ Seq(
@@ -128,7 +148,8 @@ object WorkerExchange:
"-classpath",
fullCp.mkString(File.pathSeparator),
classOf[WorkerMain].getCanonicalName,
) ++ connArgs
) ++ connArgs ++
(if persistent then Seq("--persistent_worker") else Nil)
val onStdoutLine: String => Unit = connectionType match
case WorkerConnection.Stdio => notifyListeners
case _ => (line) => scala.Console.out.println(line)
@@ -158,6 +179,66 @@ object WorkerExchange:
Files.deleteIfExists(path)
path
/**
* `maxPoolSize` bounds the total number of idle workers kept across every key combined (not
* per key); once a checkin would exceed it, the globally least-recently-used idle worker is
* closed to make room, regardless of which key it belongs to.
*/
def withWorker[A1](fo: ForkOptions, extraCp: Seq[File], persistent: Boolean, maxPoolSize: Int)(
f: WorkerProxy => A1
): A1 =
val ct =
if persistent then WorkerConnection.Tcp
else if supportsUnixDomainSockets(fo.javaHome) then WorkerConnection.Ipc(newIpcSocketPath())
else WorkerConnection.Stdio
val key = WorkerKey(fo)
def checkout: Option[WorkerProxy] =
val it = idleWorkers.iterator()
Iterator
.continually(if it.hasNext() then Some(it.next()) else None)
.takeWhile(_.isDefined)
.flatten
.filter((k, _) => k == key)
.map: (_, w) =>
it.remove(); w
.find: w =>
val alive = w.process.isAlive()
if !alive then
try w.close()
catch case NonFatal(_) => ()
alive
def checkin(worker: WorkerProxy): Unit =
if worker.process.isAlive() then
idleWorkers.add(key -> worker)
evictExcess(maxPoolSize)
val w = (if persistent then checkout else None)
.getOrElse(startWorker(fo, extraCp, ct, persistent))
try f(w)
finally
if persistent && w.process.isAlive() then checkin(w)
else w.close()
end withWorker
/** Asks the worker to shut itself down before closing the connection and reaping the process. */
private def shutdownWorker(w: WorkerProxy): Unit =
try w.println("""{ "jsonrpc": "2.0", "method": "bye", "params": {}, "id": 0 }""")
catch case NonFatal(_) => ()
try w.close()
catch case NonFatal(_) => ()
w.process.destroy()
private def evictExcess(maxPoolSize: Int): Unit =
while idleWorkers.size() > maxPoolSize do
Option(idleWorkers.poll()).foreach((_, w) => shutdownWorker(w))
/** Close every idle worker. Workers currently checked out are unaffected. */
def closeIdleWorkers(): Unit =
Iterator
.continually(Option(idleWorkers.poll()))
.takeWhile(_.isDefined)
.flatten
.foreach((_, w) => shutdownWorker(w))
def registerListener(listener: WorkerResponseListener): Unit =
synchronized:
listeners.append(listener)
+51 -1
View File
@@ -194,6 +194,7 @@ object Defaults extends BuildCommon with DefExtra:
testForkedParallel :== true,
testForkedParallelism :== None,
workerMaxInstances :== SysProp.workerMaxInstances,
testPersistentWorker :== SysProp.testPersistentWorker,
javaOptions :== Nil,
sbtPlugin :== false,
isMetaBuild :== false,
@@ -1244,6 +1245,10 @@ object Defaults extends BuildCommon with DefExtra:
testTaskOptions(testSelected),
testTaskOptions(testQuick),
testDefaults,
baseDirectory := {
if testPersistentWorker.value then (ThisBuild / baseDirectory).value
else baseDirectory.value
},
testLoader := Def.uncached(ClassLoaders.testTask.value),
loadedTestFrameworks := Def.uncached {
val loader = testLoader.value
@@ -1264,6 +1269,11 @@ object Defaults extends BuildCommon with DefExtra:
executeTests := Def.uncached(Def.taskDyn {
import sbt.TupleSyntax.*
val fpm = testForkedParallelism.value
val pw = testPersistentWorker.value
// Reuse the ForkedTestGroup concurrency limit as the persistent worker pool cap: never keep
// more idle worker JVMs around than the number of forked test groups allowed to run at once.
val pwMax =
Tags.effectiveLimit(concurrentRestrictions.value, Tags.ForkedTestGroup, 12)
(
test / streams,
loadedTestFrameworks,
@@ -1289,7 +1299,9 @@ object Defaults extends BuildCommon with DefExtra:
jo,
clls,
s"${Util.quoteIfNotScalaId(thisProj.id)} / ",
c
c,
pw,
pwMax,
)
}
}.value),
@@ -1558,6 +1570,9 @@ object Defaults extends BuildCommon with DefExtra:
classLoaderLayeringStrategy.value,
projectId = s"${Util.quoteIfNotScalaId(thisProject.value.id)} / ",
converter = fileConverter.value,
persistentWorker = testPersistentWorker.value,
persistentWorkerPoolMax =
Tags.effectiveLimit(concurrentRestrictions.value, Tags.ForkedTestGroup, 12),
)
val taskName = display.show(resolvedScoped.value)
val trl = testResultLogger.value
@@ -1697,6 +1712,39 @@ object Defaults extends BuildCommon with DefExtra:
strategy: ClassLoaderLayeringStrategy,
projectId: String,
converter: FileConverter,
): Task[Tests.Output] =
allTestGroupsTask(
s,
frameworks,
loader,
groups,
config,
cp,
forkedParallelExecution,
forkedParallelism,
javaOptions,
strategy,
projectId,
converter,
persistentWorker = false,
persistentWorkerPoolMax = 1,
)
private[sbt] def allTestGroupsTask(
s: TaskStreams,
frameworks: Map[TestFramework, Framework],
loader: ClassLoader,
groups: Seq[Tests.Group],
config: Tests.Execution,
cp: Classpath,
forkedParallelExecution: Boolean,
forkedParallelism: Option[Int],
javaOptions: Seq[String],
strategy: ClassLoaderLayeringStrategy,
projectId: String,
converter: FileConverter,
persistentWorker: Boolean,
persistentWorkerPoolMax: Int,
): Task[Tests.Output] =
val processedOptions: Map[Tests.Group, Tests.ProcessedOptions] =
groups
@@ -1741,6 +1789,8 @@ object Defaults extends BuildCommon with DefExtra:
s.log,
forkedParallelism,
strategy != ClassLoaderLayeringStrategy.Raw,
persistentWorker,
persistentWorkerPoolMax,
Vector((Tags.ForkedTestGroup, 1)) ++ innerTags ++ group.tags*
)
case Tests.InProcess =>
+1
View File
@@ -403,6 +403,7 @@ object Keys {
@transient
val testListeners = taskKey[Seq[TestReportListener]]("Defines test listeners.").withRank(DTask)
val testPersistentWorker = settingKey[Boolean]("Whether forked test worker JVMs are pooled and reused across separate test task executions instead of forked fresh every time. Default false.")
val testForkedParallel = settingKey[Boolean]("Whether forked tests should be executed in parallel").withRank(CTask)
val testForkedParallelism = settingKey[Option[Int]]("Maximum number of parallel test threads when using testForkedParallel. Default: 2.").withRank(CTask)
val workerMaxInstances = settingKey[Int]("Maximum number of test workers. Default: 2")
+8
View File
@@ -76,6 +76,14 @@ object Tags:
def getInt(m: TagMap, tag: Tag): Int = m.getOrElse(tag, 0)
def effectiveLimit(rules: Seq[Rule], tag: Tag, bound: Int): Int =
val pred = predicate(rules)
@tailrec def loop(n: Int): Int =
if n > bound then bound
else if pred(Map(tag -> n, All -> n)) then loop(n + 1)
else n - 1
loop(1) max 1
/**
* Constructs a custom Rule from the predicate `f`.
* The input represents the weighted tags of a set of tasks.
@@ -116,6 +116,8 @@ object SysProp:
def workerMaxInstances: Int = int("sbt.worker_max_instances", 2)
def testPersistentWorker: Boolean = getOrFalse("sbt.test_persistent_worker")
def analysisCacheMaxCount: Int = int("sbt.local_cache.analysis_count", 20)
/**
+9
View File
@@ -47,5 +47,14 @@ object TagsTest extends Properties("Tags"):
excl(etag)(tm)
}
property("effectiveLimit recovers the max a single limit() rule was constructed with") =
forAll(tag, Gen.choose(1, 500)) { (t: Tag, max: Int) =>
effectiveLimit(limit(t, max) :: Nil, t, bound = 512) == max
}
property("effectiveLimit is capped at bound when the tag is unrestricted") = forAll { (t: Tag) =>
effectiveLimit(Nil, t, bound = 7) == 7
}
private def excl(tag: Tag): TagMap => Boolean = predicate(exclusive(tag) :: Nil)
end TagsTest
@@ -0,0 +1,34 @@
import Tests._
import Defaults._
scalaVersion := "3.8.4"
organization := "com.example"
val check = TaskKey[Unit]("check", "Check that the two runs shared the same worker JVM")
val checkDistinct = TaskKey[Unit]("checkDistinct", "Check that the two runs used different worker JVMs")
val clearPids = TaskKey[Unit]("clearPids", "Delete the pids marker file")
Test / fork := true
Global / concurrentRestrictions += Tags.limit(Tags.ForkedTestGroup, 4)
libraryDependencies += "org.scalameta" %% "munit" % "1.0.4" % Test
check := Def.uncached {
val lines = IO.readLines(file("pids")).filter(_.nonEmpty)
if lines.size != 2 then
sys.error(s"Expected exactly 2 recorded runs, saw ${lines.size}: $lines")
if lines(0) != lines(1) then
sys.error(s"Expected the same worker JVM to be reused, but saw ${lines(0)} then ${lines(1)}")
}
checkDistinct := Def.uncached {
val lines = IO.readLines(file("pids")).filter(_.nonEmpty)
if lines.size != 2 then
sys.error(s"Expected exactly 2 recorded runs, saw ${lines.size}: $lines")
if lines(0) == lines(1) then
sys.error(s"Expected a fresh worker JVM per run, but both used ${lines(0)}")
}
clearPids := Def.uncached {
IO.delete(file("pids"))
}
@@ -0,0 +1,13 @@
package example
import java.io.{ File, FileWriter }
import java.lang.management.ManagementFactory
import munit.FunSuite
class Marker extends FunSuite:
test("mark") {
val pid = ManagementFactory.getRuntimeMXBean.getName
val w = new FileWriter(new File("pids"), true)
try w.write(pid + "\n")
finally w.close()
}
@@ -0,0 +1,11 @@
> set testPersistentWorker := true
> testFull
> testFull
> check
> clearPids
> set testPersistentWorker := false
> testFull
> testFull
> checkDistinct
> clearPids
@@ -1,6 +1,7 @@
package sbt.internal.worker1;
import java.net.URI;
import java.util.Objects;
public class FilePath {
public URI path;
@@ -10,4 +11,17 @@ public class FilePath {
this.path = path;
this.digest = digest;
}
@Override
public boolean equals(Object o) {
if (this == o) return true;
if (!(o instanceof FilePath)) return false;
FilePath other = (FilePath) o;
return Objects.equals(path, other.path) && Objects.equals(digest, other.digest);
}
@Override
public int hashCode() {
return Objects.hash(path, digest);
}
}
@@ -9,6 +9,7 @@
package sbt.internal.worker1;
import java.io.ByteArrayOutputStream;
import java.io.IOException;
import java.io.InputStream;
import java.io.PrintStream;
import java.lang.reflect.Method;
@@ -21,7 +22,13 @@ import java.nio.channels.SocketChannel;
import java.nio.file.Files;
import java.nio.file.Path;
import java.nio.file.Paths;
import java.util.ArrayList;
import java.util.Arrays;
import java.util.Collections;
import java.util.HashSet;
import java.util.List;
import java.util.Scanner;
import java.util.Set;
import org.scalasbt.shadedgson.com.google.gson.Gson;
import org.scalasbt.shadedgson.com.google.gson.GsonBuilder;
import org.scalasbt.shadedgson.com.google.gson.JsonElement;
@@ -72,10 +79,11 @@ public final class WorkerMain {
WorkerMain app = new WorkerMain();
app.argFileWork(Paths.get(args[0].substring(1)));
System.exit(0);
} else if (args.length == 2 && args[0].equals("--tcp")) {
} else if (args.length >= 2 && args[0].equals("--tcp")) {
WorkerMain app = new WorkerMain();
int serverPort = Integer.parseInt(args[1]);
app.socketWork(serverPort);
boolean persistentWorker = Arrays.asList(args).contains("--persistent_worker");
app.socketWork(serverPort, persistentWorker);
System.exit(0);
} else if (args.length == 2 && args[0].equals("--ipc")) {
WorkerMain app = new WorkerMain();
@@ -115,14 +123,18 @@ public final class WorkerMain {
process(line);
}
void socketWork(int serverPort) throws Exception {
void socketWork(int serverPort, boolean persistentWorker) throws Exception {
InetAddress loopback = InetAddress.getByName(null);
Socket client = new Socket(loopback, serverPort);
this.jsonOut = new PrintStream(client.getOutputStream(), true, "UTF-8");
this.inScanner = new Scanner(client.getInputStream(), "UTF-8");
if (this.inScanner.hasNextLine()) {
boolean keepGoing = true;
while (keepGoing && this.inScanner.hasNextLine()) {
String line = this.inScanner.nextLine();
process(line);
keepGoing = process(line) && persistentWorker;
if (keepGoing) {
client.setSoTimeout(30 * 60 * 1000);
}
}
}
@@ -137,7 +149,7 @@ public final class WorkerMain {
}
/** This processes single request of supposed JSON line. */
void process(String json) throws Exception {
boolean process(String json) throws Exception {
JsonElement elem = JsonParser.parseString(json);
JsonObject o = elem.getAsJsonObject();
if (!o.has("jsonrpc")) {
@@ -161,13 +173,14 @@ public final class WorkerMain {
case "console":
ConsoleInfo consoleInfo = g.fromJson(params, ConsoleInfo.class);
console(id, consoleInfo);
return;
return false;
case "bye":
break;
}
String response = String.format("{ \"jsonrpc\": \"2.0\", \"result\": 0, \"id\": %d }", id);
this.jsonOut.println(response);
this.jsonOut.flush();
return !method.equals("bye");
} catch (Throwable e) {
WorkerError err = new WorkerError(1, e.getMessage());
String errMessage = g.toJson(err, err.getClass());
@@ -176,6 +189,7 @@ public final class WorkerMain {
this.jsonOut.println(errJson);
this.jsonOut.flush();
e.printStackTrace();
return false;
}
}
@@ -211,6 +225,50 @@ public final class WorkerMain {
}
}
private Set<FilePath> stableLayerEntries = Collections.emptySet();
private URLClassLoader stableLayer;
private URLClassLoader topLayer;
/** Caches non-build-output (library) entries as a parent layer; rebuilds only "target" output. */
private URLClassLoader classLoaderFor(RunInfo.JvmRunInfo info, ClassLoader parent)
throws IOException {
List<FilePath> stable = new ArrayList<>();
List<FilePath> changed = new ArrayList<>();
for (FilePath fp : info.classpath) {
(isBuildOutput(fp) ? changed : stable).add(fp);
}
Set<FilePath> stableSet = new HashSet<>(stable);
if (stableLayer == null || !stableLayerEntries.equals(stableSet)) {
stableLayer = urlClassLoaderOf(stable, parent);
stableLayerEntries = stableSet;
}
if (topLayer != null) topLayer.close();
topLayer = changed.isEmpty() ? null : urlClassLoaderOf(changed, stableLayer);
return topLayer != null ? topLayer : stableLayer;
}
private static boolean isBuildOutput(FilePath fp) {
String path = fp.path.getPath();
return path != null && (path.contains("/target/") || path.contains("\\target\\"));
}
private URLClassLoader urlClassLoaderOf(List<FilePath> entries, ClassLoader parent) {
URL[] urls =
entries.stream()
.map(
filePath -> {
try {
return filePath.path.toURL();
} catch (MalformedURLException e) {
throw new RuntimeException(e);
}
})
.toArray(URL[]::new);
return new URLClassLoader(urls, parent);
}
void console(long id, ConsoleInfo info) throws Exception {
ForkConsoleMain.main(id, info);
return;
@@ -27,6 +27,6 @@ object WorkerTest extends verify.BasicTestSuite:
val runInfo =
s"""{ "jvm": true, "jvmRunInfo": { "args": ["hi"], "classpath": $cp, "mainClass": "example.Hello" } }"""
val json = s"""{ "jsonrpc": "2.0", "id": 1, "method": "run", "params": $runInfo }"""
main.process(json)
val _ = main.process(json)
}
end WorkerTest