feat: add stable agent event message identity (#10716)

This commit is contained in:
Lifei Zhou
2026-07-29 15:32:55 +10:00
committed by GitHub
parent 5b547350e2
commit 8b73e1a1b6
14 changed files with 1166 additions and 198 deletions
@@ -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);
@@ -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 <function=not-a-tool"},"index":0,"finish_reason":null}],"object":"chat.completion.chunk","created":123}
data: {"id":"ollama-source-id","model":"test-model","choices":[{"delta":{"content":" should remain text"},"index":0,"finish_reason":"stop"}],"object":"chat.completion.chunk","created":124}
data: [DONE]"#;
let lines = response_lines.lines().map(|s| Ok(s.to_string()));
let response_stream = tokio_stream::iter(lines);
let mut messages = std::pin::pin!(response_to_streaming_message_ollama(response_stream));
let (message, usage) = messages
.next()
.await
.expect("expected invalid XML fallback message")?;
assert!(usage.is_none());
let message = message.expect("expected invalid XML fallback message");
assert_eq!(message.role, Role::Assistant);
assert_eq!(message.content.len(), 1);
let MessageContentBlock::Text(text) = &message.content[0] else {
panic!("expected invalid XML fallback to remain text-only");
};
assert_eq!(text.text, "literal <function=not-a-tool should remain text");
let message_id = message
.id
.as_deref()
.expect("invalid XML fallback message should have an ID");
assert!(message_id.starts_with("msg_"));
assert!(messages.next().await.is_none());
Ok(())
}
+70 -1
View File
@@ -595,7 +595,8 @@ impl Provider for AcpProvider {
// tool_response so downstream consumers see the rejection.
if reject_all_tools {
let message = Message::assistant()
.with_text("Tool call was denied.");
.with_text("Tool call was denied.")
.with_generated_id();
yield (Some(message), None);
} else {
let denial = vec![RmcpContent::text("Tool call was denied.")];
@@ -1868,6 +1869,74 @@ mod tests {
assert!(!provider.handoff_context_sent.load(Ordering::Acquire));
}
#[tokio::test]
async fn chat_mode_denied_tool_messages_get_distinct_provider_ids() {
use futures::StreamExt;
let (tx, mut rx) = mpsc::channel(1);
let (provider, model) = test_provider_with_tx(Some(tx));
*provider.goose_mode.lock().unwrap() = GooseMode::Chat;
let messages = vec![Message::user().with_text("inspect src/lib.rs")];
let mut stream = provider.stream(&model, "", &messages, &[]).await.unwrap();
let response_tx = match rx.recv().await.expect("expected ACP prompt request") {
ClientRequest::Prompt { response_tx, .. } => 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")
+45 -74
View File
@@ -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<Agent>,
cx: &ConnectionTo<Client>,
) -> 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]
+21 -8
View File
@@ -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]
@@ -15,14 +15,12 @@ struct ToolChainStep {
#[derive(Debug)]
struct TrackedToolChain {
message_id: String,
steps: Vec<ToolChainStep>,
}
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<ToolRequest>,
}
@@ -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");
@@ -36,13 +36,11 @@ pub(crate) fn spawn_tool_title_enrichment(
tool_call_notifier: ToolCallNotifier,
session_manager: &Arc<SessionManager>,
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
+223 -41
View File
@@ -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<Message> {
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<Message> {
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<String>, 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<CancellationToken>,
) -> Result<BoxStream<'_, Result<AgentEvent>>> {
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<CancellationToken>,
) -> Result<BoxStream<'_, Result<AgentEvent>>> {
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<String> = None;
let mut pending_turn_usage: Option<ProviderUsage> = None;
let mut preferred_turn_usage_message_id: Option<String> = 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::<Vec<_>>();
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
+265 -1
View File
@@ -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<std::time::Instant> = None;
let mut active_mergeable_assistant_id: Option<String> = 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<MessageStream, ProviderError> {
type StreamItem = Result<(Option<Message>, Option<ProviderUsage>), 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<StreamItem>)))
}
}
#[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::<Vec<_>>();
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<Message>,
}
#[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<MessageStream, ProviderError> {
type StreamItem = Result<(Option<Message>, Option<ProviderUsage>), ProviderError>;
let items = self
.messages
.iter()
.cloned()
.map(|message| Ok((Some(message), None)))
.collect::<Vec<StreamItem>>();
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();
+62 -2
View File
@@ -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;
+50 -4
View File
@@ -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>, 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<MessageContent> = 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,
+14 -27
View File
@@ -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<String> {
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<ToolChainSummary> {
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;
+24 -1
View File
@@ -1163,7 +1163,30 @@ pub async fn run_prompt_basic<C: Connection>() {
.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);
}
+343 -15
View File
@@ -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<MessageStream, ProviderError> {
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::<Vec<_>>();
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(())
}
}