diff --git a/crates/goose-provider-types/src/conversation/message.rs b/crates/goose-provider-types/src/conversation/message.rs index ce6133a11..d6bfda109 100644 --- a/crates/goose-provider-types/src/conversation/message.rs +++ b/crates/goose-provider-types/src/conversation/message.rs @@ -859,6 +859,14 @@ impl Message { self.with_id(format!("msg_{}", Uuid::new_v4())) } + pub fn with_generated_id_if_missing(self) -> Self { + if self.id.is_some() { + self + } else { + self.with_generated_id() + } + } + /// Add any MessageContent to the message pub fn with_content(mut self, content: MessageContentBlock) -> Self { self.content.push(content); diff --git a/crates/goose-provider-types/src/formats/ollama.rs b/crates/goose-provider-types/src/formats/ollama.rs index 76b666e85..97b3729fb 100644 --- a/crates/goose-provider-types/src/formats/ollama.rs +++ b/crates/goose-provider-types/src/formats/ollama.rs @@ -221,7 +221,8 @@ where Role::Assistant, chrono::Utc::now().timestamp(), vec![MessageContentBlock::text(&accumulated_text)], - ); + ) + .with_generated_id(); yield (Some(msg), last_usage); } @@ -353,8 +354,10 @@ hello } } - #[test] - fn test_response_to_message_xml_fallback() -> anyhow::Result<()> { + #[tokio::test] + async fn test_response_to_message_xml_fallback() -> anyhow::Result<()> { + use futures::StreamExt; + // Test that response_to_message falls back to XML parsing when no JSON tool_calls let response = json!({ "choices": [{ @@ -375,6 +378,32 @@ hello panic!("Expected ToolRequest content from XML parsing"); } + let response_lines = r#"data: {"id":"ollama-source-id","model":"test-model","choices":[{"delta":{"role":"assistant","content":"literal response_tx, + _ => panic!("expected ACP prompt request"), + }; + + for id in ["call-1", "call-2"] { + response_tx + .send(AcpUpdate::ToolCallStart { + id: id.to_string(), + name: "read_file".to_string(), + kind: ToolKind::Read, + raw_input: None, + }) + .await + .unwrap(); + response_tx + .send(AcpUpdate::ToolCallComplete { + id: id.to_string(), + raw_output: None, + content: None, + is_error: false, + }) + .await + .unwrap(); + } + response_tx + .send(AcpUpdate::Complete(StopReason::EndTurn, None)) + .await + .unwrap(); + + let mut messages = Vec::new(); + while let Some(item) = stream.next().await { + let (message, usage) = item.unwrap(); + assert!(usage.is_none()); + if let Some(message) = message { + messages.push(message); + } + } + + assert_eq!(messages.len(), 2); + assert_eq!(messages[0].as_concat_text(), "Tool call was denied."); + assert_eq!(messages[1].as_concat_text(), "Tool call was denied."); + + let first_id = messages[0] + .id + .as_deref() + .expect("first denial should have a provider message ID"); + let second_id = messages[1] + .id + .as_deref() + .expect("second denial should have a provider message ID"); + + assert!(first_id.starts_with("msg_")); + assert!(second_id.starts_with("msg_")); + assert_ne!(first_id, second_id); + } + #[test] fn live_acp_text_update_preserves_assistant_only_audience() { let text = TextContent::new("assistant-only") diff --git a/crates/goose/src/acp/server.rs b/crates/goose/src/acp/server.rs index 23222560b..4e959e2a7 100644 --- a/crates/goose/src/acp/server.rs +++ b/crates/goose/src/acp/server.rs @@ -45,7 +45,7 @@ use agent_client_protocol::schema::v1::{ EmbeddedResourceResource, FileSystemCapabilities, ForkSessionRequest, ForkSessionResponse, ImageContent, Implementation, InitializeRequest, InitializeResponse, ListSessionsRequest, ListSessionsResponse, LoadSessionRequest, LoadSessionResponse, McpCapabilities, McpServer, - Meta, NewSessionRequest, NewSessionResponse, PermissionOption, PermissionOptionKind, + MessageId, Meta, NewSessionRequest, NewSessionResponse, PermissionOption, PermissionOptionKind, PromptCapabilities, PromptRequest, PromptResponse, RequestPermissionOutcome, RequestPermissionRequest, ResourceLink, SessionCapabilities, SessionCloseCapabilities, SessionConfigOption, SessionId, SessionInfoUpdate, SessionListCapabilities, @@ -1001,9 +1001,12 @@ impl GooseAcpAgent { ) -> Result<(), agent_client_protocol::Error> { match content_item { MessageContent::Text(text) => { - let chunk = - ContentChunk::new(ContentBlock::Text(TextContent::new(text.text.clone()))) - .meta(message_update_meta(message_id, message_created, steer)); + let chunk = content_chunk_for_message( + ContentBlock::Text(TextContent::new(text.text.clone())), + message_id, + message_created, + steer, + ); let update = match role { Role::User => SessionUpdate::UserMessageChunk(chunk), Role::Assistant => SessionUpdate::AgentMessageChunk(chunk), @@ -1011,15 +1014,8 @@ impl GooseAcpAgent { cx.send_notification(SessionNotification::new(session_id.clone(), update))?; } MessageContent::ToolRequest(tool_request) => { - self.handle_tool_request( - tool_request, - session_id, - session_id_str, - message_id, - agent, - cx, - ) - .await?; + self.handle_tool_request(tool_request, session_id, session_id_str, agent, cx) + .await?; } MessageContent::ToolResponse(tool_response) => { self.handle_tool_response( @@ -1033,16 +1029,12 @@ impl GooseAcpAgent { MessageContent::Thinking(thinking) => { cx.send_notification(SessionNotification::new( session_id.clone(), - SessionUpdate::AgentThoughtChunk( - ContentChunk::new(ContentBlock::Text(TextContent::new( - thinking.thinking.clone(), - ))) - .meta(message_update_meta( - message_id, - message_created, - steer, - )), - ), + SessionUpdate::AgentThoughtChunk(content_chunk_for_message( + ContentBlock::Text(TextContent::new(thinking.thinking.clone())), + message_id, + message_created, + steer, + )), ))?; } MessageContent::ActionRequired(action_required) => match &action_required.data { @@ -1098,8 +1090,12 @@ impl GooseAcpAgent { ), ); } - let chunk = ContentChunk::new(ContentBlock::Image(image_content)) - .meta(message_update_meta(message_id, message_created, steer)); + let chunk = content_chunk_for_message( + ContentBlock::Image(image_content), + message_id, + message_created, + steer, + ); let update = match role { Role::User => SessionUpdate::UserMessageChunk(chunk), Role::Assistant => SessionUpdate::AgentMessageChunk(chunk), @@ -1145,7 +1141,6 @@ impl GooseAcpAgent { tool_request: &ToolRequest, session_id: &SessionId, session_id_for_persist: &str, - message_id: Option<&str>, agent: &Arc, cx: &ConnectionTo, ) -> Result<(), agent_client_protocol::Error> { @@ -1165,7 +1160,6 @@ impl GooseAcpAgent { tool_call_notifier, &self.session_manager, session_id_for_persist, - message_id, tool_request, ); } @@ -1256,18 +1250,6 @@ impl GooseAcpAgent { Ok(()) } - - fn is_builtin_agent_command(command: &str) -> bool { - let normalized = command.trim_start_matches('/'); - - crate::agents::execute_commands::list_commands() - .iter() - .any(|cmd| cmd.name == normalized) - || crate::agents::execute_commands::COMPACT_TRIGGERS - .iter() - .filter_map(|trigger| trigger.strip_prefix('/')) - .any(|trigger| trigger == normalized) - } } fn extract_client_supports_goose_custom_notifications( @@ -1423,6 +1405,20 @@ fn message_update_meta(message_id: Option<&str>, created: i64, steer: bool) -> M meta } +fn content_chunk_for_message( + content: ContentBlock, + message_id: Option<&str>, + created: i64, + steer: bool, +) -> ContentChunk { + let mut chunk = + ContentChunk::new(content).meta(message_update_meta(message_id, created, steer)); + if let Some(message_id) = message_id { + chunk = chunk.message_id(MessageId::new(message_id)); + } + chunk +} + impl GooseAcpAgent { async fn on_initialize( &self, @@ -1766,35 +1762,6 @@ impl GooseAcpAgent { let user_message = Self::convert_acp_prompt_to_message(&args.prompt); - let message_text = user_message.as_concat_text(); - if let Some(parsed) = crate::agents::execute_commands::parse_slash_command(&message_text) { - let full_command = format!("/{}", parsed.command); - - if !Self::is_builtin_agent_command(parsed.command) { - if let Some(recipe_path) = - crate::slash_commands::recipe_slash_command::get_recipe_for_command( - &full_command, - ) - { - if recipe_path.exists() { - if let Err(error) = cx.send_notification(SessionNotification::new( - args.session_id.clone(), - SessionUpdate::AgentMessageChunk(ContentChunk::new( - ContentBlock::Text(TextContent::new(format!( - "Running recipe: {}", - full_command - ))), - )), - )) { - self.clear_active_run(&session_id, &run_id).await; - let _ = Self::send_active_run_update(cx, &args.session_id, None); - return Err(error); - } - } - } - } - } - let session_config = SessionConfig { id: session_id.clone(), schedule_id: None, @@ -1871,12 +1838,7 @@ impl GooseAcpAgent { let ready_chain = match content_item { MessageContent::ToolRequest(tool_request) => { - if let Some(message_id) = stored_message_id.as_deref() { - chain_tracker.record_request( - tool_request.clone(), - message_id.to_string(), - ); - } + chain_tracker.record_request(tool_request.clone()); None } MessageContent::ToolResponse(tool_response) => { @@ -2494,6 +2456,15 @@ print(\"hello, world\") "messageId": "msg_live", })), ); + + let chunk = content_chunk_for_message( + ContentBlock::Text(TextContent::new("hello")), + Some("msg_live"), + 1_700_000_000, + true, + ); + + assert_eq!(chunk.message_id, Some(MessageId::new("msg_live"))); } #[test] diff --git a/crates/goose/src/acp/server/load_session.rs b/crates/goose/src/acp/server/load_session.rs index 5114327d4..4ef4c2cd3 100644 --- a/crates/goose/src/acp/server/load_session.rs +++ b/crates/goose/src/acp/server/load_session.rs @@ -3,7 +3,7 @@ use super::tool_calls::conversion::{ }; use super::tool_calls::enrichment::tool_chain_summary; use super::*; -use agent_client_protocol::schema::v1::ToolCall; +use agent_client_protocol::schema::v1::{MessageId, ToolCall}; fn replay_message_meta(message: &Message) -> Meta { let mut meta = serde_json::Map::new(); @@ -62,7 +62,7 @@ fn send_replay_content_chunk( message: &Message, content: ContentBlock, ) -> std::result::Result<(), agent_client_protocol::Error> { - let chunk = ContentChunk::new(content).meta(replay_message_meta(message)); + let chunk = replay_content_chunk_for_message(message, content); let update = match message.role { Role::User => SessionUpdate::UserMessageChunk(chunk), Role::Assistant => SessionUpdate::AgentMessageChunk(chunk), @@ -70,6 +70,14 @@ fn send_replay_content_chunk( cx.send_notification(SessionNotification::new(session_id.clone(), update)) } +fn replay_content_chunk_for_message(message: &Message, content: ContentBlock) -> ContentChunk { + let mut chunk = ContentChunk::new(content).meta(replay_message_meta(message)); + if let Some(message_id) = message.id.as_deref() { + chunk = chunk.message_id(MessageId::new(message_id)); + } + chunk +} + fn build_replayed_tool_call( tool_request: &ToolRequest, client_requests_tool_call_label_enrichment: bool, @@ -175,12 +183,10 @@ fn replay_conversation_to_client( MessageContent::Thinking(thinking) => { cx.send_notification(SessionNotification::new( session_id.clone(), - SessionUpdate::AgentThoughtChunk( - ContentChunk::new(ContentBlock::Text(TextContent::new( - thinking.thinking.clone(), - ))) - .meta(replay_message_meta(message)), - ), + SessionUpdate::AgentThoughtChunk(replay_content_chunk_for_message( + message, + ContentBlock::Text(TextContent::new(thinking.thinking.clone())), + )), ))?; } MessageContent::SystemNotification(_) => {} @@ -371,6 +377,13 @@ mod tests { "messageId": "msg_2", })), ); + + let chunk = replay_content_chunk_for_message( + &message, + ContentBlock::Text(TextContent::new("replayed text")), + ); + + assert_eq!(chunk.message_id, Some(MessageId::new("msg_2"))); } #[test] diff --git a/crates/goose/src/acp/server/tool_calls/chain.rs b/crates/goose/src/acp/server/tool_calls/chain.rs index bf1c7c125..207aacd72 100644 --- a/crates/goose/src/acp/server/tool_calls/chain.rs +++ b/crates/goose/src/acp/server/tool_calls/chain.rs @@ -15,14 +15,12 @@ struct ToolChainStep { #[derive(Debug)] struct TrackedToolChain { - message_id: String, steps: Vec, } impl TrackedToolChain { - fn new(request: ToolRequest, message_id: String) -> Self { + fn new(request: ToolRequest) -> Self { Self { - message_id, steps: vec![ToolChainStep { request, responded: false, @@ -61,14 +59,12 @@ impl TrackedToolChain { fn into_ready(self) -> ReadyToolChain { ReadyToolChain { - message_id: self.message_id, tool_requests: self.steps.into_iter().map(|step| step.request).collect(), } } } pub(crate) struct ReadyToolChain { - pub(crate) message_id: String, pub(crate) tool_requests: Vec, } @@ -80,11 +76,11 @@ pub(crate) struct ToolChainTracker { } impl ToolChainTracker { - pub(crate) fn record_request(&mut self, request: ToolRequest, message_id: String) { + pub(crate) fn record_request(&mut self, request: ToolRequest) { if let Some(chain) = &mut self.current_chain { chain.add_request(request); } else { - self.current_chain = Some(TrackedToolChain::new(request, message_id)); + self.current_chain = Some(TrackedToolChain::new(request)); } } @@ -155,20 +151,19 @@ mod tests { let mut tracker = ToolChainTracker::default(); for id in ["a", "b", "c"] { - tracker.record_request(request(id), format!("message-{id}")); + tracker.record_request(request(id)); assert!(tracker.record_response(id).is_none()); } let ready = tracker.close_current_chain().expect("A-B-C is ready"); assert_eq!(request_ids(&ready), ["a", "b", "c"]); - assert_eq!(ready.message_id, "message-a"); } #[test] fn closed_chain_waits_for_its_last_response() { let mut tracker = ToolChainTracker::default(); for id in ["a", "b", "c"] { - tracker.record_request(request(id), format!("message-{id}")); + tracker.record_request(request(id)); } tracker.record_response("a"); tracker.record_response("b"); @@ -183,19 +178,19 @@ mod tests { fn boundary_separates_request_runs_and_discards_singletons() { let mut tracker = ToolChainTracker::default(); for id in ["a", "b"] { - tracker.record_request(request(id), format!("message-{id}")); + tracker.record_request(request(id)); tracker.record_response(id); } let first = tracker.close_current_chain().expect("A-B is ready"); assert_eq!(request_ids(&first), ["a", "b"]); - tracker.record_request(request("c"), "message-c".to_string()); + tracker.record_request(request("c")); tracker.record_response("c"); assert!(tracker.close_current_chain().is_none()); for id in ["d", "e"] { - tracker.record_request(request(id), format!("message-{id}")); + tracker.record_request(request(id)); tracker.record_response(id); } let second = tracker.close_current_chain().expect("D-E is ready"); diff --git a/crates/goose/src/acp/server/tool_calls/enrichment.rs b/crates/goose/src/acp/server/tool_calls/enrichment.rs index 2101fee22..6104bf4c2 100644 --- a/crates/goose/src/acp/server/tool_calls/enrichment.rs +++ b/crates/goose/src/acp/server/tool_calls/enrichment.rs @@ -36,13 +36,11 @@ pub(crate) fn spawn_tool_title_enrichment( tool_call_notifier: ToolCallNotifier, session_manager: &Arc, session_id: &str, - message_id: Option<&str>, tool_request: &ToolRequest, ) { let agent = agent.clone(); let session_manager = session_manager.clone(); let session_id = session_id.to_string(); - let message_id = message_id.map(str::to_string); let tool_request = tool_request.clone(); spawn(async move { @@ -50,7 +48,6 @@ pub(crate) fn spawn_tool_title_enrichment( agent.as_ref(), session_manager.as_ref(), &session_id, - message_id.as_deref(), &tool_request, ) .await @@ -75,17 +72,13 @@ pub(crate) fn spawn_chain_summary_enrichment( let session_manager = session_manager.clone(); spawn(async move { - let ReadyToolChain { - message_id, - tool_requests, - } = chain; + let ReadyToolChain { tool_requests } = chain; let first_tool_call_id = tool_requests[0].id.clone(); let Some(summary) = generate_tool_chain_summary( agent.as_ref(), session_manager.as_ref(), &session_id.0, - &message_id, &tool_requests, ) .await diff --git a/crates/goose/src/agents/agent.rs b/crates/goose/src/agents/agent.rs index d83bdf1f4..9f1878012 100644 --- a/crates/goose/src/agents/agent.rs +++ b/crates/goose/src/agents/agent.rs @@ -8,7 +8,6 @@ use anyhow::{anyhow, Context, Result}; use futures::stream::BoxStream; use futures::{stream, FutureExt, Stream, StreamExt, TryStreamExt}; use tracing_futures::Instrument; -use uuid::Uuid; use super::container::Container; use super::final_output_tool::FinalOutputTool; @@ -276,6 +275,40 @@ pub enum AgentEvent { HistoryReplaced(Conversation), } +fn ensure_message_event_id(event: AgentEvent) -> AgentEvent { + match event { + AgentEvent::Message(message) => AgentEvent::Message(message.with_generated_id_if_missing()), + other => other, + } +} + +fn push_message_with_id(messages: &mut Conversation, message: Message) -> Message { + let message = message.with_generated_id_if_missing(); + messages.push(message.clone()); + message +} + +async fn persist_message_with_id( + session_manager: &SessionManager, + session_id: &str, + message: Message, +) -> Result { + let message = message.with_generated_id_if_missing(); + session_manager.add_message(session_id, &message).await?; + Ok(message) +} + +async fn persist_and_push_message_with_id( + session_manager: &SessionManager, + session_id: &str, + conversation: &mut Conversation, + message: Message, +) -> Result { + let message = persist_message_with_id(session_manager, session_id, message).await?; + conversation.push(message.clone()); + Ok(message) +} + fn project_message_for_user_event(message: &Message) -> Message { message.user_visible_content() } @@ -287,12 +320,22 @@ fn agent_visible_message_text(message: &Message) -> String { fn attach_turn_usage( messages: &mut Conversation, usage: &ProviderUsage, + preferred_message_id: Option<&str>, ) -> Option<(Option, MessageUsage)> { - let message = messages - .messages_mut() - .iter_mut() - .rev() - .find(|m| m.role == rmcp::model::Role::Assistant)?; + let message_index = preferred_message_id + .and_then(|preferred_message_id| { + messages.messages().iter().rposition(|message| { + message.role == rmcp::model::Role::Assistant + && message.id.as_deref() == Some(preferred_message_id) + }) + }) + .or_else(|| { + messages + .messages() + .iter() + .rposition(|message| message.role == rmcp::model::Role::Assistant) + })?; + let message = &mut messages.messages_mut()[message_index]; let has_user_visible_content = !message.user_visible_content().content.is_empty(); let message_usage = MessageUsage::from_provider_usage(usage, false); message.metadata.usage = Some(Box::new(message_usage.clone())); @@ -1552,6 +1595,22 @@ impl Agent { session_config: SessionConfig, cancel_token: Option, ) -> Result>> { + let events = self + .reply_impl(user_message, session_config, cancel_token) + .await?; + + // This is the single live-event identity boundary. Callers that intentionally stream + // multiple events for one logical message must assign their shared ID before this point. + Ok(Box::pin(events.map_ok(ensure_message_event_id))) + } + + async fn reply_impl( + &self, + user_message: Message, + session_config: SessionConfig, + cancel_token: Option, + ) -> Result>> { + let user_message = user_message.with_generated_id_if_missing(); let session_manager = self.config.session_manager.clone(); let message_text_for_trace = agent_visible_message_text(&user_message); @@ -1660,6 +1719,8 @@ impl Agent { if response.role == rmcp::model::Role::Assistant && crate::agents::execute_commands::command_starts_turn(&message_text) => { + let response = response.with_generated_id_if_missing(); + // Setting a goal/grind should immediately start a turn so the // agent begins pursuing it, rather than waiting for the next // user prompt. Record the command and its confirmation as @@ -1695,6 +1756,8 @@ impl Agent { ]; } Ok(Some(response)) if response.role == rmcp::model::Role::Assistant => { + let response = response.with_generated_id_if_missing(); + session_manager .add_message( &session_config.id, @@ -1967,8 +2030,13 @@ impl Agent { .emit(crate::hooks::HookEvent::UserPromptSubmit, ctx) .await; } - session_manager.add_message(&session_config.id, &message).await?; - conversation.push(message.clone()); + let message = persist_and_push_message_with_id( + &session_manager, + &session_config.id, + &mut conversation, + message, + ) + .await?; yield AgentEvent::Message(message); } } @@ -1979,7 +2047,9 @@ impl Agent { }; if let Some(output) = final_output { last_assistant_text = output.clone(); - let message = Message::assistant().with_text(output); + let message = Message::assistant() + .with_text(output) + .with_generated_id_if_missing(); yield AgentEvent::Message(message.clone()); session_manager.add_message(&session_config.id, &message).await?; conversation.push(message); @@ -1995,15 +2065,23 @@ impl Agent { crate::hooks::HookDecision::Deny { reason, plugin } => { consecutive_stop_hook_blocks += 1; if consecutive_stop_hook_blocks > stop_hook_block_cap { - let message = stop_hook_block_cap_warning(&plugin, stop_hook_block_cap); - session_manager.add_message(&session_config.id, &message).await?; + let message = persist_message_with_id( + &session_manager, + &session_config.id, + stop_hook_block_cap_warning(&plugin, stop_hook_block_cap), + ) + .await?; yield AgentEvent::Message(message); stop_hook_handled_for_exit = true; break; } - let message = stop_hook_denial_context_message(&plugin, &reason); - session_manager.add_message(&session_config.id, &message).await?; - conversation.push(message); + persist_and_push_message_with_id( + &session_manager, + &session_config.id, + &mut conversation, + stop_hook_denial_context_message(&plugin, &reason), + ) + .await?; yield AgentEvent::Message(stop_hook_denial_notification(&plugin)); retrying_after_stop_hook_denial = true; continue; @@ -2071,6 +2149,7 @@ impl Agent { let mut provider_produced_content = false; let mut pending_final_output: Option = None; let mut pending_turn_usage: Option = None; + let mut preferred_turn_usage_message_id: Option = None; // Track whether this provider turn has already emitted visible // thinking so a later tool-call chunk can suppress replayed @@ -2286,10 +2365,8 @@ impl Agent { match tool_item { Some((request_id, item)) => { match item { - ToolStreamItem::ActionRequired(mut msg) => { - if msg.id.is_none() { - msg = msg.with_generated_id(); - } + ToolStreamItem::ActionRequired(msg) => { + let msg = msg.with_generated_id_if_missing(); if let Err(e) = session_manager.add_message(&session_config.id, &msg).await { warn!("Failed to save elicitation message to session: {}", e); } @@ -2456,9 +2533,33 @@ impl Agent { merged }; + let response_message_id = response + .id + .as_deref() + .expect("provider stream responses have IDs"); + let has_existing_message_id_carrier = messages_to_add + .iter() + .any(|message| { + message.id.as_deref() == Some(response_message_id) + }); + let carrier_tool_call_id = if has_existing_message_id_carrier { + None + } else { + remaining_requests + .first() + .or_else(|| frontend_requests.first()) + .map(|request| request.id.as_str()) + }; + preferred_turn_usage_message_id = + Some(response_message_id.to_owned()); + for request in frontend_requests.iter().chain(remaining_requests.iter()) { - let mut request_msg = Message::assistant() - .with_id(format!("msg_{}", Uuid::new_v4())); + let mut request_msg = + if carrier_tool_call_id == Some(request.id.as_str()) { + Message::assistant().with_id(response_message_id) + } else { + Message::assistant().with_generated_id() + }; for thinking in &response_thinking { request_msg = request_msg.with_content(thinking.clone()); @@ -2706,8 +2807,10 @@ impl Agent { match final_output { Some(None) => { warn!("Final output tool has not been called yet. Continuing agent loop."); - let message = Message::user().with_text(FINAL_OUTPUT_CONTINUATION_MESSAGE); - messages_to_add.push(message.clone()); + let message = push_message_with_id( + &mut messages_to_add, + Message::user().with_text(FINAL_OUTPUT_CONTINUATION_MESSAGE), + ); yield AgentEvent::Message(message); } Some(Some(output)) => { @@ -2728,7 +2831,7 @@ impl Agent { ); let message = Message::user().with_text(&nudge) .with_visibility(false, true); - messages_to_add.push(message); + push_message_with_id(&mut messages_to_add, message); yield AgentEvent::Message( Message::assistant().with_system_notification( SystemNotificationType::InlineMessage, @@ -2746,7 +2849,7 @@ impl Agent { ); let message = Message::user().with_text(&nudge) .with_visibility(false, true); - messages_to_add.push(message); + push_message_with_id(&mut messages_to_add, message); yield AgentEvent::Message( Message::assistant().with_system_notification( SystemNotificationType::InlineMessage, @@ -2786,8 +2889,10 @@ impl Agent { } else { warn!("Provider returned an empty response after retries; ending turn"); last_assistant_text = EMPTY_TURN_MESSAGE.to_string(); - let message = Message::assistant().with_text(EMPTY_TURN_MESSAGE); - messages_to_add.push(message.clone()); + let message = push_message_with_id( + &mut messages_to_add, + Message::assistant().with_text(EMPTY_TURN_MESSAGE), + ); yield AgentEvent::Message(message); exit_chat = true; } @@ -2796,8 +2901,8 @@ impl Agent { // Surface and persist the failure message // through the normal path so recipes don't // exit silently when retries are exhausted. + let message = push_message_with_id(&mut messages_to_add, message); last_assistant_text = message.as_concat_text(); - messages_to_add.push(message.clone()); yield AgentEvent::Message(message); exit_chat = true; } @@ -2856,9 +2961,12 @@ impl Agent { } if let Some(output) = pending_final_output.take() { + preferred_turn_usage_message_id = None; last_assistant_text = output.clone(); - let message = Message::assistant().with_text(output); - messages_to_add.push(message.clone()); + let message = push_message_with_id( + &mut messages_to_add, + Message::assistant().with_text(output), + ); yield AgentEvent::Message(message); } @@ -2873,9 +2981,11 @@ impl Agent { }; if let Some(usage) = pending_turn_usage.take() { - if let Some((message_id, usage)) = - attach_turn_usage(&mut messages_to_add, &usage) - { + if let Some((message_id, usage)) = attach_turn_usage( + &mut messages_to_add, + &usage, + preferred_turn_usage_message_id.as_deref(), + ) { yield AgentEvent::MessageUsage { message_id, usage }; } } @@ -2901,15 +3011,23 @@ impl Agent { crate::hooks::HookDecision::Deny { reason, plugin } => { consecutive_stop_hook_blocks += 1; if consecutive_stop_hook_blocks > stop_hook_block_cap { - let message = stop_hook_block_cap_warning(&plugin, stop_hook_block_cap); - session_manager.add_message(&session_config.id, &message).await?; + let message = persist_message_with_id( + &session_manager, + &session_config.id, + stop_hook_block_cap_warning(&plugin, stop_hook_block_cap), + ) + .await?; yield AgentEvent::Message(message); stop_hook_handled_for_exit = true; break; } - let message = stop_hook_denial_context_message(&plugin, &reason); - session_manager.add_message(&session_config.id, &message).await?; - conversation.push(message); + persist_and_push_message_with_id( + &session_manager, + &session_config.id, + &mut conversation, + stop_hook_denial_context_message(&plugin, &reason), + ) + .await?; yield AgentEvent::Message(stop_hook_denial_notification(&plugin)); retrying_after_stop_hook_denial = true; } @@ -3503,6 +3621,34 @@ mod tests { use std::sync::atomic::{AtomicUsize, Ordering}; use tempfile::TempDir; + #[test] + fn ensure_message_event_id_assigns_missing_ids_and_preserves_existing_ids() { + let generated = + ensure_message_event_id(AgentEvent::Message(Message::assistant().with_text("hello"))); + let AgentEvent::Message(generated_message) = generated else { + panic!("expected message event"); + }; + let generated_id = generated_message + .id + .as_deref() + .expect("generated message id"); + assert!(generated_id.starts_with("msg_")); + + let preserved = ensure_message_event_id(AgentEvent::Message( + Message::assistant() + .with_id("provider-message-id") + .with_text("hello"), + )); + let AgentEvent::Message(preserved_message) = preserved else { + panic!("expected message event"); + }; + assert_eq!(preserved_message.id.as_deref(), Some("provider-message-id")); + + let non_message = + ensure_message_event_id(AgentEvent::HistoryReplaced(Conversation::empty())); + assert!(matches!(non_message, AgentEvent::HistoryReplaced(_))); + } + #[test] fn resolve_use_login_shell_path_defaults_by_platform() { assert!(resolve_use_login_shell_path( @@ -3939,8 +4085,13 @@ echo start >> "$PLUGIN_ROOT/hook.log" .reply(Message::user().with_text("hi"), session_config, None) .await?; tokio::pin!(reply_stream); + let mut emitted_refusal_id = None; while let Some(event) = reply_stream.next().await { - event?; + if let AgentEvent::Message(message) = event? { + if message.as_concat_text().contains("provider refused") { + emitted_refusal_id = message.id; + } + } } assert_eq!( @@ -3948,6 +4099,9 @@ echo start >> "$PLUGIN_ROOT/hook.log" 1, "a refused request must not be resent" ); + let emitted_refusal_id = + emitted_refusal_id.expect("refusal message should be emitted with an ID"); + assert!(emitted_refusal_id.starts_with("msg_")); Ok(()) } @@ -4162,6 +4316,34 @@ echo start >> "$PLUGIN_ROOT/hook.log" }) })); + let stored_session = agent + .config + .session_manager + .get_session(&session_id, true) + .await?; + let stored_messages = stored_session + .conversation + .expect("session should have stored conversation"); + let stop_hook_context_messages = stored_messages + .messages() + .iter() + .filter(|message| { + message.role == rmcp::model::Role::User + && !message.is_user_visible() + && message.is_agent_visible() + && message + .as_concat_text() + .contains("Address this policy hook denial") + }) + .collect::>(); + assert_eq!(stop_hook_context_messages.len(), 2); + assert!(stop_hook_context_messages.iter().all(|message| { + message + .id + .as_deref() + .is_some_and(|id| id.starts_with("msg_")) + })); + Ok(()) } @@ -4351,7 +4533,7 @@ echo start >> "$PLUGIN_ROOT/hook.log" ]); let (message_id, attached) = - attach_turn_usage(&mut conversation, &usage).expect("usage should attach"); + attach_turn_usage(&mut conversation, &usage, None).expect("usage should attach"); assert_eq!(message_id.as_deref(), Some("a2")); assert_eq!(attached.input_tokens, Some(1200)); @@ -4376,7 +4558,7 @@ echo start >> "$PLUGIN_ROOT/hook.log" let usage = ProviderUsage::new("test-model".to_string(), Usage::default()); let mut conversation = Conversation::new_unvalidated([Message::user().with_text("hi")]); - assert!(attach_turn_usage(&mut conversation, &usage).is_none()); + assert!(attach_turn_usage(&mut conversation, &usage, None).is_none()); assert!( conversation.messages()[0].metadata.usage.is_none(), "user message must stay untouched" @@ -4400,7 +4582,7 @@ echo start >> "$PLUGIN_ROOT/hook.log" .with_content(MessageContent::Text(assistant_only)), ]); - assert!(attach_turn_usage(&mut conversation, &usage).is_none()); + assert!(attach_turn_usage(&mut conversation, &usage, None).is_none()); let stored = conversation.messages()[1] .metadata diff --git a/crates/goose/src/agents/reply_parts.rs b/crates/goose/src/agents/reply_parts.rs index 78a6307fd..f9aa91fd4 100644 --- a/crates/goose/src/agents/reply_parts.rs +++ b/crates/goose/src/agents/reply_parts.rs @@ -168,6 +168,19 @@ fn message_has_timing_content(message: &Message) -> bool { .any(|content| !matches!(content, MessageContent::SystemNotification(_))) } +fn is_mergeable_assistant_chunk(message: &Message) -> bool { + message.role == rmcp::model::Role::Assistant + && !message.content.is_empty() + && message.content.iter().all(|content| { + matches!( + content, + MessageContent::Text(_) + | MessageContent::Thinking(_) + | MessageContent::RedactedThinking(_) + ) + }) +} + impl Agent { pub async fn prepare_tools_and_prompt( &self, @@ -396,13 +409,14 @@ impl Agent { if let Some(msg) = accumulated_message { let processed = toolshim_postprocess(msg, &toolshim_tools).await?; - yield (Some(processed), final_usage); + yield (Some(processed.with_generated_id_if_missing()), final_usage); } else if final_usage.is_some() { // Preserve usage-only responses (no message content) yield (None, final_usage); } } else { let mut first_content_at: Option = None; + let mut active_mergeable_assistant_id: Option = None; while let Some(result) = stream.next().await { let (message, mut usage) = result?; @@ -415,6 +429,21 @@ impl Agent { fill_stream_timing(usage, request_started, first_content_at); } + let message = message.map(|message| { + if message.id.is_some() { + active_mergeable_assistant_id = None; + message + } else if is_mergeable_assistant_chunk(&message) { + let id = active_mergeable_assistant_id + .get_or_insert_with(|| format!("msg_{}", uuid::Uuid::new_v4())) + .clone(); + message.with_id(id) + } else { + active_mergeable_assistant_id = None; + message.with_generated_id() + } + }); + yield (message, usage); } } @@ -1038,6 +1067,241 @@ mod tests { ); } + struct MixedMessageIdStreamProvider; + + #[async_trait] + impl Provider for MixedMessageIdStreamProvider { + fn get_name(&self) -> &str { + "mixed-message-id-stream" + } + + async fn stream( + &self, + _model_config: &ModelConfig, + _system: &str, + _messages: &[Message], + _tools: &[Tool], + ) -> Result { + type StreamItem = Result<(Option, Option), ProviderError>; + Ok(Box::pin(futures::stream::iter(vec![ + Ok((Some(Message::assistant().with_text("Hel")), None)), + Ok((Some(Message::assistant().with_text("lo")), None)), + Ok(( + Some(Message::assistant().with_action_required( + "permission-a", + "shell".to_string(), + object!({}), + Some("Approve A?".to_string()), + )), + None, + )), + Ok(( + Some(Message::assistant().with_action_required( + "permission-b", + "shell".to_string(), + object!({}), + Some("Approve B?".to_string()), + )), + None, + )), + Ok(( + Some( + Message::assistant() + .with_id("provider-id") + .with_text("done"), + ), + None, + )), + Ok((Some(Message::assistant().with_text("next")), None)), + Ok((Some(Message::assistant().with_text("a")), None)), + Ok(( + Some(Message::assistant().with_tool_request( + "tool-t", + Ok(rmcp::model::CallToolRequestParams::new("test_tool")), + )), + None, + )), + Ok((Some(Message::assistant().with_text("b")), None)), + ] + as Vec))) + } + } + + #[tokio::test] + async fn normal_provider_stream_groups_only_contiguous_mergeable_chunks() -> anyhow::Result<()> + { + let provider = Arc::new(MixedMessageIdStreamProvider); + let mut stream = Agent::stream_response_from_provider( + provider, + ModelConfig::new("test-model"), + "test-session", + "system", + &[Message::user().with_text("hi")], + &[], + &[], + ) + .await?; + + let mut messages = Vec::new(); + while let Some(item) = stream.next().await { + let (message, usage) = item?; + assert!(usage.is_none()); + if let Some(message) = message { + messages.push(message); + } + } + + assert_eq!(messages.len(), 9); + + let ids = messages + .iter() + .map(|message| { + message + .id + .as_deref() + .expect("streamed provider message should have an ID") + }) + .collect::>(); + + assert_eq!(messages[0].as_concat_text(), "Hel"); + assert_eq!(messages[1].as_concat_text(), "lo"); + assert_eq!(ids[0], ids[1]); + assert!(ids[0].starts_with("msg_")); + + assert!(matches!( + messages[2].content.first(), + Some(MessageContent::ActionRequired(_)) + )); + assert!(matches!( + messages[3].content.first(), + Some(MessageContent::ActionRequired(_)) + )); + assert_ne!(ids[2], ids[3]); + assert_ne!(ids[2], ids[0]); + assert_ne!(ids[3], ids[0]); + + assert_eq!(messages[4].as_concat_text(), "done"); + assert_eq!(ids[4], "provider-id"); + + assert_eq!(messages[5].as_concat_text(), "next"); + assert_eq!(messages[6].as_concat_text(), "a"); + assert_ne!(ids[5], ids[0]); + assert_eq!(ids[5], ids[6]); + + assert!(matches!( + messages[7].content.first(), + Some(MessageContent::ToolRequest(_)) + )); + assert_ne!(ids[7], ids[5]); + + assert_eq!(messages[8].as_concat_text(), "b"); + assert_ne!(ids[8], ids[5]); + assert_ne!(ids[8], ids[7]); + + Ok(()) + } + + struct ToolshimMessageIdProvider { + messages: Vec, + } + + #[async_trait] + impl Provider for ToolshimMessageIdProvider { + fn get_name(&self) -> &str { + "toolshim-message-id" + } + + async fn stream( + &self, + _model_config: &ModelConfig, + _system: &str, + _messages: &[Message], + _tools: &[Tool], + ) -> Result { + type StreamItem = Result<(Option, Option), ProviderError>; + let items = self + .messages + .iter() + .cloned() + .map(|message| Ok((Some(message), None))) + .collect::>(); + Ok(Box::pin(futures::stream::iter(items))) + } + } + + #[tokio::test] + async fn toolshim_provider_stream_assigns_missing_message_id() -> anyhow::Result<()> { + let provider = Arc::new(ToolshimMessageIdProvider { + messages: vec![ + Message::assistant().with_text("Hel"), + Message::assistant().with_text("lo"), + ], + }); + let mut stream = Agent::stream_response_from_provider( + provider, + ModelConfig::new("test-model").with_toolshim(true), + "test-session", + "system", + &[Message::user().with_text("hi")], + &[], + &[], + ) + .await?; + + let mut messages = Vec::new(); + while let Some(item) = stream.next().await { + let (message, usage) = item?; + assert!(usage.is_none()); + if let Some(message) = message { + messages.push(message); + } + } + + assert_eq!(messages.len(), 1); + assert_eq!(messages[0].as_concat_text(), "Hello"); + let id = messages[0] + .id + .as_deref() + .expect("toolshim provider message should have an ID"); + assert!(id.starts_with("msg_")); + + Ok(()) + } + + #[tokio::test] + async fn toolshim_provider_stream_preserves_provider_message_id() -> anyhow::Result<()> { + let provider = Arc::new(ToolshimMessageIdProvider { + messages: vec![Message::assistant() + .with_id("provider-toolshim-id") + .with_text("hello")], + }); + let mut stream = Agent::stream_response_from_provider( + provider, + ModelConfig::new("test-model").with_toolshim(true), + "test-session", + "system", + &[Message::user().with_text("hi")], + &[], + &[], + ) + .await?; + + let mut messages = Vec::new(); + while let Some(item) = stream.next().await { + let (message, usage) = item?; + assert!(usage.is_none()); + if let Some(message) = message { + messages.push(message); + } + } + + assert_eq!(messages.len(), 1); + assert_eq!(messages[0].as_concat_text(), "hello"); + assert_eq!(messages[0].id.as_deref(), Some("provider-toolshim-id")); + + Ok(()) + } + #[tokio::test] async fn categorize_tool_requests_keeps_thinking_when_not_previously_streamed() { let agent = crate::agents::Agent::new(); diff --git a/crates/goose/src/providers/claude_code.rs b/crates/goose/src/providers/claude_code.rs index c8cb0678a..9ae24116b 100644 --- a/crates/goose/src/providers/claude_code.rs +++ b/crates/goose/src/providers/claude_code.rs @@ -737,7 +737,7 @@ impl Provider for ClaudeCodeProvider { let blocks = self.last_user_content_blocks(messages); let ndjson_line = build_stream_json_input(&blocks, &session_id); let model_name = model_config.model_name.clone(); - let message_id = uuid::Uuid::new_v4().to_string(); + let mut current_text_message_id = uuid::Uuid::new_v4().to_string(); let pending_confirmations = Arc::clone(&self.pending_confirmations); Ok(Box::pin(try_stream! { @@ -820,7 +820,7 @@ impl Provider for ClaudeCodeProvider { vec![MessageContent::text(text)], ); partial_message.id = - Some(message_id.clone()); + Some(current_text_message_id.clone()); yield (Some(partial_message), None); } } @@ -945,6 +945,7 @@ impl Provider for ClaudeCodeProvider { request_id.clone(), tool_name, input.clone(), None, ); yield (Some(action_msg), None); + current_text_message_id = uuid::Uuid::new_v4().to_string(); let confirmation = rx.await.unwrap_or(PermissionConfirmation { principal_type: PrincipalType::Tool, @@ -1499,6 +1500,65 @@ mod tests { assert_eq!(response_data, expected_response); } + #[tokio::test] + async fn test_text_message_id_rotates_after_action_required() { + use futures::StreamExt; + + let (provider, mut stream, _stdin_reader) = stream_with_canned_stdout(&[ + r#"{"type":"control_response","response":{"subtype":"success","request_id":"req_0"}}"#, + r#"{"type":"stream_event","event":{"type":"content_block_delta","delta":{"type":"text_delta","text":"before permission"}}}"#, + r#"{"type":"control_request","request_id":"perm_1","request":{"subtype":"can_use_tool","tool_name":"Write","input":{"path":"foo.txt","content":"hello"},"tool_use_id":"tu_1"}}"#, + r#"{"type":"stream_event","event":{"type":"content_block_delta","delta":{"type":"text_delta","text":"after permission"}}}"#, + r#"{"type":"result","result":"Done","usage":{"input_tokens":10,"output_tokens":5}}"#, + ]).await; + + let (before_msg, usage) = stream.next().await.unwrap().unwrap(); + assert!(usage.is_none()); + let before_msg = before_msg.unwrap(); + assert_eq!(before_msg.role, Role::Assistant); + assert_eq!(before_msg.as_concat_text(), "before permission"); + let before_id = before_msg + .id + .as_deref() + .expect("text before permission should have a provider ID") + .to_string(); + + let (action_msg, usage) = stream.next().await.unwrap().unwrap(); + assert!(usage.is_none()); + assert!(action_msg + .unwrap() + .content + .iter() + .any(|content| content.as_action_required().is_some())); + + let handled = provider + .handle_permission_confirmation( + "perm_1", + &PermissionConfirmation { + principal_type: PrincipalType::Tool, + permission: Permission::AllowOnce, + }, + ) + .await; + assert!(handled); + + let (after_msg, usage) = stream.next().await.unwrap().unwrap(); + assert!(usage.is_none()); + let after_msg = after_msg.unwrap(); + assert_eq!(after_msg.role, Role::Assistant); + assert_eq!(after_msg.as_concat_text(), "after permission"); + let after_id = after_msg + .id + .as_deref() + .expect("text after permission should have a provider ID"); + + assert_ne!(before_id, after_id); + + while let Some(item) = stream.next().await { + item.unwrap(); + } + } + #[tokio::test] async fn test_can_use_tool_cancel_on_drop() { use futures::StreamExt; diff --git a/crates/goose/src/session/session_manager.rs b/crates/goose/src/session/session_manager.rs index 48b567c9a..efbb65a8a 100644 --- a/crates/goose/src/session/session_manager.rs +++ b/crates/goose/src/session/session_manager.rs @@ -631,16 +631,16 @@ impl SessionManager { /// Patch `tool_meta` on a specific `ToolRequest` within a stored message. /// Used to persist LLM-generated tool titles and chain summaries so they /// survive session reload. Merge-based: existing keys not in `patch` are - /// preserved. No-op if the message or tool_call_id is not found. + /// preserved. Searches the most recently inserted messages in the session + /// and is a no-op if the tool_call_id is not found. pub async fn update_tool_request_meta( &self, session_id: &str, - message_id: &str, tool_call_id: &str, patch: serde_json::Value, ) -> Result<()> { self.storage - .update_tool_request_meta(session_id, message_id, tool_call_id, patch) + .update_tool_request_meta(session_id, tool_call_id, patch) .await } } @@ -2451,12 +2451,58 @@ impl SessionStorage { Ok(()) } + async fn update_tool_request_meta( + &self, + session_id: &str, + tool_call_id: &str, + patch: serde_json::Value, + ) -> Result<()> { + use crate::conversation::message::MessageContent; + + let pool = self.pool().await?; + let rows = sqlx::query_as::<_, (Option, String)>( + "SELECT message_id, content_json FROM messages \ + WHERE session_id = ? \ + ORDER BY id DESC \ + LIMIT 100", + ) + .bind(session_id) + .fetch_all(pool) + .await?; + + for (message_id, content_json) in rows { + let content: Vec = serde_json::from_str(&content_json)?; + let contains_tool_request = content.iter().any(|block| { + matches!( + block, + MessageContent::ToolRequest(tool_request) + if tool_request.id == tool_call_id + ) + }); + if contains_tool_request { + let Some(message_id) = message_id else { + return Ok(()); + }; + return self + .update_tool_request_meta_by_message_id( + session_id, + &message_id, + tool_call_id, + patch, + ) + .await; + } + } + + Ok(()) + } + /// Patch `tool_meta` on a specific `ToolRequest` within a stored message's /// `content_json`. Finds the row(s) with matching `message_id`, scans each /// row's content for a `ToolRequest` with the given `tool_call_id`, and /// merges `patch` into its `tool_meta`. Uses `BEGIN IMMEDIATE` so /// concurrent writers serialize correctly. - async fn update_tool_request_meta( + async fn update_tool_request_meta_by_message_id( &self, session_id: &str, message_id: &str, diff --git a/crates/goose/src/tool_call_labels.rs b/crates/goose/src/tool_call_labels.rs index 90953edfe..45d3bb20b 100644 --- a/crates/goose/src/tool_call_labels.rs +++ b/crates/goose/src/tool_call_labels.rs @@ -33,7 +33,6 @@ pub(crate) async fn generate_tool_title( agent: &Agent, session_manager: &SessionManager, session_id: &str, - message_id: Option<&str>, tool_request: &ToolRequest, ) -> Option { let provider = agent.provider().await.ok()?; @@ -54,16 +53,14 @@ pub(crate) async fn generate_tool_title( .await?; let request_id = &tool_request.id; - if let Some(message_id) = message_id { - let patch = json!({ - (TOOL_META_TITLE_KEY): &title, - }); - if let Err(error) = session_manager - .update_tool_request_meta(session_id, message_id, request_id, patch) - .await - { - warn!("tool call title: persist failed for {request_id} in {message_id}: {error}",); - } + let patch = json!({ + (TOOL_META_TITLE_KEY): &title, + }); + if let Err(error) = session_manager + .update_tool_request_meta(session_id, request_id, patch) + .await + { + warn!("tool call title: persist failed for {request_id}: {error}"); } Some(title) @@ -73,7 +70,6 @@ pub(crate) async fn generate_tool_chain_summary( agent: &Agent, session_manager: &SessionManager, session_id: &str, - message_id: &str, tool_requests: &[ToolRequest], ) -> Option { let steps = prepare_tool_chain_steps(tool_requests); @@ -105,11 +101,11 @@ pub(crate) async fn generate_tool_chain_summary( (TOOL_META_CHAIN_SUMMARY_KEY): &chain_summary, }); if let Err(error) = session_manager - .update_tool_request_meta(session_id, message_id, first_tool_call_id, patch) + .update_tool_request_meta(session_id, first_tool_call_id, patch) .await { warn!( - "tool chain summary: persist failed for chain anchored at {first_tool_call_id} in {message_id}: {error}", + "tool chain summary: persist failed for chain anchored at {first_tool_call_id}: {error}", ); } @@ -492,7 +488,6 @@ mod tests { &agent, session_manager.as_ref(), &session.id, - None, &tool_request(json!({})), ) .await; @@ -531,7 +526,7 @@ mod tests { } #[tokio::test] - async fn persists_title_for_known_message_id() { + async fn persists_title_for_recent_tool_request() { let temp_dir = TempDir::new().unwrap(); let session_manager = Arc::new(SessionManager::new(temp_dir.path().join("sessions"))); let permission_manager = @@ -569,14 +564,9 @@ mod tests { .await .unwrap(); - let title = generate_tool_title( - &agent, - session_manager.as_ref(), - &session.id, - message.id.as_deref(), - &tool_request, - ) - .await; + let title = + generate_tool_title(&agent, session_manager.as_ref(), &session.id, &tool_request) + .await; assert_eq!(title.as_deref(), Some("checking project status")); let loaded = session_manager @@ -728,7 +718,6 @@ mod tests { &agent, session_manager.as_ref(), &session.id, - "message-1", &chain_tool_requests(), ) .await; @@ -774,7 +763,6 @@ mod tests { &agent, session_manager.as_ref(), &session.id, - "message-1", &[tool_request(json!({}))], ) .await; @@ -833,7 +821,6 @@ mod tests { &agent, session_manager.as_ref(), &session.id, - message.id.as_deref().unwrap(), &tool_requests, ) .await; diff --git a/crates/goose/tests/acp_common_tests/mod.rs b/crates/goose/tests/acp_common_tests/mod.rs index 059a1f90e..1df3728e5 100644 --- a/crates/goose/tests/acp_common_tests/mod.rs +++ b/crates/goose/tests/acp_common_tests/mod.rs @@ -1163,7 +1163,30 @@ pub async fn run_prompt_basic() { .await .unwrap(); assert_eq!(output.text, "2"); - assert_notifications(&session.notifications(), &[Notification::AgentMessage]); + let updates = session.session_updates(); + let (standard_message_id, goose_message_id) = updates + .iter() + .find_map(|update| { + let SessionUpdate::AgentMessageChunk(chunk) = update else { + return None; + }; + let standard_message_id = chunk.message_id.as_ref()?.0.to_string(); + let goose_message_id = chunk + .meta + .as_ref()? + .get("goose")? + .get("messageId")? + .as_str()? + .to_string(); + Some((standard_message_id, goose_message_id)) + }) + .expect("expected live agent message chunk with standard and goose message IDs"); + assert!(!standard_message_id.is_empty()); + assert_eq!(standard_message_id, goose_message_id); + assert_notifications( + &fixtures::to_notifications(&updates), + &[Notification::AgentMessage], + ); expected_session_id.assert_matches(&session.session_id().0); } diff --git a/crates/goose/tests/agent.rs b/crates/goose/tests/agent.rs index 85bc81c26..4f35c41ea 100644 --- a/crates/goose/tests/agent.rs +++ b/crates/goose/tests/agent.rs @@ -1923,7 +1923,7 @@ mod tests { } #[tokio::test] - async fn test_reasoning_preserved_on_all_tool_calls_when_thinking_in_separate_chunk( + async fn test_multi_tool_response_preserves_reasoning_and_message_id_correlation( ) -> Result<()> { use goose_providers::formats::openai::{ format_messages_with_options, OpenAiFormatOptions, @@ -1972,8 +1972,23 @@ mod tests { ) .await?; tokio::pin!(reply_stream); + let mut live_tool_message_id = None; + let mut usage_message_ids = Vec::new(); while let Some(event) = reply_stream.next().await { - event?; + match event? { + AgentEvent::Message(message) + if message + .content + .iter() + .any(|content| matches!(content, MessageContent::ToolRequest(_))) => + { + live_tool_message_id = message.id; + } + AgentEvent::MessageUsage { message_id, .. } => { + usage_message_ids.push(message_id); + } + _ => {} + } } let reloaded = session_manager.get_session(&session_id, true).await?; @@ -1983,6 +1998,45 @@ mod tests { .messages() .to_vec(); + let live_tool_message_id = + live_tool_message_id.expect("live tool message must have a generated ID"); + let persisted_tool_message_ids: Vec<&str> = messages + .iter() + .filter(|message| { + message + .content + .iter() + .any(|content| matches!(content, MessageContent::ToolRequest(_))) + }) + .map(|message| { + message + .id + .as_deref() + .expect("persisted tool message must have an ID") + }) + .collect(); + + assert_eq!(persisted_tool_message_ids.len(), 2); + assert_ne!( + persisted_tool_message_ids[0], persisted_tool_message_ids[1], + "split tool messages must keep distinct message IDs" + ); + assert_eq!( + persisted_tool_message_ids + .iter() + .copied() + .filter(|message_id| *message_id == live_tool_message_id.as_str()) + .count(), + 1, + "exactly one persisted tool message must retain the live message ID" + ); + assert!( + usage_message_ids.iter().any(|message_id| { + message_id.as_deref() == Some(live_tool_message_id.as_str()) + }), + "tool-turn usage must reference the live message ID" + ); + let spec = format_messages_with_options( &messages, &ImageFormat::OpenAi, @@ -2494,8 +2548,21 @@ mod tests { .reply(Message::user().with_text("/goal"), session_config, None) .await?; tokio::pin!(reply_stream); + + let mut emitted_user_id = None; + let mut emitted_response_id = None; while let Some(event) = reply_stream.next().await { - let _ = event?; + if let AgentEvent::Message(message) = event? { + if message.role == rmcp::model::Role::User + && message.as_concat_text() == "/goal" + { + emitted_user_id = message.id; + } else if message.role == rmcp::model::Role::Assistant + && message.as_concat_text().contains("No goal set") + { + emitted_response_id = message.id; + } + } } assert_eq!( @@ -2504,6 +2571,42 @@ mod tests { "Querying the goal should not start an agent turn" ); + let emitted_user_id = emitted_user_id.expect("User message should be emitted with ID"); + assert!(emitted_user_id.starts_with("msg_")); + let emitted_response_id = + emitted_response_id.expect("Slash command response should be emitted with ID"); + assert!(emitted_response_id.starts_with("msg_")); + + let reloaded = session_manager.get_session(&session.id, true).await?; + let conversation = reloaded + .conversation + .expect("Session should have a conversation"); + let stored_user_message = conversation + .messages() + .iter() + .find(|message| { + message.role == rmcp::model::Role::User && message.as_concat_text() == "/goal" + }) + .expect("User message should be stored"); + + assert_eq!( + stored_user_message.id.as_deref(), + Some(emitted_user_id.as_str()) + ); + let stored_response_message = conversation + .messages() + .iter() + .find(|message| { + message.role == rmcp::model::Role::Assistant + && message.as_concat_text().contains("No goal set") + }) + .expect("Slash command response should be stored"); + + assert_eq!( + stored_response_message.id.as_deref(), + Some(emitted_response_id.as_str()) + ); + Ok(()) } } @@ -2971,7 +3074,9 @@ mod tests { mod empty_turn_tests { use super::*; use async_trait::async_trait; - use goose::agents::{AgentEvent, SessionConfig}; + use goose::agents::final_output_tool::FINAL_OUTPUT_TOOL_NAME; + use goose::agents::{AgentConfig, AgentEvent, GoosePlatform, SessionConfig}; + use goose::config::permission::PermissionManager; use goose::config::GooseMode; use goose::conversation::message::{Message, MessageContent}; use goose::conversation::Conversation; @@ -2982,9 +3087,11 @@ mod tests { use goose_providers::conversation::token_usage::{ProviderUsage, Usage}; use goose_providers::errors::ProviderError; use goose_providers::model::ModelConfig; - use rmcp::model::Tool; + use rmcp::model::{CallToolRequestParams, Tool}; + use rmcp::object; use std::path::PathBuf; use std::sync::atomic::{AtomicUsize, Ordering}; + use std::sync::Arc; fn usage() -> ProviderUsage { ProviderUsage::new( @@ -3003,6 +3110,10 @@ mod tests { struct AssistantOnlyProvider; + struct FinalOutputRequestProvider { + call_count: AtomicUsize, + } + impl goose::providers::base::ProviderDescriptor for AssistantOnlyProvider { fn metadata() -> ProviderMetadata { ProviderMetadata { @@ -3073,6 +3184,14 @@ mod tests { } } + impl FinalOutputRequestProvider { + fn new() -> Self { + Self { + call_count: AtomicUsize::new(0), + } + } + } + impl goose::providers::base::ProviderDescriptor for EmptyThenTextProvider { fn metadata() -> ProviderMetadata { ProviderMetadata { @@ -3132,6 +3251,33 @@ mod tests { } } + #[async_trait] + impl Provider for FinalOutputRequestProvider { + async fn stream( + &self, + _model_config: &ModelConfig, + _system_prompt: &str, + _messages: &[Message], + _tools: &[Tool], + ) -> Result { + let call = self.call_count.fetch_add(1, Ordering::SeqCst); + if call != 0 { + panic!("unexpected provider call after final-output tool request"); + } + + let tool_call = CallToolRequestParams::new(FINAL_OUTPUT_TOOL_NAME) + .with_arguments(object!({"result": "Final answer"})); + Ok(stream_from_single_message( + Message::assistant().with_tool_request("final-output-call", Ok(tool_call)), + usage(), + )) + } + + fn get_name(&self) -> &str { + "final-output-request-mock" + } + } + /// Runs a reply to completion and returns the messages yielded to the /// caller along with the conversation persisted to the session store. async fn run_reply( @@ -3260,6 +3406,18 @@ mod tests { !persisted.iter().any(is_empty_assistant), "empty assistant turn must not be persisted alongside the fallback: {persisted:?}" ); + + let emitted_fallback_id = last + .id + .as_deref() + .expect("empty-turn fallback should be emitted with ID"); + assert!(emitted_fallback_id.starts_with("msg_")); + + let stored_fallback = persisted + .iter() + .find(|message| message.as_concat_text().contains("empty response")) + .expect("empty-turn fallback should be stored"); + assert_eq!(stored_fallback.id.as_deref(), Some(emitted_fallback_id)); Ok(()) } @@ -3337,8 +3495,15 @@ mod tests { .reply(Message::user().with_text("Hi"), session_config, None) .await?; tokio::pin!(reply_stream); + let mut emitted_steer_id = None; while let Some(event) = reply_stream.next().await { - event?; + if let AgentEvent::Message(message) = event? { + if message.role == rmcp::model::Role::User + && message.as_concat_text().contains("keep going") + { + emitted_steer_id = message.id; + } + } } let persisted = agent @@ -3360,6 +3525,14 @@ mod tests { .any(|m| m.as_concat_text().contains("keep going")), "the queued steer should have been consumed: {persisted:?}" ); + let emitted_steer_id = + emitted_steer_id.expect("queued steer should be emitted with ID"); + assert!(emitted_steer_id.starts_with("msg_")); + let stored_steer = persisted + .iter() + .find(|message| message.as_concat_text().contains("keep going")) + .expect("queued steer should be stored"); + assert_eq!(stored_steer.id.as_deref(), Some(emitted_steer_id.as_str())); Ok(()) } @@ -3400,7 +3573,7 @@ mod tests { .await; let session_config = SessionConfig { - id: session.id, + id: session.id.clone(), schedule_id: None, max_turns: Some(3), retry_config: None, @@ -3411,6 +3584,124 @@ mod tests { .await?; tokio::pin!(reply_stream); + let mut messages = Vec::new(); + let mut emitted_nudge_ids = Vec::new(); + while let Some(event) = reply_stream.next().await { + if let AgentEvent::Message(m) = event? { + if m.role == rmcp::model::Role::User + && m.as_concat_text() + .contains(FINAL_OUTPUT_CONTINUATION_MESSAGE) + { + emitted_nudge_ids.push( + m.id.clone() + .expect("Final-output nudge should be emitted with ID"), + ); + } + messages.push(m); + } + } + + let text = concat_text(&messages); + assert!( + text.contains(FINAL_OUTPUT_CONTINUATION_MESSAGE), + "expected the final-output nudge, got: {text:?}" + ); + assert!( + !text.contains("empty response"), + "empty-turn fallback must not pre-empt the final-output nudge: {text:?}" + ); + + assert!( + !emitted_nudge_ids.is_empty(), + "expected at least one emitted final-output nudge" + ); + assert!(emitted_nudge_ids.iter().all(|id| id.starts_with("msg_"))); + + let reloaded = agent + .config + .session_manager + .get_session(&session.id, true) + .await?; + let conversation = reloaded + .conversation + .expect("Session should have a conversation"); + let stored_nudge_ids = conversation + .messages() + .iter() + .filter(|message| { + message.role == rmcp::model::Role::User + && message + .as_concat_text() + .contains(FINAL_OUTPUT_CONTINUATION_MESSAGE) + }) + .map(|message| { + message + .id + .clone() + .expect("Stored final-output nudge should have ID") + }) + .collect::>(); + + assert_eq!(stored_nudge_ids, emitted_nudge_ids); + Ok(()) + } + + #[tokio::test] + async fn test_final_output_result_id_matches_persisted_message() -> Result<()> { + use goose::recipe::Response; + use goose::session::SessionManager; + use tempfile::TempDir; + + let temp_dir = TempDir::new()?; + let session_manager = Arc::new(SessionManager::new(temp_dir.path().join("data"))); + let agent = Agent::with_config(AgentConfig::new( + session_manager.clone(), + Arc::new(PermissionManager::new(temp_dir.path().join("config"))), + None, + GooseMode::Auto, + true, + GoosePlatform::GooseCli, + )); + + let session = session_manager + .create_session( + PathBuf::default(), + "final-output-result".to_string(), + SessionType::Hidden, + GooseMode::Auto, + ) + .await?; + let session_id = session.id.clone(); + let provider = Arc::new(FinalOutputRequestProvider::new()); + agent + .update_provider( + provider.clone(), + ModelConfig::new("mock-model"), + &session.id, + ) + .await?; + agent + .add_final_output_tool(Response { + json_schema: Some(serde_json::json!({ + "type": "object", + "properties": { "result": { "type": "string" } }, + "required": ["result"] + })), + }) + .await; + + let session_config = SessionConfig { + id: session.id, + schedule_id: None, + max_turns: Some(5), + retry_config: None, + }; + + let reply_stream = agent + .reply(Message::user().with_text("Hi"), session_config, None) + .await?; + tokio::pin!(reply_stream); + let mut messages = Vec::new(); while let Some(event) = reply_stream.next().await { if let AgentEvent::Message(m) = event? { @@ -3418,15 +3709,38 @@ mod tests { } } - let text = concat_text(&messages); - assert!( - text.contains(FINAL_OUTPUT_CONTINUATION_MESSAGE), - "expected the final-output nudge, got: {text:?}" - ); - assert!( - !text.contains("empty response"), - "empty-turn fallback must not pre-empt the final-output nudge: {text:?}" + let emitted_final_output = messages + .iter() + .find(|message| { + message.role == rmcp::model::Role::Assistant + && message.as_concat_text().contains("Final answer") + }) + .expect("final-output result should be emitted"); + let emitted_final_output_id = emitted_final_output + .id + .as_deref() + .expect("final-output result should be emitted with ID"); + assert!(emitted_final_output_id.starts_with("msg_")); + + let persisted = session_manager + .get_session(&session_id, true) + .await? + .conversation + .map(|c| c.messages().to_vec()) + .unwrap_or_default(); + let stored_final_output = persisted + .iter() + .find(|message| { + message.role == rmcp::model::Role::Assistant + && message.as_concat_text().contains("Final answer") + }) + .expect("final-output result should be stored"); + + assert_eq!( + stored_final_output.id.as_deref(), + Some(emitted_final_output_id) ); + assert_eq!(provider.call_count.load(Ordering::SeqCst), 1); Ok(()) } @@ -3551,6 +3865,15 @@ mod tests { text.contains("Maximum retry attempts"), "exhausted recipe retries must surface the failure message: {text:?}" ); + let emitted_failure = messages + .iter() + .find(|message| message.as_concat_text().contains("Maximum retry attempts")) + .expect("max-retry failure message should be emitted"); + let emitted_failure_id = emitted_failure + .id + .as_deref() + .expect("max-retry failure message should be emitted with ID"); + assert!(emitted_failure_id.starts_with("msg_")); let persisted = agent .config @@ -3564,6 +3887,11 @@ mod tests { concat_text(&persisted).contains("Maximum retry attempts"), "the max-retry failure message must be persisted: {persisted:?}" ); + let stored_failure = persisted + .iter() + .find(|message| message.as_concat_text().contains("Maximum retry attempts")) + .expect("max-retry failure message should be stored"); + assert_eq!(stored_failure.id.as_deref(), Some(emitted_failure_id)); Ok(()) } }