feat: add stable agent event message identity (#10716)
This commit is contained in:
@@ -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(())
|
||||
}
|
||||
|
||||
|
||||
@@ -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,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]
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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
@@ -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(())
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user