Fix: P0 fix-pass for Hermes Multi-Host architecture and event parsing
This commit is contained in:
parent
fc3ddeda65
commit
7e2512e390
11 changed files with 850 additions and 199 deletions
|
|
@ -178,40 +178,31 @@ sealed class GatewayEvent {
|
|||
private val json = Json { ignoreUnknownKeys = true; isLenient = true }
|
||||
|
||||
fun parse(root: JsonObject): GatewayEvent {
|
||||
// Find event name and data container
|
||||
var eventType = ""
|
||||
var dataObj: JsonObject = root
|
||||
val params = root["params"]?.jsonObject ?: root
|
||||
|
||||
if (root.containsKey("method")) {
|
||||
val method = root["method"]?.jsonPrimitive?.content ?: ""
|
||||
if (method == "event" && root.containsKey("params")) {
|
||||
val params = root["params"]?.jsonObject ?: JsonObject(emptyMap())
|
||||
eventType = params["event"]?.jsonPrimitive?.content
|
||||
?: params["type"]?.jsonPrimitive?.content
|
||||
?: ""
|
||||
dataObj = params["data"]?.jsonObject
|
||||
?: params["payload"]?.jsonObject
|
||||
?: params
|
||||
} else if (method.isNotEmpty()) {
|
||||
eventType = method
|
||||
dataObj = root["params"]?.jsonObject ?: root
|
||||
}
|
||||
}
|
||||
// 1. event type -> params["type"]
|
||||
val eventType = params["type"]?.jsonPrimitive?.content
|
||||
?: params["event"]?.jsonPrimitive?.content
|
||||
?: root["type"]?.jsonPrimitive?.content
|
||||
?: root["event"]?.jsonPrimitive?.content
|
||||
?: ""
|
||||
|
||||
if (eventType.isEmpty()) {
|
||||
eventType = root["event"]?.jsonPrimitive?.content
|
||||
?: root["type"]?.jsonPrimitive?.content
|
||||
?: ""
|
||||
if (root.containsKey("data") && root["data"] is JsonObject) {
|
||||
dataObj = root["data"]!!.jsonObject
|
||||
} else if (root.containsKey("payload") && root["payload"] is JsonObject) {
|
||||
dataObj = root["payload"]!!.jsonObject
|
||||
}
|
||||
}
|
||||
// 2. runtime session -> params["session_id"] (Do NOT search inside payload)
|
||||
val sessionId = params["session_id"]?.jsonPrimitive?.content
|
||||
?: params["session_key"]?.jsonPrimitive?.content
|
||||
?: root["session_id"]?.jsonPrimitive?.content
|
||||
?: root["session_key"]?.jsonPrimitive?.content
|
||||
|
||||
// 3. event body -> params["payload"]
|
||||
val payloadObj = params["payload"]?.jsonObject
|
||||
?: params["data"]?.jsonObject
|
||||
?: root["payload"]?.jsonObject
|
||||
?: root["data"]?.jsonObject
|
||||
?: params
|
||||
|
||||
fun getString(vararg keys: String): String {
|
||||
for (k in keys) {
|
||||
val v = dataObj[k]?.jsonPrimitive?.content ?: root[k]?.jsonPrimitive?.content
|
||||
val v = payloadObj[k]?.jsonPrimitive?.content
|
||||
if (v != null) return v
|
||||
}
|
||||
return ""
|
||||
|
|
@ -219,8 +210,7 @@ sealed class GatewayEvent {
|
|||
|
||||
fun getNullableString(vararg keys: String): String? {
|
||||
for (k in keys) {
|
||||
val el = dataObj[k] ?: root[k]
|
||||
val v = el?.jsonPrimitive?.content
|
||||
val v = payloadObj[k]?.jsonPrimitive?.content
|
||||
if (v != null) return v
|
||||
}
|
||||
return null
|
||||
|
|
@ -228,7 +218,7 @@ sealed class GatewayEvent {
|
|||
|
||||
fun getLong(vararg keys: String): Long {
|
||||
for (k in keys) {
|
||||
val v = (dataObj[k]?.jsonPrimitive ?: root[k]?.jsonPrimitive)?.longOrNull
|
||||
val v = payloadObj[k]?.jsonPrimitive?.longOrNull
|
||||
if (v != null) return v
|
||||
}
|
||||
return 0L
|
||||
|
|
@ -236,7 +226,7 @@ sealed class GatewayEvent {
|
|||
|
||||
fun getInt(vararg keys: String): Int {
|
||||
for (k in keys) {
|
||||
val v = (dataObj[k]?.jsonPrimitive ?: root[k]?.jsonPrimitive)?.intOrNull
|
||||
val v = payloadObj[k]?.jsonPrimitive?.intOrNull
|
||||
if (v != null) return v
|
||||
}
|
||||
return 0
|
||||
|
|
@ -244,19 +234,17 @@ sealed class GatewayEvent {
|
|||
|
||||
fun getBoolean(vararg keys: String): Boolean {
|
||||
for (k in keys) {
|
||||
val v = (dataObj[k]?.jsonPrimitive ?: root[k]?.jsonPrimitive)?.booleanOrNull
|
||||
val v = payloadObj[k]?.jsonPrimitive?.booleanOrNull
|
||||
if (v != null) return v
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
fun getStringList(key: String): List<String> {
|
||||
val array = (dataObj[key] ?: root[key]) as? JsonArray ?: return emptyList()
|
||||
val array = payloadObj[key] as? JsonArray ?: return emptyList()
|
||||
return array.mapNotNull { it.jsonPrimitive.content }
|
||||
}
|
||||
|
||||
val sessionKey = getNullableString("session_id", "session_key", "sessionKey", "sessionId")
|
||||
|
||||
return when (eventType) {
|
||||
"gateway.ready" -> GatewayReadyEvent(
|
||||
version = getString("version", "server_version"),
|
||||
|
|
@ -266,69 +254,69 @@ sealed class GatewayEvent {
|
|||
"message.start" -> MessageStartEvent(
|
||||
messageId = getString("message_id", "id"),
|
||||
role = getString("role").ifEmpty { "assistant" },
|
||||
sessionId = sessionKey,
|
||||
sessionId = sessionId,
|
||||
rawPayload = root
|
||||
)
|
||||
"message.delta" -> MessageDeltaEvent(
|
||||
messageId = getString("message_id", "id"),
|
||||
delta = getString("delta", "text", "chunk"),
|
||||
sessionId = sessionKey,
|
||||
sessionId = sessionId,
|
||||
rawPayload = root
|
||||
)
|
||||
"message.interim" -> MessageInterimEvent(
|
||||
messageId = getString("message_id", "id"),
|
||||
content = getString("content", "text"),
|
||||
sessionId = sessionKey,
|
||||
sessionId = sessionId,
|
||||
rawPayload = root
|
||||
)
|
||||
"message.complete" -> MessageCompleteEvent(
|
||||
messageId = getString("message_id", "id"),
|
||||
content = getString("content", "text"),
|
||||
sessionId = sessionKey,
|
||||
sessionId = sessionId,
|
||||
rawPayload = root
|
||||
)
|
||||
"thinking.delta" -> ThinkingDeltaEvent(
|
||||
messageId = getString("message_id", "id"),
|
||||
delta = getString("delta", "text", "chunk"),
|
||||
sessionId = sessionKey,
|
||||
sessionId = sessionId,
|
||||
rawPayload = root
|
||||
)
|
||||
"reasoning.delta" -> ReasoningDeltaEvent(
|
||||
messageId = getString("message_id", "id"),
|
||||
delta = getString("delta", "text", "chunk"),
|
||||
sessionId = sessionKey,
|
||||
sessionId = sessionId,
|
||||
rawPayload = root
|
||||
)
|
||||
"reasoning.available" -> ReasoningAvailableEvent(
|
||||
messageId = getString("message_id", "id"),
|
||||
reasoning = getString("reasoning", "content"),
|
||||
sessionId = sessionKey,
|
||||
sessionId = sessionId,
|
||||
rawPayload = root
|
||||
)
|
||||
"tool.start" -> ToolStartEvent(
|
||||
toolId = getString("tool_id", "id"),
|
||||
name = getString("name", "tool_name"),
|
||||
input = dataObj["input"] ?: root["input"],
|
||||
sessionId = sessionKey,
|
||||
input = payloadObj["input"],
|
||||
sessionId = sessionId,
|
||||
rawPayload = root
|
||||
)
|
||||
"tool.progress" -> ToolProgressEvent(
|
||||
toolId = getString("tool_id", "id"),
|
||||
progress = getString("progress", "message"),
|
||||
sessionId = sessionKey,
|
||||
sessionId = sessionId,
|
||||
rawPayload = root
|
||||
)
|
||||
"tool.generating" -> ToolGeneratingEvent(
|
||||
toolId = getString("tool_id", "id"),
|
||||
name = getString("name", "tool_name"),
|
||||
sessionId = sessionKey,
|
||||
sessionId = sessionId,
|
||||
rawPayload = root
|
||||
)
|
||||
"tool.complete" -> ToolCompleteEvent(
|
||||
toolId = getString("tool_id", "id"),
|
||||
result = getString("result", "output"),
|
||||
isError = getBoolean("is_error", "error"),
|
||||
sessionId = sessionKey,
|
||||
sessionId = sessionId,
|
||||
rawPayload = root
|
||||
)
|
||||
"approval.request" -> {
|
||||
|
|
@ -338,8 +326,8 @@ sealed class GatewayEvent {
|
|||
command = getNullableString("command"),
|
||||
description = getNullableString("description", "prompt"),
|
||||
choices = if (choices.isNotEmpty()) choices else listOf("once", "deny"),
|
||||
sessionKey = sessionKey,
|
||||
sessionId = sessionKey,
|
||||
sessionKey = sessionId,
|
||||
sessionId = sessionId,
|
||||
rawPayload = root
|
||||
)
|
||||
}
|
||||
|
|
@ -348,32 +336,32 @@ sealed class GatewayEvent {
|
|||
questionId = getNullableString("question_id", "questionId"),
|
||||
question = getString("question", "prompt"),
|
||||
promptType = ClarifyType.CLARIFY,
|
||||
sessionId = sessionKey,
|
||||
sessionId = sessionId,
|
||||
rawPayload = root
|
||||
)
|
||||
"sudo.request" -> SudoRequestEvent(
|
||||
requestId = getString("request_id", "id"),
|
||||
question = getString("question", "prompt").ifEmpty { "Administrator password required:" },
|
||||
sessionId = sessionKey,
|
||||
sessionId = sessionId,
|
||||
rawPayload = root
|
||||
)
|
||||
"secret.request" -> SecretRequestEvent(
|
||||
requestId = getString("request_id", "id"),
|
||||
question = getString("question", "prompt").ifEmpty { "Secret / Token required:" },
|
||||
sessionId = sessionKey,
|
||||
sessionId = sessionId,
|
||||
rawPayload = root
|
||||
)
|
||||
"status.update" -> StatusUpdateEvent(
|
||||
status = getString("status"),
|
||||
message = getNullableString("message"),
|
||||
sessionId = sessionKey,
|
||||
sessionId = sessionId,
|
||||
rawPayload = root
|
||||
)
|
||||
"session.usage" -> SessionUsageEvent(
|
||||
inputTokens = getLong("input_tokens", "prompt_tokens"),
|
||||
outputTokens = getLong("output_tokens", "completion_tokens"),
|
||||
totalTokens = getLong("total_tokens"),
|
||||
sessionId = sessionKey,
|
||||
sessionId = sessionId,
|
||||
rawPayload = root
|
||||
)
|
||||
"session.info" -> SessionInfoEvent(
|
||||
|
|
@ -384,24 +372,24 @@ sealed class GatewayEvent {
|
|||
branch = getNullableString("branch"),
|
||||
project = getNullableString("project")
|
||||
),
|
||||
sessionId = sessionKey,
|
||||
sessionId = sessionId,
|
||||
rawPayload = root
|
||||
)
|
||||
"background.complete" -> BackgroundCompleteEvent(
|
||||
taskId = getString("task_id", "id"),
|
||||
result = getNullableString("result"),
|
||||
sessionId = sessionKey,
|
||||
sessionId = sessionId,
|
||||
rawPayload = root
|
||||
)
|
||||
"error" -> ErrorEvent(
|
||||
code = getInt("code"),
|
||||
message = getString("message").ifEmpty { "Unknown error" },
|
||||
sessionId = sessionKey,
|
||||
sessionId = sessionId,
|
||||
rawPayload = root
|
||||
)
|
||||
else -> UnknownGatewayEvent(
|
||||
eventType = eventType.ifEmpty { "unknown" },
|
||||
sessionId = sessionKey,
|
||||
sessionId = sessionId,
|
||||
rawPayload = root
|
||||
)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -5,6 +5,8 @@ import kotlinx.serialization.json.JsonElement
|
|||
import kotlinx.serialization.json.JsonObject
|
||||
import kotlinx.serialization.json.buildJsonObject
|
||||
|
||||
import java.io.IOException
|
||||
|
||||
@Serializable
|
||||
data class JsonRpcRequest(
|
||||
val jsonrpc: String = "2.0",
|
||||
|
|
@ -27,3 +29,10 @@ data class JsonRpcError(
|
|||
val message: String,
|
||||
val data: JsonElement? = null
|
||||
)
|
||||
|
||||
open class JsonRpcException(
|
||||
val code: Int,
|
||||
val errorMessage: String,
|
||||
val data: JsonElement? = null
|
||||
) : IOException("RPC Error [$code]: $errorMessage")
|
||||
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ import app.hermes.mobile.core.model.CreateSessionResult
|
|||
import app.hermes.mobile.core.model.DurableSessionId
|
||||
import app.hermes.mobile.core.model.GatewayEvent
|
||||
import app.hermes.mobile.core.model.JsonRpcError
|
||||
import app.hermes.mobile.core.model.JsonRpcException
|
||||
import app.hermes.mobile.core.model.JsonRpcRequest
|
||||
import app.hermes.mobile.core.model.JsonRpcResponse
|
||||
import app.hermes.mobile.core.model.PromptSubmitResult
|
||||
|
|
@ -76,7 +77,7 @@ class JsonRpcGatewayClient(
|
|||
private val _connectionState = MutableStateFlow<ConnectionState>(ConnectionState.Disconnected)
|
||||
val connectionState: StateFlow<ConnectionState> = _connectionState.asStateFlow()
|
||||
|
||||
private val _events = MutableSharedFlow<GatewayEvent>(extraBufferCapacity = 64)
|
||||
private val _events = MutableSharedFlow<GatewayEvent>(replay = 1, extraBufferCapacity = 64)
|
||||
val events: SharedFlow<GatewayEvent> = _events.asSharedFlow()
|
||||
|
||||
private fun nextId(): String = "a${reqCounter.incrementAndGet()}"
|
||||
|
|
@ -250,7 +251,7 @@ class JsonRpcGatewayClient(
|
|||
val params = buildJsonObject { put("limit", limit) }
|
||||
val response = sendRequest("session.list", params)
|
||||
if (response.error != null) {
|
||||
throw IOException("RPC Error [${response.error.code}]: ${response.error.message}")
|
||||
throw JsonRpcException(response.error.code, response.error.message, response.error.data)
|
||||
}
|
||||
val result = response.result ?: return emptyList()
|
||||
|
||||
|
|
@ -294,7 +295,7 @@ class JsonRpcGatewayClient(
|
|||
}
|
||||
val response = sendRequest("session.create", params)
|
||||
if (response.error != null) {
|
||||
throw IOException("RPC Error [${response.error.code}]: ${response.error.message}")
|
||||
throw JsonRpcException(response.error.code, response.error.message, response.error.data)
|
||||
}
|
||||
val result = response.result as? JsonObject
|
||||
?: throw IOException("Invalid response format for session.create")
|
||||
|
|
@ -324,7 +325,7 @@ class JsonRpcGatewayClient(
|
|||
}
|
||||
val response = sendRequest("session.resume", params)
|
||||
if (response.error != null) {
|
||||
throw IOException("RPC Error [${response.error.code}]: ${response.error.message}")
|
||||
throw JsonRpcException(response.error.code, response.error.message, response.error.data)
|
||||
}
|
||||
val result = response.result as? JsonObject
|
||||
?: throw IOException("Invalid response format for session.resume")
|
||||
|
|
@ -352,7 +353,7 @@ class JsonRpcGatewayClient(
|
|||
}
|
||||
val response = sendRequest("prompt.submit", params)
|
||||
if (response.error != null) {
|
||||
throw IOException("RPC Error [${response.error.code}]: ${response.error.message}")
|
||||
throw JsonRpcException(response.error.code, response.error.message, response.error.data)
|
||||
}
|
||||
val result = response.result as? JsonObject
|
||||
val turnId = result?.get("turn_id")?.jsonPrimitive?.content
|
||||
|
|
|
|||
|
|
@ -44,8 +44,8 @@ class UnifiedSessionRepository(
|
|||
private val _activeClarify = MutableStateFlow<HostAttributedClarify?>(null)
|
||||
val activeClarify: StateFlow<HostAttributedClarify?> = _activeClarify.asStateFlow()
|
||||
|
||||
// Mapping from runtimeSessionId to (sessionId, hostId)
|
||||
private val runtimeToSessionMap = ConcurrentHashMap<String, Pair<UnifiedSessionId, HermesHostId>>()
|
||||
// Mapping from (hostId, runtimeSessionId) to sessionId
|
||||
private val runtimeToSessionMap = ConcurrentHashMap<Pair<HermesHostId, String>, UnifiedSessionId>()
|
||||
|
||||
// In-memory active session messages cache for reactive streaming updates
|
||||
private val sessionMessagesState = ConcurrentHashMap<UnifiedSessionId, MutableStateFlow<List<UnifiedMessage>>>()
|
||||
|
|
@ -55,17 +55,6 @@ class UnifiedSessionRepository(
|
|||
private val sessionExecutingState = ConcurrentHashMap<UnifiedSessionId, MutableStateFlow<Boolean>>()
|
||||
|
||||
init {
|
||||
scope.launch {
|
||||
sessions.collect { list ->
|
||||
for (s in list) {
|
||||
for ((hId, b) in s.bindings) {
|
||||
if (b.runtimeSessionId.value.isNotEmpty()) {
|
||||
runtimeToSessionMap[b.runtimeSessionId.value] = Pair(s.id, hId)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
scope.launch {
|
||||
connectionManager.allEvents.collect { hostEvent ->
|
||||
handleHostGatewayEvent(hostEvent)
|
||||
|
|
@ -142,12 +131,12 @@ class UnifiedSessionRepository(
|
|||
sessionMessagesState.remove(sessionId)
|
||||
sessionExecutingState.remove(sessionId)
|
||||
hostExecutingState.entries.removeIf { it.key.first == sessionId }
|
||||
runtimeToSessionMap.entries.removeIf { it.value.first == sessionId }
|
||||
runtimeToSessionMap.entries.removeIf { it.value == sessionId }
|
||||
}
|
||||
|
||||
fun registerRuntimeBinding(sessionId: UnifiedSessionId, hostId: HermesHostId, runtimeSessionId: RuntimeSessionId) {
|
||||
if (runtimeSessionId.value.isNotEmpty()) {
|
||||
runtimeToSessionMap[runtimeSessionId.value] = Pair(sessionId, hostId)
|
||||
runtimeToSessionMap[Pair(hostId, runtimeSessionId.value)] = sessionId
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -175,21 +164,26 @@ class UnifiedSessionRepository(
|
|||
syncedAt = null
|
||||
)
|
||||
sessionDao.insertOrUpdateBinding(binding.toEntity(sessionId.value))
|
||||
runtimeToSessionMap[createRes.runtimeId.value] = Pair(sessionId, targetHostId)
|
||||
runtimeToSessionMap[Pair(targetHostId, createRes.runtimeId.value)] = sessionId
|
||||
return binding
|
||||
}
|
||||
|
||||
// We have an existing durableSessionId.
|
||||
// Check if current runtimeSessionId is already registered and valid, or if we need to resume
|
||||
// Check if current runtimeSessionId is already registered and valid in memory in the current process
|
||||
val currentRuntimeId = binding.runtimeSessionId.value
|
||||
val isRegistered = currentRuntimeId.isNotEmpty() && runtimeToSessionMap.containsKey(currentRuntimeId)
|
||||
val isRegistered = currentRuntimeId.isNotEmpty() && runtimeToSessionMap.containsKey(Pair(targetHostId, currentRuntimeId))
|
||||
|
||||
if (!isRegistered || binding.state == BindingState.NOT_CREATED || binding.state == BindingState.OFFLINE || binding.state == BindingState.ERROR) {
|
||||
val resumeRes = try {
|
||||
runtime.gatewayClient.resumeSession(binding.durableSessionId, source = "android")
|
||||
} catch (_: Exception) {
|
||||
val createRes = runtime.gatewayClient.createSession(source = "android")
|
||||
ResumeSessionResult(createRes.durableId, createRes.runtimeId)
|
||||
} catch (e: Exception) {
|
||||
if (isDefinitivelyMissingSession(e)) {
|
||||
val createRes = runtime.gatewayClient.createSession(source = "android")
|
||||
ResumeSessionResult(createRes.durableId, createRes.runtimeId)
|
||||
} else {
|
||||
// For transient errors during resume (timeout, network, auth), throw without creating a new session or destroying the binding
|
||||
throw e
|
||||
}
|
||||
}
|
||||
binding = binding.copy(
|
||||
durableSessionId = resumeRes.durableId,
|
||||
|
|
@ -198,14 +192,31 @@ class UnifiedSessionRepository(
|
|||
state = BindingState.READY
|
||||
)
|
||||
sessionDao.insertOrUpdateBinding(binding.toEntity(sessionId.value))
|
||||
runtimeToSessionMap[resumeRes.runtimeId.value] = Pair(sessionId, targetHostId)
|
||||
runtimeToSessionMap[Pair(targetHostId, resumeRes.runtimeId.value)] = sessionId
|
||||
} else {
|
||||
runtimeToSessionMap[currentRuntimeId] = Pair(sessionId, targetHostId)
|
||||
runtimeToSessionMap[Pair(targetHostId, currentRuntimeId)] = sessionId
|
||||
}
|
||||
|
||||
return binding
|
||||
}
|
||||
|
||||
private fun isDefinitivelyMissingSession(e: Throwable): Boolean {
|
||||
if (e is JsonRpcException) {
|
||||
if (e.code == 404 || e.code == -32004) return true
|
||||
val msg = e.errorMessage.lowercase()
|
||||
if (msg.contains("not found") || msg.contains("does not exist") ||
|
||||
msg.contains("invalid session") || msg.contains("no such session") ||
|
||||
msg.contains("session destroyed") || msg.contains("unrecoverable")) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
val msg = e.message?.lowercase() ?: ""
|
||||
if (msg.contains("404") || msg.contains("session not found") || msg.contains("session does not exist")) {
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
suspend fun sendPrompt(sessionId: UnifiedSessionId, text: String): String {
|
||||
val details = sessionDao.getSessionWithDetails(sessionId.value)
|
||||
?: throw IllegalArgumentException("Session not found: ${sessionId.value}")
|
||||
|
|
@ -452,18 +463,11 @@ class UnifiedSessionRepository(
|
|||
messageId: String? = null,
|
||||
toolId: String? = null
|
||||
): UnifiedSessionId? {
|
||||
// 1. Exact match via runtimeSessionId in runtimeToSessionMap
|
||||
// 1. Exact match via (hostId, runtimeSessionId) in runtimeToSessionMap
|
||||
if (!sessionIdFromEvent.isNullOrEmpty()) {
|
||||
val mapped = runtimeToSessionMap[sessionIdFromEvent]
|
||||
if (mapped != null && mapped.second == hostId) {
|
||||
return mapped.first
|
||||
}
|
||||
for (session in sessions.value) {
|
||||
val b = session.bindings[hostId]
|
||||
if (b != null && b.runtimeSessionId.value == sessionIdFromEvent) {
|
||||
runtimeToSessionMap[sessionIdFromEvent] = Pair(session.id, hostId)
|
||||
return session.id
|
||||
}
|
||||
val mapped = runtimeToSessionMap[Pair(hostId, sessionIdFromEvent)]
|
||||
if (mapped != null) {
|
||||
return mapped
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -44,7 +44,7 @@ class HermesConnectionManager(
|
|||
private val _activeHostId = MutableStateFlow<HermesHostId?>(null)
|
||||
val activeHostId: StateFlow<HermesHostId?> = _activeHostId.asStateFlow()
|
||||
|
||||
private val _allEvents = MutableSharedFlow<HostGatewayEvent>(extraBufferCapacity = 128)
|
||||
private val _allEvents = MutableSharedFlow<HostGatewayEvent>(replay = 1, extraBufferCapacity = 128)
|
||||
val allEvents: SharedFlow<HostGatewayEvent> = _allEvents.asSharedFlow()
|
||||
|
||||
init {
|
||||
|
|
|
|||
|
|
@ -89,7 +89,7 @@ class EndToEndContractScenarioTest {
|
|||
serverWs = webSocket
|
||||
wsConnectedLatch.countDown()
|
||||
// Send gateway.ready
|
||||
webSocket.send("""{"event":"gateway.ready","data":{"version":"1.0.0","session_count":1}}""")
|
||||
webSocket.send("""{"jsonrpc":"2.0","method":"event","params":{"type":"gateway.ready","payload":{"version":"1.0.0","session_count":1}}}""")
|
||||
}
|
||||
|
||||
override fun onMessage(webSocket: WebSocket, text: String) {
|
||||
|
|
@ -100,16 +100,16 @@ class EndToEndContractScenarioTest {
|
|||
} else if (text.contains("prompt.submit")) {
|
||||
webSocket.send("""{"jsonrpc":"2.0","id":"a3","result":{"turn_id":"turn_001"}}""")
|
||||
// Emit streaming events
|
||||
webSocket.send("""{"event":"message.start","data":{"message_id":"msg_resp_1","role":"assistant"}}""")
|
||||
webSocket.send("""{"event":"message.delta","data":{"message_id":"msg_resp_1","delta":"Sure, I can "}}""")
|
||||
webSocket.send("""{"event":"message.delta","data":{"message_id":"msg_resp_1","delta":"run that tool."}}""")
|
||||
webSocket.send("""{"event":"tool.start","data":{"tool_id":"t_exec","name":"run_command"}}""")
|
||||
webSocket.send("""{"event":"tool.progress","data":{"tool_id":"t_exec","progress":"Executing ls..."}}""")
|
||||
webSocket.send("""{"event":"tool.complete","data":{"tool_id":"t_exec","result":"file1.txt\nfile2.txt","is_error":false}}""")
|
||||
webSocket.send("""{"event":"approval.request","data":{"request_id":"app_req_1","command":"git status","description":"Run git status","choices":["once","deny"]}}""")
|
||||
webSocket.send("""{"jsonrpc":"2.0","method":"event","params":{"type":"message.start","session_id":"runtime_202","payload":{"message_id":"msg_resp_1","role":"assistant"}}}""")
|
||||
webSocket.send("""{"jsonrpc":"2.0","method":"event","params":{"type":"message.delta","session_id":"runtime_202","payload":{"message_id":"msg_resp_1","delta":"Sure, I can "}}}""")
|
||||
webSocket.send("""{"jsonrpc":"2.0","method":"event","params":{"type":"message.delta","session_id":"runtime_202","payload":{"message_id":"msg_resp_1","delta":"run that tool."}}}""")
|
||||
webSocket.send("""{"jsonrpc":"2.0","method":"event","params":{"type":"tool.start","session_id":"runtime_202","payload":{"tool_id":"t_exec","name":"run_command"}}}""")
|
||||
webSocket.send("""{"jsonrpc":"2.0","method":"event","params":{"type":"tool.progress","session_id":"runtime_202","payload":{"tool_id":"t_exec","progress":"Executing ls..."}}}""")
|
||||
webSocket.send("""{"jsonrpc":"2.0","method":"event","params":{"type":"tool.complete","session_id":"runtime_202","payload":{"tool_id":"t_exec","result":"file1.txt\nfile2.txt","is_error":false}}}""")
|
||||
webSocket.send("""{"jsonrpc":"2.0","method":"event","params":{"type":"approval.request","session_id":"runtime_202","payload":{"request_id":"app_req_1","command":"git status","description":"Run git status","choices":["once","deny"]}}}""")
|
||||
} else if (text.contains("approval.respond")) {
|
||||
webSocket.send("""{"jsonrpc":"2.0","id":"a4","result":{"accepted":true}}""")
|
||||
webSocket.send("""{"event":"message.complete","data":{"message_id":"msg_resp_1","content":"Sure, I can run that tool. Done!"}}""")
|
||||
webSocket.send("""{"jsonrpc":"2.0","method":"event","params":{"type":"message.complete","session_id":"runtime_202","payload":{"message_id":"msg_resp_1","content":"Sure, I can run that tool. Done!"}}}""")
|
||||
}
|
||||
}
|
||||
})
|
||||
|
|
@ -239,7 +239,7 @@ class EndToEndContractScenarioTest {
|
|||
path.startsWith("/api/ws") || path.startsWith("/ws") -> {
|
||||
MockResponse().withWebSocketUpgrade(object : WebSocketListener() {
|
||||
override fun onOpen(webSocket: WebSocket, response: Response) {
|
||||
webSocket.send("""{"event":"gateway.ready","data":{"version":"1.0.0","session_count":0}}""")
|
||||
webSocket.send("""{"jsonrpc":"2.0","method":"event","params":{"type":"gateway.ready","payload":{"version":"1.0.0","session_count":0}}}""")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
|
|
|||
|
|
@ -47,7 +47,7 @@ class JsonRpcGatewayClientTest {
|
|||
MockResponse().withWebSocketUpgrade(object : WebSocketListener() {
|
||||
override fun onOpen(webSocket: WebSocket, response: Response) {
|
||||
serverWebSocket = webSocket
|
||||
webSocket.send("""{"event":"gateway.ready","data":{"version":"1.0.0","session_count":0}}""")
|
||||
webSocket.send("""{"jsonrpc":"2.0","method":"event","params":{"type":"gateway.ready","payload":{"version":"1.0.0","session_count":0}}}""")
|
||||
}
|
||||
|
||||
override fun onMessage(webSocket: WebSocket, text: String) {
|
||||
|
|
@ -98,7 +98,7 @@ class JsonRpcGatewayClientTest {
|
|||
assertEquals(ConnectionState.Connecting, client.connectionState.value)
|
||||
|
||||
// Send gateway.ready
|
||||
serverWebSocket?.send("""{"event":"gateway.ready","data":{"version":"1.0.0","session_count":1}}""")
|
||||
serverWebSocket?.send("""{"jsonrpc":"2.0","method":"event","params":{"type":"gateway.ready","payload":{"version":"1.0.0","session_count":1}}}""")
|
||||
|
||||
client.awaitGatewayReady(5000)
|
||||
assertEquals(ConnectionState.Connected, client.connectionState.value)
|
||||
|
|
@ -115,23 +115,26 @@ class JsonRpcGatewayClientTest {
|
|||
MockResponse().withWebSocketUpgrade(object : WebSocketListener() {
|
||||
override fun onOpen(webSocket: WebSocket, response: Response) {
|
||||
serverWebSocket = webSocket
|
||||
webSocket.send("""{"jsonrpc":"2.0","method":"event","params":{"type":"gateway.ready","payload":{"version":"1.0.0"}}}""")
|
||||
// Send an incoming server notification/event
|
||||
webSocket.send("""{"event":"message.delta","data":{"message_id":"m100","delta":"Streaming token"}}""")
|
||||
webSocket.send("""{"jsonrpc":"2.0","method":"event","params":{"type":"message.delta","session_id":"rt_test","payload":{"message_id":"m100","delta":"Streaming token"}}}""")
|
||||
}
|
||||
})
|
||||
)
|
||||
|
||||
val wsUrl = "ws://${server.hostName}:${server.port}/api/ws"
|
||||
client.connect(wsUrl, allowCleartext = true)
|
||||
client.awaitGatewayReady(5000)
|
||||
|
||||
val event = withTimeout(5000) {
|
||||
client.events.first()
|
||||
client.events.first { it is GatewayEvent.MessageDeltaEvent }
|
||||
}
|
||||
|
||||
assertTrue(event is GatewayEvent.MessageDeltaEvent)
|
||||
val deltaEvent = event as GatewayEvent.MessageDeltaEvent
|
||||
assertEquals("m100", deltaEvent.messageId)
|
||||
assertEquals("Streaming token", deltaEvent.delta)
|
||||
assertEquals("rt_test", deltaEvent.sessionId)
|
||||
|
||||
serverWebSocket?.close(1000, "done")
|
||||
client.disconnect()
|
||||
|
|
@ -146,7 +149,7 @@ class JsonRpcGatewayClientTest {
|
|||
MockResponse().withWebSocketUpgrade(object : WebSocketListener() {
|
||||
override fun onOpen(webSocket: WebSocket, response: Response) {
|
||||
serverWebSocket = webSocket
|
||||
webSocket.send("""{"event":"gateway.ready","data":{"version":"1.0.0","session_count":0}}""")
|
||||
webSocket.send("""{"jsonrpc":"2.0","method":"event","params":{"type":"gateway.ready","payload":{"version":"1.0.0","session_count":0}}}""")
|
||||
}
|
||||
|
||||
override fun onMessage(webSocket: WebSocket, text: String) {
|
||||
|
|
|
|||
|
|
@ -82,81 +82,309 @@ class JsonRpcWireFormatTest {
|
|||
@Test
|
||||
fun testAllGatewayEventParsers() {
|
||||
// 1. Gateway Ready
|
||||
val readyJson = json.decodeFromString<JsonObject>("""{"event":"gateway.ready","data":{"version":"1.0.0","session_count":3}}""")
|
||||
val readyJson = json.decodeFromString<JsonObject>("""{"jsonrpc":"2.0","method":"event","params":{"type":"gateway.ready","payload":{"version":"1.0.0","session_count":3}}}""")
|
||||
val readyEvent = GatewayEvent.parse(readyJson) as GatewayEvent.GatewayReadyEvent
|
||||
assertEquals("1.0.0", readyEvent.version)
|
||||
assertEquals(3, readyEvent.sessionCount)
|
||||
|
||||
// 2. Message Start
|
||||
val msgStartJson = json.decodeFromString<JsonObject>("""{"event":"message.start","data":{"message_id":"msg_1","role":"assistant"}}""")
|
||||
val msgStartJson = json.decodeFromString<JsonObject>("""{"jsonrpc":"2.0","method":"event","params":{"type":"message.start","session_id":"rt_test","payload":{"message_id":"msg_1","role":"assistant"}}}""")
|
||||
val msgStart = GatewayEvent.parse(msgStartJson) as GatewayEvent.MessageStartEvent
|
||||
assertEquals("msg_1", msgStart.messageId)
|
||||
assertEquals("assistant", msgStart.role)
|
||||
assertEquals("rt_test", msgStart.sessionId)
|
||||
|
||||
// 3. Message Delta
|
||||
val msgDeltaJson = json.decodeFromString<JsonObject>("""{"event":"message.delta","data":{"message_id":"msg_1","delta":"Hello world"}}""")
|
||||
val msgDeltaJson = json.decodeFromString<JsonObject>("""{"jsonrpc":"2.0","method":"event","params":{"type":"message.delta","session_id":"rt_test","payload":{"message_id":"msg_1","delta":"Hello world"}}}""")
|
||||
val msgDelta = GatewayEvent.parse(msgDeltaJson) as GatewayEvent.MessageDeltaEvent
|
||||
assertEquals("msg_1", msgDelta.messageId)
|
||||
assertEquals("Hello world", msgDelta.delta)
|
||||
assertEquals("rt_test", msgDelta.sessionId)
|
||||
|
||||
// 4. Message Complete
|
||||
val msgCompleteJson = json.decodeFromString<JsonObject>("""{"event":"message.complete","data":{"message_id":"msg_1","content":"Final answer"}}""")
|
||||
val msgCompleteJson = json.decodeFromString<JsonObject>("""{"jsonrpc":"2.0","method":"event","params":{"type":"message.complete","session_id":"rt_test","payload":{"message_id":"msg_1","content":"Final answer"}}}""")
|
||||
val msgComplete = GatewayEvent.parse(msgCompleteJson) as GatewayEvent.MessageCompleteEvent
|
||||
assertEquals("msg_1", msgComplete.messageId)
|
||||
assertEquals("Final answer", msgComplete.content)
|
||||
assertEquals("rt_test", msgComplete.sessionId)
|
||||
|
||||
// 5. Thinking Delta
|
||||
val thinkJson = json.decodeFromString<JsonObject>("""{"event":"thinking.delta","data":{"message_id":"msg_1","delta":"Analyzing requirements..."}}""")
|
||||
val thinkJson = json.decodeFromString<JsonObject>("""{"jsonrpc":"2.0","method":"event","params":{"type":"thinking.delta","session_id":"rt_test","payload":{"message_id":"msg_1","delta":"Analyzing requirements..."}}}""")
|
||||
val think = GatewayEvent.parse(thinkJson) as GatewayEvent.ThinkingDeltaEvent
|
||||
assertEquals("Analyzing requirements...", think.delta)
|
||||
assertEquals("rt_test", think.sessionId)
|
||||
|
||||
// 6. Tool Lifecycle
|
||||
val toolStartJson = json.decodeFromString<JsonObject>("""{"event":"tool.start","data":{"tool_id":"t1","name":"exec_command"}}""")
|
||||
val toolStartJson = json.decodeFromString<JsonObject>("""{"jsonrpc":"2.0","method":"event","params":{"type":"tool.start","session_id":"rt_test","payload":{"tool_id":"t1","name":"exec_command"}}}""")
|
||||
val toolStart = GatewayEvent.parse(toolStartJson) as GatewayEvent.ToolStartEvent
|
||||
assertEquals("t1", toolStart.toolId)
|
||||
assertEquals("exec_command", toolStart.name)
|
||||
assertEquals("rt_test", toolStart.sessionId)
|
||||
|
||||
val toolProgressJson = json.decodeFromString<JsonObject>("""{"event":"tool.progress","data":{"tool_id":"t1","progress":"Running build..."}}""")
|
||||
val toolProgressJson = json.decodeFromString<JsonObject>("""{"jsonrpc":"2.0","method":"event","params":{"type":"tool.progress","session_id":"rt_test","payload":{"tool_id":"t1","progress":"Running build..."}}}""")
|
||||
val toolProgress = GatewayEvent.parse(toolProgressJson) as GatewayEvent.ToolProgressEvent
|
||||
assertEquals("Running build...", toolProgress.progress)
|
||||
assertEquals("rt_test", toolProgress.sessionId)
|
||||
|
||||
val toolCompleteJson = json.decodeFromString<JsonObject>("""{"event":"tool.complete","data":{"tool_id":"t1","result":"Success","is_error":false}}""")
|
||||
val toolCompleteJson = json.decodeFromString<JsonObject>("""{"jsonrpc":"2.0","method":"event","params":{"type":"tool.complete","session_id":"rt_test","payload":{"tool_id":"t1","result":"Success","is_error":false}}}""")
|
||||
val toolComplete = GatewayEvent.parse(toolCompleteJson) as GatewayEvent.ToolCompleteEvent
|
||||
assertEquals("Success", toolComplete.result)
|
||||
assertEquals(false, toolComplete.isError)
|
||||
assertEquals("rt_test", toolComplete.sessionId)
|
||||
|
||||
// 7. Approval Request
|
||||
val approvalJson = json.decodeFromString<JsonObject>("""{"event":"approval.request","data":{"request_id":"req_app","command":"rm -rf /tmp/cache","description":"Clear cache directory","choices":["once","deny"]}}""")
|
||||
val approvalJson = json.decodeFromString<JsonObject>("""{"jsonrpc":"2.0","method":"event","params":{"type":"approval.request","session_id":"rt_test","payload":{"request_id":"req_app","command":"rm -rf /tmp/cache","description":"Clear cache directory","choices":["once","deny"]}}}""")
|
||||
val approval = GatewayEvent.parse(approvalJson) as GatewayEvent.ApprovalRequestEvent
|
||||
assertEquals("req_app", approval.requestId)
|
||||
assertEquals("rm -rf /tmp/cache", approval.command)
|
||||
assertEquals(2, approval.choices.size)
|
||||
assertEquals("rt_test", approval.sessionId)
|
||||
|
||||
// 8. Clarify, Sudo, Secret
|
||||
val clarifyJson = json.decodeFromString<JsonObject>("""{"event":"clarify.request","data":{"request_id":"c1","question":"Which port?"}}""")
|
||||
val clarifyJson = json.decodeFromString<JsonObject>("""{"jsonrpc":"2.0","method":"event","params":{"type":"clarify.request","session_id":"rt_test","payload":{"request_id":"c1","question":"Which port?"}}}""")
|
||||
val clarify = GatewayEvent.parse(clarifyJson) as GatewayEvent.ClarifyRequestEvent
|
||||
assertEquals("Which port?", clarify.question)
|
||||
assertEquals(ClarifyType.CLARIFY, clarify.promptType)
|
||||
assertEquals("rt_test", clarify.sessionId)
|
||||
|
||||
val sudoJson = json.decodeFromString<JsonObject>("""{"event":"sudo.request","data":{"request_id":"s1","question":"Root password required:"}}""")
|
||||
val sudoJson = json.decodeFromString<JsonObject>("""{"jsonrpc":"2.0","method":"event","params":{"type":"sudo.request","session_id":"rt_test","payload":{"request_id":"s1","question":"Root password required:"}}}""")
|
||||
val sudo = GatewayEvent.parse(sudoJson) as GatewayEvent.SudoRequestEvent
|
||||
assertEquals("Root password required:", sudo.question)
|
||||
assertEquals("rt_test", sudo.sessionId)
|
||||
|
||||
val secretJson = json.decodeFromString<JsonObject>("""{"event":"secret.request","data":{"request_id":"sec1","question":"OpenAI API Key:"}}""")
|
||||
val secretJson = json.decodeFromString<JsonObject>("""{"jsonrpc":"2.0","method":"event","params":{"type":"secret.request","session_id":"rt_test","payload":{"request_id":"sec1","question":"OpenAI API Key:"}}}""")
|
||||
val secret = GatewayEvent.parse(secretJson) as GatewayEvent.SecretRequestEvent
|
||||
assertEquals("OpenAI API Key:", secret.question)
|
||||
assertEquals("rt_test", secret.sessionId)
|
||||
|
||||
// 9. Session Info & Usage
|
||||
val infoJson = json.decodeFromString<JsonObject>("""{"event":"session.info","data":{"model":"claude-3-5-sonnet","provider":"anthropic","branch":"main"}}""")
|
||||
val infoJson = json.decodeFromString<JsonObject>("""{"jsonrpc":"2.0","method":"event","params":{"type":"session.info","session_id":"rt_test","payload":{"model":"claude-3-5-sonnet","provider":"anthropic","branch":"main"}}}""")
|
||||
val info = GatewayEvent.parse(infoJson) as GatewayEvent.SessionInfoEvent
|
||||
assertEquals("claude-3-5-sonnet", info.info.model)
|
||||
assertEquals("main", info.info.branch)
|
||||
assertEquals("rt_test", info.sessionId)
|
||||
|
||||
val usageJson = json.decodeFromString<JsonObject>("""{"event":"session.usage","data":{"input_tokens":1200,"output_tokens":350,"total_tokens":1550}}""")
|
||||
val usageJson = json.decodeFromString<JsonObject>("""{"jsonrpc":"2.0","method":"event","params":{"type":"session.usage","session_id":"rt_test","payload":{"input_tokens":1200,"output_tokens":350,"total_tokens":1550}}}""")
|
||||
val usage = GatewayEvent.parse(usageJson) as GatewayEvent.SessionUsageEvent
|
||||
assertEquals(1200L, usage.inputTokens)
|
||||
assertEquals(350L, usage.outputTokens)
|
||||
assertEquals(1550L, usage.totalTokens)
|
||||
assertEquals("rt_test", usage.sessionId)
|
||||
}
|
||||
|
||||
@Test
|
||||
fun testUpstreamHermesContractMessageStart() {
|
||||
val raw = """
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
"method": "event",
|
||||
"params": {
|
||||
"type": "message.start",
|
||||
"session_id": "rt-session-001",
|
||||
"payload": {
|
||||
"message_id": "msg-start-001",
|
||||
"role": "assistant"
|
||||
}
|
||||
}
|
||||
}
|
||||
""".trimIndent()
|
||||
val root = json.decodeFromString<JsonObject>(raw)
|
||||
val event = GatewayEvent.parse(root)
|
||||
assertTrue(event is GatewayEvent.MessageStartEvent)
|
||||
val startEvent = event as GatewayEvent.MessageStartEvent
|
||||
assertEquals("msg-start-001", startEvent.messageId)
|
||||
assertEquals("assistant", startEvent.role)
|
||||
assertEquals("rt-session-001", startEvent.sessionId)
|
||||
}
|
||||
|
||||
@Test
|
||||
fun testUpstreamHermesContractMessageDelta() {
|
||||
val raw = """
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
"method": "event",
|
||||
"params": {
|
||||
"type": "message.delta",
|
||||
"session_id": "rt-session-001",
|
||||
"payload": {
|
||||
"message_id": "msg-delta-001",
|
||||
"delta": "Hello from Hermes"
|
||||
}
|
||||
}
|
||||
}
|
||||
""".trimIndent()
|
||||
val root = json.decodeFromString<JsonObject>(raw)
|
||||
val event = GatewayEvent.parse(root)
|
||||
assertTrue(event is GatewayEvent.MessageDeltaEvent)
|
||||
val deltaEvent = event as GatewayEvent.MessageDeltaEvent
|
||||
assertEquals("msg-delta-001", deltaEvent.messageId)
|
||||
assertEquals("Hello from Hermes", deltaEvent.delta)
|
||||
assertEquals("rt-session-001", deltaEvent.sessionId)
|
||||
}
|
||||
|
||||
@Test
|
||||
fun testUpstreamHermesContractMessageComplete() {
|
||||
val raw = """
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
"method": "event",
|
||||
"params": {
|
||||
"type": "message.complete",
|
||||
"session_id": "rt-session-001",
|
||||
"payload": {
|
||||
"message_id": "msg-complete-001",
|
||||
"content": "Execution completed successfully."
|
||||
}
|
||||
}
|
||||
}
|
||||
""".trimIndent()
|
||||
val root = json.decodeFromString<JsonObject>(raw)
|
||||
val event = GatewayEvent.parse(root)
|
||||
assertTrue(event is GatewayEvent.MessageCompleteEvent)
|
||||
val completeEvent = event as GatewayEvent.MessageCompleteEvent
|
||||
assertEquals("msg-complete-001", completeEvent.messageId)
|
||||
assertEquals("Execution completed successfully.", completeEvent.content)
|
||||
assertEquals("rt-session-001", completeEvent.sessionId)
|
||||
}
|
||||
|
||||
@Test
|
||||
fun testUpstreamHermesContractToolStart() {
|
||||
val raw = """
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
"method": "event",
|
||||
"params": {
|
||||
"type": "tool.start",
|
||||
"session_id": "rt-session-001",
|
||||
"payload": {
|
||||
"tool_id": "tool-call-101",
|
||||
"name": "bash_execution",
|
||||
"input": {
|
||||
"command": "ls -la"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
""".trimIndent()
|
||||
val root = json.decodeFromString<JsonObject>(raw)
|
||||
val event = GatewayEvent.parse(root)
|
||||
assertTrue(event is GatewayEvent.ToolStartEvent)
|
||||
val toolStart = event as GatewayEvent.ToolStartEvent
|
||||
assertEquals("tool-call-101", toolStart.toolId)
|
||||
assertEquals("bash_execution", toolStart.name)
|
||||
assertNotNull(toolStart.input)
|
||||
assertEquals("rt-session-001", toolStart.sessionId)
|
||||
}
|
||||
|
||||
@Test
|
||||
fun testUpstreamHermesContractToolProgress() {
|
||||
val raw = """
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
"method": "event",
|
||||
"params": {
|
||||
"type": "tool.progress",
|
||||
"session_id": "rt-session-001",
|
||||
"payload": {
|
||||
"tool_id": "tool-call-101",
|
||||
"progress": "Downloading dependencies..."
|
||||
}
|
||||
}
|
||||
}
|
||||
""".trimIndent()
|
||||
val root = json.decodeFromString<JsonObject>(raw)
|
||||
val event = GatewayEvent.parse(root)
|
||||
assertTrue(event is GatewayEvent.ToolProgressEvent)
|
||||
val toolProgress = event as GatewayEvent.ToolProgressEvent
|
||||
assertEquals("tool-call-101", toolProgress.toolId)
|
||||
assertEquals("Downloading dependencies...", toolProgress.progress)
|
||||
assertEquals("rt-session-001", toolProgress.sessionId)
|
||||
}
|
||||
|
||||
@Test
|
||||
fun testUpstreamHermesContractToolComplete() {
|
||||
val raw = """
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
"method": "event",
|
||||
"params": {
|
||||
"type": "tool.complete",
|
||||
"session_id": "rt-session-001",
|
||||
"payload": {
|
||||
"tool_id": "tool-call-101",
|
||||
"result": "Total 14 files found",
|
||||
"is_error": false
|
||||
}
|
||||
}
|
||||
}
|
||||
""".trimIndent()
|
||||
val root = json.decodeFromString<JsonObject>(raw)
|
||||
val event = GatewayEvent.parse(root)
|
||||
assertTrue(event is GatewayEvent.ToolCompleteEvent)
|
||||
val toolComplete = event as GatewayEvent.ToolCompleteEvent
|
||||
assertEquals("tool-call-101", toolComplete.toolId)
|
||||
assertEquals("Total 14 files found", toolComplete.result)
|
||||
assertEquals(false, toolComplete.isError)
|
||||
assertEquals("rt-session-001", toolComplete.sessionId)
|
||||
}
|
||||
|
||||
@Test
|
||||
fun testUpstreamHermesContractApprovalRequest() {
|
||||
val raw = """
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
"method": "event",
|
||||
"params": {
|
||||
"type": "approval.request",
|
||||
"session_id": "rt-session-001",
|
||||
"payload": {
|
||||
"request_id": "appr-req-999",
|
||||
"command": "rm -rf /var/log/*",
|
||||
"description": "Clean logs directory",
|
||||
"choices": ["once", "deny", "always"]
|
||||
}
|
||||
}
|
||||
}
|
||||
""".trimIndent()
|
||||
val root = json.decodeFromString<JsonObject>(raw)
|
||||
val event = GatewayEvent.parse(root)
|
||||
assertTrue(event is GatewayEvent.ApprovalRequestEvent)
|
||||
val approval = event as GatewayEvent.ApprovalRequestEvent
|
||||
assertEquals("appr-req-999", approval.requestId)
|
||||
assertEquals("rm -rf /var/log/*", approval.command)
|
||||
assertEquals("Clean logs directory", approval.description)
|
||||
assertEquals(3, approval.choices.size)
|
||||
assertEquals(listOf("once", "deny", "always"), approval.choices)
|
||||
assertEquals("rt-session-001", approval.sessionId)
|
||||
assertEquals("rt-session-001", approval.sessionKey)
|
||||
}
|
||||
|
||||
@Test
|
||||
fun testUpstreamHermesContractClarifyRequest() {
|
||||
val raw = """
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
"method": "event",
|
||||
"params": {
|
||||
"type": "clarify.request",
|
||||
"session_id": "rt-session-001",
|
||||
"payload": {
|
||||
"request_id": "clarify-req-555",
|
||||
"question_id": "q-port-target",
|
||||
"question": "Which HTTP port should the mock server bind to?"
|
||||
}
|
||||
}
|
||||
}
|
||||
""".trimIndent()
|
||||
val root = json.decodeFromString<JsonObject>(raw)
|
||||
val event = GatewayEvent.parse(root)
|
||||
assertTrue(event is GatewayEvent.ClarifyRequestEvent)
|
||||
val clarify = event as GatewayEvent.ClarifyRequestEvent
|
||||
assertEquals("clarify-req-555", clarify.requestId)
|
||||
assertEquals("q-port-target", clarify.questionId)
|
||||
assertEquals("Which HTTP port should the mock server bind to?", clarify.question)
|
||||
assertEquals(ClarifyType.CLARIFY, clarify.promptType)
|
||||
assertEquals("rt-session-001", clarify.sessionId)
|
||||
}
|
||||
|
||||
@Test
|
||||
|
|
|
|||
|
|
@ -89,13 +89,16 @@ class ApprovalRoutingTest {
|
|||
|
||||
// Simulate incoming approval request with specific runtime session ID
|
||||
val prodEventJson = buildJsonObject {
|
||||
put("jsonrpc", "2.0")
|
||||
put("method", "event")
|
||||
put("params", buildJsonObject {
|
||||
put("event", "approval.request")
|
||||
put("request_id", "req_prod_1")
|
||||
put("session_key", "runtime_session_prod_99")
|
||||
put("command", "systemctl restart nginx")
|
||||
put("description", "Restart web server")
|
||||
put("type", "approval.request")
|
||||
put("session_id", "runtime_session_prod_99")
|
||||
put("payload", buildJsonObject {
|
||||
put("request_id", "req_prod_1")
|
||||
put("command", "systemctl restart nginx")
|
||||
put("description", "Restart web server")
|
||||
})
|
||||
})
|
||||
}
|
||||
runtime1?.gatewayClient?.handleIncomingMessage(prodEventJson.toString())
|
||||
|
|
@ -135,7 +138,7 @@ class ApprovalRoutingTest {
|
|||
path.startsWith("/api/ws") || path.startsWith("/ws") -> {
|
||||
MockResponse().withWebSocketUpgrade(object : WebSocketListener() {
|
||||
override fun onOpen(webSocket: WebSocket, response: Response) {
|
||||
webSocket.send("""{"event":"gateway.ready","data":{"version":"1.0.0"}}""")
|
||||
webSocket.send("""{"jsonrpc":"2.0","method":"event","params":{"type":"gateway.ready","payload":{"version":"1.0.0"}}}""")
|
||||
}
|
||||
|
||||
override fun onMessage(webSocket: WebSocket, text: String) {
|
||||
|
|
@ -178,12 +181,15 @@ class ApprovalRoutingTest {
|
|||
|
||||
// Incoming approval
|
||||
val prodEventJson = buildJsonObject {
|
||||
put("jsonrpc", "2.0")
|
||||
put("method", "event")
|
||||
put("params", buildJsonObject {
|
||||
put("event", "approval.request")
|
||||
put("request_id", "req_prod_1")
|
||||
put("type", "approval.request")
|
||||
put("session_id", "runtime_session_prod_99")
|
||||
put("command", "systemctl restart nginx")
|
||||
put("payload", buildJsonObject {
|
||||
put("request_id", "req_prod_1")
|
||||
put("command", "systemctl restart nginx")
|
||||
})
|
||||
})
|
||||
}
|
||||
runtime1.gatewayClient.handleIncomingMessage(prodEventJson.toString())
|
||||
|
|
|
|||
|
|
@ -110,43 +110,55 @@ class MultiHostConcurrencyExecutionTest {
|
|||
)
|
||||
testScheduler.advanceUntilIdle()
|
||||
|
||||
// Stream from Host A into Session 1
|
||||
// Stream from Host A into Session 1 using upstream Hermes envelope
|
||||
val eventA1 = buildJsonObject {
|
||||
put("jsonrpc", "2.0")
|
||||
put("method", "event")
|
||||
put("params", buildJsonObject {
|
||||
put("event", "message.start")
|
||||
put("type", "message.start")
|
||||
put("session_id", "rt_win_1")
|
||||
put("message_id", "msg_a_1")
|
||||
put("role", "assistant")
|
||||
put("payload", buildJsonObject {
|
||||
put("message_id", "msg_a_1")
|
||||
put("role", "assistant")
|
||||
})
|
||||
})
|
||||
}
|
||||
val eventA2 = buildJsonObject {
|
||||
put("jsonrpc", "2.0")
|
||||
put("method", "event")
|
||||
put("params", buildJsonObject {
|
||||
put("event", "message.delta")
|
||||
put("type", "message.delta")
|
||||
put("session_id", "rt_win_1")
|
||||
put("message_id", "msg_a_1")
|
||||
put("delta", "Windows output chunk")
|
||||
put("payload", buildJsonObject {
|
||||
put("message_id", "msg_a_1")
|
||||
put("delta", "Windows output chunk")
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
// Stream from Host B into Session 2 concurrently
|
||||
// Stream from Host B into Session 2 concurrently using upstream Hermes envelope
|
||||
val eventB1 = buildJsonObject {
|
||||
put("jsonrpc", "2.0")
|
||||
put("method", "event")
|
||||
put("params", buildJsonObject {
|
||||
put("event", "message.start")
|
||||
put("type", "message.start")
|
||||
put("session_id", "rt_lin_1")
|
||||
put("message_id", "msg_b_1")
|
||||
put("role", "assistant")
|
||||
put("payload", buildJsonObject {
|
||||
put("message_id", "msg_b_1")
|
||||
put("role", "assistant")
|
||||
})
|
||||
})
|
||||
}
|
||||
val eventB2 = buildJsonObject {
|
||||
put("jsonrpc", "2.0")
|
||||
put("method", "event")
|
||||
put("params", buildJsonObject {
|
||||
put("event", "message.delta")
|
||||
put("type", "message.delta")
|
||||
put("session_id", "rt_lin_1")
|
||||
put("message_id", "msg_b_1")
|
||||
put("delta", "Linux output chunk")
|
||||
put("payload", buildJsonObject {
|
||||
put("message_id", "msg_b_1")
|
||||
put("delta", "Linux output chunk")
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
|
|
@ -215,12 +227,15 @@ class MultiHostConcurrencyExecutionTest {
|
|||
// Host A streams message
|
||||
runtimeA!!.gatewayClient.handleIncomingMessage(
|
||||
buildJsonObject {
|
||||
put("jsonrpc", "2.0")
|
||||
put("method", "event")
|
||||
put("params", buildJsonObject {
|
||||
put("event", "message.start")
|
||||
put("type", "message.start")
|
||||
put("session_id", "rt_win_dual")
|
||||
put("message_id", "msg_win")
|
||||
put("role", "assistant")
|
||||
put("payload", buildJsonObject {
|
||||
put("message_id", "msg_win")
|
||||
put("role", "assistant")
|
||||
})
|
||||
})
|
||||
}.toString()
|
||||
)
|
||||
|
|
@ -228,12 +243,15 @@ class MultiHostConcurrencyExecutionTest {
|
|||
|
||||
runtimeA.gatewayClient.handleIncomingMessage(
|
||||
buildJsonObject {
|
||||
put("jsonrpc", "2.0")
|
||||
put("method", "event")
|
||||
put("params", buildJsonObject {
|
||||
put("event", "message.delta")
|
||||
put("type", "message.delta")
|
||||
put("session_id", "rt_win_dual")
|
||||
put("message_id", "msg_win")
|
||||
put("delta", "Windows result")
|
||||
put("payload", buildJsonObject {
|
||||
put("message_id", "msg_win")
|
||||
put("delta", "Windows result")
|
||||
})
|
||||
})
|
||||
}.toString()
|
||||
)
|
||||
|
|
@ -242,12 +260,15 @@ class MultiHostConcurrencyExecutionTest {
|
|||
// Host B concurrently streams tool
|
||||
runtimeB!!.gatewayClient.handleIncomingMessage(
|
||||
buildJsonObject {
|
||||
put("jsonrpc", "2.0")
|
||||
put("method", "event")
|
||||
put("params", buildJsonObject {
|
||||
put("event", "tool.start")
|
||||
put("type", "tool.start")
|
||||
put("session_id", "rt_lin_dual")
|
||||
put("tool_id", "tool_lin")
|
||||
put("name", "bash_exec")
|
||||
put("payload", buildJsonObject {
|
||||
put("tool_id", "tool_lin")
|
||||
put("name", "bash_exec")
|
||||
})
|
||||
})
|
||||
}.toString()
|
||||
)
|
||||
|
|
@ -255,12 +276,15 @@ class MultiHostConcurrencyExecutionTest {
|
|||
|
||||
runtimeB.gatewayClient.handleIncomingMessage(
|
||||
buildJsonObject {
|
||||
put("jsonrpc", "2.0")
|
||||
put("method", "event")
|
||||
put("params", buildJsonObject {
|
||||
put("event", "tool.complete")
|
||||
put("type", "tool.complete")
|
||||
put("session_id", "rt_lin_dual")
|
||||
put("tool_id", "tool_lin")
|
||||
put("result", "Linux command completed")
|
||||
put("payload", buildJsonObject {
|
||||
put("tool_id", "tool_lin")
|
||||
put("result", "Linux command completed")
|
||||
})
|
||||
})
|
||||
}.toString()
|
||||
)
|
||||
|
|
@ -294,21 +318,27 @@ class MultiHostConcurrencyExecutionTest {
|
|||
|
||||
// Both hosts have same requestId "req_shared_1"
|
||||
val approvalEventA = buildJsonObject {
|
||||
put("jsonrpc", "2.0")
|
||||
put("method", "event")
|
||||
put("params", buildJsonObject {
|
||||
put("event", "approval.request")
|
||||
put("type", "approval.request")
|
||||
put("session_id", "rt_win_shared")
|
||||
put("request_id", "req_shared_1")
|
||||
put("command", "powershell.exe -Command Get-Process")
|
||||
put("payload", buildJsonObject {
|
||||
put("request_id", "req_shared_1")
|
||||
put("command", "powershell.exe -Command Get-Process")
|
||||
})
|
||||
})
|
||||
}
|
||||
val approvalEventB = buildJsonObject {
|
||||
put("jsonrpc", "2.0")
|
||||
put("method", "event")
|
||||
put("params", buildJsonObject {
|
||||
put("event", "approval.request")
|
||||
put("type", "approval.request")
|
||||
put("session_id", "rt_lin_shared")
|
||||
put("request_id", "req_shared_1")
|
||||
put("command", "ps aux")
|
||||
put("payload", buildJsonObject {
|
||||
put("request_id", "req_shared_1")
|
||||
put("command", "ps aux")
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
|
|
@ -330,6 +360,68 @@ class MultiHostConcurrencyExecutionTest {
|
|||
assertEquals(RuntimeSessionId("rt_lin_shared"), appB?.runtimeSessionId)
|
||||
}
|
||||
|
||||
@Test
|
||||
fun testTwoHermesHostsWithIdenticalRuntimeSessionIdDoNotCrossRouteEvents() = runTest(testDispatcher) {
|
||||
val hostA = HermesHost(id = host1Id, displayName = "Windows PC", baseUrl = "http://pc:9119")
|
||||
val hostB = HermesHost(id = host2Id, displayName = "Linux Server", baseUrl = "http://linux:9119")
|
||||
connectionManager.addHost(hostA)
|
||||
connectionManager.addHost(hostB)
|
||||
testScheduler.advanceUntilIdle()
|
||||
|
||||
val session1 = sessionRepo.createUnifiedSession(title = "Host 1 Session", initialHostId = host1Id)
|
||||
val session2 = sessionRepo.createUnifiedSession(title = "Host 2 Session", initialHostId = host2Id)
|
||||
testScheduler.advanceUntilIdle()
|
||||
|
||||
val runtimeA = connectionManager.getRuntime(host1Id)!!
|
||||
val runtimeB = connectionManager.getRuntime(host2Id)!!
|
||||
|
||||
// Both hosts independently return the EXACT SAME runtimeSessionId "s1"
|
||||
sessionRepo.registerRuntimeBinding(session1.id, host1Id, RuntimeSessionId("s1"))
|
||||
sessionRepo.registerRuntimeBinding(session2.id, host2Id, RuntimeSessionId("s1"))
|
||||
|
||||
// Host 1 sends delta for session_id "s1"
|
||||
val eventHost1 = buildJsonObject {
|
||||
put("jsonrpc", "2.0")
|
||||
put("method", "event")
|
||||
put("params", buildJsonObject {
|
||||
put("type", "message.delta")
|
||||
put("session_id", "s1")
|
||||
put("payload", buildJsonObject {
|
||||
put("message_id", "m_shared")
|
||||
put("delta", "Chunk From Host 1")
|
||||
})
|
||||
})
|
||||
}
|
||||
// Host 2 sends delta for session_id "s1"
|
||||
val eventHost2 = buildJsonObject {
|
||||
put("jsonrpc", "2.0")
|
||||
put("method", "event")
|
||||
put("params", buildJsonObject {
|
||||
put("type", "message.delta")
|
||||
put("session_id", "s1")
|
||||
put("payload", buildJsonObject {
|
||||
put("message_id", "m_shared")
|
||||
put("delta", "Chunk From Host 2")
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
runtimeA.gatewayClient.handleIncomingMessage(eventHost1.toString())
|
||||
runtimeB.gatewayClient.handleIncomingMessage(eventHost2.toString())
|
||||
testScheduler.advanceUntilIdle()
|
||||
|
||||
val session1Messages = sessionRepo.getSessionMessages(session1.id).value
|
||||
val session2Messages = sessionRepo.getSessionMessages(session2.id).value
|
||||
|
||||
assertEquals(1, session1Messages.size)
|
||||
assertEquals("Chunk From Host 1", session1Messages[0].content)
|
||||
assertEquals(host1Id, session1Messages[0].hostId)
|
||||
|
||||
assertEquals(1, session2Messages.size)
|
||||
assertEquals("Chunk From Host 2", session2Messages[0].content)
|
||||
assertEquals(host2Id, session2Messages[0].hostId)
|
||||
}
|
||||
|
||||
@Test
|
||||
fun testAppRestartRestoresBindingViaDurableId() = runBlocking(Dispatchers.Default) {
|
||||
val server = MockWebServer()
|
||||
|
|
@ -344,7 +436,7 @@ class MultiHostConcurrencyExecutionTest {
|
|||
path.startsWith("/api/ws") || path.startsWith("/ws") -> {
|
||||
MockResponse().withWebSocketUpgrade(object : WebSocketListener() {
|
||||
override fun onOpen(webSocket: WebSocket, response: Response) {
|
||||
webSocket.send("""{"event":"gateway.ready","data":{"version":"1.0.0"}}""")
|
||||
webSocket.send("""{"jsonrpc":"2.0","method":"event","params":{"type":"gateway.ready","payload":{"version":"1.0.0"}}}""")
|
||||
}
|
||||
|
||||
override fun onMessage(webSocket: WebSocket, text: String) {
|
||||
|
|
@ -399,6 +491,9 @@ class MultiHostConcurrencyExecutionTest {
|
|||
scope = CoroutineScope(Dispatchers.Default)
|
||||
)
|
||||
|
||||
val host1 = HermesHost(id = host1Id, displayName = "Prod Server", baseUrl = wsUrl, allowCleartext = true, lastKnownStatus = HostStatus.ONLINE)
|
||||
freshConnectionManager.addHost(host1)
|
||||
|
||||
val runtime = freshConnectionManager.getRuntime(host1Id)
|
||||
assertNotNull(runtime)
|
||||
runtime!!.connect()
|
||||
|
|
@ -418,6 +513,105 @@ class MultiHostConcurrencyExecutionTest {
|
|||
server.shutdown()
|
||||
}
|
||||
|
||||
@Test
|
||||
fun testAppRestartWithStateReadyCallsSessionResume() = runBlocking(Dispatchers.Default) {
|
||||
val server = MockWebServer()
|
||||
var sessionResumeCalled = false
|
||||
var resumedDurableId = ""
|
||||
|
||||
server.dispatcher = object : Dispatcher() {
|
||||
override fun dispatch(request: RecordedRequest): MockResponse {
|
||||
val path = request.path ?: ""
|
||||
return when {
|
||||
path == "/api/status" -> {
|
||||
MockResponse().setResponseCode(200).setBody("""{"status":"ok","auth_required":false,"version":"1.0.0"}""")
|
||||
}
|
||||
path.startsWith("/api/ws") || path.startsWith("/ws") -> {
|
||||
MockResponse().withWebSocketUpgrade(object : WebSocketListener() {
|
||||
override fun onOpen(webSocket: WebSocket, response: Response) {
|
||||
webSocket.send("""{"jsonrpc":"2.0","method":"event","params":{"type":"gateway.ready","payload":{"version":"1.0.0"}}}""")
|
||||
}
|
||||
|
||||
override fun onMessage(webSocket: WebSocket, text: String) {
|
||||
if (text.contains("session.resume")) {
|
||||
sessionResumeCalled = true
|
||||
if (text.contains("valid_durable_888")) {
|
||||
resumedDurableId = "valid_durable_888"
|
||||
}
|
||||
webSocket.send("""{"jsonrpc":"2.0","id":"a1","result":{"stored_session_id":"valid_durable_888","session_id":"fresh_runtime_777"}}""")
|
||||
} else if (text.contains("prompt.submit")) {
|
||||
webSocket.send("""{"jsonrpc":"2.0","id":"a2","result":{"turn_id":"turn_ready_restart"}}""")
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
else -> MockResponse().setResponseCode(404)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
server.start()
|
||||
val wsUrl = server.url("").toString().removeSuffix("/")
|
||||
val testHostDao = FakeHostDao()
|
||||
val testSessionDao = FakeUnifiedSessionDao()
|
||||
val testTokenVault = InMemoryTokenVault()
|
||||
|
||||
testHostDao.insertOrUpdateHost(HostEntity(id = host1Id.value, displayName = "Server", baseUrl = wsUrl, allowCleartext = true, lastKnownStatus = "ONLINE"))
|
||||
|
||||
val sessionId = UnifiedSessionId("session_ready_restart")
|
||||
testSessionDao.insertSession(
|
||||
UnifiedSessionEntity(
|
||||
id = sessionId.value,
|
||||
title = "Ready Restart Test",
|
||||
activeHostId = host1Id.value
|
||||
)
|
||||
)
|
||||
// Binding in Room has state = READY, but stale runtimeSessionId from previous process run
|
||||
testSessionDao.insertOrUpdateBinding(
|
||||
HostBindingEntity(
|
||||
sessionId = sessionId.value,
|
||||
hostId = host1Id.value,
|
||||
durableSessionId = "valid_durable_888",
|
||||
runtimeSessionId = "stale_dead_runtime_999",
|
||||
state = BindingState.READY.name
|
||||
)
|
||||
)
|
||||
|
||||
// New process simulation
|
||||
val freshConnectionManager = HermesConnectionManager(
|
||||
hostDao = testHostDao,
|
||||
tokenVault = testTokenVault,
|
||||
scope = CoroutineScope(Dispatchers.Default)
|
||||
)
|
||||
val freshRepo = UnifiedSessionRepository(
|
||||
connectionManager = freshConnectionManager,
|
||||
sessionDao = testSessionDao,
|
||||
scope = CoroutineScope(Dispatchers.Default)
|
||||
)
|
||||
|
||||
val host1 = HermesHost(id = host1Id, displayName = "Server", baseUrl = wsUrl, allowCleartext = true, lastKnownStatus = HostStatus.ONLINE)
|
||||
freshConnectionManager.addHost(host1)
|
||||
|
||||
val runtime = freshConnectionManager.getRuntime(host1Id)!!
|
||||
runtime.connect()
|
||||
runtime.gatewayClient.awaitGatewayReady(5000)
|
||||
|
||||
// Calling sendPrompt must trigger session.resume since in-memory attachment does not exist in new process
|
||||
val turnId = freshRepo.sendPrompt(sessionId, "Hello after clean restart")
|
||||
assertEquals("turn_ready_restart", turnId)
|
||||
|
||||
assertTrue("session.resume MUST be called even if persisted state was READY", sessionResumeCalled)
|
||||
assertEquals("valid_durable_888", resumedDurableId)
|
||||
|
||||
val updatedBinding = testSessionDao.getBindingsForSession(sessionId.value).find { it.hostId == host1Id.value }
|
||||
assertNotNull(updatedBinding)
|
||||
assertEquals("valid_durable_888", updatedBinding?.durableSessionId)
|
||||
assertEquals("fresh_runtime_777", updatedBinding?.runtimeSessionId)
|
||||
|
||||
runtime.disconnect()
|
||||
server.shutdown()
|
||||
}
|
||||
|
||||
@Test
|
||||
fun testHostReconnectMintsNewRuntimeId() = runBlocking(Dispatchers.Default) {
|
||||
val server = MockWebServer()
|
||||
|
|
@ -433,7 +627,7 @@ class MultiHostConcurrencyExecutionTest {
|
|||
path.startsWith("/api/ws") || path.startsWith("/ws") -> {
|
||||
MockResponse().withWebSocketUpgrade(object : WebSocketListener() {
|
||||
override fun onOpen(webSocket: WebSocket, response: Response) {
|
||||
webSocket.send("""{"event":"gateway.ready","data":{"version":"1.0.0"}}""")
|
||||
webSocket.send("""{"jsonrpc":"2.0","method":"event","params":{"type":"gateway.ready","payload":{"version":"1.0.0"}}}""")
|
||||
}
|
||||
|
||||
override fun onMessage(webSocket: WebSocket, text: String) {
|
||||
|
|
@ -510,13 +704,15 @@ class MultiHostConcurrencyExecutionTest {
|
|||
path.startsWith("/api/ws") || path.startsWith("/ws") -> {
|
||||
MockResponse().withWebSocketUpgrade(object : WebSocketListener() {
|
||||
override fun onOpen(webSocket: WebSocket, response: Response) {
|
||||
webSocket.send("""{"event":"gateway.ready","data":{"version":"1.0.0"}}""")
|
||||
webSocket.send("""{"jsonrpc":"2.0","method":"event","params":{"type":"gateway.ready","payload":{"version":"1.0.0"}}}""")
|
||||
}
|
||||
|
||||
override fun onMessage(webSocket: WebSocket, text: String) {
|
||||
if (text.contains("prompt.submit")) {
|
||||
if (text.contains("session.resume")) {
|
||||
webSocket.send("""{"jsonrpc":"2.0","id":"a1","result":{"stored_session_id":"dur_fail_test","session_id":"rt_fail_test"}}""")
|
||||
} else if (text.contains("prompt.submit")) {
|
||||
// Fail prompt submission with an RPC error
|
||||
webSocket.send("""{"jsonrpc":"2.0","id":"a1","error":{"code":-32000,"message":"Model overloaded"}}""")
|
||||
webSocket.send("""{"jsonrpc":"2.0","id":"a2","error":{"code":-32000,"message":"Model overloaded"}}""")
|
||||
}
|
||||
}
|
||||
})
|
||||
|
|
@ -545,6 +741,9 @@ class MultiHostConcurrencyExecutionTest {
|
|||
scope = CoroutineScope(Dispatchers.Default)
|
||||
)
|
||||
|
||||
val host1 = HermesHost(id = host1Id, displayName = "Server", baseUrl = wsUrl, allowCleartext = true, lastKnownStatus = HostStatus.ONLINE)
|
||||
testConnectionManager.addHost(host1)
|
||||
|
||||
val session = testRepo.createUnifiedSession(title = "Cursor Test", initialHostId = host1Id)
|
||||
|
||||
// Pre-populate binding with syncedThroughMessageId = "msg_baseline"
|
||||
|
|
@ -582,6 +781,195 @@ class MultiHostConcurrencyExecutionTest {
|
|||
server.shutdown()
|
||||
}
|
||||
|
||||
@Test
|
||||
fun testTransientResumeErrorDoesNotDestroyBindingOrCallCreate() = runBlocking(Dispatchers.Default) {
|
||||
val server = MockWebServer()
|
||||
var createCalled = false
|
||||
|
||||
server.dispatcher = object : Dispatcher() {
|
||||
override fun dispatch(request: RecordedRequest): MockResponse {
|
||||
val path = request.path ?: ""
|
||||
return when {
|
||||
path == "/api/status" -> {
|
||||
MockResponse().setResponseCode(200).setBody("""{"status":"ok","auth_required":false,"version":"1.0.0"}""")
|
||||
}
|
||||
path.startsWith("/api/ws") || path.startsWith("/ws") -> {
|
||||
MockResponse().withWebSocketUpgrade(object : WebSocketListener() {
|
||||
override fun onOpen(webSocket: WebSocket, response: Response) {
|
||||
webSocket.send("""{"jsonrpc":"2.0","method":"event","params":{"type":"gateway.ready","payload":{"version":"1.0.0"}}}""")
|
||||
}
|
||||
|
||||
override fun onMessage(webSocket: WebSocket, text: String) {
|
||||
if (text.contains("session.resume")) {
|
||||
// Transient failure: 500 / -32000 Server overloaded
|
||||
webSocket.send("""{"jsonrpc":"2.0","id":"a1","error":{"code":-32000,"message":"Transient server error"}}""")
|
||||
} else if (text.contains("session.create")) {
|
||||
createCalled = true
|
||||
webSocket.send("""{"jsonrpc":"2.0","id":"a2","result":{"stored_session_id":"dur_forbidden","session_id":"rt_forbidden"}}""")
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
else -> MockResponse().setResponseCode(404)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
server.start()
|
||||
val wsUrl = server.url("").toString().removeSuffix("/")
|
||||
val testHostDao = FakeHostDao()
|
||||
val testSessionDao = FakeUnifiedSessionDao()
|
||||
val testTokenVault = InMemoryTokenVault()
|
||||
|
||||
testHostDao.insertOrUpdateHost(HostEntity(id = host1Id.value, displayName = "Server", baseUrl = wsUrl, allowCleartext = true, lastKnownStatus = "ONLINE"))
|
||||
|
||||
val sessionId = UnifiedSessionId("session_transient_test")
|
||||
testSessionDao.insertSession(
|
||||
UnifiedSessionEntity(
|
||||
id = sessionId.value,
|
||||
title = "Transient Resume Test",
|
||||
activeHostId = host1Id.value
|
||||
)
|
||||
)
|
||||
testSessionDao.insertOrUpdateBinding(
|
||||
HostBindingEntity(
|
||||
sessionId = sessionId.value,
|
||||
hostId = host1Id.value,
|
||||
durableSessionId = "dur_preserved_123",
|
||||
runtimeSessionId = "rt_stale",
|
||||
state = BindingState.READY.name
|
||||
)
|
||||
)
|
||||
|
||||
val testConnectionManager = HermesConnectionManager(
|
||||
hostDao = testHostDao,
|
||||
tokenVault = testTokenVault,
|
||||
scope = CoroutineScope(Dispatchers.Default)
|
||||
)
|
||||
val testRepo = UnifiedSessionRepository(
|
||||
connectionManager = testConnectionManager,
|
||||
sessionDao = testSessionDao,
|
||||
scope = CoroutineScope(Dispatchers.Default)
|
||||
)
|
||||
|
||||
val host1 = HermesHost(id = host1Id, displayName = "Server", baseUrl = wsUrl, allowCleartext = true, lastKnownStatus = HostStatus.ONLINE)
|
||||
testConnectionManager.addHost(host1)
|
||||
|
||||
val runtime = testConnectionManager.getRuntime(host1Id)!!
|
||||
runtime.connect()
|
||||
runtime.gatewayClient.awaitGatewayReady(5000)
|
||||
|
||||
var threw = false
|
||||
try {
|
||||
testRepo.ensureAttachedRuntimeSession(sessionId, host1Id, runtime)
|
||||
} catch (_: Exception) {
|
||||
threw = true
|
||||
}
|
||||
|
||||
assertTrue("Expected transient resume error to throw", threw)
|
||||
assertFalse("session.create MUST NOT be called on transient resume error", createCalled)
|
||||
|
||||
// Verify binding was not destroyed
|
||||
val binding = testSessionDao.getBindingsForSession(sessionId.value).find { it.hostId == host1Id.value }
|
||||
assertEquals("dur_preserved_123", binding?.durableSessionId)
|
||||
|
||||
runtime.disconnect()
|
||||
server.shutdown()
|
||||
}
|
||||
|
||||
@Test
|
||||
fun testUnrecoverableResumeErrorRecreatesSession() = runBlocking(Dispatchers.Default) {
|
||||
val server = MockWebServer()
|
||||
var createCalled = false
|
||||
|
||||
server.dispatcher = object : Dispatcher() {
|
||||
override fun dispatch(request: RecordedRequest): MockResponse {
|
||||
val path = request.path ?: ""
|
||||
return when {
|
||||
path == "/api/status" -> {
|
||||
MockResponse().setResponseCode(200).setBody("""{"status":"ok","auth_required":false,"version":"1.0.0"}""")
|
||||
}
|
||||
path.startsWith("/api/ws") || path.startsWith("/ws") -> {
|
||||
MockResponse().withWebSocketUpgrade(object : WebSocketListener() {
|
||||
override fun onOpen(webSocket: WebSocket, response: Response) {
|
||||
webSocket.send("""{"jsonrpc":"2.0","method":"event","params":{"type":"gateway.ready","payload":{"version":"1.0.0"}}}""")
|
||||
}
|
||||
|
||||
override fun onMessage(webSocket: WebSocket, text: String) {
|
||||
if (text.contains("session.resume")) {
|
||||
// Definitively unrecoverable: 404 Session not found
|
||||
webSocket.send("""{"jsonrpc":"2.0","id":"a1","error":{"code":404,"message":"Session not found"}}""")
|
||||
} else if (text.contains("session.create")) {
|
||||
createCalled = true
|
||||
webSocket.send("""{"jsonrpc":"2.0","id":"a2","result":{"stored_session_id":"dur_new_fresh_999","session_id":"rt_new_fresh_999"}}""")
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
else -> MockResponse().setResponseCode(404)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
server.start()
|
||||
val wsUrl = server.url("").toString().removeSuffix("/")
|
||||
val testHostDao = FakeHostDao()
|
||||
val testSessionDao = FakeUnifiedSessionDao()
|
||||
val testTokenVault = InMemoryTokenVault()
|
||||
|
||||
testHostDao.insertOrUpdateHost(HostEntity(id = host1Id.value, displayName = "Server", baseUrl = wsUrl, allowCleartext = true, lastKnownStatus = "ONLINE"))
|
||||
|
||||
val sessionId = UnifiedSessionId("session_unrecoverable_test")
|
||||
testSessionDao.insertSession(
|
||||
UnifiedSessionEntity(
|
||||
id = sessionId.value,
|
||||
title = "Unrecoverable Resume Test",
|
||||
activeHostId = host1Id.value
|
||||
)
|
||||
)
|
||||
testSessionDao.insertOrUpdateBinding(
|
||||
HostBindingEntity(
|
||||
sessionId = sessionId.value,
|
||||
hostId = host1Id.value,
|
||||
durableSessionId = "dur_dead_deleted",
|
||||
runtimeSessionId = "rt_dead",
|
||||
state = BindingState.READY.name
|
||||
)
|
||||
)
|
||||
|
||||
val testConnectionManager = HermesConnectionManager(
|
||||
hostDao = testHostDao,
|
||||
tokenVault = testTokenVault,
|
||||
scope = CoroutineScope(Dispatchers.Default)
|
||||
)
|
||||
val testRepo = UnifiedSessionRepository(
|
||||
connectionManager = testConnectionManager,
|
||||
sessionDao = testSessionDao,
|
||||
scope = CoroutineScope(Dispatchers.Default)
|
||||
)
|
||||
|
||||
val host1 = HermesHost(id = host1Id, displayName = "Server", baseUrl = wsUrl, allowCleartext = true, lastKnownStatus = HostStatus.ONLINE)
|
||||
testConnectionManager.addHost(host1)
|
||||
|
||||
val runtime = testConnectionManager.getRuntime(host1Id)!!
|
||||
runtime.connect()
|
||||
runtime.gatewayClient.awaitGatewayReady(5000)
|
||||
|
||||
val attachedBinding = testRepo.ensureAttachedRuntimeSession(sessionId, host1Id, runtime)
|
||||
|
||||
assertTrue("session.create MUST be called on unrecoverable 404 resume error", createCalled)
|
||||
assertEquals(DurableSessionId("dur_new_fresh_999"), attachedBinding.durableSessionId)
|
||||
assertEquals(RuntimeSessionId("rt_new_fresh_999"), attachedBinding.runtimeSessionId)
|
||||
|
||||
// Verify Room DB binding was updated
|
||||
val binding = testSessionDao.getBindingsForSession(sessionId.value).find { it.hostId == host1Id.value }
|
||||
assertEquals("dur_new_fresh_999", binding?.durableSessionId)
|
||||
assertEquals("rt_new_fresh_999", binding?.runtimeSessionId)
|
||||
|
||||
runtime.disconnect()
|
||||
server.shutdown()
|
||||
}
|
||||
|
||||
@Test
|
||||
fun testStopHostADoesNotStopHostB() = runTest(testDispatcher) {
|
||||
val hostA = HermesHost(id = host1Id, displayName = "Windows PC", baseUrl = "http://pc:9119")
|
||||
|
|
@ -619,24 +1007,30 @@ class MultiHostConcurrencyExecutionTest {
|
|||
)
|
||||
testScheduler.advanceUntilIdle()
|
||||
|
||||
// Start execution on both hosts
|
||||
// Start execution on both hosts using upstream Hermes envelope
|
||||
runtimeA!!.gatewayClient.handleIncomingMessage(
|
||||
buildJsonObject {
|
||||
put("jsonrpc", "2.0")
|
||||
put("method", "event")
|
||||
put("params", buildJsonObject {
|
||||
put("event", "message.start")
|
||||
put("type", "message.start")
|
||||
put("session_id", "rt_a")
|
||||
put("message_id", "msg_a")
|
||||
put("payload", buildJsonObject {
|
||||
put("message_id", "msg_a")
|
||||
})
|
||||
})
|
||||
}.toString()
|
||||
)
|
||||
runtimeB!!.gatewayClient.handleIncomingMessage(
|
||||
buildJsonObject {
|
||||
put("jsonrpc", "2.0")
|
||||
put("method", "event")
|
||||
put("params", buildJsonObject {
|
||||
put("event", "message.start")
|
||||
put("type", "message.start")
|
||||
put("session_id", "rt_b")
|
||||
put("message_id", "msg_b")
|
||||
put("payload", buildJsonObject {
|
||||
put("message_id", "msg_b")
|
||||
})
|
||||
})
|
||||
}.toString()
|
||||
)
|
||||
|
|
|
|||
|
|
@ -114,12 +114,15 @@ class UnifiedSessionRepositoryTest {
|
|||
|
||||
// Stream start event
|
||||
val msgStart = buildJsonObject {
|
||||
put("jsonrpc", "2.0")
|
||||
put("method", "event")
|
||||
put("params", buildJsonObject {
|
||||
put("event", "message.start")
|
||||
put("type", "message.start")
|
||||
put("session_id", "rt_stream_1")
|
||||
put("message_id", "msg_stream_1")
|
||||
put("role", "assistant")
|
||||
put("payload", buildJsonObject {
|
||||
put("message_id", "msg_stream_1")
|
||||
put("role", "assistant")
|
||||
})
|
||||
})
|
||||
}
|
||||
runtimeA?.gatewayClient?.handleIncomingMessage(msgStart.toString())
|
||||
|
|
@ -127,12 +130,15 @@ class UnifiedSessionRepositoryTest {
|
|||
|
||||
// Stream delta 1
|
||||
val msgDelta1 = buildJsonObject {
|
||||
put("jsonrpc", "2.0")
|
||||
put("method", "event")
|
||||
put("params", buildJsonObject {
|
||||
put("event", "message.delta")
|
||||
put("type", "message.delta")
|
||||
put("session_id", "rt_stream_1")
|
||||
put("message_id", "msg_stream_1")
|
||||
put("delta", "Hello ")
|
||||
put("payload", buildJsonObject {
|
||||
put("message_id", "msg_stream_1")
|
||||
put("delta", "Hello ")
|
||||
})
|
||||
})
|
||||
}
|
||||
runtimeA?.gatewayClient?.handleIncomingMessage(msgDelta1.toString())
|
||||
|
|
@ -140,12 +146,15 @@ class UnifiedSessionRepositoryTest {
|
|||
|
||||
// Stream delta 2
|
||||
val msgDelta2 = buildJsonObject {
|
||||
put("jsonrpc", "2.0")
|
||||
put("method", "event")
|
||||
put("params", buildJsonObject {
|
||||
put("event", "message.delta")
|
||||
put("type", "message.delta")
|
||||
put("session_id", "rt_stream_1")
|
||||
put("message_id", "msg_stream_1")
|
||||
put("delta", "from Multi-Hermes!")
|
||||
put("payload", buildJsonObject {
|
||||
put("message_id", "msg_stream_1")
|
||||
put("delta", "from Multi-Hermes!")
|
||||
})
|
||||
})
|
||||
}
|
||||
runtimeA?.gatewayClient?.handleIncomingMessage(msgDelta2.toString())
|
||||
|
|
@ -153,12 +162,15 @@ class UnifiedSessionRepositoryTest {
|
|||
|
||||
// Stream complete
|
||||
val msgComplete = buildJsonObject {
|
||||
put("jsonrpc", "2.0")
|
||||
put("method", "event")
|
||||
put("params", buildJsonObject {
|
||||
put("event", "message.complete")
|
||||
put("type", "message.complete")
|
||||
put("session_id", "rt_stream_1")
|
||||
put("message_id", "msg_stream_1")
|
||||
put("content", "Hello from Multi-Hermes!")
|
||||
put("payload", buildJsonObject {
|
||||
put("message_id", "msg_stream_1")
|
||||
put("content", "Hello from Multi-Hermes!")
|
||||
})
|
||||
})
|
||||
}
|
||||
runtimeA?.gatewayClient?.handleIncomingMessage(msgComplete.toString())
|
||||
|
|
@ -201,12 +213,15 @@ class UnifiedSessionRepositoryTest {
|
|||
|
||||
// Host A starts long tool operation
|
||||
val toolStart = buildJsonObject {
|
||||
put("jsonrpc", "2.0")
|
||||
put("method", "event")
|
||||
put("params", buildJsonObject {
|
||||
put("event", "tool.start")
|
||||
put("type", "tool.start")
|
||||
put("session_id", "rt_bg_1")
|
||||
put("tool_id", "tool_bg_1")
|
||||
put("name", "heavy_build_task")
|
||||
put("payload", buildJsonObject {
|
||||
put("tool_id", "tool_bg_1")
|
||||
put("name", "heavy_build_task")
|
||||
})
|
||||
})
|
||||
}
|
||||
runtimeA?.gatewayClient?.handleIncomingMessage(toolStart.toString())
|
||||
|
|
@ -218,13 +233,16 @@ class UnifiedSessionRepositoryTest {
|
|||
|
||||
// Host A finishes tool in background
|
||||
val toolComplete = buildJsonObject {
|
||||
put("jsonrpc", "2.0")
|
||||
put("method", "event")
|
||||
put("params", buildJsonObject {
|
||||
put("event", "tool.complete")
|
||||
put("type", "tool.complete")
|
||||
put("session_id", "rt_bg_1")
|
||||
put("tool_id", "tool_bg_1")
|
||||
put("result", "Build successful in 42s")
|
||||
put("is_error", false)
|
||||
put("payload", buildJsonObject {
|
||||
put("tool_id", "tool_bg_1")
|
||||
put("result", "Build successful in 42s")
|
||||
put("is_error", false)
|
||||
})
|
||||
})
|
||||
}
|
||||
runtimeA?.gatewayClient?.handleIncomingMessage(toolComplete.toString())
|
||||
|
|
|
|||
Loading…
Reference in a new issue