Fix critical Multi-Hermes defects (P0/P1) and add concurrency test suite
This commit is contained in:
parent
879b834e51
commit
fc3ddeda65
12 changed files with 1397 additions and 255 deletions
|
|
@ -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
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -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 {
|
||||||
|
if (binding.runtimeSessionId.isNotEmpty()) {
|
||||||
runtime.gatewayClient.interruptSession(RuntimeSessionId(binding.runtimeSessionId))
|
runtime.gatewayClient.interruptSession(RuntimeSessionId(binding.runtimeSessionId))
|
||||||
|
} else {
|
||||||
|
false
|
||||||
|
}
|
||||||
} catch (_: Exception) {
|
} catch (_: Exception) {
|
||||||
|
false
|
||||||
}
|
}
|
||||||
|
setHostExecuting(sessionId, hostId, false)
|
||||||
|
sessionDao.updateBindingState(sessionId.value, hostId.value, BindingState.READY.name)
|
||||||
|
return success
|
||||||
}
|
}
|
||||||
}
|
|
||||||
setExecuting(sessionId, false)
|
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
|
||||||
|
if (!messageId.isNullOrEmpty()) {
|
||||||
for ((sessionId, flow) in sessionMessagesState) {
|
for ((sessionId, flow) in sessionMessagesState) {
|
||||||
if (flow.value.any { it.id == messageId }) {
|
if (flow.value.any { it.id == messageId && (it.hostId == hostId || it.hostId == null) }) {
|
||||||
return sessionId
|
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)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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()
|
||||||
|
|
|
||||||
|
|
@ -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 ---")
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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,6 +98,7 @@ fun ChatScreen(
|
||||||
|
|
||||||
Scaffold(
|
Scaffold(
|
||||||
topBar = {
|
topBar = {
|
||||||
|
Column {
|
||||||
TopAppBar(
|
TopAppBar(
|
||||||
title = {
|
title = {
|
||||||
Column {
|
Column {
|
||||||
|
|
@ -190,6 +192,32 @@ fun ChatScreen(
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// Multi-host Status Strip showing simultaneous statuses (e.g. PC1 Running, Linux Active, PC3 Offline)
|
||||||
|
if (hosts.isNotEmpty()) {
|
||||||
|
LazyRow(
|
||||||
|
modifier = Modifier
|
||||||
|
.fillMaxWidth()
|
||||||
|
.background(MaterialTheme.colorScheme.surfaceVariant.copy(alpha = 0.5f))
|
||||||
|
.padding(horizontal = 16.dp, vertical = 6.dp),
|
||||||
|
horizontalArrangement = Arrangement.spacedBy(8.dp),
|
||||||
|
verticalAlignment = Alignment.CenterVertically
|
||||||
|
) {
|
||||||
|
items(hosts, key = { it.id.value }) { host ->
|
||||||
|
val isRunning by viewModel.getHostExecuting(host.id).collectAsState()
|
||||||
|
val isActive = host.id == currentSession?.activeHostId
|
||||||
|
|
||||||
|
HostStatusChip(
|
||||||
|
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(
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
when (req.promptType) {
|
val success = 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)
|
||||||
}
|
}
|
||||||
} catch (e: Exception) {
|
if (!success) {
|
||||||
_uiState.value = _uiState.value.copy(
|
_uiState.value = _uiState.value.copy(
|
||||||
error = e.localizedMessage ?: "Failed to respond to clarification"
|
error = "Failed to submit clarification response"
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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")
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue