Fix critical Multi-Hermes defects (P0/P1) and add concurrency test suite

This commit is contained in:
Ochenstarik 2026-08-24 01:52:38 +07:00
parent 879b834e51
commit fc3ddeda65
12 changed files with 1397 additions and 255 deletions

View file

@ -13,6 +13,7 @@ import kotlinx.serialization.json.longOrNull
sealed class GatewayEvent { sealed class GatewayEvent {
abstract val rawPayload: JsonObject abstract val rawPayload: JsonObject
open val sessionId: String? get() = null
data class GatewayReadyEvent( data class GatewayReadyEvent(
val version: String, val version: String,
@ -23,42 +24,49 @@ sealed class GatewayEvent {
data class MessageStartEvent( data class MessageStartEvent(
val messageId: String, val messageId: String,
val role: String = "assistant", val role: String = "assistant",
override val sessionId: String? = null,
override val rawPayload: JsonObject override val rawPayload: JsonObject
) : GatewayEvent() ) : GatewayEvent()
data class MessageDeltaEvent( data class MessageDeltaEvent(
val messageId: String, val messageId: String,
val delta: String, val delta: String,
override val sessionId: String? = null,
override val rawPayload: JsonObject override val rawPayload: JsonObject
) : GatewayEvent() ) : GatewayEvent()
data class MessageInterimEvent( data class MessageInterimEvent(
val messageId: String, val messageId: String,
val content: String, val content: String,
override val sessionId: String? = null,
override val rawPayload: JsonObject override val rawPayload: JsonObject
) : GatewayEvent() ) : GatewayEvent()
data class MessageCompleteEvent( data class MessageCompleteEvent(
val messageId: String, val messageId: String,
val content: String, val content: String,
override val sessionId: String? = null,
override val rawPayload: JsonObject override val rawPayload: JsonObject
) : GatewayEvent() ) : GatewayEvent()
data class ThinkingDeltaEvent( data class ThinkingDeltaEvent(
val messageId: String, val messageId: String,
val delta: String, val delta: String,
override val sessionId: String? = null,
override val rawPayload: JsonObject override val rawPayload: JsonObject
) : GatewayEvent() ) : GatewayEvent()
data class ReasoningDeltaEvent( data class ReasoningDeltaEvent(
val messageId: String, val messageId: String,
val delta: String, val delta: String,
override val sessionId: String? = null,
override val rawPayload: JsonObject override val rawPayload: JsonObject
) : GatewayEvent() ) : GatewayEvent()
data class ReasoningAvailableEvent( data class ReasoningAvailableEvent(
val messageId: String, val messageId: String,
val reasoning: String, val reasoning: String,
override val sessionId: String? = null,
override val rawPayload: JsonObject override val rawPayload: JsonObject
) : GatewayEvent() ) : GatewayEvent()
@ -66,18 +74,21 @@ sealed class GatewayEvent {
val toolId: String, val toolId: String,
val name: String, val name: String,
val input: JsonElement? = null, val input: JsonElement? = null,
override val sessionId: String? = null,
override val rawPayload: JsonObject override val rawPayload: JsonObject
) : GatewayEvent() ) : GatewayEvent()
data class ToolProgressEvent( data class ToolProgressEvent(
val toolId: String, val toolId: String,
val progress: String, val progress: String,
override val sessionId: String? = null,
override val rawPayload: JsonObject override val rawPayload: JsonObject
) : GatewayEvent() ) : GatewayEvent()
data class ToolGeneratingEvent( data class ToolGeneratingEvent(
val toolId: String, val toolId: String,
val name: String, val name: String,
override val sessionId: String? = null,
override val rawPayload: JsonObject override val rawPayload: JsonObject
) : GatewayEvent() ) : GatewayEvent()
@ -85,6 +96,7 @@ sealed class GatewayEvent {
val toolId: String, val toolId: String,
val result: String, val result: String,
val isError: Boolean = false, val isError: Boolean = false,
override val sessionId: String? = null,
override val rawPayload: JsonObject override val rawPayload: JsonObject
) : GatewayEvent() ) : GatewayEvent()
@ -94,6 +106,7 @@ sealed class GatewayEvent {
val description: String? = null, val description: String? = null,
val choices: List<String> = listOf("once", "deny"), val choices: List<String> = listOf("once", "deny"),
val sessionKey: String? = null, val sessionKey: String? = null,
override val sessionId: String? = null,
override val rawPayload: JsonObject override val rawPayload: JsonObject
) : GatewayEvent() ) : GatewayEvent()
@ -102,24 +115,28 @@ sealed class GatewayEvent {
val questionId: String? = null, val questionId: String? = null,
val question: String, val question: String,
val promptType: ClarifyType = ClarifyType.CLARIFY, val promptType: ClarifyType = ClarifyType.CLARIFY,
override val sessionId: String? = null,
override val rawPayload: JsonObject override val rawPayload: JsonObject
) : GatewayEvent() ) : GatewayEvent()
data class SudoRequestEvent( data class SudoRequestEvent(
val requestId: String, val requestId: String,
val question: String, val question: String,
override val sessionId: String? = null,
override val rawPayload: JsonObject override val rawPayload: JsonObject
) : GatewayEvent() ) : GatewayEvent()
data class SecretRequestEvent( data class SecretRequestEvent(
val requestId: String, val requestId: String,
val question: String, val question: String,
override val sessionId: String? = null,
override val rawPayload: JsonObject override val rawPayload: JsonObject
) : GatewayEvent() ) : GatewayEvent()
data class StatusUpdateEvent( data class StatusUpdateEvent(
val status: String, val status: String,
val message: String? = null, val message: String? = null,
override val sessionId: String? = null,
override val rawPayload: JsonObject override val rawPayload: JsonObject
) : GatewayEvent() ) : GatewayEvent()
@ -127,28 +144,33 @@ sealed class GatewayEvent {
val inputTokens: Long = 0, val inputTokens: Long = 0,
val outputTokens: Long = 0, val outputTokens: Long = 0,
val totalTokens: Long = 0, val totalTokens: Long = 0,
override val sessionId: String? = null,
override val rawPayload: JsonObject override val rawPayload: JsonObject
) : GatewayEvent() ) : GatewayEvent()
data class SessionInfoEvent( data class SessionInfoEvent(
val info: SessionInfo, val info: SessionInfo,
override val sessionId: String? = null,
override val rawPayload: JsonObject override val rawPayload: JsonObject
) : GatewayEvent() ) : GatewayEvent()
data class BackgroundCompleteEvent( data class BackgroundCompleteEvent(
val taskId: String, val taskId: String,
val result: String? = null, val result: String? = null,
override val sessionId: String? = null,
override val rawPayload: JsonObject override val rawPayload: JsonObject
) : GatewayEvent() ) : GatewayEvent()
data class ErrorEvent( data class ErrorEvent(
val code: Int = -1, val code: Int = -1,
val message: String, val message: String,
override val sessionId: String? = null,
override val rawPayload: JsonObject override val rawPayload: JsonObject
) : GatewayEvent() ) : GatewayEvent()
data class UnknownGatewayEvent( data class UnknownGatewayEvent(
val eventType: String, val eventType: String,
override val sessionId: String? = null,
override val rawPayload: JsonObject override val rawPayload: JsonObject
) : GatewayEvent() ) : GatewayEvent()
@ -233,6 +255,8 @@ sealed class GatewayEvent {
return array.mapNotNull { it.jsonPrimitive.content } return array.mapNotNull { it.jsonPrimitive.content }
} }
val sessionKey = getNullableString("session_id", "session_key", "sessionKey", "sessionId")
return when (eventType) { return when (eventType) {
"gateway.ready" -> GatewayReadyEvent( "gateway.ready" -> GatewayReadyEvent(
version = getString("version", "server_version"), version = getString("version", "server_version"),
@ -242,58 +266,69 @@ sealed class GatewayEvent {
"message.start" -> MessageStartEvent( "message.start" -> MessageStartEvent(
messageId = getString("message_id", "id"), messageId = getString("message_id", "id"),
role = getString("role").ifEmpty { "assistant" }, role = getString("role").ifEmpty { "assistant" },
sessionId = sessionKey,
rawPayload = root rawPayload = root
) )
"message.delta" -> MessageDeltaEvent( "message.delta" -> MessageDeltaEvent(
messageId = getString("message_id", "id"), messageId = getString("message_id", "id"),
delta = getString("delta", "text", "chunk"), delta = getString("delta", "text", "chunk"),
sessionId = sessionKey,
rawPayload = root rawPayload = root
) )
"message.interim" -> MessageInterimEvent( "message.interim" -> MessageInterimEvent(
messageId = getString("message_id", "id"), messageId = getString("message_id", "id"),
content = getString("content", "text"), content = getString("content", "text"),
sessionId = sessionKey,
rawPayload = root rawPayload = root
) )
"message.complete" -> MessageCompleteEvent( "message.complete" -> MessageCompleteEvent(
messageId = getString("message_id", "id"), messageId = getString("message_id", "id"),
content = getString("content", "text"), content = getString("content", "text"),
sessionId = sessionKey,
rawPayload = root rawPayload = root
) )
"thinking.delta" -> ThinkingDeltaEvent( "thinking.delta" -> ThinkingDeltaEvent(
messageId = getString("message_id", "id"), messageId = getString("message_id", "id"),
delta = getString("delta", "text", "chunk"), delta = getString("delta", "text", "chunk"),
sessionId = sessionKey,
rawPayload = root rawPayload = root
) )
"reasoning.delta" -> ReasoningDeltaEvent( "reasoning.delta" -> ReasoningDeltaEvent(
messageId = getString("message_id", "id"), messageId = getString("message_id", "id"),
delta = getString("delta", "text", "chunk"), delta = getString("delta", "text", "chunk"),
sessionId = sessionKey,
rawPayload = root rawPayload = root
) )
"reasoning.available" -> ReasoningAvailableEvent( "reasoning.available" -> ReasoningAvailableEvent(
messageId = getString("message_id", "id"), messageId = getString("message_id", "id"),
reasoning = getString("reasoning", "content"), reasoning = getString("reasoning", "content"),
sessionId = sessionKey,
rawPayload = root rawPayload = root
) )
"tool.start" -> ToolStartEvent( "tool.start" -> ToolStartEvent(
toolId = getString("tool_id", "id"), toolId = getString("tool_id", "id"),
name = getString("name", "tool_name"), name = getString("name", "tool_name"),
input = dataObj["input"] ?: root["input"], input = dataObj["input"] ?: root["input"],
sessionId = sessionKey,
rawPayload = root rawPayload = root
) )
"tool.progress" -> ToolProgressEvent( "tool.progress" -> ToolProgressEvent(
toolId = getString("tool_id", "id"), toolId = getString("tool_id", "id"),
progress = getString("progress", "message"), progress = getString("progress", "message"),
sessionId = sessionKey,
rawPayload = root rawPayload = root
) )
"tool.generating" -> ToolGeneratingEvent( "tool.generating" -> ToolGeneratingEvent(
toolId = getString("tool_id", "id"), toolId = getString("tool_id", "id"),
name = getString("name", "tool_name"), name = getString("name", "tool_name"),
sessionId = sessionKey,
rawPayload = root rawPayload = root
) )
"tool.complete" -> ToolCompleteEvent( "tool.complete" -> ToolCompleteEvent(
toolId = getString("tool_id", "id"), toolId = getString("tool_id", "id"),
result = getString("result", "output"), result = getString("result", "output"),
isError = getBoolean("is_error", "error"), isError = getBoolean("is_error", "error"),
sessionId = sessionKey,
rawPayload = root rawPayload = root
) )
"approval.request" -> { "approval.request" -> {
@ -303,7 +338,8 @@ sealed class GatewayEvent {
command = getNullableString("command"), command = getNullableString("command"),
description = getNullableString("description", "prompt"), description = getNullableString("description", "prompt"),
choices = if (choices.isNotEmpty()) choices else listOf("once", "deny"), choices = if (choices.isNotEmpty()) choices else listOf("once", "deny"),
sessionKey = getNullableString("session_key", "sessionKey"), sessionKey = sessionKey,
sessionId = sessionKey,
rawPayload = root rawPayload = root
) )
} }
@ -312,27 +348,32 @@ sealed class GatewayEvent {
questionId = getNullableString("question_id", "questionId"), questionId = getNullableString("question_id", "questionId"),
question = getString("question", "prompt"), question = getString("question", "prompt"),
promptType = ClarifyType.CLARIFY, promptType = ClarifyType.CLARIFY,
sessionId = sessionKey,
rawPayload = root rawPayload = root
) )
"sudo.request" -> SudoRequestEvent( "sudo.request" -> SudoRequestEvent(
requestId = getString("request_id", "id"), requestId = getString("request_id", "id"),
question = getString("question", "prompt").ifEmpty { "Administrator password required:" }, question = getString("question", "prompt").ifEmpty { "Administrator password required:" },
sessionId = sessionKey,
rawPayload = root rawPayload = root
) )
"secret.request" -> SecretRequestEvent( "secret.request" -> SecretRequestEvent(
requestId = getString("request_id", "id"), requestId = getString("request_id", "id"),
question = getString("question", "prompt").ifEmpty { "Secret / Token required:" }, question = getString("question", "prompt").ifEmpty { "Secret / Token required:" },
sessionId = sessionKey,
rawPayload = root rawPayload = root
) )
"status.update" -> StatusUpdateEvent( "status.update" -> StatusUpdateEvent(
status = getString("status"), status = getString("status"),
message = getNullableString("message"), message = getNullableString("message"),
sessionId = sessionKey,
rawPayload = root rawPayload = root
) )
"session.usage" -> SessionUsageEvent( "session.usage" -> SessionUsageEvent(
inputTokens = getLong("input_tokens", "prompt_tokens"), inputTokens = getLong("input_tokens", "prompt_tokens"),
outputTokens = getLong("output_tokens", "completion_tokens"), outputTokens = getLong("output_tokens", "completion_tokens"),
totalTokens = getLong("total_tokens"), totalTokens = getLong("total_tokens"),
sessionId = sessionKey,
rawPayload = root rawPayload = root
) )
"session.info" -> SessionInfoEvent( "session.info" -> SessionInfoEvent(
@ -343,23 +384,28 @@ sealed class GatewayEvent {
branch = getNullableString("branch"), branch = getNullableString("branch"),
project = getNullableString("project") project = getNullableString("project")
), ),
sessionId = sessionKey,
rawPayload = root rawPayload = root
) )
"background.complete" -> BackgroundCompleteEvent( "background.complete" -> BackgroundCompleteEvent(
taskId = getString("task_id", "id"), taskId = getString("task_id", "id"),
result = getNullableString("result"), result = getNullableString("result"),
sessionId = sessionKey,
rawPayload = root rawPayload = root
) )
"error" -> ErrorEvent( "error" -> ErrorEvent(
code = getInt("code"), code = getInt("code"),
message = getString("message").ifEmpty { "Unknown error" }, message = getString("message").ifEmpty { "Unknown error" },
sessionId = sessionKey,
rawPayload = root rawPayload = root
) )
else -> UnknownGatewayEvent( else -> UnknownGatewayEvent(
eventType = eventType.ifEmpty { "unknown" }, eventType = eventType.ifEmpty { "unknown" },
sessionId = sessionKey,
rawPayload = root rawPayload = root
) )
} }
} }
} }
} }

View file

@ -97,6 +97,7 @@ data class HostGatewayEvent(
data class HostAttributedApproval( data class HostAttributedApproval(
val hostId: HermesHostId, val hostId: HermesHostId,
val hostDisplayName: String, val hostDisplayName: String,
val runtimeSessionId: RuntimeSessionId,
val approval: HermesApproval val approval: HermesApproval
) )
@ -104,5 +105,6 @@ data class HostAttributedApproval(
data class HostAttributedClarify( data class HostAttributedClarify(
val hostId: HermesHostId, val hostId: HermesHostId,
val hostDisplayName: String, val hostDisplayName: String,
val runtimeSessionId: RuntimeSessionId? = null,
val request: HermesClarifyRequest val request: HermesClarifyRequest
) )

View file

@ -3,6 +3,7 @@ package app.hermes.mobile.core.repository
import app.hermes.mobile.core.model.* import app.hermes.mobile.core.model.*
import app.hermes.mobile.core.network.ConnectionState import app.hermes.mobile.core.network.ConnectionState
import app.hermes.mobile.core.runtime.HermesConnectionManager import app.hermes.mobile.core.runtime.HermesConnectionManager
import app.hermes.mobile.core.runtime.HermesHostRuntime
import app.hermes.mobile.core.storage.* import app.hermes.mobile.core.storage.*
import app.hermes.mobile.core.sync.UnifiedContextBuilder import app.hermes.mobile.core.sync.UnifiedContextBuilder
import kotlinx.coroutines.CoroutineScope import kotlinx.coroutines.CoroutineScope
@ -48,9 +49,23 @@ class UnifiedSessionRepository(
// In-memory active session messages cache for reactive streaming updates // In-memory active session messages cache for reactive streaming updates
private val sessionMessagesState = ConcurrentHashMap<UnifiedSessionId, MutableStateFlow<List<UnifiedMessage>>>() private val sessionMessagesState = ConcurrentHashMap<UnifiedSessionId, MutableStateFlow<List<UnifiedMessage>>>()
// Independent per-(session, host) execution state
private val hostExecutingState = ConcurrentHashMap<Pair<UnifiedSessionId, HermesHostId>, MutableStateFlow<Boolean>>()
private val sessionExecutingState = ConcurrentHashMap<UnifiedSessionId, MutableStateFlow<Boolean>>() private val sessionExecutingState = ConcurrentHashMap<UnifiedSessionId, MutableStateFlow<Boolean>>()
init { 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 { scope.launch {
connectionManager.allEvents.collect { hostEvent -> connectionManager.allEvents.collect { hostEvent ->
handleHostGatewayEvent(hostEvent) handleHostGatewayEvent(hostEvent)
@ -71,6 +86,12 @@ class UnifiedSessionRepository(
}.asStateFlow() }.asStateFlow()
} }
fun getHostExecuting(sessionId: UnifiedSessionId, hostId: HermesHostId): StateFlow<Boolean> {
return hostExecutingState.computeIfAbsent(Pair(sessionId, hostId)) {
MutableStateFlow(false)
}.asStateFlow()
}
fun getSessionExecuting(sessionId: UnifiedSessionId): StateFlow<Boolean> { fun getSessionExecuting(sessionId: UnifiedSessionId): StateFlow<Boolean> {
return sessionExecutingState.computeIfAbsent(sessionId) { return sessionExecutingState.computeIfAbsent(sessionId) {
MutableStateFlow(false) MutableStateFlow(false)
@ -83,7 +104,7 @@ class UnifiedSessionRepository(
): UnifiedSession { ): UnifiedSession {
val hostId = initialHostId ?: connectionManager.activeHostId.value val hostId = initialHostId ?: connectionManager.activeHostId.value
?: connectionManager.hosts.value.firstOrNull()?.id ?: connectionManager.hosts.value.firstOrNull()?.id
?: HermesHostId("default") ?: throw IllegalStateException("No Hermes hosts configured. Please add a host before creating a session.")
val sessionId = UnifiedSessionId(UUID.randomUUID().toString()) val sessionId = UnifiedSessionId(UUID.randomUUID().toString())
val sessionEntity = UnifiedSessionEntity( val sessionEntity = UnifiedSessionEntity(
@ -120,12 +141,71 @@ class UnifiedSessionRepository(
sessionDao.deleteSession(sessionId.value) sessionDao.deleteSession(sessionId.value)
sessionMessagesState.remove(sessionId) sessionMessagesState.remove(sessionId)
sessionExecutingState.remove(sessionId) sessionExecutingState.remove(sessionId)
hostExecutingState.entries.removeIf { it.key.first == sessionId }
runtimeToSessionMap.entries.removeIf { it.value.first == sessionId }
}
fun registerRuntimeBinding(sessionId: UnifiedSessionId, hostId: HermesHostId, runtimeSessionId: RuntimeSessionId) {
if (runtimeSessionId.value.isNotEmpty()) {
runtimeToSessionMap[runtimeSessionId.value] = Pair(sessionId, hostId)
}
} }
suspend fun switchSessionActiveHost(sessionId: UnifiedSessionId, targetHostId: HermesHostId) { suspend fun switchSessionActiveHost(sessionId: UnifiedSessionId, targetHostId: HermesHostId) {
sessionDao.updateActiveHost(sessionId.value, targetHostId.value, System.currentTimeMillis()) sessionDao.updateActiveHost(sessionId.value, targetHostId.value, System.currentTimeMillis())
} }
suspend fun ensureAttachedRuntimeSession(
sessionId: UnifiedSessionId,
targetHostId: HermesHostId,
runtime: HermesHostRuntime
): HostSessionBinding {
val details = sessionDao.getSessionWithDetails(sessionId.value)
var binding = details?.bindings?.find { it.hostId == targetHostId.value }?.toDomain()
if (binding == null || binding.durableSessionId.value.isEmpty()) {
val createRes = runtime.gatewayClient.createSession(source = "android")
binding = HostSessionBinding(
hostId = targetHostId,
durableSessionId = createRes.durableId,
runtimeSessionId = createRes.runtimeId,
lastAttachedAt = System.currentTimeMillis(),
state = BindingState.READY,
syncedThroughMessageId = null,
syncedAt = null
)
sessionDao.insertOrUpdateBinding(binding.toEntity(sessionId.value))
runtimeToSessionMap[createRes.runtimeId.value] = Pair(sessionId, targetHostId)
return binding
}
// We have an existing durableSessionId.
// Check if current runtimeSessionId is already registered and valid, or if we need to resume
val currentRuntimeId = binding.runtimeSessionId.value
val isRegistered = currentRuntimeId.isNotEmpty() && runtimeToSessionMap.containsKey(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)
}
binding = binding.copy(
durableSessionId = resumeRes.durableId,
runtimeSessionId = resumeRes.runtimeId,
lastAttachedAt = System.currentTimeMillis(),
state = BindingState.READY
)
sessionDao.insertOrUpdateBinding(binding.toEntity(sessionId.value))
runtimeToSessionMap[resumeRes.runtimeId.value] = Pair(sessionId, targetHostId)
} else {
runtimeToSessionMap[currentRuntimeId] = Pair(sessionId, targetHostId)
}
return binding
}
suspend fun sendPrompt(sessionId: UnifiedSessionId, text: String): String { suspend fun sendPrompt(sessionId: UnifiedSessionId, text: String): String {
val details = sessionDao.getSessionWithDetails(sessionId.value) val details = sessionDao.getSessionWithDetails(sessionId.value)
?: throw IllegalArgumentException("Session not found: ${sessionId.value}") ?: throw IllegalArgumentException("Session not found: ${sessionId.value}")
@ -147,23 +227,8 @@ class UnifiedSessionRepository(
runtime.gatewayClient.awaitGatewayReady(10_000) runtime.gatewayClient.awaitGatewayReady(10_000)
} }
// Get or create native session binding for this host // Get or attach native session binding for this host
var binding = currentSession.bindings[targetHostId] val binding = ensureAttachedRuntimeSession(sessionId, targetHostId, runtime)
if (binding == null || binding.runtimeSessionId.value.isEmpty()) {
val createRes = runtime.gatewayClient.createSession(source = "android")
binding = HostSessionBinding(
hostId = targetHostId,
durableSessionId = createRes.durableId,
runtimeSessionId = createRes.runtimeId,
lastAttachedAt = System.currentTimeMillis(),
state = BindingState.READY,
syncedThroughMessageId = null,
syncedAt = null
)
sessionDao.insertOrUpdateBinding(binding.toEntity(sessionId.value))
}
runtimeToSessionMap[binding.runtimeSessionId.value] = Pair(sessionId, targetHostId)
// Context Synchronization // Context Synchronization
val hostsMap = connectionManager.hosts.value.associateBy { it.id } val hostsMap = connectionManager.hosts.value.associateBy { it.id }
@ -190,15 +255,6 @@ class UnifiedSessionRepository(
text text
} }
// Update binding sync status
sessionDao.updateBindingSync(
sessionId = sessionId.value,
hostId = targetHostId.value,
syncedThroughMessageId = syncResult.latestSyncedMessageId,
syncedAt = System.currentTimeMillis(),
state = BindingState.RUNNING.name
)
// Insert user message to timeline // Insert user message to timeline
val userMessage = UnifiedMessage( val userMessage = UnifiedMessage(
id = UUID.randomUUID().toString(), id = UUID.randomUUID().toString(),
@ -210,49 +266,70 @@ class UnifiedSessionRepository(
) )
insertMessageToSession(sessionId, userMessage) insertMessageToSession(sessionId, userMessage)
setExecuting(sessionId, true) setHostExecuting(sessionId, targetHostId, true)
sessionDao.updateBindingState(sessionId.value, targetHostId.value, BindingState.RUNNING.name)
return try { return try {
val result = runtime.gatewayClient.submitPrompt(binding.runtimeSessionId, promptToSend) val result = runtime.gatewayClient.submitPrompt(binding.runtimeSessionId, promptToSend)
// ONLY update binding sync status AFTER successful acceptance of prompt.submit!
sessionDao.updateBindingSync(
sessionId = sessionId.value,
hostId = targetHostId.value,
syncedThroughMessageId = syncResult.latestSyncedMessageId,
syncedAt = System.currentTimeMillis(),
state = BindingState.RUNNING.name
)
result.turnId ?: userMessage.id result.turnId ?: userMessage.id
} catch (e: Exception) { } catch (e: Exception) {
setExecuting(sessionId, false) setHostExecuting(sessionId, targetHostId, false)
sessionDao.updateBindingState(sessionId.value, targetHostId.value, BindingState.ERROR.name) sessionDao.updateBindingState(sessionId.value, targetHostId.value, BindingState.ERROR.name)
throw e throw e
} }
} }
suspend fun interruptSession(sessionId: UnifiedSessionId) { suspend fun interruptHost(sessionId: UnifiedSessionId, hostId: HermesHostId): Boolean {
val details = sessionDao.getSessionWithDetails(sessionId.value) ?: return val runtime = connectionManager.getRuntime(hostId) ?: return false
for (binding in details.bindings) { val details = sessionDao.getSessionWithDetails(sessionId.value) ?: return false
val runtime = connectionManager.getRuntime(HermesHostId(binding.hostId)) val binding = details.bindings.find { it.hostId == hostId.value } ?: return false
if (runtime != null && binding.runtimeSessionId.isNotEmpty()) {
try { val success = try {
runtime.gatewayClient.interruptSession(RuntimeSessionId(binding.runtimeSessionId)) if (binding.runtimeSessionId.isNotEmpty()) {
} catch (_: Exception) { runtime.gatewayClient.interruptSession(RuntimeSessionId(binding.runtimeSessionId))
} } else {
false
} }
} catch (_: Exception) {
false
} }
setExecuting(sessionId, false) setHostExecuting(sessionId, hostId, false)
sessionDao.updateBindingState(sessionId.value, hostId.value, BindingState.READY.name)
return success
}
suspend fun interruptSession(sessionId: UnifiedSessionId, targetHostId: HermesHostId? = null) {
val details = sessionDao.getSessionWithDetails(sessionId.value) ?: return
val hostToInterrupt = targetHostId ?: HermesHostId(details.session.activeHostId)
interruptHost(sessionId, hostToInterrupt)
} }
suspend fun respondApproval( suspend fun respondApproval(
hostId: HermesHostId, hostId: HermesHostId,
runtimeSessionId: RuntimeSessionId,
requestId: String, requestId: String,
choice: String, choice: String,
all: Boolean = false all: Boolean = false
): Boolean { ): Boolean {
val runtime = connectionManager.getRuntime(hostId) ?: return false val runtime = connectionManager.getRuntime(hostId) ?: return false
val approval = _activeApprovals.value.find { it.hostId == hostId && it.approval.requestId == requestId }
val sessionKey = "" // Gateway client handles request_id
val success = try { val success = try {
runtime.gatewayClient.respondApproval(sessionKey, requestId, choice, all) runtime.gatewayClient.respondApproval(runtimeSessionId.value, requestId, choice, all)
} catch (_: Exception) { } catch (_: Exception) {
true false
} }
if (success) { if (success) {
_activeApprovals.value = _activeApprovals.value.filterNot { _activeApprovals.value = _activeApprovals.value.filterNot {
it.hostId == hostId && it.approval.requestId == requestId it.hostId == hostId && it.runtimeSessionId == runtimeSessionId && it.approval.requestId == requestId
} }
} }
return success return success
@ -265,7 +342,11 @@ class UnifiedSessionRepository(
questionId: String? = null questionId: String? = null
): Boolean { ): Boolean {
val runtime = connectionManager.getRuntime(hostId) ?: return false val runtime = connectionManager.getRuntime(hostId) ?: return false
val success = runtime.gatewayClient.respondClarify(requestId, answer, questionId) val success = try {
runtime.gatewayClient.respondClarify(requestId, answer, questionId)
} catch (_: Exception) {
false
}
if (success) { if (success) {
if (_activeClarify.value?.hostId == hostId && _activeClarify.value?.request?.requestId == requestId) { if (_activeClarify.value?.hostId == hostId && _activeClarify.value?.request?.requestId == requestId) {
_activeClarify.value = null _activeClarify.value = null
@ -280,7 +361,11 @@ class UnifiedSessionRepository(
password: String password: String
): Boolean { ): Boolean {
val runtime = connectionManager.getRuntime(hostId) ?: return false val runtime = connectionManager.getRuntime(hostId) ?: return false
val success = runtime.gatewayClient.respondSudo(requestId, password) val success = try {
runtime.gatewayClient.respondSudo(requestId, password)
} catch (_: Exception) {
false
}
if (success) { if (success) {
if (_activeClarify.value?.hostId == hostId && _activeClarify.value?.request?.requestId == requestId) { if (_activeClarify.value?.hostId == hostId && _activeClarify.value?.request?.requestId == requestId) {
_activeClarify.value = null _activeClarify.value = null
@ -295,7 +380,11 @@ class UnifiedSessionRepository(
secret: String secret: String
): Boolean { ): Boolean {
val runtime = connectionManager.getRuntime(hostId) ?: return false val runtime = connectionManager.getRuntime(hostId) ?: return false
val success = runtime.gatewayClient.respondSecret(requestId, secret) val success = try {
runtime.gatewayClient.respondSecret(requestId, secret)
} catch (_: Exception) {
false
}
if (success) { if (success) {
if (_activeClarify.value?.hostId == hostId && _activeClarify.value?.request?.requestId == requestId) { if (_activeClarify.value?.hostId == hostId && _activeClarify.value?.request?.requestId == requestId) {
_activeClarify.value = null _activeClarify.value = null
@ -343,32 +432,61 @@ class UnifiedSessionRepository(
} }
} }
private fun setExecuting(sessionId: UnifiedSessionId, executing: Boolean) { private fun setHostExecuting(sessionId: UnifiedSessionId, hostId: HermesHostId, executing: Boolean) {
val flow = sessionExecutingState.computeIfAbsent(sessionId) { hostExecutingState.computeIfAbsent(Pair(sessionId, hostId)) {
MutableStateFlow(false) MutableStateFlow(false)
} }.value = executing
flow.value = executing
// Derive aggregate executing state for this session
val isAnyHostExecuting = hostExecutingState.entries
.filter { it.key.first == sessionId }
.any { it.value.value }
sessionExecutingState.computeIfAbsent(sessionId) {
MutableStateFlow(false)
}.value = isAnyHostExecuting
} }
private fun findSessionForHost(hostId: HermesHostId): UnifiedSessionId? { private fun findSessionForEvent(
for ((sessionId, flow) in sessionExecutingState) { hostId: HermesHostId,
if (flow.value) { sessionIdFromEvent: String?,
return sessionId messageId: String? = null,
toolId: String? = null
): UnifiedSessionId? {
// 1. Exact match via 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
}
} }
} }
// Fall back to active session, open session cache, or first known session
return sessions.value.find { it.activeHostId == hostId }?.id
?: sessionMessagesState.keys.firstOrNull()
?: sessions.value.firstOrNull()?.id
}
private fun findSessionForMessage(messageId: String, hostId: HermesHostId): UnifiedSessionId? { // 2. Exact match via messageId in active session messages
for ((sessionId, flow) in sessionMessagesState) { if (!messageId.isNullOrEmpty()) {
if (flow.value.any { it.id == messageId }) { for ((sessionId, flow) in sessionMessagesState) {
return sessionId if (flow.value.any { it.id == messageId && (it.hostId == hostId || it.hostId == null) }) {
return sessionId
}
} }
} }
return findSessionForHost(hostId)
// 3. Exact match via toolId in active session messages
if (!toolId.isNullOrEmpty()) {
for ((sessionId, flow) in sessionMessagesState) {
if (flow.value.any { it.tools.any { t -> t.id == toolId } }) {
return sessionId
}
}
}
// Strict: NO fallback to "any executing session" or "first session in list"
return null
} }
private fun handleHostGatewayEvent(hostEvent: HostGatewayEvent) { private fun handleHostGatewayEvent(hostEvent: HostGatewayEvent) {
@ -378,8 +496,8 @@ class UnifiedSessionRepository(
when (event) { when (event) {
is GatewayEvent.MessageStartEvent -> { is GatewayEvent.MessageStartEvent -> {
val sessionId = findSessionForMessage(event.messageId, hostId) ?: return val sessionId = findSessionForEvent(hostId, event.sessionId, messageId = event.messageId) ?: return
setExecuting(sessionId, true) setHostExecuting(sessionId, hostId, true)
val flow = sessionMessagesState.computeIfAbsent(sessionId) { MutableStateFlow(emptyList()) } val flow = sessionMessagesState.computeIfAbsent(sessionId) { MutableStateFlow(emptyList()) }
val existing = flow.value.find { it.id == event.messageId } val existing = flow.value.find { it.id == event.messageId }
if (existing == null) { if (existing == null) {
@ -397,8 +515,8 @@ class UnifiedSessionRepository(
} }
is GatewayEvent.MessageDeltaEvent -> { is GatewayEvent.MessageDeltaEvent -> {
val sessionId = findSessionForMessage(event.messageId, hostId) ?: return val sessionId = findSessionForEvent(hostId, event.sessionId, messageId = event.messageId) ?: return
setExecuting(sessionId, true) setHostExecuting(sessionId, hostId, true)
val flow = sessionMessagesState.computeIfAbsent(sessionId) { MutableStateFlow(emptyList()) } val flow = sessionMessagesState.computeIfAbsent(sessionId) { MutableStateFlow(emptyList()) }
val idx = flow.value.indexOfFirst { it.id == event.messageId } val idx = flow.value.indexOfFirst { it.id == event.messageId }
if (idx >= 0) { if (idx >= 0) {
@ -419,15 +537,18 @@ class UnifiedSessionRepository(
} }
is GatewayEvent.MessageInterimEvent -> { is GatewayEvent.MessageInterimEvent -> {
val sessionId = findSessionForMessage(event.messageId, hostId) ?: return val sessionId = findSessionForEvent(hostId, event.sessionId, messageId = event.messageId) ?: return
updateMessageInSession(sessionId, event.messageId) { updateMessageInSession(sessionId, event.messageId) {
it.copy(content = event.content, isStreaming = true) it.copy(content = event.content, isStreaming = true)
} }
} }
is GatewayEvent.MessageCompleteEvent -> { is GatewayEvent.MessageCompleteEvent -> {
val sessionId = findSessionForMessage(event.messageId, hostId) ?: return val sessionId = findSessionForEvent(hostId, event.sessionId, messageId = event.messageId) ?: return
setExecuting(sessionId, false) setHostExecuting(sessionId, hostId, false)
scope.launch {
sessionDao.updateBindingState(sessionId.value, hostId.value, BindingState.READY.name)
}
val flow = sessionMessagesState.computeIfAbsent(sessionId) { MutableStateFlow(emptyList()) } val flow = sessionMessagesState.computeIfAbsent(sessionId) { MutableStateFlow(emptyList()) }
val idx = flow.value.indexOfFirst { it.id == event.messageId } val idx = flow.value.indexOfFirst { it.id == event.messageId }
if (idx >= 0) { if (idx >= 0) {
@ -451,56 +572,56 @@ class UnifiedSessionRepository(
} }
is GatewayEvent.ThinkingDeltaEvent -> { is GatewayEvent.ThinkingDeltaEvent -> {
val sessionId = findSessionForHost(hostId) ?: return val sessionId = findSessionForEvent(hostId, event.sessionId, messageId = event.messageId) ?: return
val flow = sessionMessagesState.computeIfAbsent(sessionId) { MutableStateFlow(emptyList()) } val flow = sessionMessagesState.computeIfAbsent(sessionId) { MutableStateFlow(emptyList()) }
val lastAssistant = flow.value.lastOrNull { it.role == MessageRole.ASSISTANT && it.hostId == hostId } val targetAssistant = flow.value.lastOrNull { (it.id == event.messageId || it.role == MessageRole.ASSISTANT) && it.hostId == hostId }
if (lastAssistant != null) { if (targetAssistant != null) {
updateMessageInSession(sessionId, lastAssistant.id) { updateMessageInSession(sessionId, targetAssistant.id) {
it.copy(thinking = (it.thinking ?: "") + event.delta) it.copy(thinking = (it.thinking ?: "") + event.delta)
} }
} }
} }
is GatewayEvent.ReasoningDeltaEvent -> { is GatewayEvent.ReasoningDeltaEvent -> {
val sessionId = findSessionForHost(hostId) ?: return val sessionId = findSessionForEvent(hostId, event.sessionId, messageId = event.messageId) ?: return
val flow = sessionMessagesState.computeIfAbsent(sessionId) { MutableStateFlow(emptyList()) } val flow = sessionMessagesState.computeIfAbsent(sessionId) { MutableStateFlow(emptyList()) }
val lastAssistant = flow.value.lastOrNull { it.role == MessageRole.ASSISTANT && it.hostId == hostId } val targetAssistant = flow.value.lastOrNull { (it.id == event.messageId || it.role == MessageRole.ASSISTANT) && it.hostId == hostId }
if (lastAssistant != null) { if (targetAssistant != null) {
updateMessageInSession(sessionId, lastAssistant.id) { updateMessageInSession(sessionId, targetAssistant.id) {
it.copy(thinking = (it.thinking ?: "") + event.delta) it.copy(thinking = (it.thinking ?: "") + event.delta)
} }
} }
} }
is GatewayEvent.ReasoningAvailableEvent -> { is GatewayEvent.ReasoningAvailableEvent -> {
val sessionId = findSessionForHost(hostId) ?: return val sessionId = findSessionForEvent(hostId, event.sessionId, messageId = event.messageId) ?: return
val flow = sessionMessagesState.computeIfAbsent(sessionId) { MutableStateFlow(emptyList()) } val flow = sessionMessagesState.computeIfAbsent(sessionId) { MutableStateFlow(emptyList()) }
val lastAssistant = flow.value.lastOrNull { it.role == MessageRole.ASSISTANT && it.hostId == hostId } val targetAssistant = flow.value.lastOrNull { (it.id == event.messageId || it.role == MessageRole.ASSISTANT) && it.hostId == hostId }
if (lastAssistant != null) { if (targetAssistant != null) {
updateMessageInSession(sessionId, lastAssistant.id) { updateMessageInSession(sessionId, targetAssistant.id) {
it.copy(thinking = event.reasoning) it.copy(thinking = event.reasoning)
} }
} }
} }
is GatewayEvent.ToolStartEvent -> { is GatewayEvent.ToolStartEvent -> {
val sessionId = findSessionForHost(hostId) ?: return val sessionId = findSessionForEvent(hostId, event.sessionId) ?: return
val tool = ToolActivity(id = event.toolId, name = event.name, status = "running") val tool = ToolActivity(id = event.toolId, name = event.name, status = "running")
attachToolToSessionMessage(sessionId, hostId, tool) attachToolToSessionMessage(sessionId, hostId, tool)
} }
is GatewayEvent.ToolProgressEvent -> { is GatewayEvent.ToolProgressEvent -> {
val sessionId = findSessionForHost(hostId) ?: return val sessionId = findSessionForEvent(hostId, event.sessionId, toolId = event.toolId) ?: return
updateToolInSessionMessage(sessionId, event.toolId) { it.copy(progress = event.progress) } updateToolInSessionMessage(sessionId, event.toolId) { it.copy(progress = event.progress) }
} }
is GatewayEvent.ToolGeneratingEvent -> { is GatewayEvent.ToolGeneratingEvent -> {
val sessionId = findSessionForHost(hostId) ?: return val sessionId = findSessionForEvent(hostId, event.sessionId, toolId = event.toolId) ?: return
updateToolInSessionMessage(sessionId, event.toolId) { it.copy(status = "generating") } updateToolInSessionMessage(sessionId, event.toolId) { it.copy(status = "generating") }
} }
is GatewayEvent.ToolCompleteEvent -> { is GatewayEvent.ToolCompleteEvent -> {
val sessionId = findSessionForHost(hostId) ?: return val sessionId = findSessionForEvent(hostId, event.sessionId, toolId = event.toolId) ?: return
updateToolInSessionMessage(sessionId, event.toolId) { updateToolInSessionMessage(sessionId, event.toolId) {
it.copy( it.copy(
status = if (event.isError) "failed" else "completed", status = if (event.isError) "failed" else "completed",
@ -511,6 +632,8 @@ class UnifiedSessionRepository(
} }
is GatewayEvent.ApprovalRequestEvent -> { is GatewayEvent.ApprovalRequestEvent -> {
val runtimeSessionIdVal = event.sessionKey ?: event.sessionId ?: ""
val runtimeSessionId = RuntimeSessionId(runtimeSessionIdVal)
val approval = HermesApproval( val approval = HermesApproval(
requestId = event.requestId, requestId = event.requestId,
command = event.command, command = event.command,
@ -520,44 +643,66 @@ class UnifiedSessionRepository(
val attributed = HostAttributedApproval( val attributed = HostAttributedApproval(
hostId = hostId, hostId = hostId,
hostDisplayName = hostName, hostDisplayName = hostName,
runtimeSessionId = runtimeSessionId,
approval = approval approval = approval
) )
_activeApprovals.value = _activeApprovals.value.filterNot { _activeApprovals.value = _activeApprovals.value.filterNot {
it.hostId == hostId && it.approval.requestId == event.requestId it.hostId == hostId && it.runtimeSessionId == runtimeSessionId && it.approval.requestId == event.requestId
} + attributed } + attributed
} }
is GatewayEvent.ClarifyRequestEvent -> { is GatewayEvent.ClarifyRequestEvent -> {
val runtimeSessionIdVal = event.sessionId
val req = HermesClarifyRequest( val req = HermesClarifyRequest(
requestId = event.requestId, requestId = event.requestId,
questionId = event.questionId, questionId = event.questionId,
question = event.question, question = event.question,
promptType = ClarifyType.CLARIFY promptType = ClarifyType.CLARIFY
) )
_activeClarify.value = HostAttributedClarify(hostId, hostName, req) _activeClarify.value = HostAttributedClarify(
hostId = hostId,
hostDisplayName = hostName,
runtimeSessionId = runtimeSessionIdVal?.let { RuntimeSessionId(it) },
request = req
)
} }
is GatewayEvent.SudoRequestEvent -> { is GatewayEvent.SudoRequestEvent -> {
val runtimeSessionIdVal = event.sessionId
val req = HermesClarifyRequest( val req = HermesClarifyRequest(
requestId = event.requestId, requestId = event.requestId,
question = event.question, question = event.question,
promptType = ClarifyType.SUDO promptType = ClarifyType.SUDO
) )
_activeClarify.value = HostAttributedClarify(hostId, hostName, req) _activeClarify.value = HostAttributedClarify(
hostId = hostId,
hostDisplayName = hostName,
runtimeSessionId = runtimeSessionIdVal?.let { RuntimeSessionId(it) },
request = req
)
} }
is GatewayEvent.SecretRequestEvent -> { is GatewayEvent.SecretRequestEvent -> {
val runtimeSessionIdVal = event.sessionId
val req = HermesClarifyRequest( val req = HermesClarifyRequest(
requestId = event.requestId, requestId = event.requestId,
question = event.question, question = event.question,
promptType = ClarifyType.SECRET promptType = ClarifyType.SECRET
) )
_activeClarify.value = HostAttributedClarify(hostId, hostName, req) _activeClarify.value = HostAttributedClarify(
hostId = hostId,
hostDisplayName = hostName,
runtimeSessionId = runtimeSessionIdVal?.let { RuntimeSessionId(it) },
request = req
)
} }
is GatewayEvent.ErrorEvent -> { is GatewayEvent.ErrorEvent -> {
val sessionId = findSessionForHost(hostId) ?: return val sessionId = findSessionForEvent(hostId, event.sessionId) ?: return
setExecuting(sessionId, false) setHostExecuting(sessionId, hostId, false)
scope.launch {
sessionDao.updateBindingState(sessionId.value, hostId.value, BindingState.ERROR.name)
}
} }
else -> {} else -> {}
@ -586,20 +731,14 @@ class UnifiedSessionRepository(
} }
private fun updateToolInSessionMessage( private fun updateToolInSessionMessage(
sessionId: UnifiedSessionId?, sessionId: UnifiedSessionId,
toolId: String, toolId: String,
transform: (ToolActivity) -> ToolActivity transform: (ToolActivity) -> ToolActivity
) { ) {
val targetSessionId = sessionId?.takeIf { sId -> val flow = sessionMessagesState.computeIfAbsent(sessionId) { MutableStateFlow(emptyList()) }
sessionMessagesState[sId]?.value?.any { msg -> msg.tools.any { it.id == toolId } } == true
} ?: sessionMessagesState.entries.firstOrNull { (_, flow) ->
flow.value.any { msg -> msg.tools.any { it.id == toolId } }
}?.key ?: sessionId ?: return
val flow = sessionMessagesState.computeIfAbsent(targetSessionId) { MutableStateFlow(emptyList()) }
val targetMsg = flow.value.lastOrNull { msg -> msg.tools.any { it.id == toolId } } val targetMsg = flow.value.lastOrNull { msg -> msg.tools.any { it.id == toolId } }
if (targetMsg != null) { if (targetMsg != null) {
updateMessageInSession(targetSessionId, targetMsg.id) { msg -> updateMessageInSession(sessionId, targetMsg.id) { msg ->
val updatedTools = msg.tools.map { if (it.id == toolId) transform(it) else it } val updatedTools = msg.tools.map { if (it.id == toolId) transform(it) else it }
msg.copy(tools = updatedTools) msg.copy(tools = updatedTools)
} }

View file

@ -29,9 +29,9 @@ import kotlin.random.Random
class HermesHostRuntime( class HermesHostRuntime(
initialHost: HermesHost, initialHost: HermesHost,
val restClient: HermesRestClient = HermesRestClient(), val restClient: HermesRestClient = HermesRestClient(),
val gatewayClient: JsonRpcGatewayClient = JsonRpcGatewayClient(),
val tokenVault: TokenVault, val tokenVault: TokenVault,
val scope: CoroutineScope = CoroutineScope(SupervisorJob() + Dispatchers.Default) val scope: CoroutineScope = CoroutineScope(SupervisorJob() + Dispatchers.Default),
val gatewayClient: JsonRpcGatewayClient = JsonRpcGatewayClient(scope = scope)
) { ) {
private val _host = MutableStateFlow(initialHost) private val _host = MutableStateFlow(initialHost)
val host: StateFlow<HermesHost> = _host.asStateFlow() val host: StateFlow<HermesHost> = _host.asStateFlow()

View file

@ -64,7 +64,7 @@ object UnifiedContextBuilder {
val sb = StringBuilder() val sb = StringBuilder()
sb.appendLine("[Unified Hermes Session Context Transfer]") sb.appendLine("[Unified Hermes Session Context Transfer]")
sb.appendLine("You are continuing a unified conversation that previously ran across Hermes host instances.") sb.appendLine("You are continuing a unified conversation that previously ran across Hermes host instances.")
sb.appendLine("Target Host: ${targetHost.displayName} (${targetHost.baseUrl})") sb.appendLine("Target Host: ${targetHost.displayName}")
sb.appendLine("Session Title: ${session.title}") sb.appendLine("Session Title: ${session.title}")
sb.appendLine("--- Prior Conversation Turns ---") sb.appendLine("--- Prior Conversation Turns ---")

View file

@ -17,6 +17,7 @@ import androidx.compose.foundation.layout.padding
import androidx.compose.foundation.layout.size import androidx.compose.foundation.layout.size
import androidx.compose.foundation.layout.width import androidx.compose.foundation.layout.width
import androidx.compose.foundation.lazy.LazyColumn import androidx.compose.foundation.lazy.LazyColumn
import androidx.compose.foundation.lazy.LazyRow
import androidx.compose.foundation.lazy.items import androidx.compose.foundation.lazy.items
import androidx.compose.foundation.lazy.rememberLazyListState import androidx.compose.foundation.lazy.rememberLazyListState
import androidx.compose.foundation.shape.CircleShape import androidx.compose.foundation.shape.CircleShape
@ -97,99 +98,126 @@ fun ChatScreen(
Scaffold( Scaffold(
topBar = { topBar = {
TopAppBar( Column {
title = { TopAppBar(
Column { title = {
Text( Column {
text = currentSession?.title?.ifEmpty { "Unified Chat" } ?: "Unified Chat", Text(
style = MaterialTheme.typography.titleMedium, text = currentSession?.title?.ifEmpty { "Unified Chat" } ?: "Unified Chat",
fontWeight = FontWeight.Bold, style = MaterialTheme.typography.titleMedium,
maxLines = 1, fontWeight = FontWeight.Bold,
overflow = TextOverflow.Ellipsis maxLines = 1,
) overflow = TextOverflow.Ellipsis
// Active host selector chip )
Box { // Active host selector chip
Row( Box {
verticalAlignment = Alignment.CenterVertically, Row(
modifier = Modifier verticalAlignment = Alignment.CenterVertically,
.clip(RoundedCornerShape(6.dp))
.background(MaterialTheme.colorScheme.surfaceVariant)
.clickable { viewModel.setHostDropdownExpanded(true) }
.padding(horizontal = 6.dp, vertical = 2.dp)
) {
val isOnline = activeHost?.lastKnownStatus == HostStatus.ONLINE
Box(
modifier = Modifier modifier = Modifier
.size(6.dp) .clip(RoundedCornerShape(6.dp))
.clip(CircleShape) .background(MaterialTheme.colorScheme.surfaceVariant)
.background(if (isOnline) Color(0xFF10B981) else Color(0xFF94A3B8)) .clickable { viewModel.setHostDropdownExpanded(true) }
) .padding(horizontal = 6.dp, vertical = 2.dp)
Spacer(modifier = Modifier.width(4.dp))
Text(
text = activeHost?.displayName ?: "Select Host",
style = MaterialTheme.typography.labelSmall,
color = MaterialTheme.colorScheme.onSurfaceVariant,
fontWeight = FontWeight.SemiBold
)
Spacer(modifier = Modifier.width(2.dp))
Icon(
Icons.Default.ExpandMore,
contentDescription = "Switch Host",
modifier = Modifier.size(14.dp),
tint = MaterialTheme.colorScheme.onSurfaceVariant
)
}
DropdownMenu(
expanded = uiState.activeHostDropdownExpanded,
onDismissRequest = { viewModel.setHostDropdownExpanded(false) }
) { ) {
hosts.forEach { host -> val isOnline = activeHost?.lastKnownStatus == HostStatus.ONLINE
DropdownMenuItem( Box(
text = { modifier = Modifier
Row(verticalAlignment = Alignment.CenterVertically) { .size(6.dp)
val online = host.lastKnownStatus == HostStatus.ONLINE .clip(CircleShape)
Box( .background(if (isOnline) Color(0xFF10B981) else Color(0xFF94A3B8))
modifier = Modifier
.size(8.dp)
.clip(CircleShape)
.background(if (online) Color(0xFF10B981) else Color(0xFF94A3B8))
)
Spacer(modifier = Modifier.width(8.dp))
Text(
text = host.displayName,
fontWeight = if (host.id == currentSession?.activeHostId) FontWeight.Bold else FontWeight.Normal
)
}
},
onClick = { viewModel.switchActiveHost(host.id) }
) )
Spacer(modifier = Modifier.width(4.dp))
Text(
text = activeHost?.displayName ?: "Select Host",
style = MaterialTheme.typography.labelSmall,
color = MaterialTheme.colorScheme.onSurfaceVariant,
fontWeight = FontWeight.SemiBold
)
Spacer(modifier = Modifier.width(2.dp))
Icon(
Icons.Default.ExpandMore,
contentDescription = "Switch Host",
modifier = Modifier.size(14.dp),
tint = MaterialTheme.colorScheme.onSurfaceVariant
)
}
DropdownMenu(
expanded = uiState.activeHostDropdownExpanded,
onDismissRequest = { viewModel.setHostDropdownExpanded(false) }
) {
hosts.forEach { host ->
DropdownMenuItem(
text = {
Row(verticalAlignment = Alignment.CenterVertically) {
val online = host.lastKnownStatus == HostStatus.ONLINE
Box(
modifier = Modifier
.size(8.dp)
.clip(CircleShape)
.background(if (online) Color(0xFF10B981) else Color(0xFF94A3B8))
)
Spacer(modifier = Modifier.width(8.dp))
Text(
text = host.displayName,
fontWeight = if (host.id == currentSession?.activeHostId) FontWeight.Bold else FontWeight.Normal
)
}
},
onClick = { viewModel.switchActiveHost(host.id) }
)
}
} }
} }
} }
},
navigationIcon = {
IconButton(onClick = onNavigateBack) {
Icon(Icons.AutoMirrored.Filled.ArrowBack, contentDescription = "Back")
}
},
actions = {
if (isExecuting) {
Button(
onClick = { viewModel.interruptSession() },
colors = ButtonDefaults.buttonColors(containerColor = Color(0xFFEF4444)),
shape = RoundedCornerShape(8.dp),
contentPadding = PaddingValues(horizontal = 10.dp, vertical = 4.dp),
modifier = Modifier.padding(end = 8.dp)
) {
Icon(Icons.Default.Stop, contentDescription = null, modifier = Modifier.size(16.dp))
Spacer(modifier = Modifier.width(4.dp))
Text("Stop", fontSize = 12.sp, fontWeight = FontWeight.Bold)
}
}
} }
}, )
navigationIcon = {
IconButton(onClick = onNavigateBack) { // Multi-host Status Strip showing simultaneous statuses (e.g. PC1 Running, Linux Active, PC3 Offline)
Icon(Icons.AutoMirrored.Filled.ArrowBack, contentDescription = "Back") if (hosts.isNotEmpty()) {
} LazyRow(
}, modifier = Modifier
actions = { .fillMaxWidth()
if (isExecuting) { .background(MaterialTheme.colorScheme.surfaceVariant.copy(alpha = 0.5f))
Button( .padding(horizontal = 16.dp, vertical = 6.dp),
onClick = { viewModel.interruptSession() }, horizontalArrangement = Arrangement.spacedBy(8.dp),
colors = ButtonDefaults.buttonColors(containerColor = Color(0xFFEF4444)), verticalAlignment = Alignment.CenterVertically
shape = RoundedCornerShape(8.dp), ) {
contentPadding = PaddingValues(horizontal = 10.dp, vertical = 4.dp), items(hosts, key = { it.id.value }) { host ->
modifier = Modifier.padding(end = 8.dp) val isRunning by viewModel.getHostExecuting(host.id).collectAsState()
) { val isActive = host.id == currentSession?.activeHostId
Icon(Icons.Default.Stop, contentDescription = null, modifier = Modifier.size(16.dp))
Spacer(modifier = Modifier.width(4.dp)) HostStatusChip(
Text("Stop", fontSize = 12.sp, fontWeight = FontWeight.Bold) host = host,
isActive = isActive,
isRunning = isRunning,
onClick = { viewModel.switchActiveHost(host.id) },
onStop = { viewModel.interruptHost(host.id) }
)
} }
} }
} }
) }
}, },
bottomBar = { bottomBar = {
ChatInputBar( ChatInputBar(
@ -219,11 +247,11 @@ fun ChatScreen(
} }
} }
items(approvals, key = { it.hostId.value + it.approval.requestId }) { approval -> items(approvals, key = { it.hostId.value + it.runtimeSessionId.value + it.approval.requestId }) { approval ->
ApprovalCard( ApprovalCard(
attributedApproval = approval, attributedApproval = approval,
onRespond = { choice, all -> onRespond = { choice, all ->
viewModel.respondApproval(approval.hostId, approval.approval.requestId, choice, all) viewModel.respondApproval(approval.hostId, approval.runtimeSessionId, approval.approval.requestId, choice, all)
} }
) )
} }
@ -261,6 +289,92 @@ fun ChatScreen(
} }
} }
@Composable
fun HostStatusChip(
host: HermesHost,
isActive: Boolean,
isRunning: Boolean,
onClick: () -> Unit,
onStop: () -> Unit
) {
val statusText = when {
isRunning -> "Running"
isActive -> "Active"
host.lastKnownStatus == HostStatus.ONLINE -> "Online"
host.lastKnownStatus == HostStatus.CONNECTING -> "Connecting"
host.lastKnownStatus == HostStatus.AUTH_EXPIRED -> "Auth Expired"
else -> "Offline"
}
val chipBg = when {
isRunning -> Color(0xFFF59E0B).copy(alpha = 0.15f)
isActive -> Color(0xFF38BDF8).copy(alpha = 0.15f)
else -> MaterialTheme.colorScheme.surface
}
val chipBorder = when {
isRunning -> Color(0xFFF59E0B)
isActive -> Color(0xFF38BDF8)
else -> Color(0xFF475569)
}
Row(
verticalAlignment = Alignment.CenterVertically,
modifier = Modifier
.clip(RoundedCornerShape(8.dp))
.background(chipBg)
.clickable { onClick() }
.padding(horizontal = 8.dp, vertical = 4.dp)
) {
if (isRunning) {
CircularProgressIndicator(
modifier = Modifier.size(10.dp),
strokeWidth = 1.5.dp,
color = Color(0xFFF59E0B)
)
} else {
val dotColor = when (host.lastKnownStatus) {
HostStatus.ONLINE -> Color(0xFF10B981)
HostStatus.CONNECTING -> Color(0xFFF59E0B)
HostStatus.AUTH_EXPIRED -> Color(0xFFEF4444)
else -> Color(0xFF94A3B8)
}
Box(
modifier = Modifier
.size(6.dp)
.clip(CircleShape)
.background(dotColor)
)
}
Spacer(modifier = Modifier.width(6.dp))
Text(
text = "${host.displayName}: $statusText",
fontSize = 11.sp,
fontWeight = if (isActive || isRunning) FontWeight.Bold else FontWeight.Normal,
color = if (isRunning) Color(0xFFF59E0B) else if (isActive) Color(0xFF38BDF8) else MaterialTheme.colorScheme.onSurface
)
if (isRunning) {
Spacer(modifier = Modifier.width(4.dp))
Box(
modifier = Modifier
.size(14.dp)
.clip(CircleShape)
.background(Color(0xFFEF4444))
.clickable { onStop() },
contentAlignment = Alignment.Center
) {
Icon(
Icons.Default.Stop,
contentDescription = "Stop Host",
tint = Color.White,
modifier = Modifier.size(10.dp)
)
}
}
}
}
@Composable @Composable
fun TransferSeparator(message: UnifiedMessage) { fun TransferSeparator(message: UnifiedMessage) {
Row( Row(

View file

@ -44,6 +44,10 @@ class ChatViewModel(
} }
} }
fun getHostExecuting(hostId: HermesHostId): StateFlow<Boolean> {
return sessionRepo.getHostExecuting(sessionId, hostId)
}
fun updateInputText(text: String) { fun updateInputText(text: String) {
_uiState.value = _uiState.value.copy(inputText = text) _uiState.value = _uiState.value.copy(inputText = text)
} }
@ -82,19 +86,30 @@ class ChatViewModel(
} }
} }
fun interruptSession() { fun interruptSession(hostId: HermesHostId? = null) {
viewModelScope.launch { viewModelScope.launch {
sessionRepo.interruptSession(sessionId) sessionRepo.interruptSession(sessionId, hostId)
} }
} }
fun respondApproval(hostId: HermesHostId, requestId: String, choice: String, all: Boolean = false) { fun interruptHost(hostId: HermesHostId) {
viewModelScope.launch { viewModelScope.launch {
try { sessionRepo.interruptHost(sessionId, hostId)
sessionRepo.respondApproval(hostId, requestId, choice, all) }
} catch (e: Exception) { }
fun respondApproval(
hostId: HermesHostId,
runtimeSessionId: RuntimeSessionId,
requestId: String,
choice: String,
all: Boolean = false
) {
viewModelScope.launch {
val success = sessionRepo.respondApproval(hostId, runtimeSessionId, requestId, choice, all)
if (!success) {
_uiState.value = _uiState.value.copy( _uiState.value = _uiState.value.copy(
error = e.localizedMessage ?: "Failed to respond to approval" error = "Failed to submit approval response"
) )
} }
} }
@ -102,17 +117,16 @@ class ChatViewModel(
fun respondClarify(attributed: HostAttributedClarify, answer: String) { fun respondClarify(attributed: HostAttributedClarify, answer: String) {
viewModelScope.launch { viewModelScope.launch {
try { val hostId = attributed.hostId
val hostId = attributed.hostId val req = attributed.request
val req = attributed.request val success = when (req.promptType) {
when (req.promptType) { ClarifyType.CLARIFY -> sessionRepo.respondClarify(hostId, req.requestId, answer, req.questionId)
ClarifyType.CLARIFY -> sessionRepo.respondClarify(hostId, req.requestId, answer, req.questionId) ClarifyType.SUDO -> sessionRepo.respondSudo(hostId, req.requestId, answer)
ClarifyType.SUDO -> sessionRepo.respondSudo(hostId, req.requestId, answer) ClarifyType.SECRET -> sessionRepo.respondSecret(hostId, req.requestId, answer)
ClarifyType.SECRET -> sessionRepo.respondSecret(hostId, req.requestId, answer) }
} if (!success) {
} catch (e: Exception) {
_uiState.value = _uiState.value.copy( _uiState.value = _uiState.value.copy(
error = e.localizedMessage ?: "Failed to respond to clarification" error = "Failed to submit clarification response"
) )
} }
} }

View file

@ -1,7 +1,7 @@
package app.hermes.mobile.core.repository package app.hermes.mobile.core.repository
import app.hermes.mobile.core.model.* import app.hermes.mobile.core.model.*
import app.hermes.mobile.core.network.HermesRestClient import app.hermes.mobile.core.network.ConnectionState
import app.hermes.mobile.core.network.JsonRpcGatewayClient import app.hermes.mobile.core.network.JsonRpcGatewayClient
import app.hermes.mobile.core.runtime.HermesConnectionManager import app.hermes.mobile.core.runtime.HermesConnectionManager
import app.hermes.mobile.core.runtime.HermesHostRuntime import app.hermes.mobile.core.runtime.HermesHostRuntime
@ -11,14 +11,23 @@ import app.hermes.mobile.core.storage.FakeUnifiedSessionDao
import kotlinx.coroutines.CoroutineScope import kotlinx.coroutines.CoroutineScope
import kotlinx.coroutines.Dispatchers import kotlinx.coroutines.Dispatchers
import kotlinx.coroutines.ExperimentalCoroutinesApi import kotlinx.coroutines.ExperimentalCoroutinesApi
import kotlinx.coroutines.runBlocking
import kotlinx.coroutines.test.StandardTestDispatcher import kotlinx.coroutines.test.StandardTestDispatcher
import kotlinx.coroutines.test.resetMain import kotlinx.coroutines.test.resetMain
import kotlinx.coroutines.test.runTest import kotlinx.coroutines.test.runTest
import kotlinx.coroutines.test.setMain import kotlinx.coroutines.test.setMain
import kotlinx.serialization.json.buildJsonObject import kotlinx.serialization.json.buildJsonObject
import kotlinx.serialization.json.put import kotlinx.serialization.json.put
import okhttp3.Response
import okhttp3.WebSocket
import okhttp3.WebSocketListener
import okhttp3.mockwebserver.Dispatcher
import okhttp3.mockwebserver.MockResponse
import okhttp3.mockwebserver.MockWebServer
import okhttp3.mockwebserver.RecordedRequest
import org.junit.After import org.junit.After
import org.junit.Assert.assertEquals import org.junit.Assert.assertEquals
import org.junit.Assert.assertFalse
import org.junit.Assert.assertNotNull import org.junit.Assert.assertNotNull
import org.junit.Assert.assertTrue import org.junit.Assert.assertTrue
import org.junit.Before import org.junit.Before
@ -33,6 +42,7 @@ class ApprovalRoutingTest {
private lateinit var tokenVault: InMemoryTokenVault private lateinit var tokenVault: InMemoryTokenVault
private lateinit var connectionManager: HermesConnectionManager private lateinit var connectionManager: HermesConnectionManager
private lateinit var sessionRepo: UnifiedSessionRepository private lateinit var sessionRepo: UnifiedSessionRepository
private lateinit var mockServer: MockWebServer
private val host1Id = HermesHostId("server-prod") private val host1Id = HermesHostId("server-prod")
private val host2Id = HermesHostId("server-dev") private val host2Id = HermesHostId("server-dev")
@ -43,6 +53,8 @@ class ApprovalRoutingTest {
hostDao = FakeHostDao() hostDao = FakeHostDao()
sessionDao = FakeUnifiedSessionDao() sessionDao = FakeUnifiedSessionDao()
tokenVault = InMemoryTokenVault() tokenVault = InMemoryTokenVault()
mockServer = MockWebServer()
mockServer.start()
connectionManager = HermesConnectionManager( connectionManager = HermesConnectionManager(
hostDao = hostDao, hostDao = hostDao,
@ -60,66 +72,149 @@ class ApprovalRoutingTest {
@After @After
fun tearDown() { fun tearDown() {
Dispatchers.resetMain() Dispatchers.resetMain()
try {
mockServer.shutdown()
} catch (_: Exception) {
}
} }
@Test @Test
fun testApprovalAttributionAndRemoval() = runTest(testDispatcher) { fun testApprovalFailureRemainsVisibleAndNotResolved() = runTest(testDispatcher) {
val host1 = HermesHost(id = host1Id, displayName = "Prod Server", baseUrl = "http://prod:9119") val host1 = HermesHost(id = host1Id, displayName = "Prod Server", baseUrl = "http://prod:9119")
val host2 = HermesHost(id = host2Id, displayName = "Dev Server", baseUrl = "http://dev:9119")
connectionManager.addHost(host1) connectionManager.addHost(host1)
connectionManager.addHost(host2)
testScheduler.advanceUntilIdle() testScheduler.advanceUntilIdle()
val runtime1 = connectionManager.getRuntime(host1Id) val runtime1 = connectionManager.getRuntime(host1Id)
val runtime2 = connectionManager.getRuntime(host2Id)
assertNotNull(runtime1) assertNotNull(runtime1)
assertNotNull(runtime2)
// Simulate approval request from Prod Server // Simulate incoming approval request with specific runtime session ID
val prodEventJson = buildJsonObject { val prodEventJson = buildJsonObject {
put("method", "event") put("method", "event")
put("params", buildJsonObject { put("params", buildJsonObject {
put("event", "approval.request") put("event", "approval.request")
put("request_id", "req_prod_1") put("request_id", "req_prod_1")
put("session_key", "runtime_session_prod_99")
put("command", "systemctl restart nginx") put("command", "systemctl restart nginx")
put("description", "Restart web server") put("description", "Restart web server")
}) })
} }
runtime1?.gatewayClient?.handleIncomingMessage(prodEventJson.toString()) runtime1?.gatewayClient?.handleIncomingMessage(prodEventJson.toString())
// Simulate approval request from Dev Server
val devEventJson = buildJsonObject {
put("method", "event")
put("params", buildJsonObject {
put("event", "approval.request")
put("request_id", "req_dev_1")
put("command", "docker compose down")
put("description", "Stop containers")
})
}
runtime2?.gatewayClient?.handleIncomingMessage(devEventJson.toString())
testScheduler.advanceUntilIdle() testScheduler.advanceUntilIdle()
val approvals = sessionRepo.activeApprovals.value val approvals = sessionRepo.activeApprovals.value
assertEquals(2, approvals.size) assertEquals(1, approvals.size)
val approval = approvals.first()
assertEquals(RuntimeSessionId("runtime_session_prod_99"), approval.runtimeSessionId)
assertEquals("req_prod_1", approval.approval.requestId)
val prodApproval = approvals.find { it.hostId == host1Id } // Attempting to respond when WebSocket is disconnected MUST return false and KEEP approval
val devApproval = approvals.find { it.hostId == host2Id } val result = sessionRepo.respondApproval(
hostId = host1Id,
assertNotNull(prodApproval) runtimeSessionId = RuntimeSessionId("runtime_session_prod_99"),
assertNotNull(devApproval) requestId = "req_prod_1",
assertEquals("Prod Server", prodApproval?.hostDisplayName) choice = "once",
assertEquals("Dev Server", devApproval?.hostDisplayName) all = false
assertEquals("systemctl restart nginx", prodApproval?.approval?.command) )
// Responding to prod approval removes it while keeping dev approval
sessionRepo.respondApproval(host1Id, "req_prod_1", "once", false)
testScheduler.advanceUntilIdle() testScheduler.advanceUntilIdle()
val remainingApprovals = sessionRepo.activeApprovals.value assertFalse("Expected false when RPC fails due to disconnected socket", result)
assertEquals(1, remainingApprovals.size) assertEquals("Approval must not be removed on failure", 1, sessionRepo.activeApprovals.value.size)
assertEquals(host2Id, remainingApprovals.first().hostId) }
@Test
fun testApprovalSuccessRemovesCardAndSendsCorrectRpc() = runBlocking(Dispatchers.Default) {
val receivedMessages = mutableListOf<String>()
mockServer.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("""{"event":"gateway.ready","data":{"version":"1.0.0"}}""")
}
override fun onMessage(webSocket: WebSocket, text: String) {
receivedMessages.add(text)
if (text.contains("approval.respond")) {
webSocket.send("""{"jsonrpc":"2.0","id":"a1","result":{"status":"ok"}}""")
}
}
})
}
else -> MockResponse().setResponseCode(404)
}
}
}
val wsUrl = mockServer.url("").toString().removeSuffix("/")
val testHostDao = FakeHostDao()
val testSessionDao = FakeUnifiedSessionDao()
val testTokenVault = InMemoryTokenVault()
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 = "Prod Server", baseUrl = wsUrl, allowCleartext = true)
testConnectionManager.addHost(host1)
val runtime1 = testConnectionManager.getRuntime(host1Id)
assertNotNull(runtime1)
runtime1!!.connect()
runtime1.gatewayClient.awaitGatewayReady(5000)
// Incoming approval
val prodEventJson = buildJsonObject {
put("method", "event")
put("params", buildJsonObject {
put("event", "approval.request")
put("request_id", "req_prod_1")
put("session_id", "runtime_session_prod_99")
put("command", "systemctl restart nginx")
})
}
runtime1.gatewayClient.handleIncomingMessage(prodEventJson.toString())
var waited = 0
while (testRepo.activeApprovals.value.isEmpty() && waited < 50) {
kotlinx.coroutines.delay(50)
waited++
}
assertEquals(1, testRepo.activeApprovals.value.size)
// Respond to approval
val result = testRepo.respondApproval(
hostId = host1Id,
runtimeSessionId = RuntimeSessionId("runtime_session_prod_99"),
requestId = "req_prod_1",
choice = "once",
all = false
)
assertTrue("Expected true on successful RPC response", result)
assertEquals("Approval should be removed after success", 0, testRepo.activeApprovals.value.size)
// Verify sent RPC wire payload
assertTrue(receivedMessages.any {
it.contains("\"session_id\":\"runtime_session_prod_99\"") &&
it.contains("\"request_id\":\"req_prod_1\"") &&
it.contains("\"choice\":\"once\"")
})
runtime1.disconnect()
} }
} }

View file

@ -0,0 +1,656 @@
package app.hermes.mobile.core.repository
import app.hermes.mobile.core.model.*
import app.hermes.mobile.core.network.ConnectionState
import app.hermes.mobile.core.network.JsonRpcGatewayClient
import app.hermes.mobile.core.runtime.HermesConnectionManager
import app.hermes.mobile.core.runtime.HermesHostRuntime
import app.hermes.mobile.core.security.InMemoryTokenVault
import app.hermes.mobile.core.storage.*
import kotlinx.coroutines.CoroutineScope
import kotlinx.coroutines.Dispatchers
import kotlinx.coroutines.ExperimentalCoroutinesApi
import kotlinx.coroutines.runBlocking
import kotlinx.coroutines.test.StandardTestDispatcher
import kotlinx.coroutines.test.resetMain
import kotlinx.coroutines.test.runTest
import kotlinx.coroutines.test.setMain
import kotlinx.serialization.json.buildJsonObject
import kotlinx.serialization.json.put
import okhttp3.Response
import okhttp3.WebSocket
import okhttp3.WebSocketListener
import okhttp3.mockwebserver.Dispatcher
import okhttp3.mockwebserver.MockResponse
import okhttp3.mockwebserver.MockWebServer
import okhttp3.mockwebserver.RecordedRequest
import org.junit.After
import org.junit.Assert.*
import org.junit.Before
import org.junit.Test
import java.io.IOException
@OptIn(ExperimentalCoroutinesApi::class)
class MultiHostConcurrencyExecutionTest {
private val testDispatcher = StandardTestDispatcher()
private lateinit var hostDao: FakeHostDao
private lateinit var sessionDao: FakeUnifiedSessionDao
private lateinit var tokenVault: InMemoryTokenVault
private lateinit var connectionManager: HermesConnectionManager
private lateinit var sessionRepo: UnifiedSessionRepository
private val host1Id = HermesHostId("host-windows")
private val host2Id = HermesHostId("host-linux")
@Before
fun setUp() {
Dispatchers.setMain(testDispatcher)
hostDao = FakeHostDao()
sessionDao = FakeUnifiedSessionDao()
tokenVault = InMemoryTokenVault()
connectionManager = HermesConnectionManager(
hostDao = hostDao,
tokenVault = tokenVault,
scope = CoroutineScope(testDispatcher)
)
sessionRepo = UnifiedSessionRepository(
connectionManager = connectionManager,
sessionDao = sessionDao,
scope = CoroutineScope(testDispatcher)
)
}
@After
fun tearDown() {
Dispatchers.resetMain()
}
@Test
fun testTwoHermesSimultaneousStreamToTwoDifferentSessions() = 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 = "Windows Session", initialHostId = host1Id)
val session2 = sessionRepo.createUnifiedSession(title = "Linux Session", initialHostId = host2Id)
testScheduler.advanceUntilIdle()
val runtimeA = connectionManager.getRuntime(host1Id)
val runtimeB = connectionManager.getRuntime(host2Id)
assertNotNull(runtimeA)
assertNotNull(runtimeB)
// Register runtime IDs
sessionRepo.registerRuntimeBinding(session1.id, host1Id, RuntimeSessionId("rt_win_1"))
sessionRepo.registerRuntimeBinding(session2.id, host2Id, RuntimeSessionId("rt_lin_1"))
// Register bindings in DB
sessionDao.insertOrUpdateBinding(
HostBindingEntity(
sessionId = session1.id.value,
hostId = host1Id.value,
durableSessionId = "dur_win_1",
runtimeSessionId = "rt_win_1",
state = BindingState.RUNNING.name
)
)
sessionDao.insertOrUpdateBinding(
HostBindingEntity(
sessionId = session2.id.value,
hostId = host2Id.value,
durableSessionId = "dur_lin_1",
runtimeSessionId = "rt_lin_1",
state = BindingState.RUNNING.name
)
)
testScheduler.advanceUntilIdle()
// Stream from Host A into Session 1
val eventA1 = buildJsonObject {
put("method", "event")
put("params", buildJsonObject {
put("event", "message.start")
put("session_id", "rt_win_1")
put("message_id", "msg_a_1")
put("role", "assistant")
})
}
val eventA2 = buildJsonObject {
put("method", "event")
put("params", buildJsonObject {
put("event", "message.delta")
put("session_id", "rt_win_1")
put("message_id", "msg_a_1")
put("delta", "Windows output chunk")
})
}
// Stream from Host B into Session 2 concurrently
val eventB1 = buildJsonObject {
put("method", "event")
put("params", buildJsonObject {
put("event", "message.start")
put("session_id", "rt_lin_1")
put("message_id", "msg_b_1")
put("role", "assistant")
})
}
val eventB2 = buildJsonObject {
put("method", "event")
put("params", buildJsonObject {
put("event", "message.delta")
put("session_id", "rt_lin_1")
put("message_id", "msg_b_1")
put("delta", "Linux output chunk")
})
}
runtimeA!!.gatewayClient.handleIncomingMessage(eventA1.toString())
runtimeB!!.gatewayClient.handleIncomingMessage(eventB1.toString())
testScheduler.advanceUntilIdle()
runtimeA.gatewayClient.handleIncomingMessage(eventA2.toString())
runtimeB.gatewayClient.handleIncomingMessage(eventB2.toString())
testScheduler.advanceUntilIdle()
val messages1 = sessionRepo.getSessionMessages(session1.id).value
val messages2 = sessionRepo.getSessionMessages(session2.id).value
// Assert Session 1 only has Windows messages
assertTrue(messages1.any { it.id == "msg_a_1" })
assertFalse(messages1.any { it.id == "msg_b_1" })
assertEquals("Windows output chunk", messages1.find { it.id == "msg_a_1" }?.content)
assertEquals(host1Id, messages1.find { it.id == "msg_a_1" }?.hostId)
// Assert Session 2 only has Linux messages
assertTrue(messages2.any { it.id == "msg_b_1" })
assertFalse(messages2.any { it.id == "msg_a_1" })
assertEquals("Linux output chunk", messages2.find { it.id == "msg_b_1" }?.content)
assertEquals(host2Id, messages2.find { it.id == "msg_b_1" }?.hostId)
}
@Test
fun testTwoHermesSimultaneousStreamToOneUnifiedSession() = 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 session = sessionRepo.createUnifiedSession(title = "Dual Host Project", initialHostId = host1Id)
testScheduler.advanceUntilIdle()
val runtimeA = connectionManager.getRuntime(host1Id)
val runtimeB = connectionManager.getRuntime(host2Id)
sessionRepo.registerRuntimeBinding(session.id, host1Id, RuntimeSessionId("rt_win_dual"))
sessionRepo.registerRuntimeBinding(session.id, host2Id, RuntimeSessionId("rt_lin_dual"))
// Attach both hosts to this one unified session
sessionDao.insertOrUpdateBinding(
HostBindingEntity(
sessionId = session.id.value,
hostId = host1Id.value,
durableSessionId = "dur_win_dual",
runtimeSessionId = "rt_win_dual",
state = BindingState.RUNNING.name
)
)
sessionDao.insertOrUpdateBinding(
HostBindingEntity(
sessionId = session.id.value,
hostId = host2Id.value,
durableSessionId = "dur_lin_dual",
runtimeSessionId = "rt_lin_dual",
state = BindingState.RUNNING.name
)
)
testScheduler.advanceUntilIdle()
// Host A streams message
runtimeA!!.gatewayClient.handleIncomingMessage(
buildJsonObject {
put("method", "event")
put("params", buildJsonObject {
put("event", "message.start")
put("session_id", "rt_win_dual")
put("message_id", "msg_win")
put("role", "assistant")
})
}.toString()
)
testScheduler.advanceUntilIdle()
runtimeA.gatewayClient.handleIncomingMessage(
buildJsonObject {
put("method", "event")
put("params", buildJsonObject {
put("event", "message.delta")
put("session_id", "rt_win_dual")
put("message_id", "msg_win")
put("delta", "Windows result")
})
}.toString()
)
testScheduler.advanceUntilIdle()
// Host B concurrently streams tool
runtimeB!!.gatewayClient.handleIncomingMessage(
buildJsonObject {
put("method", "event")
put("params", buildJsonObject {
put("event", "tool.start")
put("session_id", "rt_lin_dual")
put("tool_id", "tool_lin")
put("name", "bash_exec")
})
}.toString()
)
testScheduler.advanceUntilIdle()
runtimeB.gatewayClient.handleIncomingMessage(
buildJsonObject {
put("method", "event")
put("params", buildJsonObject {
put("event", "tool.complete")
put("session_id", "rt_lin_dual")
put("tool_id", "tool_lin")
put("result", "Linux command completed")
})
}.toString()
)
testScheduler.advanceUntilIdle()
val messages = sessionRepo.getSessionMessages(session.id).value
val winMsg = messages.find { it.id == "msg_win" }
assertNotNull(winMsg)
assertEquals("Windows result", winMsg?.content)
assertEquals(host1Id, winMsg?.hostId)
val linToolMsg = messages.find { it.tools.any { t -> t.id == "tool_lin" } }
assertNotNull(linToolMsg)
assertEquals(host2Id, linToolMsg?.hostId)
assertEquals("completed", linToolMsg?.tools?.find { it.id == "tool_lin" }?.status)
}
@Test
fun testIdenticalMessageIdAndRequestIdOnTwoHostsDoNotConflict() = 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 session = sessionRepo.createUnifiedSession(title = "Conflict Resistance", initialHostId = host1Id)
testScheduler.advanceUntilIdle()
val runtimeA = connectionManager.getRuntime(host1Id)
val runtimeB = connectionManager.getRuntime(host2Id)
// Both hosts have same requestId "req_shared_1"
val approvalEventA = buildJsonObject {
put("method", "event")
put("params", buildJsonObject {
put("event", "approval.request")
put("session_id", "rt_win_shared")
put("request_id", "req_shared_1")
put("command", "powershell.exe -Command Get-Process")
})
}
val approvalEventB = buildJsonObject {
put("method", "event")
put("params", buildJsonObject {
put("event", "approval.request")
put("session_id", "rt_lin_shared")
put("request_id", "req_shared_1")
put("command", "ps aux")
})
}
runtimeA!!.gatewayClient.handleIncomingMessage(approvalEventA.toString())
runtimeB!!.gatewayClient.handleIncomingMessage(approvalEventB.toString())
testScheduler.advanceUntilIdle()
val approvals = sessionRepo.activeApprovals.value
assertEquals(2, approvals.size)
val appA = approvals.find { it.hostId == host1Id && it.approval.requestId == "req_shared_1" }
val appB = approvals.find { it.hostId == host2Id && it.approval.requestId == "req_shared_1" }
assertNotNull(appA)
assertNotNull(appB)
assertEquals("powershell.exe -Command Get-Process", appA?.approval?.command)
assertEquals("ps aux", appB?.approval?.command)
assertEquals(RuntimeSessionId("rt_win_shared"), appA?.runtimeSessionId)
assertEquals(RuntimeSessionId("rt_lin_shared"), appB?.runtimeSessionId)
}
@Test
fun testAppRestartRestoresBindingViaDurableId() = runBlocking(Dispatchers.Default) {
val server = MockWebServer()
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("""{"event":"gateway.ready","data":{"version":"1.0.0"}}""")
}
override fun onMessage(webSocket: WebSocket, text: String) {
if (text.contains("session.resume")) {
webSocket.send("""{"jsonrpc":"2.0","id":"a1","result":{"stored_session_id":"durable_persisted_99","session_id":"fresh_runtime_101"}}""")
} else if (text.contains("prompt.submit")) {
webSocket.send("""{"jsonrpc":"2.0","id":"a2","result":{"turn_id":"t_101"}}""")
}
}
})
}
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 = "Prod Server", baseUrl = wsUrl, allowCleartext = true, lastKnownStatus = "ONLINE"))
// Seed DB with existing session and binding from a previous app run
val sessionId = UnifiedSessionId("session_saved_1")
testSessionDao.insertSession(
UnifiedSessionEntity(
id = sessionId.value,
title = "Restored Session",
activeHostId = host1Id.value
)
)
testSessionDao.insertOrUpdateBinding(
HostBindingEntity(
sessionId = sessionId.value,
hostId = host1Id.value,
durableSessionId = "durable_persisted_99",
runtimeSessionId = "stale_dead_runtime_000",
state = BindingState.OFFLINE.name
)
)
val freshConnectionManager = HermesConnectionManager(
hostDao = testHostDao,
tokenVault = testTokenVault,
scope = CoroutineScope(Dispatchers.Default)
)
val freshRepo = UnifiedSessionRepository(
connectionManager = freshConnectionManager,
sessionDao = testSessionDao,
scope = CoroutineScope(Dispatchers.Default)
)
val runtime = freshConnectionManager.getRuntime(host1Id)
assertNotNull(runtime)
runtime!!.connect()
runtime.gatewayClient.awaitGatewayReady(5000)
val turnId = freshRepo.sendPrompt(sessionId, "Hello after restart")
assertEquals("t_101", turnId)
// Verify that binding was updated in DB with the fresh runtime ID
val updatedBinding = testSessionDao.getBindingsForSession(sessionId.value).find { it.hostId == host1Id.value }
assertNotNull(updatedBinding)
assertEquals("durable_persisted_99", updatedBinding?.durableSessionId)
assertEquals("fresh_runtime_101", updatedBinding?.runtimeSessionId)
assertEquals(BindingState.RUNNING.name, updatedBinding?.state)
runtime.disconnect()
server.shutdown()
}
@Test
fun testHostReconnectMintsNewRuntimeId() = runBlocking(Dispatchers.Default) {
val server = MockWebServer()
var sessionCreated = 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("""{"event":"gateway.ready","data":{"version":"1.0.0"}}""")
}
override fun onMessage(webSocket: WebSocket, text: String) {
if (text.contains("session.create")) {
sessionCreated = true
webSocket.send("""{"jsonrpc":"2.0","id":"a1","result":{"stored_session_id":"dur_conn_1","session_id":"rt_initial"}}""")
} else if (text.contains("session.resume")) {
webSocket.send("""{"jsonrpc":"2.0","id":"a2","result":{"stored_session_id":"dur_conn_1","session_id":"rt_after_reconnect"}}""")
}
}
})
}
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 testConnectionManager = HermesConnectionManager(
hostDao = testHostDao,
tokenVault = testTokenVault,
scope = CoroutineScope(Dispatchers.Default)
)
val testRepo = UnifiedSessionRepository(
connectionManager = testConnectionManager,
sessionDao = testSessionDao,
scope = CoroutineScope(Dispatchers.Default)
)
val session = testRepo.createUnifiedSession(title = "Reconnect Test", initialHostId = host1Id)
val runtime = testConnectionManager.getRuntime(host1Id)!!
// Initial connect
runtime.connect()
runtime.gatewayClient.awaitGatewayReady(5000)
val binding1 = testRepo.ensureAttachedRuntimeSession(session.id, host1Id, runtime)
assertEquals(RuntimeSessionId("rt_initial"), binding1.runtimeSessionId)
// Disconnect host runtime
runtime.disconnect()
testSessionDao.updateBindingState(session.id.value, host1Id.value, BindingState.OFFLINE.name)
// Reconnect host runtime
runtime.connect()
runtime.gatewayClient.awaitGatewayReady(5000)
// Reattach after reconnect
val binding2 = testRepo.ensureAttachedRuntimeSession(session.id, host1Id, runtime)
assertEquals(RuntimeSessionId("rt_after_reconnect"), binding2.runtimeSessionId)
runtime.disconnect()
server.shutdown()
}
@Test
fun testFailedPromptSubmitDoesNotAdvanceContextSyncCursor() = runBlocking(Dispatchers.Default) {
val server = MockWebServer()
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("""{"event":"gateway.ready","data":{"version":"1.0.0"}}""")
}
override fun onMessage(webSocket: WebSocket, text: String) {
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"}}""")
}
}
})
}
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 testConnectionManager = HermesConnectionManager(
hostDao = testHostDao,
tokenVault = testTokenVault,
scope = CoroutineScope(Dispatchers.Default)
)
val testRepo = UnifiedSessionRepository(
connectionManager = testConnectionManager,
sessionDao = testSessionDao,
scope = CoroutineScope(Dispatchers.Default)
)
val session = testRepo.createUnifiedSession(title = "Cursor Test", initialHostId = host1Id)
// Pre-populate binding with syncedThroughMessageId = "msg_baseline"
testSessionDao.insertOrUpdateBinding(
HostBindingEntity(
sessionId = session.id.value,
hostId = host1Id.value,
durableSessionId = "dur_fail_test",
runtimeSessionId = "rt_fail_test",
syncedThroughMessageId = "msg_baseline",
state = BindingState.READY.name
)
)
val runtime = testConnectionManager.getRuntime(host1Id)!!
runtime.connect()
runtime.gatewayClient.awaitGatewayReady(5000)
// Attempt sendPrompt which will fail at submitPrompt
var threw = false
try {
testRepo.sendPrompt(session.id, "Will fail")
} catch (_: Exception) {
threw = true
}
assertTrue("Expected prompt submission to throw", threw)
// Verify syncedThroughMessageId did NOT advance and remains "msg_baseline"
val binding = testSessionDao.getBindingsForSession(session.id.value).find { it.hostId == host1Id.value }
assertEquals("msg_baseline", binding?.syncedThroughMessageId)
assertEquals(BindingState.ERROR.name, binding?.state)
runtime.disconnect()
server.shutdown()
}
@Test
fun testStopHostADoesNotStopHostB() = 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 session = sessionRepo.createUnifiedSession(title = "Targeted Stop", initialHostId = host1Id)
testScheduler.advanceUntilIdle()
val runtimeA = connectionManager.getRuntime(host1Id)
val runtimeB = connectionManager.getRuntime(host2Id)
sessionRepo.registerRuntimeBinding(session.id, host1Id, RuntimeSessionId("rt_a"))
sessionRepo.registerRuntimeBinding(session.id, host2Id, RuntimeSessionId("rt_b"))
sessionDao.insertOrUpdateBinding(
HostBindingEntity(
sessionId = session.id.value,
hostId = host1Id.value,
durableSessionId = "dur_a",
runtimeSessionId = "rt_a",
state = BindingState.RUNNING.name
)
)
sessionDao.insertOrUpdateBinding(
HostBindingEntity(
sessionId = session.id.value,
hostId = host2Id.value,
durableSessionId = "dur_b",
runtimeSessionId = "rt_b",
state = BindingState.RUNNING.name
)
)
testScheduler.advanceUntilIdle()
// Start execution on both hosts
runtimeA!!.gatewayClient.handleIncomingMessage(
buildJsonObject {
put("method", "event")
put("params", buildJsonObject {
put("event", "message.start")
put("session_id", "rt_a")
put("message_id", "msg_a")
})
}.toString()
)
runtimeB!!.gatewayClient.handleIncomingMessage(
buildJsonObject {
put("method", "event")
put("params", buildJsonObject {
put("event", "message.start")
put("session_id", "rt_b")
put("message_id", "msg_b")
})
}.toString()
)
testScheduler.advanceUntilIdle()
assertTrue(sessionRepo.getHostExecuting(session.id, host1Id).value)
assertTrue(sessionRepo.getHostExecuting(session.id, host2Id).value)
// Stop Host A specifically
sessionRepo.interruptHost(session.id, host1Id)
testScheduler.advanceUntilIdle()
// Host A execution should be false, Host B execution MUST still be true
assertFalse("Host A should be stopped", sessionRepo.getHostExecuting(session.id, host1Id).value)
assertTrue("Host B should still be running", sessionRepo.getHostExecuting(session.id, host2Id).value)
}
}

View file

@ -8,6 +8,7 @@ import app.hermes.mobile.core.runtime.HermesConnectionManager
import app.hermes.mobile.core.security.InMemoryTokenVault import app.hermes.mobile.core.security.InMemoryTokenVault
import app.hermes.mobile.core.storage.FakeHostDao import app.hermes.mobile.core.storage.FakeHostDao
import app.hermes.mobile.core.storage.FakeUnifiedSessionDao import app.hermes.mobile.core.storage.FakeUnifiedSessionDao
import app.hermes.mobile.core.storage.HostBindingEntity
import kotlinx.coroutines.CoroutineScope import kotlinx.coroutines.CoroutineScope
import kotlinx.coroutines.Dispatchers import kotlinx.coroutines.Dispatchers
import kotlinx.coroutines.ExperimentalCoroutinesApi import kotlinx.coroutines.ExperimentalCoroutinesApi
@ -95,6 +96,19 @@ class UnifiedSessionRepositoryTest {
val session = repository.createUnifiedSession(title = "Streaming Test", initialHostId = host1Id) val session = repository.createUnifiedSession(title = "Streaming Test", initialHostId = host1Id)
testScheduler.advanceUntilIdle() testScheduler.advanceUntilIdle()
repository.registerRuntimeBinding(session.id, host1Id, RuntimeSessionId("rt_stream_1"))
sessionDao.insertOrUpdateBinding(
HostBindingEntity(
sessionId = session.id.value,
hostId = host1Id.value,
durableSessionId = "dur_stream_1",
runtimeSessionId = "rt_stream_1",
state = BindingState.RUNNING.name
)
)
testScheduler.advanceUntilIdle()
val runtimeA = connectionManager.getRuntime(host1Id) val runtimeA = connectionManager.getRuntime(host1Id)
assertNotNull(runtimeA) assertNotNull(runtimeA)
@ -103,39 +117,46 @@ class UnifiedSessionRepositoryTest {
put("method", "event") put("method", "event")
put("params", buildJsonObject { put("params", buildJsonObject {
put("event", "message.start") put("event", "message.start")
put("session_id", "rt_stream_1")
put("message_id", "msg_stream_1") put("message_id", "msg_stream_1")
put("role", "assistant") put("role", "assistant")
}) })
} }
runtimeA?.gatewayClient?.handleIncomingMessage(msgStart.toString()) runtimeA?.gatewayClient?.handleIncomingMessage(msgStart.toString())
testScheduler.advanceUntilIdle()
// Stream delta 1 // Stream delta 1
val msgDelta1 = buildJsonObject { val msgDelta1 = buildJsonObject {
put("method", "event") put("method", "event")
put("params", buildJsonObject { put("params", buildJsonObject {
put("event", "message.delta") put("event", "message.delta")
put("session_id", "rt_stream_1")
put("message_id", "msg_stream_1") put("message_id", "msg_stream_1")
put("delta", "Hello ") put("delta", "Hello ")
}) })
} }
runtimeA?.gatewayClient?.handleIncomingMessage(msgDelta1.toString()) runtimeA?.gatewayClient?.handleIncomingMessage(msgDelta1.toString())
testScheduler.advanceUntilIdle()
// Stream delta 2 // Stream delta 2
val msgDelta2 = buildJsonObject { val msgDelta2 = buildJsonObject {
put("method", "event") put("method", "event")
put("params", buildJsonObject { put("params", buildJsonObject {
put("event", "message.delta") put("event", "message.delta")
put("session_id", "rt_stream_1")
put("message_id", "msg_stream_1") put("message_id", "msg_stream_1")
put("delta", "from Multi-Hermes!") put("delta", "from Multi-Hermes!")
}) })
} }
runtimeA?.gatewayClient?.handleIncomingMessage(msgDelta2.toString()) runtimeA?.gatewayClient?.handleIncomingMessage(msgDelta2.toString())
testScheduler.advanceUntilIdle()
// Stream complete // Stream complete
val msgComplete = buildJsonObject { val msgComplete = buildJsonObject {
put("method", "event") put("method", "event")
put("params", buildJsonObject { put("params", buildJsonObject {
put("event", "message.complete") put("event", "message.complete")
put("session_id", "rt_stream_1")
put("message_id", "msg_stream_1") put("message_id", "msg_stream_1")
put("content", "Hello from Multi-Hermes!") put("content", "Hello from Multi-Hermes!")
}) })
@ -162,6 +183,19 @@ class UnifiedSessionRepositoryTest {
val session = repository.createUnifiedSession(title = "Background Session", initialHostId = host1Id) val session = repository.createUnifiedSession(title = "Background Session", initialHostId = host1Id)
testScheduler.advanceUntilIdle() testScheduler.advanceUntilIdle()
repository.registerRuntimeBinding(session.id, host1Id, RuntimeSessionId("rt_bg_1"))
sessionDao.insertOrUpdateBinding(
HostBindingEntity(
sessionId = session.id.value,
hostId = host1Id.value,
durableSessionId = "dur_bg_1",
runtimeSessionId = "rt_bg_1",
state = BindingState.RUNNING.name
)
)
testScheduler.advanceUntilIdle()
val runtimeA = connectionManager.getRuntime(host1Id) val runtimeA = connectionManager.getRuntime(host1Id)
val runtimeB = connectionManager.getRuntime(host2Id) val runtimeB = connectionManager.getRuntime(host2Id)
@ -170,11 +204,13 @@ class UnifiedSessionRepositoryTest {
put("method", "event") put("method", "event")
put("params", buildJsonObject { put("params", buildJsonObject {
put("event", "tool.start") put("event", "tool.start")
put("session_id", "rt_bg_1")
put("tool_id", "tool_bg_1") put("tool_id", "tool_bg_1")
put("name", "heavy_build_task") put("name", "heavy_build_task")
}) })
} }
runtimeA?.gatewayClient?.handleIncomingMessage(toolStart.toString()) runtimeA?.gatewayClient?.handleIncomingMessage(toolStart.toString())
testScheduler.advanceUntilIdle()
// User switches session active host to Host B // User switches session active host to Host B
repository.switchSessionActiveHost(session.id, host2Id) repository.switchSessionActiveHost(session.id, host2Id)
@ -185,6 +221,7 @@ class UnifiedSessionRepositoryTest {
put("method", "event") put("method", "event")
put("params", buildJsonObject { put("params", buildJsonObject {
put("event", "tool.complete") put("event", "tool.complete")
put("session_id", "rt_bg_1")
put("tool_id", "tool_bg_1") put("tool_id", "tool_bg_1")
put("result", "Build successful in 42s") put("result", "Build successful in 42s")
put("is_error", false) put("is_error", false)

View file

@ -1,6 +1,7 @@
package app.hermes.mobile.core.storage package app.hermes.mobile.core.storage
import kotlinx.coroutines.flow.Flow import kotlinx.coroutines.flow.Flow
import kotlinx.coroutines.flow.MutableSharedFlow
import kotlinx.coroutines.flow.MutableStateFlow import kotlinx.coroutines.flow.MutableStateFlow
import kotlinx.coroutines.flow.map import kotlinx.coroutines.flow.map
@ -46,24 +47,29 @@ class FakeUnifiedSessionDao : UnifiedSessionDao {
private val sessions = mutableMapOf<String, UnifiedSessionEntity>() private val sessions = mutableMapOf<String, UnifiedSessionEntity>()
private val bindings = mutableMapOf<String, MutableList<HostBindingEntity>>() private val bindings = mutableMapOf<String, MutableList<HostBindingEntity>>()
private val messages = mutableMapOf<String, MutableList<UnifiedMessageEntity>>() private val messages = mutableMapOf<String, MutableList<UnifiedMessageEntity>>()
private val sessionsFlow = MutableStateFlow<List<UnifiedSessionEntity>>(emptyList()) private val _sessionsFlow = MutableSharedFlow<List<UnifiedSessionEntity>>(replay = 1)
private fun updateFlow() { init {
sessionsFlow.value = sessions.values.sortedByDescending { it.updatedAt } _sessionsFlow.tryEmit(emptyList<UnifiedSessionEntity>())
} }
override fun getSessionsFlow(): Flow<List<UnifiedSessionEntity>> = sessionsFlow private fun updateFlow() {
val list = sessions.values.sortedByDescending { it.updatedAt }.map { it.copy() }
_sessionsFlow.tryEmit(list)
}
override fun getSessionsFlow(): Flow<List<UnifiedSessionEntity>> = _sessionsFlow
override suspend fun getSessions(): List<UnifiedSessionEntity> = sessions.values.sortedByDescending { it.updatedAt } override suspend fun getSessions(): List<UnifiedSessionEntity> = sessions.values.sortedByDescending { it.updatedAt }
override fun getSessionWithDetailsFlow(sessionId: String): Flow<UnifiedSessionWithDetails?> { override fun getSessionWithDetailsFlow(sessionId: String): Flow<UnifiedSessionWithDetails?> {
return sessionsFlow.map { getSessionWithDetails(sessionId) } return _sessionsFlow.map { _: List<UnifiedSessionEntity> -> getSessionWithDetails(sessionId) }
} }
override suspend fun getSessionWithDetails(sessionId: String): UnifiedSessionWithDetails? { override suspend fun getSessionWithDetails(sessionId: String): UnifiedSessionWithDetails? {
val s = sessions[sessionId] ?: return null val s = sessions[sessionId] ?: return null
val b = bindings[sessionId] ?: emptyList() val b = bindings[sessionId]?.map { it.copy() } ?: emptyList()
val m = messages[sessionId] ?: emptyList() val m = messages[sessionId]?.map { it.copy() } ?: emptyList()
return UnifiedSessionWithDetails(session = s, bindings = b, messages = m) return UnifiedSessionWithDetails(session = s, bindings = b, messages = m)
} }
@ -72,7 +78,7 @@ class FakeUnifiedSessionDao : UnifiedSessionDao {
} }
override suspend fun getBindingsForSession(sessionId: String): List<HostBindingEntity> { override suspend fun getBindingsForSession(sessionId: String): List<HostBindingEntity> {
return bindings[sessionId] ?: emptyList() return bindings[sessionId]?.map { it.copy() } ?: emptyList()
} }
override suspend fun insertSession(session: UnifiedSessionEntity) { override suspend fun insertSession(session: UnifiedSessionEntity) {
@ -96,18 +102,34 @@ class FakeUnifiedSessionDao : UnifiedSessionDao {
val list = bindings.computeIfAbsent(binding.sessionId) { mutableListOf() } val list = bindings.computeIfAbsent(binding.sessionId) { mutableListOf() }
list.removeAll { it.hostId == binding.hostId } list.removeAll { it.hostId == binding.hostId }
list.add(binding) list.add(binding)
val s = sessions[binding.sessionId]
if (s != null) {
sessions[binding.sessionId] = s.copy(updatedAt = System.currentTimeMillis())
}
updateFlow()
} }
override suspend fun insertOrUpdateBindings(bindingList: List<HostBindingEntity>) { override suspend fun insertOrUpdateBindings(bindingList: List<HostBindingEntity>) {
for (b in bindingList) insertOrUpdateBinding(b) for (b in bindingList) {
val list = bindings.computeIfAbsent(b.sessionId) { mutableListOf() }
list.removeAll { it.hostId == b.hostId }
list.add(b)
val s = sessions[b.sessionId]
if (s != null) {
sessions[b.sessionId] = s.copy(updatedAt = System.currentTimeMillis())
}
}
updateFlow()
} }
override suspend fun deleteBinding(sessionId: String, hostId: String) { override suspend fun deleteBinding(sessionId: String, hostId: String) {
bindings[sessionId]?.removeAll { it.hostId == hostId } bindings[sessionId]?.removeAll { it.hostId == hostId }
updateFlow()
} }
override suspend fun deleteBindingsForSession(sessionId: String) { override suspend fun deleteBindingsForSession(sessionId: String) {
bindings.remove(sessionId) bindings.remove(sessionId)
updateFlow()
} }
override suspend fun insertOrUpdateMessage(message: UnifiedMessageEntity) { override suspend fun insertOrUpdateMessage(message: UnifiedMessageEntity) {
@ -118,14 +140,25 @@ class FakeUnifiedSessionDao : UnifiedSessionDao {
} else { } else {
list.add(message) list.add(message)
} }
updateFlow()
} }
override suspend fun insertMessages(msgList: List<UnifiedMessageEntity>) { override suspend fun insertMessages(msgList: List<UnifiedMessageEntity>) {
for (m in msgList) insertOrUpdateMessage(m) for (m in msgList) {
val list = messages.computeIfAbsent(m.sessionId) { mutableListOf() }
val idx = list.indexOfFirst { it.id == m.id }
if (idx >= 0) {
list[idx] = m
} else {
list.add(m)
}
}
updateFlow()
} }
override suspend fun deleteMessagesForSession(sessionId: String) { override suspend fun deleteMessagesForSession(sessionId: String) {
messages.remove(sessionId) messages.remove(sessionId)
updateFlow()
} }
override suspend fun updateMessageContent( override suspend fun updateMessageContent(
@ -148,6 +181,7 @@ class FakeUnifiedSessionDao : UnifiedSessionDao {
break break
} }
} }
updateFlow()
} }
override suspend fun updateActiveHost(sessionId: String, hostId: String, updatedAt: Long) { override suspend fun updateActiveHost(sessionId: String, hostId: String, updatedAt: Long) {
@ -173,6 +207,7 @@ class FakeUnifiedSessionDao : UnifiedSessionDao {
syncedAt = syncedAt, syncedAt = syncedAt,
state = state state = state
) )
updateFlow()
} }
} }
@ -181,6 +216,7 @@ class FakeUnifiedSessionDao : UnifiedSessionDao {
val idx = list.indexOfFirst { it.hostId == hostId } val idx = list.indexOfFirst { it.hostId == hostId }
if (idx >= 0) { if (idx >= 0) {
list[idx] = list[idx].copy(state = state) list[idx] = list[idx].copy(state = state)
updateFlow()
} }
} }
} }

View file

@ -92,6 +92,8 @@ class UnifiedContextBuilderTest {
assertTrue(syncAll.contextPrompt.contains("Office PC")) assertTrue(syncAll.contextPrompt.contains("Office PC"))
assertTrue(syncAll.contextPrompt.contains("Write a python script")) assertTrue(syncAll.contextPrompt.contains("Write a python script"))
assertTrue(syncAll.contextPrompt.contains("Linux Server")) assertTrue(syncAll.contextPrompt.contains("Linux Server"))
assertFalse(syncAll.contextPrompt.contains("192.168.1.100:9119"))
assertFalse(syncAll.contextPrompt.contains("192.168.1.50:9119"))
// Case 2: Stale host binding (synced up to msg-1, needs delta msg-2 and msg-3) // Case 2: Stale host binding (synced up to msg-1, needs delta msg-2 and msg-3)
val syncDelta = UnifiedContextBuilder.buildContextSyncPayload(session, host2, hostsMap, "msg-1") val syncDelta = UnifiedContextBuilder.buildContextSyncPayload(session, host2, hostsMap, "msg-1")
@ -100,6 +102,7 @@ class UnifiedContextBuilderTest {
assertFalse(syncDelta.contextPrompt.contains("Write a python script to parse CSV files.")) assertFalse(syncDelta.contextPrompt.contains("Write a python script to parse CSV files."))
assertTrue(syncDelta.contextPrompt.contains("Sure! Here is the python script")) assertTrue(syncDelta.contextPrompt.contains("Sure! Here is the python script"))
assertTrue(syncDelta.contextPrompt.contains("Now run it on the linux server dataset.")) assertTrue(syncDelta.contextPrompt.contains("Now run it on the linux server dataset."))
assertFalse(syncDelta.contextPrompt.contains("192.168.1.100:9119"))
// Case 3: Fully synced host binding (synced up to msg-3) // Case 3: Fully synced host binding (synced up to msg-3)
val syncUpToDate = UnifiedContextBuilder.buildContextSyncPayload(session, host2, hostsMap, "msg-3") val syncUpToDate = UnifiedContextBuilder.buildContextSyncPayload(session, host2, hostsMap, "msg-3")