From b6fdd9df726390d79078054494c36121f9bb8bee Mon Sep 17 00:00:00 2001 From: Bryan Helmkamp Date: Mon, 3 Aug 2026 16:56:54 -0400 Subject: [PATCH 1/2] Keep task reminders transactional across interrupts --- lib/components/fabro-agent/src/session.rs | 161 ++++++++++++++++++++-- 1 file changed, 151 insertions(+), 10 deletions(-) diff --git a/lib/components/fabro-agent/src/session.rs b/lib/components/fabro-agent/src/session.rs index c72d33683..f6ae3541c 100644 --- a/lib/components/fabro-agent/src/session.rs +++ b/lib/components/fabro-agent/src/session.rs @@ -1532,10 +1532,13 @@ impl Session { compaction_failed = self.compact_if_needed().await; } - self.inject_task_reminder_if_needed(); + // Keep generated directives local to the round until its assistant + // response commits. An interrupted round must not leave a system + // message behind for later steering to follow. + let pending_task_reminder = self.task_reminder_if_needed(); // Build request - let built_request = self.build_request(); + let built_request = self.build_request(pending_task_reminder.as_deref()); let local_context_window = built_request.context_window.clone(); let request = built_request.request; @@ -1887,6 +1890,12 @@ impl Session { *usage_accumulator += usage.clone(); UsdMicros::accumulate(cost_accumulator, response.cost_usd.map(UsdMicros::from_usd)); + if let Some(reminder) = pending_task_reminder { + self.history.push(Message::System { + content: reminder, + timestamp: SystemTime::now(), + }); + } self.history.push(Message::Assistant { content: text.clone(), tool_calls: tool_calls.clone(), @@ -2123,12 +2132,15 @@ impl Session { } } - fn build_request(&self) -> BuiltRequest { + fn build_request(&self, pending_task_reminder: Option<&str>) -> BuiltRequest { let mut messages = Vec::new(); if !self.system_prompt.trim().is_empty() { messages.push(LlmMessage::system(self.system_prompt.clone())); } messages.extend(self.history.convert_to_messages()); + if let Some(reminder) = pending_task_reminder { + messages.push(LlmMessage::system(reminder)); + } let tools_with_source = self.effective_tools(); let tools: Vec<_> = tools_with_source @@ -2180,19 +2192,14 @@ impl Session { } } - fn inject_task_reminder_if_needed(&mut self) { + fn task_reminder_if_needed(&self) -> Option { let tools: Vec<_> = self .effective_tools() .into_iter() .map(|tool| tool.definition) .collect(); let tool_names: Vec<&str> = tools.iter().map(|tool| tool.name.as_str()).collect(); - if let Some(reminder) = task_reminder::maybe_reminder(&self.history, &tool_names) { - self.history.push(Message::System { - content: reminder, - timestamp: SystemTime::now(), - }); - } + task_reminder::maybe_reminder(&self.history, &tool_names) } } @@ -2543,6 +2550,56 @@ mod tests { } } + struct BlockingAfterFirstOutputProvider { + requests: Mutex>, + response: Response, + call_index: AtomicUsize, + } + + impl BlockingAfterFirstOutputProvider { + fn new(response: Response) -> Self { + Self { + requests: Mutex::new(Vec::new()), + response, + call_index: AtomicUsize::new(0), + } + } + } + + #[async_trait::async_trait] + impl ProviderAdapter for BlockingAfterFirstOutputProvider { + fn name(&self) -> &'static str { + "mock" + } + + async fn complete(&self, _request: &Request) -> Result { + Err(LlmError::Configuration { + message: "BlockingAfterFirstOutputProvider does not implement complete()".into(), + source: None, + }) + } + + async fn stream(&self, request: &Request) -> Result { + self.requests + .lock() + .expect("request capture lock poisoned") + .push(request.clone()); + if self.call_index.fetch_add(1, Ordering::SeqCst) == 0 { + let first_output = StreamEvent::ToolCallStart { + tool_call: ToolCall::new( + "call_1", + "TaskUpdate", + serde_json::json!({"taskId": "1", "status": "completed"}), + ), + }; + return Ok(Box::pin( + stream::iter([Ok(first_output)]).chain(stream::pending()), + )); + } + Ok(response_to_stream(self.response.clone())) + } + } + async fn make_session_with_provider(provider: Arc) -> Session { make_session_with_provider_and_manager(provider, None).await } @@ -2922,6 +2979,90 @@ mod tests { assert!(!control.is_waiting_for_steer()); } + #[tokio::test] + async fn interrupt_after_task_reminder_keeps_resumed_request_order_valid() { + let provider = Arc::new(BlockingAfterFirstOutputProvider::new(text_response( + "resumed", + ))); + let client = make_client(provider.clone()).await; + let mut registry = ToolRegistry::new(); + registry.register(make_named_noop_tool("TaskCreate")); + registry.register(make_named_noop_tool("TaskUpdate")); + let profile = Arc::new(TestProfile::with_tools(registry)); + let env = Arc::new(MockSandbox::default()); + let mut session = Session::new(client, profile, env, SessionOptions::default(), None); + for index in 0..10 { + session.history.push(Message::User { + content: format!("turn {index}"), + timestamp: SystemTime::now(), + }); + session.history.push(Message::Assistant { + content: "done".into(), + tool_calls: Vec::new(), + provider_parts: Vec::new(), + usage: Box::::default(), + response_id: format!("response_{index}"), + timestamp: SystemTime::now(), + }); + } + + let control = session.control_handle(); + let mut events = session.subscribe(); + let control_for_controller = control.clone(); + let controller = tokio::spawn(async move { + wait_for_agent_event(&mut events, |event| { + matches!(event, AgentEvent::LlmFirstOutput { + kind: LlmOutputKind::ToolCall, + }) + }) + .await; + control_for_controller.interrupt(None); + wait_for_agent_event(&mut events, |event| { + matches!(event, AgentEvent::RoundInterrupted { generation: 1 }) + }) + .await; + control_for_controller.steer("wrap up now".into(), None); + }); + + timeout(Duration::from_secs(1), session.process_input("continue")) + .await + .expect("interrupted session should resume after steering") + .unwrap(); + controller.await.unwrap(); + + let requests = provider + .requests + .lock() + .expect("request capture lock poisoned"); + let interrupted = requests + .first() + .expect("the interrupted request should be captured"); + assert!( + matches!(interrupted.messages.last(), Some(message) + if message.role == Role::System + && message.text().contains("")), + "the interrupted request should include the staged task reminder" + ); + let resumed = requests + .get(1) + .expect("steering should trigger a second provider request"); + assert!( + matches!(resumed.messages.as_slice(), [.., steering, reminder] + if steering.role == Role::User + && steering.text() == "wrap up now" + && reminder.role == Role::System + && reminder.text().contains("")), + "the resumed request should place steering before a newly staged reminder" + ); + drop(requests); + + assert!( + matches!(session.history.turns(), [.., Message::System { content: reminder, .. }, Message::Assistant { content, .. }] + if reminder.contains("") && content == "resumed"), + "the reminder should commit with the successful assistant turn" + ); + } + #[tokio::test] async fn interrupt_during_tool_settles_once_after_balancing_tool_result() { let blocking_tool = RegisteredTool { From 9c152ccddf2262e00a8150cfd444f1deba03fe53 Mon Sep 17 00:00:00 2001 From: Bryan Helmkamp Date: Tue, 4 Aug 2026 14:20:32 -0400 Subject: [PATCH 2/2] refactor(agent): simplify task reminder staging and test fixtures Stage the pending task reminder as a Message and add Message::to_llm_message so durable history and the round-staged turn share one turn-to-wire conversion. Replace the one-off BlockingAfterFirstOutputProvider with request capture and an EventsThenPending variant on ScriptedStreamProvider, add a shared make_session_with_provider_and_tools helper, and assert the reminder tests against task_reminder::TASK_REMINDER_TEXT instead of a substring. Co-Authored-By: Claude Fable 5 --- lib/components/fabro-agent/src/history.rs | 56 +----- lib/components/fabro-agent/src/session.rs | 174 ++++++++---------- .../fabro-agent/src/test_support.rs | 7 + lib/components/fabro-agent/src/types.rs | 59 +++++- 4 files changed, 140 insertions(+), 156 deletions(-) diff --git a/lib/components/fabro-agent/src/history.rs b/lib/components/fabro-agent/src/history.rs index b42cf1993..23f8ab7a7 100644 --- a/lib/components/fabro-agent/src/history.rs +++ b/lib/components/fabro-agent/src/history.rs @@ -1,6 +1,6 @@ use std::collections::HashSet; -use fabro_llm::types::{ContentPart, Message as LlmMessage, Role, TokenCounts}; +use fabro_llm::types::{Message as LlmMessage, TokenCounts}; use fabro_types::SessionMessage; use crate::types::Message; @@ -94,57 +94,7 @@ impl History { #[must_use] pub fn convert_to_messages(&self) -> Vec { - self.turns - .iter() - .map(|turn| match turn { - Message::User { content, .. } => LlmMessage::user(content), - Message::Assistant { - content, - tool_calls, - provider_parts, - .. - } => { - let mut parts: Vec = Vec::new(); - // Provider-specific opaque parts (e.g. OpenAI reasoning items, - // Anthropic thinking blocks with signatures) must precede - // function calls for correct round-tripping. - parts.extend(provider_parts.iter().cloned()); - if !content.is_empty() { - parts.push(ContentPart::text(content)); - } - for tc in tool_calls { - parts.push(ContentPart::ToolCall(tc.clone())); - } - LlmMessage { - role: Role::Assistant, - content: parts, - name: None, - tool_call_id: None, - } - } - Message::ToolResults { results, .. } => { - let content: Vec = results - .iter() - .map(|r| ContentPart::ToolResult(r.clone())) - .collect(); - // Use the first result's tool_call_id if available - let tool_call_id = results.first().map(|r| r.tool_call_id.clone()); - LlmMessage { - role: Role::Tool, - content, - name: None, - tool_call_id, - } - } - Message::System { content, .. } => LlmMessage::system(content), - Message::Steering { content, .. } => LlmMessage { - role: Role::User, - content: vec![ContentPart::text(content)], - name: None, - tool_call_id: None, - }, - }) - .collect() + self.turns.iter().map(Message::to_llm_message).collect() } } @@ -214,7 +164,7 @@ fn add_tool_result_call_ids<'a>(turns: &'a [Message], call_ids: &mut HashSet<&'a mod tests { use std::time::SystemTime; - use fabro_llm::types::{ThinkingData, TokenCounts, ToolCall, ToolResult}; + use fabro_llm::types::{ContentPart, Role, ThinkingData, TokenCounts, ToolCall, ToolResult}; use super::*; diff --git a/lib/components/fabro-agent/src/session.rs b/lib/components/fabro-agent/src/session.rs index f6ae3541c..3d0cf6ef0 100644 --- a/lib/components/fabro-agent/src/session.rs +++ b/lib/components/fabro-agent/src/session.rs @@ -1538,7 +1538,7 @@ impl Session { let pending_task_reminder = self.task_reminder_if_needed(); // Build request - let built_request = self.build_request(pending_task_reminder.as_deref()); + let built_request = self.build_request(pending_task_reminder.as_ref()); let local_context_window = built_request.context_window.clone(); let request = built_request.request; @@ -1891,10 +1891,7 @@ impl Session { UsdMicros::accumulate(cost_accumulator, response.cost_usd.map(UsdMicros::from_usd)); if let Some(reminder) = pending_task_reminder { - self.history.push(Message::System { - content: reminder, - timestamp: SystemTime::now(), - }); + self.history.push(reminder); } self.history.push(Message::Assistant { content: text.clone(), @@ -2132,14 +2129,14 @@ impl Session { } } - fn build_request(&self, pending_task_reminder: Option<&str>) -> BuiltRequest { + fn build_request(&self, pending_task_reminder: Option<&Message>) -> BuiltRequest { let mut messages = Vec::new(); if !self.system_prompt.trim().is_empty() { messages.push(LlmMessage::system(self.system_prompt.clone())); } messages.extend(self.history.convert_to_messages()); if let Some(reminder) = pending_task_reminder { - messages.push(LlmMessage::system(reminder)); + messages.push(reminder.to_llm_message()); } let tools_with_source = self.effective_tools(); @@ -2192,14 +2189,16 @@ impl Session { } } - fn task_reminder_if_needed(&self) -> Option { - let tools: Vec<_> = self - .effective_tools() - .into_iter() - .map(|tool| tool.definition) + fn task_reminder_if_needed(&self) -> Option { + let tools = self.effective_tools(); + let tool_names: Vec<&str> = tools + .iter() + .map(|tool| tool.definition.name.as_str()) .collect(); - let tool_names: Vec<&str> = tools.iter().map(|tool| tool.name.as_str()).collect(); - task_reminder::maybe_reminder(&self.history, &tool_names) + task_reminder::maybe_reminder(&self.history, &tool_names).map(|content| Message::System { + content, + timestamp: SystemTime::now(), + }) } } @@ -2396,11 +2395,14 @@ mod tests { enum ScriptedStreamCall { Response(Box), Events(Vec>), + /// Emit the events, then hang until the round is cancelled. + EventsThenPending(Vec>), Error(LlmError), } struct ScriptedStreamProvider { calls: Vec, + requests: Mutex>, call_index: AtomicUsize, } @@ -2412,6 +2414,7 @@ mod tests { ); Self { calls, + requests: Mutex::new(Vec::new()), call_index: AtomicUsize::new(0), } } @@ -2453,7 +2456,11 @@ mod tests { }) } - async fn stream(&self, _request: &Request) -> Result { + async fn stream(&self, request: &Request) -> Result { + self.requests + .lock() + .expect("request capture lock poisoned") + .push(request.clone()); let idx = self.call_index.fetch_add(1, Ordering::SeqCst); let scripted = if idx < self.calls.len() { self.calls[idx].clone() @@ -2466,6 +2473,9 @@ mod tests { Ok(Box::pin(stream::iter(Self::events_for_response(*response)))) } ScriptedStreamCall::Events(events) => Ok(Box::pin(stream::iter(events))), + ScriptedStreamCall::EventsThenPending(events) => { + Ok(Box::pin(stream::iter(events).chain(stream::pending()))) + } ScriptedStreamCall::Error(err) => Err(err), } } @@ -2550,56 +2560,6 @@ mod tests { } } - struct BlockingAfterFirstOutputProvider { - requests: Mutex>, - response: Response, - call_index: AtomicUsize, - } - - impl BlockingAfterFirstOutputProvider { - fn new(response: Response) -> Self { - Self { - requests: Mutex::new(Vec::new()), - response, - call_index: AtomicUsize::new(0), - } - } - } - - #[async_trait::async_trait] - impl ProviderAdapter for BlockingAfterFirstOutputProvider { - fn name(&self) -> &'static str { - "mock" - } - - async fn complete(&self, _request: &Request) -> Result { - Err(LlmError::Configuration { - message: "BlockingAfterFirstOutputProvider does not implement complete()".into(), - source: None, - }) - } - - async fn stream(&self, request: &Request) -> Result { - self.requests - .lock() - .expect("request capture lock poisoned") - .push(request.clone()); - if self.call_index.fetch_add(1, Ordering::SeqCst) == 0 { - let first_output = StreamEvent::ToolCallStart { - tool_call: ToolCall::new( - "call_1", - "TaskUpdate", - serde_json::json!({"taskId": "1", "status": "completed"}), - ), - }; - return Ok(Box::pin( - stream::iter([Ok(first_output)]).chain(stream::pending()), - )); - } - Ok(response_to_stream(self.response.clone())) - } - } - async fn make_session_with_provider(provider: Arc) -> Session { make_session_with_provider_and_manager(provider, None).await } @@ -2980,17 +2940,21 @@ mod tests { } #[tokio::test] - async fn interrupt_after_task_reminder_keeps_resumed_request_order_valid() { - let provider = Arc::new(BlockingAfterFirstOutputProvider::new(text_response( - "resumed", - ))); - let client = make_client(provider.clone()).await; + async fn interrupted_round_does_not_commit_task_reminder() { + let provider = Arc::new(ScriptedStreamProvider::new(vec![ + ScriptedStreamCall::EventsThenPending(vec![Ok(StreamEvent::ToolCallStart { + tool_call: ToolCall::new( + "call_1", + "TaskUpdate", + serde_json::json!({"taskId": "1", "status": "completed"}), + ), + })]), + ScriptedStreamCall::Response(Box::new(text_response("resumed"))), + ])); let mut registry = ToolRegistry::new(); registry.register(make_named_noop_tool("TaskCreate")); registry.register(make_named_noop_tool("TaskUpdate")); - let profile = Arc::new(TestProfile::with_tools(registry)); - let env = Arc::new(MockSandbox::default()); - let mut session = Session::new(client, profile, env, SessionOptions::default(), None); + let mut session = make_session_with_provider_and_tools(provider.clone(), registry).await; for index in 0..10 { session.history.push(Message::User { content: format!("turn {index}"), @@ -3037,30 +3001,42 @@ mod tests { let interrupted = requests .first() .expect("the interrupted request should be captured"); - assert!( - matches!(interrupted.messages.last(), Some(message) - if message.role == Role::System - && message.text().contains("")), - "the interrupted request should include the staged task reminder" - ); + let staged = interrupted + .messages + .last() + .expect("the interrupted request should not be empty"); + assert_eq!(staged.role, Role::System); + assert_eq!(staged.text(), task_reminder::TASK_REMINDER_TEXT); + let resumed = requests .get(1) .expect("steering should trigger a second provider request"); - assert!( - matches!(resumed.messages.as_slice(), [.., steering, reminder] - if steering.role == Role::User - && steering.text() == "wrap up now" - && reminder.role == Role::System - && reminder.text().contains("")), - "the resumed request should place steering before a newly staged reminder" - ); - drop(requests); - - assert!( - matches!(session.history.turns(), [.., Message::System { content: reminder, .. }, Message::Assistant { content, .. }] - if reminder.contains("") && content == "resumed"), - "the reminder should commit with the successful assistant turn" - ); + let [.., steering, reminder] = resumed.messages.as_slice() else { + panic!( + "the resumed request should end with steering and a restaged reminder: {:?}", + resumed.messages + ); + }; + assert_eq!(steering.role, Role::User); + assert_eq!(steering.text(), "wrap up now"); + assert_eq!(reminder.role, Role::System); + assert_eq!(reminder.text(), task_reminder::TASK_REMINDER_TEXT); + + let [ + .., + Message::System { + content: committed, .. + }, + Message::Assistant { content, .. }, + ] = session.history.turns() + else { + panic!( + "the reminder should commit with the successful assistant turn: {:?}", + session.history.turns() + ); + }; + assert_eq!(committed, task_reminder::TASK_REMINDER_TEXT); + assert_eq!(content, "resumed"); } #[tokio::test] @@ -4032,13 +4008,10 @@ mod tests { async fn request_injects_task_reminder_after_ten_unused_assistant_turns() { let provider = Arc::new(CapturingLlmProvider::new()); let provider_ref = provider.clone(); - let client = make_client(provider as Arc).await; let mut registry = ToolRegistry::new(); registry.register(make_named_noop_tool("TaskCreate")); registry.register(make_named_noop_tool("TaskUpdate")); - let profile = Arc::new(TestProfile::with_tools(registry)); - let env = Arc::new(MockSandbox::default()); - let mut session = Session::new(client, profile, env, SessionOptions::default(), None); + let mut session = make_session_with_provider_and_tools(provider, registry).await; for index in 0..10 { session @@ -4054,10 +4027,7 @@ mod tests { .expect("request should have been captured"); assert!( request.messages.iter().any(|message| { - message.role == Role::System - && message.text().contains("") - && message.text().contains("TaskCreate") - && message.text().contains("TaskUpdate") + message.role == Role::System && message.text() == task_reminder::TASK_REMINDER_TEXT }), "request should include task reminder system message" ); diff --git a/lib/components/fabro-agent/src/test_support.rs b/lib/components/fabro-agent/src/test_support.rs index 9c8cdbd69..b68c9ed8a 100644 --- a/lib/components/fabro-agent/src/test_support.rs +++ b/lib/components/fabro-agent/src/test_support.rs @@ -212,6 +212,13 @@ pub async fn make_session(responses: Vec) -> Session { pub async fn make_session_with_tools(responses: Vec, registry: ToolRegistry) -> Session { let provider = Arc::new(MockLlmProvider::new(responses)); + make_session_with_provider_and_tools(provider, registry).await +} + +pub async fn make_session_with_provider_and_tools( + provider: Arc, + registry: ToolRegistry, +) -> Session { let client = make_client(provider).await; let profile = Arc::new(TestProfile::with_tools(registry)); let env = Arc::new(MockSandbox::default()); diff --git a/lib/components/fabro-agent/src/types.rs b/lib/components/fabro-agent/src/types.rs index 4cbd85736..de3d0434a 100644 --- a/lib/components/fabro-agent/src/types.rs +++ b/lib/components/fabro-agent/src/types.rs @@ -2,7 +2,9 @@ use std::time::SystemTime; use chrono::{DateTime, Utc}; use fabro_llm::Error as LlmError; -use fabro_llm::types::{ContentPart, ThinkingData, TokenCounts, ToolCall, ToolResult}; +use fabro_llm::types::{ + ContentPart, Message as LlmMessage, Role, ThinkingData, TokenCounts, ToolCall, ToolResult, +}; use fabro_model::{CostSource, ModelRef}; use fabro_types::{ CommandTermination, ExecOutputTail, LlmOutputKind, LlmRetryPhase, ReasoningOutput, @@ -93,6 +95,61 @@ impl Message { }) } + /// Convert this turn into the wire message sent to the provider. Durable + /// history and round-staged turns must share this conversion so a staged + /// turn produces the same wire shape it will have once committed. + #[must_use] + pub fn to_llm_message(&self) -> LlmMessage { + match self { + Self::User { content, .. } => LlmMessage::user(content), + Self::Assistant { + content, + tool_calls, + provider_parts, + .. + } => { + let mut parts: Vec = Vec::new(); + // Provider-specific opaque parts (e.g. OpenAI reasoning items, + // Anthropic thinking blocks with signatures) must precede + // function calls for correct round-tripping. + parts.extend(provider_parts.iter().cloned()); + if !content.is_empty() { + parts.push(ContentPart::text(content)); + } + for tc in tool_calls { + parts.push(ContentPart::ToolCall(tc.clone())); + } + LlmMessage { + role: Role::Assistant, + content: parts, + name: None, + tool_call_id: None, + } + } + Self::ToolResults { results, .. } => { + let content: Vec = results + .iter() + .map(|r| ContentPart::ToolResult(r.clone())) + .collect(); + // Use the first result's tool_call_id if available + let tool_call_id = results.first().map(|r| r.tool_call_id.clone()); + LlmMessage { + role: Role::Tool, + content, + name: None, + tool_call_id, + } + } + Self::System { content, .. } => LlmMessage::system(content), + Self::Steering { content, .. } => LlmMessage { + role: Role::User, + content: vec![ContentPart::text(content)], + name: None, + tool_call_id: None, + }, + } + } + #[must_use] pub fn to_session_message(&self) -> SessionMessage { match self {