use crate::builtin_agent_roles::PLAN_AGENT_ID;
use crate::gateway::{
AppRuntimeFactory, AppRuntimeFactoryError, ResolvedSessionRuntime, SessionRecord,
SessionRuntimeBuildInput, SessionServiceError,
};
use agent_contracts::backend::OperationBackend;
use agent_contracts::{ChannelFileSender, InteractionHandle, LoopEventSink};
use agent_types::common::ids::AgentId;
use agent_types::events::{LoopEndSummary, ToolResultEvent};
use agent_types::ReasoningEffort;
use memory::{MemoryManager, MemorySnapshot};
use std::sync::Arc;
use tokio_util::sync::CancellationToken;
use tool::ToolSpecSnapshot;
use xiaoo_api::runtime::{
LoopStateSnapshot, LoopStopRule, RuntimeInput, RuntimeOutput, RuntimeState,
};
pub(crate) struct SessionWorkerInput {
pub runtime_input: SessionRuntimeBuildInput,
pub resolved_runtime: ResolvedSessionRuntime,
pub session: SessionRecord,
pub agent_id: AgentId,
pub operation_backend: Arc<dyn OperationBackend>,
pub user_message: String,
pub append_user_message: bool,
pub reasoning_effort: ReasoningEffort,
pub loop_event_sink_override: Option<Arc<dyn LoopEventSink>>,
pub interaction_handle_override: Option<Arc<dyn InteractionHandle>>,
pub channel_file_sender_override: Option<Arc<dyn ChannelFileSender>>,
pub loop_state: Option<LoopStateSnapshot>,
pub memory_snapshot: Option<MemorySnapshot>,
pub tool_manifest: Option<Vec<ToolSpecSnapshot>>,
pub cancellation_token: Option<CancellationToken>,
pub command_context: Option<agent_types::chat::CommandContext>,
}
pub(crate) struct SessionWorkerResult {
pub loop_result: RuntimeOutput,
pub loop_state: LoopStateSnapshot,
pub memory_snapshot: MemorySnapshot,
pub tool_manifest: Vec<ToolSpecSnapshot>,
}
pub(crate) struct SessionWorker;
impl SessionWorker {
pub async fn run(
input: SessionWorkerInput,
) -> Result<SessionWorkerResult, SessionServiceError> {
let is_root_lane = input.agent_id == input.session.runtime.agent_id;
let mut resolved = input.resolved_runtime;
if !is_root_lane {
resolved.bindings.interaction_handle = input.interaction_handle_override.clone();
resolved.bindings.pending_user_messages = None;
if let Some(override_sender) = input.channel_file_sender_override.clone() {
resolved.bindings.channel_file_sender = Some(override_sender);
}
resolved.bindings.loop_event_sink = merge_loop_event_sinks(
resolved.bindings.loop_event_sink.clone(),
input.loop_event_sink_override.clone(),
);
} else {
if let Some(override_handle) = input.interaction_handle_override.clone() {
resolved.bindings.interaction_handle = Some(override_handle);
}
if let Some(override_sender) = input.channel_file_sender_override.clone() {
resolved.bindings.channel_file_sender = Some(override_sender);
}
resolved.bindings.loop_event_sink = merge_loop_event_sinks(
resolved.bindings.loop_event_sink.clone(),
input.loop_event_sink_override.clone(),
);
}
let loop_session_id = input
.loop_state
.as_ref()
.map(|snapshot| snapshot.session_id.clone())
.unwrap_or_else(|| input.session.session_id.clone());
let cancel = input
.cancellation_token
.clone()
.unwrap_or_else(CancellationToken::new);
let mut loop_state = input
.loop_state
.clone()
.map(|snapshot| RuntimeState::from_snapshot(snapshot, cancel.clone()))
.unwrap_or_else(|| RuntimeState::new_with_cancel(loop_session_id, cancel));
let messages = loop_state.messages_arc();
let assembly = AppRuntimeFactory::build(
&resolved,
&input.session,
messages,
input.tool_manifest.clone(),
input.operation_backend.clone(),
)
.await?;
let tool_manifest = assembly.tool_manifest.clone();
let mut memory_manager = match input.memory_snapshot.clone() {
Some(snapshot) => MemoryManager::from_snapshot(snapshot),
None => {
let memory_session_id = if is_root_lane {
input.session.session_id.clone()
} else {
input.agent_id.0.clone()
};
MemoryManager::new(memory_session_id, current_time_ms()).map_err(|error| {
SessionServiceError::Memory {
message: error.to_string(),
}
})?
}
};
let mut loop_input = RuntimeInput::new(input.user_message)
.with_agent_id(input.agent_id.clone())
.with_visible_tools(assembly.visible_tools.clone())
.with_reasoning_effort(input.reasoning_effort);
if !input.append_user_message {
loop_input = loop_input.resume_without_user_message();
}
loop_input.command_context = input.command_context.clone();
if input.runtime_input.entry.runtime_profile_id.as_deref() == Some(PLAN_AGENT_ID) {
loop_input = loop_input.with_stop_rules([LoopStopRule::AfterSuccessfulTool {
tool_name: "todo_write".to_string(),
}]);
}
if let Some(loop_event_sink) = resolved.bindings.loop_event_sink.clone() {
loop_input = loop_input.with_event_sink(loop_event_sink);
}
if let Some(runtime_view) = assembly.runtime_view.clone() {
loop_input = loop_input.with_runtime_view(runtime_view);
}
if let Some(pending_user_messages) = resolved.bindings.pending_user_messages.clone() {
loop_input = loop_input.with_pending_user_messages(pending_user_messages);
}
let loop_result = assembly.runtime.run(&mut loop_state, loop_input).await;
let shutdown_result = assembly.shutdown().await;
let loop_result = match loop_result {
Ok(loop_result) => loop_result,
Err(error) => {
tracing::warn!(
session_id = %input.session.session_id,
agent_id = %input.agent_id,
error = %error,
messages_count = loop_state.messages.read().len(),
"agent loop failed, preserving partial state for recovery"
);
memory_manager.sync_from_loop_state(&loop_state.messages.read(), current_time_ms());
if let Err(shutdown_error) = shutdown_result {
tracing::warn!(
session_id = %input.session.session_id,
agent_id = %input.agent_id,
shutdown_error = %shutdown_error,
"runtime shutdown failed after loop error"
);
}
return Err(SessionServiceError::CoreRunWithState {
message: error.to_string(),
partial_loop_state: loop_state.to_snapshot(),
partial_memory_snapshot: memory_manager.snapshot().clone(),
tool_manifest,
});
}
};
if let Err(error) = shutdown_result {
tracing::warn!(
session_id = %input.session.session_id,
agent_id = %input.agent_id,
shutdown_error = %error,
messages_count = loop_state.messages.read().len(),
"runtime shutdown failed after successful loop, preserving state for recovery"
);
memory_manager.sync_from_loop_state(&loop_state.messages.read(), current_time_ms());
return Err(SessionServiceError::CoreRunWithState {
message: format!("runtime shutdown failed: {error}"),
partial_loop_state: loop_state.to_snapshot(),
partial_memory_snapshot: memory_manager.snapshot().clone(),
tool_manifest,
});
}
memory_manager.sync_from_loop_state(&loop_state.messages.read(), current_time_ms());
Ok(SessionWorkerResult {
loop_result,
loop_state: loop_state.to_snapshot(),
memory_snapshot: memory_manager.snapshot().clone(),
tool_manifest,
})
}
}
#[derive(Clone)]
struct FanoutLoopEventSink {
sinks: Vec<Arc<dyn LoopEventSink>>,
}
impl LoopEventSink for FanoutLoopEventSink {
fn on_turn_start(&self, agent_id: &AgentId, turn: u32) {
for sink in &self.sinks {
sink.on_turn_start(agent_id, turn);
}
}
fn on_assistant_message(&self, agent_id: &AgentId, text: &str) {
for sink in &self.sinks {
sink.on_assistant_message(agent_id, text);
}
}
fn on_assistant_reasoning(&self, agent_id: &AgentId, text: &str) {
for sink in &self.sinks {
sink.on_assistant_reasoning(agent_id, text);
}
}
fn on_tool_result(&self, agent_id: &AgentId, event: &ToolResultEvent) {
for sink in &self.sinks {
sink.on_tool_result(agent_id, event);
}
}
fn on_loop_end(&self, agent_id: &AgentId, summary: &LoopEndSummary) {
for sink in &self.sinks {
sink.on_loop_end(agent_id, summary);
}
}
}
pub(crate) fn merge_loop_event_sinks(
primary: Option<Arc<dyn LoopEventSink>>,
secondary: Option<Arc<dyn LoopEventSink>>,
) -> Option<Arc<dyn LoopEventSink>> {
match (primary, secondary) {
(None, None) => None,
(Some(sink), None) | (None, Some(sink)) => Some(sink),
(Some(primary), Some(secondary)) => Some(Arc::new(FanoutLoopEventSink {
sinks: vec![primary, secondary],
})),
}
}
impl From<AppRuntimeFactoryError> for SessionServiceError {
fn from(value: AppRuntimeFactoryError) -> Self {
Self::RuntimeBuild {
message: value.to_string(),
}
}
}
fn current_time_ms() -> u64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|duration| duration.as_millis() as u64)
.unwrap_or(0)
}