mirror of
https://github.com/fabro-sh/fabro.git
synced 2026-08-28 05:27:41 +00:00
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 <noreply@anthropic.com>
This commit is contained in:
parent
b6fdd9df72
commit
9c152ccddf
4 changed files with 139 additions and 155 deletions
|
|
@ -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::*;
|
||||
|
||||
|
|
|
|||
|
|
@ -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<String> {
|
||||
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();
|
||||
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<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,
|
||||
}
|
||||
|
||||
|
|
@ -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<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()
|
||||
|
|
@ -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<Vec<Request>>,
|
||||
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<Response, LlmError> {
|
||||
Err(LlmError::Configuration {
|
||||
message: "BlockingAfterFirstOutputProvider does not implement complete()".into(),
|
||||
source: None,
|
||||
})
|
||||
}
|
||||
|
||||
async fn stream(&self, request: &Request) -> Result<StreamEventStream, LlmError> {
|
||||
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<dyn ProviderAdapter>) -> 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("<system-reminder>")),
|
||||
"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("<system-reminder>")),
|
||||
"the resumed request should place steering before a newly staged reminder"
|
||||
);
|
||||
drop(requests);
|
||||
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);
|
||||
|
||||
assert!(
|
||||
matches!(session.history.turns(), [.., Message::System { content: reminder, .. }, Message::Assistant { content, .. }]
|
||||
if reminder.contains("<system-reminder>") && content == "resumed"),
|
||||
"the reminder should commit with the successful assistant turn"
|
||||
);
|
||||
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<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
|
||||
|
|
@ -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("<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"
|
||||
);
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue