Merge pull request #721 from fabro-sh/fix/interrupt-steering-task-reminder

Keep task reminders transactional across interrupts
This commit is contained in:
Bryan Helmkamp 2026-08-04 14:59:19 -04:00 committed by GitHub
commit 646d7e8a29
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 203 additions and 78 deletions

View file

@ -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<LlmMessage> {
self.turns
.iter()
.map(|turn| match turn {
Message::User { content, .. } => LlmMessage::user(content),
Message::Assistant {
content,
tool_calls,
provider_parts,
..
} => {
let mut parts: Vec<ContentPart> = 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<ContentPart> = 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::*;

View file

@ -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_ref());
let local_context_window = built_request.context_window.clone();
let request = built_request.request;
@ -1887,6 +1890,9 @@ 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(reminder);
}
self.history.push(Message::Assistant {
content: text.clone(),
tool_calls: tool_calls.clone(),
@ -2123,12 +2129,15 @@ impl Session {
}
}
fn build_request(&self) -> 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(reminder.to_llm_message());
}
let tools_with_source = self.effective_tools();
let tools: Vec<_> = tools_with_source
@ -2180,19 +2189,16 @@ impl Session {
}
}
fn inject_task_reminder_if_needed(&mut self) {
let tools: Vec<_> = self
.effective_tools()
.into_iter()
.map(|tool| tool.definition)
fn task_reminder_if_needed(&self) -> Option<Message> {
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();
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).map(|content| Message::System {
content,
timestamp: SystemTime::now(),
})
}
}
@ -2389,11 +2395,14 @@ mod tests {
enum ScriptedStreamCall {
Response(Box<Response>),
Events(Vec<Result<StreamEvent, LlmError>>),
/// Emit the events, then hang until the round is cancelled.
EventsThenPending(Vec<Result<StreamEvent, LlmError>>),
Error(LlmError),
}
struct ScriptedStreamProvider {
calls: Vec<ScriptedStreamCall>,
requests: Mutex<Vec<Request>>,
call_index: AtomicUsize,
}
@ -2405,6 +2414,7 @@ mod tests {
);
Self {
calls,
requests: Mutex::new(Vec::new()),
call_index: AtomicUsize::new(0),
}
}
@ -2446,7 +2456,11 @@ mod tests {
})
}
async fn stream(&self, _request: &Request) -> Result<StreamEventStream, LlmError> {
async fn stream(&self, request: &Request) -> Result<StreamEventStream, LlmError> {
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()
@ -2459,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),
}
}
@ -2922,6 +2939,106 @@ mod tests {
assert!(!control.is_waiting_for_steer());
}
#[tokio::test]
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 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}"),
timestamp: SystemTime::now(),
});
session.history.push(Message::Assistant {
content: "done".into(),
tool_calls: Vec::new(),
provider_parts: Vec::new(),
usage: Box::<TokenCounts>::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");
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");
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]
async fn interrupt_during_tool_settles_once_after_balancing_tool_result() {
let blocking_tool = RegisteredTool {
@ -3891,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<dyn ProviderAdapter>).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
@ -3913,10 +4027,7 @@ mod tests {
.expect("request should have been captured");
assert!(
request.messages.iter().any(|message| {
message.role == Role::System
&& message.text().contains("<system-reminder>")
&& 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"
);

View file

@ -212,6 +212,13 @@ pub async fn make_session(responses: Vec<Response>) -> Session {
pub async fn make_session_with_tools(responses: Vec<Response>, 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<dyn ProviderAdapter>,
registry: ToolRegistry,
) -> Session {
let client = make_client(provider).await;
let profile = Arc::new(TestProfile::with_tools(registry));
let env = Arc::new(MockSandbox::default());

View file

@ -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<ContentPart> = 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<ContentPart> = 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 {