diff --git a/src/agent/agent_loop.rs b/src/agent/agent_loop.rs index b2b41c1..1d8a494 100644 --- a/src/agent/agent_loop.rs +++ b/src/agent/agent_loop.rs @@ -1,11 +1,15 @@ use crate::agent::context_compressor::estimate_tokens; use crate::agent::media_handler::MediaHandlerRegistry; use crate::agent::system_prompt::build_system_prompt; +use crate::agent::turn_event::{AgentTurnContext, TurnEvent}; use crate::bus::message::ContentBlock; use crate::bus::{ChatMessage, MediaRef}; use crate::config::LLMProviderConfig; use crate::observability::{Observer, ObserverEvent, ToolExecutionOutcome, truncate_args}; -use crate::providers::{ChatCompletionRequest, LLMProvider, Message, ToolCall, create_provider}; +use crate::providers::{ + ChatCompletionRequest, ChatCompletionResponse, LLMProvider, Message, ProviderChunk, + ProviderResponseAccumulator, ToolCall, create_provider, +}; use crate::tools::ToolRegistry; use std::collections::VecDeque; use std::hash::{Hash, Hasher}; @@ -13,11 +17,14 @@ use std::path::PathBuf; use std::sync::Arc; use std::time::Instant; +use futures_util::StreamExt; + /// Maximum characters in a tool result before truncation. /// Prevents context overflow from large tool outputs. const MAX_TOOL_RESULT_CHARS: usize = 16_000; /// Minimum characters to keep when truncating const TRUNCATION_SUFFIX_LEN: usize = 200; +const TOOL_PREVIEW_CHARS: usize = 1_000; enum MediaOrigin<'a> { User, @@ -164,6 +171,17 @@ fn truncate_tool_result(output: &str) -> String { } } +fn tool_result_preview(output: &str) -> String { + if output.len() <= TOOL_PREVIEW_CHARS { + output.to_string() + } else { + format!( + "{}…", + &output[..output.floor_char_boundary(TOOL_PREVIEW_CHARS)] + ) + } +} + /// Loop detection result. #[derive(Debug, Clone, PartialEq, Eq)] enum LoopDetectionResult { @@ -296,6 +314,32 @@ pub struct AgentProcessResult { pub final_response: ChatMessage, pub emitted_messages: Vec, pub total_tokens: Option, + pub usage: Option, +} + +fn merge_usage(total: &mut crate::providers::Usage, next: &crate::providers::Usage) { + total.prompt_tokens = total.prompt_tokens.saturating_add(next.prompt_tokens); + total.completion_tokens = total + .completion_tokens + .saturating_add(next.completion_tokens); + total.total_tokens = total.total_tokens.saturating_add(next.total_tokens); + total.cached_tokens = sum_optional_tokens(total.cached_tokens, next.cached_tokens); + total.cache_read_input_tokens = + sum_optional_tokens(total.cache_read_input_tokens, next.cache_read_input_tokens); + total.cache_creation_input_tokens = sum_optional_tokens( + total.cache_creation_input_tokens, + next.cache_creation_input_tokens, + ); +} + +fn sum_optional_tokens(left: Option, right: Option) -> Option { + match (left, right) { + (None, None) => None, + (left, right) => Some( + left.unwrap_or_default() + .saturating_add(right.unwrap_or_default()), + ), + } } impl AgentLoop { @@ -451,6 +495,60 @@ impl AgentLoop { &self.tools } + async fn stream_completion( + &self, + request: ChatCompletionRequest, + iteration: u32, + turn: Option<&AgentTurnContext>, + ) -> Result { + let mut provider_stream = self.provider.stream(request).await.map_err(|error| { + tracing::error!(error = %error, "LLM request failed"); + AgentError::LlmError(error.to_string()) + })?; + let mut accumulator = ProviderResponseAccumulator::default(); + while let Some(chunk) = provider_stream.next().await { + let chunk = chunk.map_err(|error| { + tracing::error!(error = %error, "LLM stream failed"); + AgentError::LlmError(error.to_string()) + })?; + if let Some(turn) = turn { + let event = match &chunk { + ProviderChunk::Reasoning(delta) => Some(TurnEvent::ReasoningDelta { + iteration, + delta: delta.clone(), + }), + ProviderChunk::Text(delta) => Some(TurnEvent::TextDelta { + iteration, + delta: delta.clone(), + }), + _ => None, + }; + if let Some(event) = event { + turn.emitter.emit(event).map_err(|error| { + AgentError::Other(format!("turn event rejected: {error}")) + })?; + } + } + accumulator.push(chunk); + } + Ok(accumulator.finish()) + } + + fn annotate_message( + message: &mut ChatMessage, + turn: Option<&AgentTurnContext>, + iteration: u32, + final_response: bool, + ) { + message.iteration = Some(iteration); + if let Some(turn) = turn { + message.turn_id = Some(turn.turn_id.clone()); + if final_response { + message.id = turn.message_id.clone(); + } + } + } + fn chat_message_to_llm_message(&self, m: &ChatMessage, include_media: bool) -> Message { let content = if m.media_refs.is_empty() || !include_media { vec![ContentBlock::text(&m.content)] @@ -498,8 +596,24 @@ impl AgentLoop { /// - The LLM returns no more tool calls (final response) /// - Maximum iterations are reached pub async fn process( + &self, + messages: Vec, + ) -> Result { + self.process_inner(messages, None).await + } + + pub async fn process_streaming( + &self, + messages: Vec, + turn: AgentTurnContext, + ) -> Result { + self.process_inner(messages, Some(turn)).await + } + + async fn process_inner( &self, mut messages: Vec, + turn: Option, ) -> Result { #[cfg(debug_assertions)] tracing::debug!( @@ -522,6 +636,7 @@ impl AgentLoop { let mut loop_detector = LoopDetector::new(LoopDetectorConfig::default()); let mut emitted_messages = Vec::new(); let mut accumulated_tokens: u32 = 0; + let mut accumulated_usage = crate::providers::Usage::default(); for iteration in 0..self.max_iterations { #[cfg(debug_assertions)] @@ -564,12 +679,14 @@ impl AgentLoop { }; // Call LLM - let response = (*self.provider).chat(request).await.map_err(|e| { - tracing::error!(error = %e, "LLM request failed"); - AgentError::LlmError(e.to_string()) - })?; + let iteration = u32::try_from(iteration) + .map_err(|_| AgentError::Other("tool iteration exceeds u32".to_string()))?; + let response = self + .stream_completion(request, iteration, turn.as_ref()) + .await?; - accumulated_tokens += response.usage.total_tokens; + accumulated_tokens = accumulated_tokens.saturating_add(response.usage.total_tokens); + merge_usage(&mut accumulated_usage, &response.usage); #[cfg(debug_assertions)] tracing::debug!( @@ -583,14 +700,23 @@ impl AgentLoop { if response.tool_calls.is_empty() { let mut assistant_message = ChatMessage::assistant(response.content); assistant_message.reasoning_content = response.reasoning_content; + assistant_message.provider_state = response.provider_state; + Self::annotate_message(&mut assistant_message, turn.as_ref(), iteration, true); emitted_messages.push(assistant_message.clone()); return Ok(AgentProcessResult { final_response: assistant_message, emitted_messages, total_tokens: Some(accumulated_tokens), + usage: Some(accumulated_usage), }); } + if let Some(turn) = turn.as_ref() { + turn.emitter + .emit(TurnEvent::TextSegmentFinished { iteration }) + .map_err(|error| AgentError::Other(format!("turn event rejected: {error}")))?; + } + // Execute tool calls — log and notify immediately { let tools_info: Vec = response @@ -614,11 +740,15 @@ impl AgentLoop { response.tool_calls.clone(), ); assistant_message.reasoning_content = response.reasoning_content; + assistant_message.provider_state = response.provider_state; + Self::annotate_message(&mut assistant_message, turn.as_ref(), iteration, false); messages.push(assistant_message.clone()); emitted_messages.push(assistant_message); // Execute tools and add results to messages - let tool_results = self.execute_tools(&response.tool_calls).await; + let tool_results = self + .execute_tools(&response.tool_calls, iteration, turn.as_ref()) + .await?; for (tool_call, result) in response.tool_calls.iter().zip(tool_results.iter()) { // Log function call with name and arguments @@ -644,22 +774,24 @@ impl AgentLoop { "Loop warning: {}", msg ); - let tool_message = ChatMessage::tool_with_media( + let mut tool_message = ChatMessage::tool_with_media( tool_call.id.clone(), tool_call.name.clone(), format!("{}\n\n[上一条结果]\n{}", msg, truncated_output), result.media_refs.clone(), ); + Self::annotate_message(&mut tool_message, turn.as_ref(), iteration, false); messages.push(tool_message.clone()); emitted_messages.push(tool_message); } LoopDetectionResult::Ok => { - let tool_message = ChatMessage::tool_with_media( + let mut tool_message = ChatMessage::tool_with_media( tool_call.id.clone(), tool_call.name.clone(), truncated_output, result.media_refs.clone(), ); + Self::annotate_message(&mut tool_message, turn.as_ref(), iteration, false); messages.push(tool_message.clone()); emitted_messages.push(tool_message); } @@ -695,25 +827,56 @@ impl AgentLoop { tools: None, // No tools in final summary call }; - match (*self.provider).chat(request).await { + let summary_iteration = u32::try_from(self.max_iterations) + .map_err(|_| AgentError::Other("tool iteration exceeds u32".to_string()))?; + match self + .stream_completion(request, summary_iteration, turn.as_ref()) + .await + { Ok(response) => { - accumulated_tokens += response.usage.total_tokens; + accumulated_tokens = accumulated_tokens.saturating_add(response.usage.total_tokens); + merge_usage(&mut accumulated_usage, &response.usage); let mut assistant_message = ChatMessage::assistant(response.content); assistant_message.reasoning_content = response.reasoning_content; + assistant_message.provider_state = response.provider_state; + Self::annotate_message( + &mut assistant_message, + turn.as_ref(), + summary_iteration, + true, + ); emitted_messages.push(assistant_message.clone()); Ok(AgentProcessResult { final_response: assistant_message, emitted_messages, total_tokens: Some(accumulated_tokens), + usage: Some(accumulated_usage), }) } Err(e) => { // Fallback if summary call fails tracing::error!(error = %e, "Failed to get summary from LLM"); - let final_message = ChatMessage::assistant(format!( + let fallback = format!( "I reached the maximum number of tool call iterations ({}) without completing the task. The work done so far has been lost due to an error. Please try breaking the task into smaller steps.", self.max_iterations - )); + ); + if let Some(turn) = turn.as_ref() { + turn.emitter + .emit(TurnEvent::TextSegmentFinished { + iteration: summary_iteration, + }) + .and_then(|()| { + turn.emitter.emit(TurnEvent::TextDelta { + iteration: summary_iteration, + delta: fallback.clone(), + }) + }) + .map_err(|error| { + AgentError::Other(format!("turn event rejected: {error}")) + })?; + } + let mut final_message = ChatMessage::assistant(fallback); + Self::annotate_message(&mut final_message, turn.as_ref(), summary_iteration, true); emitted_messages.push(final_message.clone()); Ok(AgentProcessResult { final_response: final_message, @@ -723,6 +886,7 @@ impl AgentLoop { } else { None }, + usage: (accumulated_usage.total_tokens > 0).then_some(accumulated_usage), }) } } @@ -753,42 +917,76 @@ impl AgentLoop { } /// Execute multiple tool calls, choosing parallel or sequential based on conditions. - async fn execute_tools(&self, tool_calls: &[ToolCall]) -> Vec { + async fn execute_tools( + &self, + tool_calls: &[ToolCall], + iteration: u32, + turn: Option<&AgentTurnContext>, + ) -> Result, AgentError> { if self.should_execute_in_parallel(tool_calls) { tracing::debug!("Executing {} tools in parallel", tool_calls.len()); - self.execute_tools_parallel(tool_calls).await + self.execute_tools_parallel(tool_calls, iteration, turn) + .await } else { tracing::debug!("Executing {} tools sequentially", tool_calls.len()); - self.execute_tools_sequential(tool_calls).await + self.execute_tools_sequential(tool_calls, iteration, turn) + .await } } /// Execute tools in parallel using join_all. - async fn execute_tools_parallel(&self, tool_calls: &[ToolCall]) -> Vec { + async fn execute_tools_parallel( + &self, + tool_calls: &[ToolCall], + iteration: u32, + turn: Option<&AgentTurnContext>, + ) -> Result, AgentError> { let futures: Vec<_> = tool_calls .iter() - .map(|tc| self.execute_one_tool(tc)) + .map(|tool_call| self.execute_one_tool(tool_call, iteration, turn)) .collect(); - futures_util::future::join_all(futures).await + futures_util::future::join_all(futures) + .await + .into_iter() + .collect() } /// Execute tools sequentially. - async fn execute_tools_sequential(&self, tool_calls: &[ToolCall]) -> Vec { + async fn execute_tools_sequential( + &self, + tool_calls: &[ToolCall], + iteration: u32, + turn: Option<&AgentTurnContext>, + ) -> Result, AgentError> { let mut outcomes = Vec::with_capacity(tool_calls.len()); for tool_call in tool_calls { - outcomes.push(self.execute_one_tool(tool_call).await); + outcomes.push(self.execute_one_tool(tool_call, iteration, turn).await?); } - outcomes + Ok(outcomes) } /// Execute a single tool and return the outcome with event tracking. - async fn execute_one_tool(&self, tool_call: &ToolCall) -> ToolExecutionOutcome { + async fn execute_one_tool( + &self, + tool_call: &ToolCall, + iteration: u32, + turn: Option<&AgentTurnContext>, + ) -> Result { let start = Instant::now(); let tool_name = tool_call.name.clone(); + if let Some(turn) = turn { + turn.emitter + .emit(TurnEvent::ToolStarted { + iteration, + call: tool_call.clone(), + }) + .map_err(|error| AgentError::Other(format!("turn event rejected: {error}")))?; + } + // Record ToolCallStart event if let Some(ref observer) = self.observer { observer.record_event(&ObserverEvent::ToolCallStart { @@ -800,6 +998,17 @@ impl AgentLoop { let result = self.execute_tool_internal(tool_call).await; let duration = start.elapsed(); + if let Some(turn) = turn { + turn.emitter + .emit(TurnEvent::ToolFinished { + iteration, + call_id: tool_call.id.clone(), + success: result.success, + preview: Some(tool_result_preview(&truncate_tool_result(&result.output))), + }) + .map_err(|error| AgentError::Other(format!("turn event rejected: {error}")))?; + } + // Record ToolCall event if let Some(ref observer) = self.observer { observer.record_event(&ObserverEvent::ToolCall { @@ -810,7 +1019,7 @@ impl AgentLoop { } // Apply duration - ToolExecutionOutcome { duration, ..result } + Ok(ToolExecutionOutcome { duration, ..result }) } /// Internal tool execution without event tracking. @@ -851,13 +1060,115 @@ impl AgentLoop { mod tests { use super::*; use crate::observability::{MultiObserver, Observer}; - use crate::providers::{ChatCompletionResponse, Usage}; + use crate::providers::{ + ChatCompletionResponse, FinishReason, ProviderChunk, ProviderStream, Usage, + }; + use crate::session::{TurnBlock, TurnController}; use crate::tools::FileReadTool; struct TestObserver { events: std::sync::Mutex>, } + struct StreamingTextProvider; + + #[async_trait::async_trait] + impl LLMProvider for StreamingTextProvider { + async fn stream( + &self, + _request: ChatCompletionRequest, + ) -> Result { + let chunks = vec![ + ProviderChunk::Metadata { + id: "response".into(), + model: "streaming-test".into(), + }, + ProviderChunk::Reasoning("because ".into()), + ProviderChunk::Reasoning("facts".into()), + ProviderChunk::Text("hello ".into()), + ProviderChunk::Text("world".into()), + ProviderChunk::ProviderState(crate::bus::ProviderReasoningState { + provider: "test".into(), + payload: serde_json::json!({"opaque":"state"}), + }), + ProviderChunk::Usage(Usage { + prompt_tokens: 2, + completion_tokens: 3, + total_tokens: 5, + ..Usage::default() + }), + ProviderChunk::Done(FinishReason::Stop), + ]; + Ok(Box::pin(futures_util::stream::iter( + chunks.into_iter().map(Ok), + ))) + } + + fn ptype(&self) -> &str { + "test" + } + + fn name(&self) -> &str { + "streaming-test" + } + + fn model_id(&self) -> &str { + "streaming-test" + } + } + + #[tokio::test] + async fn process_streaming_emits_turn_blocks_and_stamps_durable_message() { + let agent = AgentLoop::with_provider( + Arc::new(StreamingTextProvider), + 1, + "streaming-test".into(), + PathBuf::from("."), + Vec::new(), + ); + let (controller, emitter, _) = TurnController::start("session", "assistant-id"); + let initial = controller.snapshot(); + let context = + AgentTurnContext::new(initial.id.0.clone(), initial.message_id.clone(), emitter); + + let result = agent + .process_streaming(vec![ChatMessage::user("hi")], context) + .await + .unwrap(); + + assert_eq!(result.final_response.id, "assistant-id"); + assert_eq!(result.final_response.content, "hello world"); + assert_eq!( + result.final_response.reasoning_content.as_deref(), + Some("because facts") + ); + assert_eq!(result.final_response.turn_id, Some(initial.id.0.clone())); + assert_eq!(result.final_response.iteration, Some(0)); + assert_eq!( + result + .final_response + .provider_state + .as_ref() + .map(|state| state.provider.as_str()), + Some("test") + ); + assert_eq!( + result.usage.as_ref().map(|usage| usage.total_tokens), + Some(5) + ); + + let snapshot = controller.snapshot(); + assert_eq!(snapshot.blocks.len(), 2); + assert!(matches!( + &snapshot.blocks[0], + TurnBlock::Reasoning { text, .. } if text == "because facts" + )); + assert!(matches!( + &snapshot.blocks[1], + TurnBlock::Assistant { text, .. } if text == "hello world" + )); + } + impl TestObserver { fn new() -> Self { Self { @@ -973,8 +1284,13 @@ mod tests { vec!["text".to_string(), "image".to_string()], ); + let (controller, emitter, _) = TurnController::start("session", "assistant-id"); + let turn = controller.snapshot(); let result = agent - .process(vec![ChatMessage::user("inspect the image")]) + .process_streaming( + vec![ChatMessage::user("inspect the image")], + AgentTurnContext::new(turn.id.0.clone(), turn.message_id.clone(), emitter), + ) .await .unwrap(); @@ -999,6 +1315,14 @@ mod tests { .iter() .any(|media| media.media_type == "image") })); + assert!(controller.snapshot().blocks.iter().any(|block| matches!( + block, + TurnBlock::Tool { + id, + status: crate::session::ToolStatus::Completed, + .. + } if id == "call-image" + ))); } #[test] diff --git a/src/agent/mod.rs b/src/agent/mod.rs index 6e825b2..dc75e49 100644 --- a/src/agent/mod.rs +++ b/src/agent/mod.rs @@ -15,4 +15,4 @@ pub use system_prompt::{ PromptContext, PromptSection, SystemPromptBuilder, build_sub_agent_system_prompt, build_system_prompt, }; -pub use turn_event::{TurnEmitError, TurnEmitter, TurnEvent}; +pub use turn_event::{AgentTurnContext, TurnEmitError, TurnEmitter, TurnEvent}; diff --git a/src/agent/turn_event.rs b/src/agent/turn_event.rs index 7f7c1cc..da10f92 100644 --- a/src/agent/turn_event.rs +++ b/src/agent/turn_event.rs @@ -1,4 +1,4 @@ -use std::sync::Arc; +use std::sync::{Arc, Mutex}; use thiserror::Error; @@ -49,6 +49,29 @@ type EmitFn = dyn Fn(TurnEvent) -> Result<(), TurnEmitError> + Send + Sync; #[derive(Clone)] pub struct TurnEmitter { emit: Arc, + enabled: Arc>, +} + +/// Session-owned identity and emitter for one AgentLoop execution. +#[derive(Clone)] +pub struct AgentTurnContext { + pub turn_id: String, + pub message_id: String, + pub emitter: TurnEmitter, +} + +impl AgentTurnContext { + pub fn new( + turn_id: impl Into, + message_id: impl Into, + emitter: TurnEmitter, + ) -> Self { + Self { + turn_id: turn_id.into(), + message_id: message_id.into(), + emitter, + } + } } impl TurnEmitter { @@ -58,10 +81,27 @@ impl TurnEmitter { { Self { emit: Arc::new(emit), + enabled: Arc::new(Mutex::new(true)), } } pub fn emit(&self, event: TurnEvent) -> Result<(), TurnEmitError> { + let enabled = self + .enabled + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + if !*enabled { + return Ok(()); + } (self.emit)(event) } + + /// Stop forwarding new presentation facts without changing durable or + /// terminal Turn status. Session uses this before invalidating a worker. + pub fn deactivate(&self) { + *self + .enabled + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()) = false; + } } diff --git a/src/providers/mod.rs b/src/providers/mod.rs index c86abf4..9ebff78 100644 --- a/src/providers/mod.rs +++ b/src/providers/mod.rs @@ -8,8 +8,8 @@ pub use self::openai::OpenAIProvider; use crate::config::LLMProviderConfig; pub use stream::{ - DynProviderError, FinishReason, ProviderChunk, ProviderStream, ProviderStreamItem, - collect_provider_stream, provider_stream_from_response, + DynProviderError, FinishReason, ProviderChunk, ProviderResponseAccumulator, ProviderStream, + ProviderStreamItem, collect_provider_stream, provider_stream_from_response, }; pub use traits::{ ChatCompletionRequest, ChatCompletionResponse, LLMProvider, Message, Tool, ToolCall, diff --git a/src/providers/stream.rs b/src/providers/stream.rs index 6de3ed0..7236eca 100644 --- a/src/providers/stream.rs +++ b/src/providers/stream.rs @@ -64,34 +64,35 @@ struct PartialToolCall { arguments: String, } -pub async fn collect_provider_stream( - mut provider_stream: ProviderStream, -) -> Result { - let mut id = String::new(); - let mut model = String::new(); - let mut content = String::new(); - let mut reasoning_content = String::new(); - let mut provider_state = None; - let mut usage = Usage::default(); - let mut tool_calls = BTreeMap::::new(); +#[derive(Default)] +pub struct ProviderResponseAccumulator { + id: String, + model: String, + content: String, + reasoning_content: String, + provider_state: Option, + usage: Usage, + tool_calls: BTreeMap, +} - while let Some(chunk) = provider_stream.next().await { - match chunk? { +impl ProviderResponseAccumulator { + pub fn push(&mut self, chunk: ProviderChunk) { + match chunk { ProviderChunk::Metadata { id: response_id, model: response_model, } => { - id = response_id; - model = response_model; + self.id = response_id; + self.model = response_model; } - ProviderChunk::Text(delta) => content.push_str(&delta), - ProviderChunk::Reasoning(delta) => reasoning_content.push_str(&delta), + ProviderChunk::Text(delta) => self.content.push_str(&delta), + ProviderChunk::Reasoning(delta) => self.reasoning_content.push_str(&delta), ProviderChunk::ToolCallStart { index, id: call_id, name, } => { - let partial = tool_calls.entry(index).or_default(); + let partial = self.tool_calls.entry(index).or_default(); if call_id.is_some() { partial.id = call_id; } @@ -100,36 +101,51 @@ pub async fn collect_provider_stream( } } ProviderChunk::ToolCallArguments { index, delta } => { - tool_calls + self.tool_calls .entry(index) .or_default() .arguments .push_str(&delta); } - ProviderChunk::ProviderState(state) => provider_state = Some(state), - ProviderChunk::Usage(value) => usage = value, + ProviderChunk::ProviderState(state) => self.provider_state = Some(state), + ProviderChunk::Usage(value) => self.usage = value, ProviderChunk::Done(_) => {} } } - let tool_calls = tool_calls - .into_iter() - .map(|(index, partial)| ToolCall { - id: partial.id.unwrap_or_else(|| format!("tool_call_{index}")), - name: partial.name.unwrap_or_default(), - arguments: serde_json::from_str(&partial.arguments).unwrap_or(serde_json::Value::Null), - }) - .collect(); + pub fn finish(self) -> ChatCompletionResponse { + let tool_calls = self + .tool_calls + .into_iter() + .map(|(index, partial)| ToolCall { + id: partial.id.unwrap_or_else(|| format!("tool_call_{index}")), + name: partial.name.unwrap_or_default(), + arguments: serde_json::from_str(&partial.arguments) + .unwrap_or(serde_json::Value::Null), + }) + .collect(); - Ok(ChatCompletionResponse { - id, - model, - content, - reasoning_content: (!reasoning_content.is_empty()).then_some(reasoning_content), - provider_state, - tool_calls, - usage, - }) + ChatCompletionResponse { + id: self.id, + model: self.model, + content: self.content, + reasoning_content: (!self.reasoning_content.is_empty()) + .then_some(self.reasoning_content), + provider_state: self.provider_state, + tool_calls, + usage: self.usage, + } + } +} + +pub async fn collect_provider_stream( + mut provider_stream: ProviderStream, +) -> Result { + let mut accumulator = ProviderResponseAccumulator::default(); + while let Some(chunk) = provider_stream.next().await { + accumulator.push(chunk?); + } + Ok(accumulator.finish()) } /// Adapt a complete response to the stream-first provider contract. diff --git a/src/session/persistence.rs b/src/session/persistence.rs index 303925a..4de586f 100644 --- a/src/session/persistence.rs +++ b/src/session/persistence.rs @@ -1,10 +1,12 @@ use std::sync::Arc; +use std::{fmt::Display, future::Future}; use tokio::sync::Mutex; use super::session::{MessagePersistSnapshot, Session}; use crate::bus::ChatMessage; use crate::storage::StorageError; +use crate::{providers::Usage, session::TurnController}; async fn persist_added_messages( snapshots: Vec>, @@ -63,3 +65,74 @@ pub(super) async fn append_persisted_messages( } Ok(()) } + +/// Publish `Completed` only after the supplied durable write succeeds. +/// +/// Keeping this ordering in one helper makes the user-visible terminal status +/// impossible to publish optimistically before SQLite commits. +pub(super) async fn finalize_turn_after_persistence( + controller: &TurnController, + usage: Option, + persistence: F, +) -> Result +where + F: Future>, + E: Display, +{ + controller.begin_finalizing(); + match persistence.await { + Ok(value) => { + controller.complete(usage); + Ok(value) + } + Err(error) => { + controller.fail(format!("failed to persist turn: {error}")); + Err(error) + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::session::{TurnController, TurnStatus}; + + #[tokio::test] + async fn completed_is_published_only_after_persistence_succeeds() { + let (controller, _emitter, receiver) = TurnController::start("session", "message"); + let persisted = Arc::new(std::sync::atomic::AtomicBool::new(false)); + let persisted_in_future = persisted.clone(); + + let result: Result<(), String> = + finalize_turn_after_persistence(&controller, None, async move { + assert_eq!(receiver.borrow().status, TurnStatus::Running); + assert_eq!( + receiver.borrow().phase, + crate::session::TurnPhase::Finalizing + ); + persisted_in_future.store(true, std::sync::atomic::Ordering::SeqCst); + Ok(()) + }) + .await; + + assert!(result.is_ok()); + assert!(persisted.load(std::sync::atomic::Ordering::SeqCst)); + assert_eq!(controller.snapshot().status, TurnStatus::Completed); + } + + #[tokio::test] + async fn persistence_failure_never_publishes_completed() { + let (controller, _emitter, _receiver) = TurnController::start("session", "message"); + let result: Result<(), &str> = + finalize_turn_after_persistence(&controller, None, async { Err("database down") }) + .await; + + assert_eq!(result, Err("database down")); + let snapshot = controller.snapshot(); + assert_eq!(snapshot.status, TurnStatus::Failed); + assert_eq!( + snapshot.error.as_deref(), + Some("failed to persist turn: database down") + ); + } +} diff --git a/src/session/session.rs b/src/session/session.rs index ee17eca..e767e40 100644 --- a/src/session/session.rs +++ b/src/session/session.rs @@ -3,8 +3,11 @@ use std::sync::Arc; use tokio::sync::{Mutex, mpsc, oneshot}; -use super::persistence::append_persisted_messages; -use crate::bus::{ChatMessage, MediaItem, MediaRef, MessageSource, OutboundMessage, SourceKind}; +use super::persistence::{append_persisted_messages, finalize_turn_after_persistence}; +use super::turn::{TurnBlock, TurnController, TurnSnapshot}; +use crate::bus::{ + ChatMessage, CompletionStatus, MediaItem, MediaRef, MessageSource, OutboundMessage, SourceKind, +}; use crate::mcp::get_mcp_status; use crate::storage::{Storage, StorageError}; use std::sync::Arc as StdArc; @@ -37,7 +40,7 @@ pub enum HandleResult { } use crate::agent::context_compressor::ContextCompressionConfig; use crate::agent::system_prompt::{build_runtime_context, build_system_prompt}; -use crate::agent::{AgentError, AgentLoop, ContextCompressor}; +use crate::agent::{AgentError, AgentLoop, AgentTurnContext, ContextCompressor, TurnEmitter}; use crate::channels::slash_command::parse_slash_command; use crate::config::BrowserConfig; use crate::config::LLMProviderConfig; @@ -53,6 +56,115 @@ fn is_context_overflow_error(msg: &str) -> bool { || lower.contains("prompt is too long") || lower.contains("input is too long") } + +fn partial_assistant_message( + snapshot: &TurnSnapshot, + completion_status: CompletionStatus, +) -> Option { + let mut assistant_segments = Vec::new(); + let mut reasoning_segments = Vec::new(); + let mut last_iteration = None; + for block in &snapshot.blocks { + match block { + TurnBlock::Assistant { + iteration, text, .. + } if !text.is_empty() => { + assistant_segments.push(text.as_str()); + last_iteration = Some(*iteration); + } + TurnBlock::Reasoning { text, .. } if !text.is_empty() => { + reasoning_segments.push(text.as_str()); + } + _ => {} + } + } + if assistant_segments.is_empty() { + return None; + } + + let mut message = ChatMessage::assistant(assistant_segments.join("\n\n")); + message.id = snapshot.message_id.clone(); + message.turn_id = Some(snapshot.id.0.clone()); + message.iteration = last_iteration; + message.completion_status = completion_status; + message.reasoning_content = + (!reasoning_segments.is_empty()).then(|| reasoning_segments.join("\n\n")); + Some(message) +} + +async fn fail_turn_with_partial( + controller: &TurnController, + session: &Arc>, + error: String, +) { + let snapshot = controller.snapshot(); + let partial = partial_assistant_message(&snapshot, CompletionStatus::Interrupted); + if let Some(partial) = partial { + controller.begin_finalizing(); + if let Err(persistence_error) = append_persisted_messages(session, vec![partial]).await { + controller.fail(format!( + "{error}; failed to persist interrupted turn: {persistence_error}" + )); + return; + } + } + controller.fail(error); +} + +#[cfg(test)] +mod cancelled_partial_tests { + use super::*; + use crate::agent::TurnEvent; + + #[test] + fn visible_partial_text_becomes_cancelled_persisted_message() { + let (controller, emitter, _) = TurnController::start("session", "message-id"); + emitter + .emit(TurnEvent::ReasoningDelta { + iteration: 0, + delta: "reason".into(), + }) + .unwrap(); + emitter + .emit(TurnEvent::TextDelta { + iteration: 0, + delta: "first".into(), + }) + .unwrap(); + emitter + .emit(TurnEvent::TextSegmentFinished { iteration: 0 }) + .unwrap(); + emitter + .emit(TurnEvent::TextDelta { + iteration: 1, + delta: "second".into(), + }) + .unwrap(); + + let message = + partial_assistant_message(&controller.snapshot(), CompletionStatus::Cancelled).unwrap(); + assert_eq!(message.id, "message-id"); + assert_eq!(message.content, "first\n\nsecond"); + assert_eq!(message.reasoning_content.as_deref(), Some("reason")); + assert_eq!(message.iteration, Some(1)); + assert_eq!(message.completion_status, CompletionStatus::Cancelled); + } + + #[test] + fn reasoning_only_cancel_does_not_create_assistant_history() { + let (controller, emitter, _) = TurnController::start("session", "message-id"); + emitter + .emit(TurnEvent::ReasoningDelta { + iteration: 0, + delta: "private".into(), + }) + .unwrap(); + assert!( + partial_assistant_message(&controller.snapshot(), CompletionStatus::Cancelled,) + .is_none() + ); + } +} use crate::bus::MessageBus; use crate::providers::{LLMProvider, create_provider}; use crate::session::events::DialogInfo; @@ -93,6 +205,7 @@ pub struct Session { agent_tx: Option>, /// Cancel signal for the currently executing agent task current_cancel: Option>, + active_turn_emitter: Option, /// Monotonic counter to detect stale workers worker_generation: u64, /// Monotonic counter for in-memory session mutations. @@ -108,6 +221,11 @@ pub struct Session { pub(super) persistence_lock: Arc>, } +struct ActiveTurnEmitter { + turn_id: String, + emitter: TurnEmitter, +} + /// A task to be processed by the per-session agent worker struct AgentTask { channel: String, @@ -178,6 +296,7 @@ impl Session { memory_manager, agent_tx: None, current_cancel: None, + active_turn_emitter: None, worker_generation: 0, state_version: 0, persistence_lock: Arc::new(Mutex::new(())), @@ -366,6 +485,7 @@ impl Session { memory_manager, agent_tx: None, current_cancel: None, + active_turn_emitter: None, worker_generation: 0, state_version: 0, persistence_lock: Arc::new(Mutex::new(())), @@ -1473,6 +1593,9 @@ impl SessionManager { if guard.current_cancel.take().is_some() { msgs.push("当前任务已发送停止信号。".to_string()); } + if let Some(active_turn) = guard.active_turn_emitter.take() { + active_turn.emitter.deactivate(); + } if guard.agent_tx.take().is_some() { msgs.push("消息队列已清空。".to_string()); } @@ -2391,12 +2514,40 @@ fn spawn_agent_worker( Session::append_runtime_context_to_user_message(last_msg, &runtime_context); } + let (turn_controller, turn_emitter, _turn_receiver) = TurnController::start( + unified_str.clone(), + uuid::Uuid::new_v4().to_string(), + ); + let initial_turn = turn_controller.snapshot(); + let active_turn_id = initial_turn.id.0.clone(); + { + let mut guard = session.lock().await; + if guard.worker_generation != worker_gen || guard.state_version != base_version { + turn_emitter.deactivate(); + turn_controller.cancel(Some( + "session changed before model execution".to_string(), + )); + guard.current_cancel = None; + continue 'tasks; + } + guard.active_turn_emitter = Some(ActiveTurnEmitter { + turn_id: initial_turn.id.0.clone(), + emitter: turn_emitter.clone(), + }); + } + let agent_turn = AgentTurnContext::new( + initial_turn.id.0.clone(), + initial_turn.message_id.clone(), + turn_emitter, + ); + // Phase 2 + 3: LLM call with cancellation let session2 = session.clone(); let bus2 = bus.clone(); let chan2 = task_chan.clone(); let cid2 = task_cid.clone(); let unified_str2 = unified_str.clone(); + let turn_lifecycle = &turn_controller; let process_future = async move { let response_session_id = unified_str2.clone(); let process_result = crate::agent::sub_agent::DELEGATE_CONTEXT.scope( @@ -2405,7 +2556,7 @@ fn spawn_agent_worker( channel: chan2.clone(), chat_id: cid2.clone(), }, - agent.process(history_out.clone()), + agent.process_streaming(history_out.clone(), agent_turn.clone()), ).await; let result = match process_result { Ok(r) => r, @@ -2435,6 +2586,12 @@ fn spawn_agent_worker( Ok(r) => r, Err(e) => { tracing::error!(error = %e, "Retry compression failed"); + fail_turn_with_partial( + turn_lifecycle, + &session2, + format!("context overflow handling failed: {e}"), + ) + .await; let err_outbound = OutboundMessage { channel: chan2, chat_id: cid2, @@ -2459,6 +2616,10 @@ fn spawn_agent_worker( session_id = %guard.id, "Session changed while retry-compressing after context overflow" ); + turn_lifecycle.cancel(Some( + "session changed during context overflow recovery" + .to_string(), + )); return; } guard.compressor.set_context_window(new_window); @@ -2490,13 +2651,22 @@ fn spawn_agent_worker( retry }; - match agent.process(retry_history).await { + match agent + .process_streaming(retry_history, agent_turn.clone()) + .await + { Ok(r) => r, Err(e) => { tracing::error!( error = %e, "Agent retry after overflow failed" ); + fail_turn_with_partial( + turn_lifecycle, + &session2, + e.to_string(), + ) + .await; let err_outbound = OutboundMessage { channel: chan2, chat_id: cid2, @@ -2513,6 +2683,12 @@ fn spawn_agent_worker( } Err(e) => { tracing::error!(error = %e, "Agent processing error"); + fail_turn_with_partial( + turn_lifecycle, + &session2, + e.to_string(), + ) + .await; let err_outbound = OutboundMessage { channel: chan2, chat_id: cid2, @@ -2529,16 +2705,36 @@ fn spawn_agent_worker( let response_content = result.final_response.content; let total_tokens = result.total_tokens; - let response = - if let Err(e) = append_persisted_messages(&session2, result.emitted_messages).await { - tracing::error!(error = %e, "Failed to atomically persist agent turn"); - None - } else { + let usage = result.usage; + { + let guard = session2.lock().await; + if guard.worker_generation != worker_gen + || guard.state_version != base_version + { + turn_lifecycle.cancel(Some( + "session changed before turn commit".to_string(), + )); + return; + } + } + let response = match finalize_turn_after_persistence( + turn_lifecycle, + usage, + append_persisted_messages(&session2, result.emitted_messages), + ) + .await + { + Ok(()) => { let mut guard = session2.lock().await; let sent_count = guard.messages.len(); guard.compressor.set_last_api_info(sent_count, total_tokens); Some(response_content) - }; + } + Err(e) => { + tracing::error!(error = %e, "Failed to atomically persist agent turn"); + None + } + }; let Some(response) = response else { let err_outbound = OutboundMessage { @@ -2575,11 +2771,39 @@ fn spawn_agent_worker( () = process_future => {} _ = cancel_rx => { // cancelled — current_cancel already taken by /stop + let snapshot = turn_controller.snapshot(); + if let Some(partial) = partial_assistant_message( + &snapshot, + CompletionStatus::Cancelled, + ) { + turn_controller.begin_finalizing(); + match append_persisted_messages(&session, vec![partial]).await { + Ok(()) => { + turn_controller.cancel(Some("stopped by user".to_string())); + } + Err(error) => { + tracing::error!(error = %error, "Failed to persist cancelled partial turn"); + turn_controller.fail(format!( + "failed to persist cancelled turn: {error}" + )); + } + } + } else { + turn_controller.cancel(Some("stopped by user".to_string())); + } } } // Clean up let mut guard = session.lock().await; + if guard + .active_turn_emitter + .as_ref() + .is_some_and(|active| active.turn_id == active_turn_id) + && let Some(active) = guard.active_turn_emitter.take() + { + active.emitter.deactivate(); + } if guard.worker_generation == worker_gen { guard.current_cancel = None; } diff --git a/src/session/turn.rs b/src/session/turn.rs index 89d042b..a102394 100644 --- a/src/session/turn.rs +++ b/src/session/turn.rs @@ -589,4 +589,18 @@ mod tests { assert_eq!(snapshot.error.as_deref(), Some("provider disconnected")); assert_eq!(snapshot.phase, TurnPhase::Finalizing); } + + #[test] + fn deactivated_emitter_drops_events_before_stale_worker_reduction() { + let (controller, emitter, _) = start(); + emitter.deactivate(); + emitter + .emit(TurnEvent::TextDelta { + iteration: 0, + delta: "late".into(), + }) + .unwrap(); + assert_eq!(controller.snapshot().revision, 0); + assert!(controller.snapshot().blocks.is_empty()); + } }