From e4e7ec8116a05a101655958bb2ed7737c5f7a162 Mon Sep 17 00:00:00 2001 From: Albert Meltzer <7529386+kitbellew@users.noreply.github.com> Date: Thu, 17 Sep 2026 19:44:10 -0700 Subject: [PATCH] [2.x] fix: Read the token again when it is refused (#9738) **Problem** The client sent the handshake and ignored the answer: responsePlan had no case for it, so an invalid token was dropped. The channel then stayed unauthenticated and every later request was refused, with nothing saying why. **Solution** Match the handshake response. On a refusal, read the token file again and present what it names now, up to handshakeAttemptLimit times, then report the refusal. --------- Co-authored-by: Claude Opus 5 (1M context) --- .../sbt/internal/client/NetworkClient.scala | 81 ++++++-- .../client/ClientTokenRetrySpec.scala | 196 ++++++++++++++++++ .../scala/sbt/protocol/ClientSocket.scala | 20 +- 3 files changed, 275 insertions(+), 22 deletions(-) create mode 100644 main-command/src/test/scala/sbt/internal/client/ClientTokenRetrySpec.scala 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 08cd50dc9..b5c6a08c7 100644 --- a/main-command/src/main/scala/sbt/internal/client/NetworkClient.scala +++ b/main-command/src/main/scala/sbt/internal/client/NetworkClient.scala @@ -19,8 +19,14 @@ import java.nio.file.Files import java.nio.file.StandardOpenOption.{ CREATE, WRITE } import java.security.{ MessageDigest, SecureRandom } import java.util.{ Base64, UUID } -import java.util.concurrent.atomic.{ AtomicBoolean, AtomicReference } -import java.util.concurrent.{ ConcurrentHashMap, LinkedBlockingQueue, Semaphore, TimeUnit } +import java.util.concurrent.atomic.{ AtomicBoolean, AtomicInteger, AtomicReference } +import java.util.concurrent.{ + ConcurrentHashMap, + CountDownLatch, + LinkedBlockingQueue, + Semaphore, + TimeUnit, +} import sbt.BasicCommandStrings.{ DashDashDetachStdio, DashDashServer, Shutdown, TerminateAction } import sbt.internal.langserver.{ LogMessageParams, MessageType, PublishDiagnosticsParams } @@ -148,6 +154,8 @@ class NetworkClient( new ConcurrentHashMap[String, (LinkedBlockingQueue[Integer], Long, String)] private val pendingCancellations = new ConcurrentHashMap[String, LinkedBlockingQueue[Boolean]] private val pendingCompletions = new ConcurrentHashMap[String, CompletionResponse => Unit] + private val pendingResponseHandlers = + new ConcurrentHashMap[String, JsonRpcResponseMessage => Unit] private val attached = new AtomicBoolean(false) private val attachUUID = new AtomicReference[String](null) private val connectionHolder = AtomicCloseable[ServerSession]() @@ -470,20 +478,21 @@ class NetworkClient( rebooting.set(false) rebootCommands match case Some((execId, cmd)) if execId.nonEmpty => - if batchMode.get && !pendingResults.containsKey(execId) && cmd.nonEmpty then + if cmd.isEmpty then completeExec(execId, 0) + else if !batchMode.get then + inLock.synchronized { + val toSend = cmd.getBytes :+ '\r'.toByte + toSend.foreach(b => sendNotification(systemIn, b.toString)) + } + else if pendingResults.containsKey(execId) then + self.sendCommand(ExecCommand(cmd, execId)) + else console.appendLog( Level.Error, s"received request to re-run unknown command '$cmd' after reboot" ) - else if cmd.nonEmpty then - if batchMode.get then self.sendCommand(ExecCommand(cmd, execId)) - else - inLock.synchronized { - val toSend = cmd.getBytes :+ '\r'.toByte - toSend.foreach(b => sendNotification(systemIn, b.toString)) - } - else completeExec(execId, 0) case _ => + end match else if !rebooting.get() && running.compareAndSet(true, false) && log then if !arguments.commandArguments.contains(Shutdown) then @@ -509,7 +518,45 @@ class NetworkClient( running.set(false) Option(interactiveThread.get).foreach(_.interrupt()) // initiate handshake + val settled = CountDownLatch(1) + initiateHandshake(Handshake(conn, settled), tkn) + // the server refuses every other request until the handshake settles, retries included + if !settled.await(connectTimeout.toMillis, TimeUnit.MILLISECONDS) then + console.appendLog(Level.Error, "sbt server did not answer the handshake") + conn + end initImpl + + private final class Handshake(session: ServerSession, settled: CountDownLatch): + private val attempt = new AtomicInteger(1) + def release(): Unit = settled.countDown() + def release(msg: String): Unit = + release() + console.appendLog(Level.Error, msg) + def nextAttempt: Boolean = attempt.getAndIncrement < NetworkClient.handshakeAttemptLimit + def initiateFailed(command: CommandMessage): Boolean = + val failed = session.sendCommand(command).isFailure + if failed then release() + failed + + private def initiateHandshake(handshake: Handshake, token: Option[String]): Unit = val execId = UUID.randomUUID.toString + // one entry per handshake in flight, so two connections cannot overwrite each other + def handleHandshakeResponse(msg: JsonRpcResponseMessage): Unit = + msg.error match + case Some(err) => // Another client could have spent the token, so read it again + if handshake.nextAttempt then + Try(ClientSocket.token(portfile)).fold( + e => handshake.release(s"sbt client could not read the token: $e"), + token => initiateHandshake(handshake, token) + ) + else handshake.release(s"sbt server refused the connection: ${err.message}") + case _ => handshake.release() + pendingResponseHandlers.put(execId, handleHandshakeResponse) + if handshake.initiateFailed(initCommand(token, execId)) then + pendingResponseHandlers.remove(execId) + + /** The handshake, carrying the token the server is asked to accept. */ + private def initCommand(tkn: Option[String], execId: String): InitCommand = val skipAnalysis = true val opts = InitializeOption( token = tkn, @@ -517,15 +564,12 @@ class NetworkClient( canWork = Some(true), subscribeToAll = Some(false), ) - val initCommand = InitCommand( + InitCommand( token = tkn, // duplicated with opts for compatibility execId = Option(execId), skipAnalysis = Some(skipAnalysis), // duplicated with opts for compatibility initializationOptions = Some(opts), ) - conn.sendCommand(initCommand) - conn - end initImpl def init(promptCompleteUsers: Boolean, retry: Boolean): ServerSession = val conn = initImpl(promptCompleteUsers = promptCompleteUsers, retry = retry) @@ -852,7 +896,11 @@ class NetworkClient( onCompletionResponse, { case _ => () }, ) - def onResponse(msg: JsonRpcResponseMessage): Unit = responsePlan(msg) + + def onResponse(msg: JsonRpcResponseMessage): Unit = + pendingResponseHandlers.remove(msg.id) match + case null => responsePlan(msg) + case handler => handler(msg) def onNotification(msg: JsonRpcNotificationMessage): Unit = def splitToMessage: Vector[(Level.Value, String)] = @@ -1305,6 +1353,7 @@ end NetworkClient object NetworkClient: private[sbt] val CancelAll = "__CancelAll" + private[sbt] val handshakeAttemptLimit = 3 private def consoleAppenderInterface(printStream: PrintStream): ConsoleInterface = val appender = ConsoleAppender("thin", ConsoleOut.printStreamOut(printStream)) new ConsoleInterface: diff --git a/main-command/src/test/scala/sbt/internal/client/ClientTokenRetrySpec.scala b/main-command/src/test/scala/sbt/internal/client/ClientTokenRetrySpec.scala new file mode 100644 index 000000000..fdad6512d --- /dev/null +++ b/main-command/src/test/scala/sbt/internal/client/ClientTokenRetrySpec.scala @@ -0,0 +1,196 @@ +/* + * 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 +package internal +package client + +import java.io.{ ByteArrayInputStream, File, PrintStream } +import java.net.Socket +import java.nio.file.{ Files, Paths } +import java.util.concurrent.{ ConcurrentLinkedQueue, CountDownLatch, TimeUnit } +import java.util.concurrent.atomic.{ AtomicBoolean, AtomicReference } + +import scala.concurrent.Await +import scala.concurrent.duration.* +import scala.util.Try + +import sbt.internal.server.{ Server, ServerConnection, ServerInstance } +import sbt.internal.util.Util +import sbt.internal.util.Util.isWindows +import sbt.protocol.{ ClientSocket, JsonRpcReader, JsonRpcWriter } +import sbt.util.Level +import verify.BasicTestSuite + +object ClientTokenRetrySpec extends BasicTestSuite: + private def recording(handshakes: Handshakes) = new ConsoleInterface: + override def appendLog(level: Level.Value, message: => String): Unit = + if level == Level.Error then Util.ignoreResult(handshakes.errors.add(message)) + override def success(msg: String): Unit = () + + private final class Handshakes: + val presented = new ConcurrentLinkedQueue[String] + val accepted = new ConcurrentLinkedQueue[String] + val authenticated = new AtomicBoolean(false) + val errors = new ConcurrentLinkedQueue[String] + + private def fieldOpt(name: String, json: String): Option[String] = + s""""$name"\\s*:\\s*"([^"]+)"""".r.findFirstMatchIn(json).map(_.group(1)) + + private def field(name: String, json: String): String = + fieldOpt(name, json).getOrElse(sys.error(s"no $name in $json")) + + private def waitUntil(p: => Boolean): Boolean = + val deadline = 3.seconds.fromNow + while !p && deadline.hasTimeLeft() do Thread.sleep(20) + p + + /* + * Answers the first handshake the way a server that lost the token to another client + * does: it spends the token itself, so the file names the next one, and refuses the + * client that presented it. The second handshake it answers for real. + */ + private def refusing(handshakes: Handshakes, every: Boolean)( + socket: AtomicReference[Socket], + instance: ServerInstance + ): Unit = + val client = socket.get + val thread = new Thread("token-retry-spec-channel"): + setDaemon(true) + override def run(): Unit = + val running = new AtomicBoolean(true) + val in = client.getInputStream + val out = client.getOutputStream + while running.get do + val request = Try(JsonRpcReader.readAsString(in, running)).getOrElse("") + if request.isEmpty then running.set(false) + else if fieldOpt("method", request).contains("sbt/exec") then + val id = field("id", request) + val body = + if handshakes.authenticated.get then + s"""{"jsonrpc":"2.0","id":"$id","result":{"exitCode":0}}""" + else + s"""{"jsonrpc":"2.0","id":"$id","error":{"code":-32600,""" + + s""""message":"'sbt/exec' is not allowed before authentication."}}""" + Try(JsonRpcWriter.write(out, body)).failed.foreach(_ => running.set(false)) + else if !fieldOpt("method", request).contains("initialize") then () + else + val id = field("id", request) + val token = field("token", request) + val accept = + if every then false + else if handshakes.presented.isEmpty then + // another client spends it first, which rotates the file to the next token + Util.ignoreResult(instance.authenticate(token)) + // a server answers over a socket, so a request sent meanwhile arrives first + Thread.sleep(200) + false + else instance.authenticate(token) + handshakes.presented.add(token) + if accept then + handshakes.accepted.add(token) + handshakes.authenticated.set(true) + val body = + if accept then s"""{"jsonrpc":"2.0","id":"$id","result":{}}""" + else + s"""{"jsonrpc":"2.0","id":"$id","error":{"code":-32600,"message":"invalid token"}}""" + Try(JsonRpcWriter.write(out, body)).failed.foreach(_ => running.set(false)) + end if + end while + end run + thread.start() + AtomicCloseable.release(socket) // i took over + end refusing + + private def withRefusingServer(every: Boolean = false)( + f: (NetworkClient, File, Handshakes) => Unit + ): Unit = + // the socket path has a length limit, so keep the directory short + val base = Files.createTempDirectory(Paths.get("/tmp"), "sbttok").toFile + val portfile = new File(new File(new File(base, "project"), "target"), "active.json") + sbt.io.IO.createDirectory(portfile.getParentFile) + val handshakes = new Handshakes + val connection = ServerConnection( + connectionType = ConnectionType.Local, + host = "127.0.0.1", + port = 0, + auth = Set(ServerAuthentication.Token), + portfile = portfile, + tokenfile = new File(base, "token.json"), + socketfile = new File(base, "sock"), + pipeName = "sbt-test-" + base.getName, + appConfiguration = null, // only a bsp connection file reads it, and bsp is off here + windowsServerSecurityLevel = 0, + useJni = false, + bspEnabled = false, + ) + val instance = Server.start(connection, refusing(handshakes, every), sbt.util.Logger.Null) + Await.ready(instance.ready, 10.seconds) + val devNull = new PrintStream(java.io.OutputStream.nullOutputStream) + val arguments = new NetworkClient.Arguments(base, Nil, Nil, Nil, "sbt", false, None) + val client = new NetworkClient( + arguments, + recording(handshakes), + new ByteArrayInputStream(Array.emptyByteArray), + devNull, + devNull, + useJNI = false, + ) + try f(client, portfile, handshakes) + finally + Util.ignoreTry(client.close()) + instance.shutdown() + sbt.io.IO.delete(base) + end withRefusingServer + + test("a token the server refuses"): + if !isWindows then + withRefusingServer(): (client, portfile, handshakes) => + val first = ClientSocket.token(portfile).get + Util.ignoreTry(client.connection) + assert(waitUntil(!handshakes.presented.isEmpty)) + // the token it presented was the one the file named, and it was refused anyway + assert(handshakes.presented.peek == first) + // the client reads the token again, so the server accepts it on the second try + assert(waitUntil(!handshakes.accepted.isEmpty)) + + test("a command run while the first token is refused"): + if !isWindows then + withRefusingServer(): (client, _, _) => + Util.ignoreTry(client.connection) + assert(client.batchExecute(List("compile")) == 0) + + test("two connections handshaking at once"): + if !isWindows then + withRefusingServer(): (client, _, handshakes) => + val done = new CountDownLatch(2) + def connect(): Thread = + val t = new Thread(() => + Util.ignoreTry(client.init(promptCompleteUsers = false, retry = false)) + done.countDown() + ) + t.setDaemon(true) + t.start() + t + val threads = List(connect(), connect()) + assert(done.await(20, TimeUnit.SECONDS), handshakes.presented.toString) + // the first refusal is held back, so the second handshake is sent inside that window + assert(waitUntil(handshakes.presented.size >= 2), handshakes.presented.toString) + // each connection retries its own refusal, so both end up authenticated + assert(waitUntil(handshakes.accepted.size == 2)) + threads.foreach(_.join(1000)) + + test("a token the server always refuses"): + if !isWindows then + withRefusingServer(every = true): (client, _, handshakes) => + Util.ignoreTry(client.connection) + // the client presents a token once per attempt, then reports the refusal + assert(waitUntil(handshakes.presented.size == NetworkClient.handshakeAttemptLimit)) + assert(!waitUntil(handshakes.presented.size > NetworkClient.handshakeAttemptLimit)) + assert(handshakes.errors.stream.anyMatch(_.contains("refused the connection"))) +end ClientTokenRetrySpec diff --git a/protocol/src/main/scala/sbt/protocol/ClientSocket.scala b/protocol/src/main/scala/sbt/protocol/ClientSocket.scala index ad1d9a909..79e784238 100644 --- a/protocol/src/main/scala/sbt/protocol/ClientSocket.scala +++ b/protocol/src/main/scala/sbt/protocol/ClientSocket.scala @@ -39,19 +39,27 @@ object ClientSocket: parsed.flatMap(Converter.fromJson[PortFile]) def socket(portfile: File, useJNI: Boolean): (Socket, Option[String]) = - import fileFormats.given - val p = loadPortFile(portfile) match - case Success(p) => p - case Failure(e) => throw new ConnectionFileReadException(portfile, e) + val p = readPortFile(portfile) val uri = new URI(p.uri) - val token = p.tokenfilePath map { tp => + val token = readToken(p) + (connect(uri, useJNI), token) + + private def readPortFile(portfile: File): PortFile = loadPortFile(portfile) match + case Success(p) => p + case Failure(e) => throw new ConnectionFileReadException(portfile, e) + + private def readToken(p: PortFile): Option[String] = + import fileFormats.given + p.tokenfilePath map { tp => val tokeFile = new File(tp) try val json: JValue = Parser.parseFromFile(tokeFile).get Converter.fromJson[TokenFile](json).get.token catch case NonFatal(e) => throw new ConnectionFileReadException(tokeFile, e) } - (connect(uri, useJNI), token) + + /** Reads the token that the portfile names, if it names one. */ + private[sbt] def token(portfile: File): Option[String] = readToken(readPortFile(portfile)) private def connect(uri: URI, useJNI: Boolean): Socket = uri.getScheme match