Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 18 additions & 3 deletions app/src/main/java/ai/sealgate/stdiod/mcp/AndroidComputerSource.kt
Original file line number Diff line number Diff line change
Expand Up @@ -166,7 +166,14 @@ class AndroidComputerSource(context: Context) : ComputerSource {
val baseline = uiStateMarker(service)
val path = Path().apply { moveTo(x.toFloat(), y.toFloat()) }
val performed = dispatchGesture(service, path, durationMillis)
finishAction(service, "tap", performed, if (performed) null else "Android rejected the tap gesture", baseline)
finishAction(
service,
"tap",
performed,
if (performed) null else "Android rejected the tap gesture",
baseline,
requireTransition = false,
)
}

override fun swipe(
Expand All @@ -185,7 +192,14 @@ class AndroidComputerSource(context: Context) : ComputerSource {
lineTo(endX.toFloat(), endY.toFloat())
}
val performed = dispatchGesture(service, path, durationMillis)
finishAction(service, "swipe", performed, if (performed) null else "Android rejected the swipe gesture", baseline)
finishAction(
service,
"swipe",
performed,
if (performed) null else "Android rejected the swipe gesture",
baseline,
requireTransition = false,
)
}

override fun globalAction(action: String): ComputerOperationResult = withService { service ->
Expand Down Expand Up @@ -263,9 +277,10 @@ class AndroidComputerSource(context: Context) : ComputerSource {
performed: Boolean,
error: String?,
baseline: UiStateMarker,
requireTransition: Boolean = true,
): ComputerOperationResult {
if (!performed) return actionFailure(action, error ?: "action was not performed")
val settle = uiSettler.awaitPostAction(baseline) { uiStateMarker(service) }
val settle = uiSettler.awaitPostAction(baseline, requireTransition) { uiStateMarker(service) }
val observation = captureObservation(service)
val payload = buildJsonObject {
put("action", buildJsonObject {
Expand Down
43 changes: 38 additions & 5 deletions app/src/main/java/ai/sealgate/stdiod/mcp/MobileCommandRouter.kt
Original file line number Diff line number Diff line change
Expand Up @@ -52,6 +52,7 @@ class MobileCommandRouter(modules: List<BaseMcpModule>) {
private val supplementLock = Any()
private val supplements = LinkedHashMap<String, MobileCommandSupplement>()
private var pendingSupplementBytes = 0L
private var latestComputerSupplementToken: String? = null

init {
val mappedTools = SPECS.map(CommandSpec::tool).toSet()
Expand Down Expand Up @@ -122,8 +123,14 @@ class MobileCommandRouter(modules: List<BaseMcpModule>) {
.joinToString("\n")
val typedContent = content.filter { it["type"]?.jsonPrimitive?.content != "text" }
val structuredContent = result["structuredContent"] as? JsonObject
val supplementToken = if (typedContent.isNotEmpty() || structuredContent != null) {
retainSupplement(MobileCommandSupplement(typedContent, structuredContent))
// Status is already emitted as complete JSON text. Retaining its duplicate
// structured payload would make polling grow an invisible side channel.
val isComputerStatus = spec.module == ComputerModule.NAME && spec.tool == "computer_status"
val supplementToken = if (typedContent.isNotEmpty() || structuredContent != null && !isComputerStatus) {
retainSupplement(
MobileCommandSupplement(typedContent, structuredContent),
replacePreviousComputerObservation = spec.module == ComputerModule.NAME,
)
} else {
null
}
Expand All @@ -138,6 +145,7 @@ class MobileCommandRouter(modules: List<BaseMcpModule>) {
fun clearSupplements() = synchronized(supplementLock) {
supplements.clear()
pendingSupplementBytes = 0L
latestComputerSupplementToken = null
}

fun availableNamespacesJson(): String = buildJsonArray {
Expand All @@ -152,18 +160,43 @@ class MobileCommandRouter(modules: List<BaseMcpModule>) {
tokens.mapNotNull(supplements::remove).also {
supplements.clear()
pendingSupplementBytes = 0L
latestComputerSupplementToken = null
}
}

private fun retainSupplement(supplement: MobileCommandSupplement): String = synchronized(supplementLock) {
check(supplements.size < MAX_PENDING_SUPPLEMENTS) { "too many pending mobile command attachments" }
private fun retainSupplement(
supplement: MobileCommandSupplement,
replacePreviousComputerObservation: Boolean,
): String = synchronized(supplementLock) {
val supplementBytes = supplement.serializedBytes()
check(supplementBytes <= MAX_PENDING_SUPPLEMENT_BYTES - pendingSupplementBytes) {
check(supplementBytes <= MAX_PENDING_SUPPLEMENT_BYTES) {
"mobile command attachment exceeds 4 MiB"
}
// A shell script can perform many observation-producing actions while
// redirecting their textual output. Typed MCP attachments do not flow
// through Bash file descriptors, so retaining every one would grow an
// invisible side channel until the script failed. Keep only the most
// recent computer observation: it represents the device state after
// the latest action and bounds loops independently of their length.
// Other typed results (for example multiple camera snapshots) retain
// their existing multi-attachment behavior.
val previousToken = latestComputerSupplementToken.takeIf { replacePreviousComputerObservation }
val previous = previousToken?.let(supplements::get)
val previousBytes = previous?.serializedBytes() ?: 0L
val projectedCount = supplements.size - if (previous == null) 0 else 1
val projectedBytes = pendingSupplementBytes - previousBytes + supplementBytes
check(projectedCount < MAX_PENDING_SUPPLEMENTS) { "too many pending mobile command attachments" }
check(projectedBytes <= MAX_PENDING_SUPPLEMENT_BYTES) {
"pending mobile command attachments exceed 4 MiB"
}
if (previousToken != null && previous != null) {
supplements.remove(previousToken)
pendingSupplementBytes -= previousBytes
}
val token = nextSupplementId.incrementAndGet().toString()
supplements[token] = supplement
pendingSupplementBytes += supplementBytes
if (replacePreviousComputerObservation) latestComputerSupplementToken = token
token
}

Expand Down
3 changes: 2 additions & 1 deletion app/src/main/java/ai/sealgate/stdiod/mcp/UiSettler.kt
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,7 @@ internal class UiSettler(
) {
fun awaitPostAction(
baseline: UiStateMarker,
requireTransition: Boolean = true,
sample: () -> UiStateMarker,
): UiSettleResult {
val startedAt = uptimeMillis()
Expand Down Expand Up @@ -61,7 +62,7 @@ internal class UiSettler(
val transitionObserved = eventObserved || windowChanged
val eventStreamIsQuiet = now - current.lastEventUptimeMillis >= quietMillis
val activeWindowIsStable = current.hasActiveWindow && now - stableSince >= quietMillis
if (transitionObserved && eventStreamIsQuiet && activeWindowIsStable) {
if ((!requireTransition || transitionObserved) && eventStreamIsQuiet && activeWindowIsStable) {
return UiSettleResult(
settled = true,
postActionEventObserved = eventObserved,
Expand Down
83 changes: 74 additions & 9 deletions app/src/main/java/ai/sealgate/stdiod/tunnel/TunnelClient.kt
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,12 @@ import okhttp3.Response
import okhttp3.WebSocket
import okhttp3.WebSocketListener
import java.util.concurrent.TimeUnit
import java.util.concurrent.ArrayBlockingQueue
import java.util.concurrent.ConcurrentHashMap
import java.util.concurrent.RejectedExecutionException
import java.util.concurrent.ThreadPoolExecutor
import java.util.concurrent.atomic.AtomicBoolean
import java.util.concurrent.atomic.AtomicReference
import kotlin.coroutines.resume
import kotlin.concurrent.thread
import kotlin.random.Random
Expand Down Expand Up @@ -60,7 +65,7 @@ class TunnelClient(
private val modulesByName: Map<String, LocalMcpModule> = modules.associateBy { it.name }

/** server_id (backend's key for `mcp_frame`s) → built-in module. */
private val modulesByServerId = HashMap<String, LocalMcpModule>()
private val modulesByServerId = ConcurrentHashMap<String, LocalMcpModule>()

private val _state = MutableStateFlow<TunnelState>(TunnelState.Disconnected)
val state: StateFlow<TunnelState> = _state
Expand All @@ -71,6 +76,7 @@ class TunnelClient(
private val modulesClosed = AtomicBoolean(false)
private val modulesCloseFinished = CompletableDeferred<Unit>()
private val moduleLock = Any()
private val activeDispatcher = AtomicReference<McpRequestDispatcher?>()

fun start() {
if (stopped.get()) return
Expand All @@ -84,6 +90,7 @@ class TunnelClient(
loopJob = null
webSocket?.close(NORMAL_CLOSURE, "client stopping")
webSocket = null
activeDispatcher.getAndSet(null)?.close()
if (modulesClosed.compareAndSet(false, true)) {
// A QuickJS evaluation may hold its runtime lock until the 60-second
// execution limit. Never make the service/main thread wait for it.
Expand Down Expand Up @@ -131,6 +138,9 @@ class TunnelClient(

/** Runs one WebSocket session to completion. Returns true if `server_hello` arrived. */
private suspend fun runOneConnection(): Boolean = suspendCancellableCoroutine { cont ->
val sessionActive = AtomicBoolean(true)
val dispatcher = McpRequestDispatcher()
activeDispatcher.getAndSet(dispatcher)?.close()
val request = Request.Builder()
.url(gatewayUrl)
.header("Authorization", "Bearer $authToken")
Expand Down Expand Up @@ -179,7 +189,24 @@ class TunnelClient(
bindServers(webSocket, frame.added + frame.updated)
frame.removed.forEach(modulesByServerId::remove)
}
is McpFrame -> routeMcpFrame(webSocket, frame)
is McpFrame -> {
// Local modules may perform gestures, screenshots, or a
// full Bash script. Never run them on OkHttp's reader
// callback: doing so prevents WebSocket control frames
// and unrelated tunnel messages from being processed.
// Resolve the binding at receipt time. A later desired-state
// update must not retroactively change an earlier request.
val module = modulesByServerId[frame.serverId] ?: modulesByName[frame.serverId]
if (!dispatcher.submit {
routeMcpFrame(webSocket, frame, module, sessionActive)
}
) {
// Backpressure is explicit: retaining an unbounded number
// of long-running requests would eventually exhaust memory.
webSocket.close(TRY_AGAIN_LATER, "MCP request queue full")
finish()
}
}
is Ping -> send(webSocket, Pong)
is Pong -> Unit
// Built-in modules have no spawn-time env/spec to store.
Expand All @@ -205,14 +232,22 @@ class TunnelClient(
}

fun finish() {
if (!sessionActive.compareAndSet(true, false)) return
dispatcher.close()
activeDispatcher.compareAndSet(dispatcher, null)
this@TunnelClient.webSocket = null
modulesByServerId.clear()
if (cont.isActive) cont.resume(sawServerHello)
}
}

val socket = httpClient.newWebSocket(request, listener)
cont.invokeOnCancellation { socket.cancel() }
cont.invokeOnCancellation {
sessionActive.set(false)
dispatcher.close()
activeDispatcher.compareAndSet(dispatcher, null)
socket.cancel()
}
}

/**
Expand Down Expand Up @@ -248,10 +283,13 @@ class TunnelClient(
}
}

private fun routeMcpFrame(webSocket: WebSocket, frame: McpFrame) {
// Accept a module addressed by bare name too, so a backend that keys
// built-ins by name (and tests) can skip the desired-state handshake.
val module = modulesByServerId[frame.serverId] ?: modulesByName[frame.serverId]
private fun routeMcpFrame(
webSocket: WebSocket,
frame: McpFrame,
module: LocalMcpModule?,
sessionActive: AtomicBoolean,
) {
if (!sessionActive.get() || stopped.get()) return
if (module == null) {
send(
webSocket,
Expand All @@ -264,10 +302,10 @@ class TunnelClient(
return
}
val response = synchronized(moduleLock) {
if (stopped.get()) return
if (!sessionActive.get() || stopped.get()) return
module.handle(frame.frame)
} ?: return
if (stopped.get()) return
if (!sessionActive.get() || stopped.get()) return
send(webSocket, McpFrame(serverId = frame.serverId, frame = response))
}

Expand All @@ -278,6 +316,7 @@ class TunnelClient(
companion object {
private const val TAG = "TunnelClient"
private const val NORMAL_CLOSURE = 1000
private const val TRY_AGAIN_LATER = 1013
private const val INITIAL_BACKOFF_MILLIS = 1_000L
private const val MAX_BACKOFF_MILLIS = 60_000L

Expand All @@ -290,3 +329,29 @@ class TunnelClient(
.build()
}
}

internal const val MCP_REQUEST_QUEUE_CAPACITY = 16

/** A bounded, session-scoped serial dispatcher for potentially slow module calls. */
internal class McpRequestDispatcher {
private val executor = ThreadPoolExecutor(
1,
1,
0L,
TimeUnit.MILLISECONDS,
ArrayBlockingQueue(MCP_REQUEST_QUEUE_CAPACITY),
{ runnable -> Thread(runnable, "mobile-mcp-requests").apply { isDaemon = true } },
ThreadPoolExecutor.AbortPolicy(),
)

fun submit(task: () -> Unit): Boolean = try {
executor.execute(task)
true
} catch (_: RejectedExecutionException) {
false
}

fun close() {
executor.shutdownNow()
}
}
60 changes: 55 additions & 5 deletions app/src/test/java/ai/sealgate/stdiod/mcp/ComputerModuleTest.kt
Original file line number Diff line number Diff line change
Expand Up @@ -66,7 +66,7 @@ class ComputerModuleTest {
}

@Test
fun routerBoundsPendingSupplementsBySerializedBytes() {
fun routerKeepsOnlyTheLatestPendingSupplement() {
val source = FakeComputerSource().apply {
screenshotData = "a".repeat((MobileCommandRouter.MAX_PENDING_SUPPLEMENT_BYTES / 2 + 1024).toInt())
}
Expand All @@ -77,26 +77,76 @@ class ComputerModuleTest {
val second = Json.parseToJsonElement(router.executeJson(request)).jsonObject

assertEquals(0, first["exitCode"]!!.jsonPrimitive.content.toInt())
assertEquals(1, second["exitCode"]!!.jsonPrimitive.content.toInt())
assertTrue(second["stderr"]!!.jsonPrimitive.content.contains("attachments exceed 4 MiB"))
assertEquals(0, second["exitCode"]!!.jsonPrimitive.content.toInt())
val firstToken = first["supplementToken"]!!.jsonPrimitive.content
val secondToken = second["supplementToken"]!!.jsonPrimitive.content
val retained = router.takeSupplements(listOf(firstToken, secondToken)).single()
assertEquals("obs_2", retained.structuredContent!!["observationId"]!!.jsonPrimitive.content)
router.clearSupplements()
}

@Test
fun repeatedComputerStatusDoesNotAccumulateOrReplaceTheLatestObservation() {
val router = MobileCommandRouter(listOf(ComputerModule(FakeComputerSource())))

val observation = router.execute("computer", listOf("observe"))
val statuses = List(65) { router.execute("computer", listOf("status")) }

assertEquals(0, observation.exitCode)
assertTrue(statuses.all { it.exitCode == 0 })
assertTrue(statuses.all { it.supplementToken == null })
val supplement = router.takeSupplements(listOf(observation.supplementToken!!)).single()
assertEquals("image", supplement.content.single()["type"]!!.jsonPrimitive.content)
assertEquals("obs_1", supplement.structuredContent!!["observationId"]!!.jsonPrimitive.content)
}

@Test
fun failedComputerReplacementPreservesThePreviousObservation() {
val source = FakeComputerSource().apply { screenshotData = "a".repeat(256 * 1024) }
val camera = object : CameraSource {
private val result = CameraOperationResult(
payload = buildJsonObject { put("lens", JsonPrimitive("back")) },
photo = CameraPhoto("a".repeat(3 * 1024 * 1024), "image/jpeg"),
)
override fun status() = result
override fun list() = result
override fun snap(options: CameraSnapOptions) = result
}
val router = MobileCommandRouter(listOf(CameraModule(camera), ComputerModule(source)))
val cameraResult = router.execute("camera", listOf("snap"))
val firstObservation = router.execute("computer", listOf("observe"))
source.screenshotData = "b".repeat(2 * 1024 * 1024)

val failedReplacement = Json.parseToJsonElement(
router.executeJson("""{"namespace":"computer","args":["observe"]}"""),
).jsonObject

assertEquals(1, failedReplacement["exitCode"]!!.jsonPrimitive.content.toInt())
val retained = router.takeSupplements(
listOf(cameraResult.supplementToken!!, firstObservation.supplementToken!!),
)
assertEquals(2, retained.size)
assertEquals("obs_1", retained.last().structuredContent!!["observationId"]!!.jsonPrimitive.content)
}

private class FakeComputerSource : ComputerSource {
var nodeId = ""
var text = ""
var tapDurationMillis = 0
var screenshotData = "aGVsbG8="
var observationNumber = 0

private fun result() = ComputerOperationResult(
payload = buildJsonObject {
put("observationId", JsonPrimitive("obs_1"))
put("observationId", JsonPrimitive("obs_${++observationNumber}"))
put("accessibilityTree", buildJsonObject { put("nodes", kotlinx.serialization.json.buildJsonArray {}) })
},
screenshot = ComputerScreenshot(screenshotData, "image/jpeg"),
)

override fun status() = result()
override fun status() = ComputerOperationResult(
payload = buildJsonObject { put("enabled", JsonPrimitive(true)) },
)
override fun observe() = result()
override fun click(nodeId: String): ComputerOperationResult = result().also { this.nodeId = nodeId }
override fun setText(nodeId: String, text: String): ComputerOperationResult = result().also {
Expand Down
Loading
Loading