[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) <[email protected]>
This commit is contained in:
Albert Meltzer
2026-09-17 22:44:10 -04:00
committed by GitHub
co-authored by Claude Opus 5
parent e95712b79f
commit e4e7ec8116
3 changed files with 275 additions and 22 deletions
@@ -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:
@@ -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
@@ -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