mirror of
https://github.com/sbt/sbt.git
synced 2026-10-06 18:14:04 +02:00
[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:
co-authored by
Claude Opus 5
parent
e95712b79f
commit
e4e7ec8116
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user