Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
56 changes: 3 additions & 53 deletions lib/components/fabro-agent/src/history.rs
Original file line number Diff line number Diff line change
@@ -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;
Expand Down Expand Up @@ -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()
}
}

Expand Down Expand Up @@ -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::*;

Expand Down
159 changes: 135 additions & 24 deletions lib/components/fabro-agent/src/session.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;

Expand Down Expand Up @@ -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(),
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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(),
})
}
}

Expand Down Expand Up @@ -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,
}
Comment thread
brynary marked this conversation as resolved.

Expand All @@ -2405,6 +2414,7 @@ mod tests {
);
Self {
calls,
requests: Mutex::new(Vec::new()),
call_index: AtomicUsize::new(0),
}
}
Expand Down Expand Up @@ -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()
Expand All @@ -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),
}
}
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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
Expand All @@ -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"
);
Expand Down
7 changes: 7 additions & 0 deletions lib/components/fabro-agent/src/test_support.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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());
Expand Down
Loading
Loading