b09a98db48
统一系统消息、工具结果与异常响应的消息构造。 补充无效响应纠正、批量工具调用和待确认状态处理。 精简图片上下文提示并完善越界引用回退规则。 补齐运行循环和提示词行为测试。
1666 lines
58 KiB
Rust
1666 lines
58 KiB
Rust
use crate::agent::Agent;
|
||
use crate::agent::LlmApiAdaptor;
|
||
use crate::error::PromptError;
|
||
use crate::hook::Hook;
|
||
use crate::memory::{AgentMemory, StagedAgentMemory, VecMemory};
|
||
use crate::prompt::INVALID_JSON_RESPONSE_REMINDER;
|
||
use crate::run::PromptOutput::{Text, Tool};
|
||
use crate::tool::{ToolCall, ToolExecutionResult, ToolFailure, ToolOutcome};
|
||
use serde::Deserialize;
|
||
use serde_json::Value;
|
||
use std::future::{Future, poll_fn};
|
||
use std::pin::Pin;
|
||
use std::task::Poll;
|
||
|
||
pub type TextOutput = String;
|
||
|
||
#[derive(Debug, Clone)]
|
||
pub struct ToolCallOutput {
|
||
pub tool_call: ToolCall,
|
||
pub output: Value,
|
||
}
|
||
|
||
#[derive(Debug, Clone)]
|
||
pub struct ToolFailureOutput {
|
||
pub tool_call: ToolCall,
|
||
pub message: String,
|
||
pub output: Value,
|
||
pub failure: ToolFailure,
|
||
}
|
||
|
||
#[derive(Debug, Clone)]
|
||
pub enum PromptOutput {
|
||
Text(TextOutput),
|
||
Tool(ToolCallOutput),
|
||
ToolFailed(ToolFailureOutput),
|
||
}
|
||
|
||
#[derive(Debug, Clone)]
|
||
pub struct PromptRunError {
|
||
pub error: PromptError,
|
||
pub partial_outputs: Vec<PromptOutput>,
|
||
}
|
||
|
||
impl PromptRunError {
|
||
pub fn new(error: PromptError, partial_outputs: Vec<PromptOutput>) -> Self {
|
||
Self {
|
||
error,
|
||
partial_outputs,
|
||
}
|
||
}
|
||
|
||
pub fn has_tool_activity(&self) -> bool {
|
||
self.partial_outputs
|
||
.iter()
|
||
.any(|output| matches!(output, PromptOutput::Tool(_) | PromptOutput::ToolFailed(_)))
|
||
}
|
||
|
||
pub fn into_parts(self) -> (PromptError, Vec<PromptOutput>) {
|
||
(self.error, self.partial_outputs)
|
||
}
|
||
}
|
||
|
||
impl std::fmt::Display for PromptRunError {
|
||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||
self.error.fmt(formatter)
|
||
}
|
||
}
|
||
|
||
impl std::error::Error for PromptRunError {
|
||
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
|
||
Some(&self.error)
|
||
}
|
||
}
|
||
|
||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||
pub enum TextFlow {
|
||
Continue,
|
||
Stop,
|
||
}
|
||
|
||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||
pub enum ToolCallFlow {
|
||
Continue,
|
||
Skip,
|
||
Stop,
|
||
}
|
||
|
||
pub fn format_tool_call_message(
|
||
tool_call_id: impl std::fmt::Display,
|
||
args: &Value,
|
||
output: &Value,
|
||
) -> Result<String, PromptError> {
|
||
let arg_json = serde_json::to_string(args)
|
||
.map_err(|error| PromptError::InternalError(error.to_string()))?;
|
||
let output_json = serde_json::to_string(output)
|
||
.map_err(|error| PromptError::InternalError(error.to_string()))?;
|
||
Ok(format!(
|
||
"[tool_call:{tool_call_id}] args: {arg_json} output: {output_json}"
|
||
))
|
||
}
|
||
|
||
pub struct PromptRequest<'a, M: LlmApiAdaptor<Message> + 'a, Message: 'a> {
|
||
agent: &'a mut Agent<M, Message>,
|
||
message: Message,
|
||
system_prompt: Option<Message>,
|
||
hooks: Vec<Box<dyn Hook + 'static>>,
|
||
max_turns: usize,
|
||
deadline: Option<PromptDeadline<'a>>,
|
||
}
|
||
|
||
struct PromptDeadline<'a> {
|
||
future: Pin<Box<dyn Future<Output = ()> + Send + 'a>>,
|
||
error: PromptError,
|
||
}
|
||
|
||
struct PromptMemoryTransaction<'a, Message: Send + Sync + 'static> {
|
||
committed: &'a mut Option<Box<dyn AgentMemory<Message>>>,
|
||
staged: Option<Box<dyn StagedAgentMemory<Message>>>,
|
||
cancellation_message: Option<Message>,
|
||
completed_tool_activity: bool,
|
||
in_flight_tool_message: Option<Message>,
|
||
finalized: bool,
|
||
}
|
||
|
||
impl<'a, Message: Send + Sync + 'static> PromptMemoryTransaction<'a, Message> {
|
||
fn new(
|
||
committed: &'a mut Option<Box<dyn AgentMemory<Message>>>,
|
||
staged: Box<dyn StagedAgentMemory<Message>>,
|
||
cancellation_message: Message,
|
||
) -> Self {
|
||
Self {
|
||
committed,
|
||
staged: Some(staged),
|
||
cancellation_message: Some(cancellation_message),
|
||
completed_tool_activity: false,
|
||
in_flight_tool_message: None,
|
||
finalized: false,
|
||
}
|
||
}
|
||
|
||
fn get_memory(&self) -> &[Message] {
|
||
self.staged
|
||
.as_ref()
|
||
.expect("staged memory exists")
|
||
.get_memory()
|
||
}
|
||
|
||
fn append_message(&mut self, message: Message) {
|
||
self.staged
|
||
.as_mut()
|
||
.expect("staged memory exists")
|
||
.append_message(message);
|
||
}
|
||
|
||
fn begin_tool(&mut self, started_message: Message) {
|
||
self.in_flight_tool_message = Some(started_message);
|
||
}
|
||
|
||
fn finish_tool(&mut self) {
|
||
self.in_flight_tool_message = None;
|
||
self.completed_tool_activity = true;
|
||
}
|
||
|
||
fn mark_tool_activity(&mut self) {
|
||
self.completed_tool_activity = true;
|
||
}
|
||
|
||
fn has_tool_activity(&self) -> bool {
|
||
self.completed_tool_activity || self.in_flight_tool_message.is_some()
|
||
}
|
||
|
||
fn commit_staged(&mut self) {
|
||
let staged = self.staged.take().expect("staged memory exists");
|
||
*self.committed = Some(staged.commit());
|
||
}
|
||
|
||
fn finalize(mut self, commit: bool, terminal_message: Option<Message>) {
|
||
self.finalized = true;
|
||
if commit {
|
||
if let Some(message) = terminal_message {
|
||
self.append_message(message);
|
||
}
|
||
self.commit_staged();
|
||
}
|
||
}
|
||
}
|
||
|
||
impl<Message: Send + Sync + 'static> Drop for PromptMemoryTransaction<'_, Message> {
|
||
fn drop(&mut self) {
|
||
if self.finalized || !self.has_tool_activity() {
|
||
return;
|
||
}
|
||
if let Some(message) = self.in_flight_tool_message.take() {
|
||
self.append_message(message);
|
||
}
|
||
if let Some(message) = self.cancellation_message.take() {
|
||
self.append_message(message);
|
||
}
|
||
self.commit_staged();
|
||
}
|
||
}
|
||
|
||
impl<'a, M, Message> PromptRequest<'a, M, Message>
|
||
where
|
||
M: LlmApiAdaptor<Message> + 'a,
|
||
Message: 'a,
|
||
{
|
||
pub fn new(agent: &'a mut Agent<M, Message>, message: Message) -> Self {
|
||
let max_turns = agent.default_max_turns;
|
||
Self {
|
||
agent,
|
||
message,
|
||
system_prompt: None,
|
||
hooks: Vec::new(),
|
||
max_turns,
|
||
deadline: None,
|
||
}
|
||
}
|
||
|
||
pub fn system_prompt(mut self, msg: Message) -> Self {
|
||
self.system_prompt = Some(msg);
|
||
self
|
||
}
|
||
|
||
pub fn add_hook(mut self, hook: impl Hook + 'static) -> Self {
|
||
self.hooks.push(Box::new(hook));
|
||
self
|
||
}
|
||
|
||
pub fn max_turns(mut self, n: usize) -> Self {
|
||
self.max_turns = n;
|
||
self
|
||
}
|
||
|
||
/// 在调用方提供的 deadline future 完成时,从 runner 内部正常收口当前执行进度。
|
||
///
|
||
/// 这与从外部 drop `PromptRequest` 不同:已完成工具会进入 `partial_outputs`,memory
|
||
/// 也会按正常失败事务边界提交或回滚。deadline future 由具体 runtime 提供,因此
|
||
/// 公共 harness 不绑定 Tokio 或其它异步运行时。
|
||
pub fn deadline(
|
||
mut self,
|
||
future: impl Future<Output = ()> + Send + 'a,
|
||
error: PromptError,
|
||
) -> Self {
|
||
self.deadline = Some(PromptDeadline {
|
||
future: Box::pin(future),
|
||
error,
|
||
});
|
||
self
|
||
}
|
||
}
|
||
|
||
async fn await_with_deadline<F>(
|
||
future: F,
|
||
deadline: &mut Option<PromptDeadline<'_>>,
|
||
) -> Result<F::Output, PromptError>
|
||
where
|
||
F: Future + Send,
|
||
{
|
||
let Some(deadline) = deadline else {
|
||
return Ok(future.await);
|
||
};
|
||
let mut future = Box::pin(future);
|
||
poll_fn(|context| {
|
||
if let Poll::Ready(()) = deadline.future.as_mut().poll(context) {
|
||
return Poll::Ready(Err(deadline.error.clone()));
|
||
}
|
||
if let Poll::Ready(output) = future.as_mut().poll(context) {
|
||
return Poll::Ready(Ok(output));
|
||
}
|
||
Poll::Pending
|
||
})
|
||
.await
|
||
}
|
||
|
||
async fn ensure_deadline_not_elapsed(
|
||
deadline: &mut Option<PromptDeadline<'_>>,
|
||
) -> Result<(), PromptError> {
|
||
await_with_deadline(std::future::ready(()), deadline).await
|
||
}
|
||
|
||
impl<'a, M, Message> IntoFuture for PromptRequest<'a, M, Message>
|
||
where
|
||
M: LlmApiAdaptor<Message> + Send + Sync + 'a,
|
||
Message: Send + Sync + Clone + 'a + 'static,
|
||
{
|
||
type Output = Result<Vec<PromptOutput>, PromptRunError>;
|
||
type IntoFuture = Pin<Box<dyn Future<Output = Self::Output> + Send + 'a>>;
|
||
|
||
fn into_future(self) -> Self::IntoFuture {
|
||
let agent = self.agent;
|
||
let message = self.message;
|
||
let request_system_prompt = self.system_prompt;
|
||
let extra_hooks = self.hooks;
|
||
let max_turns = self.max_turns;
|
||
let mut deadline = self.deadline;
|
||
|
||
Box::pin(async move {
|
||
let Agent {
|
||
model,
|
||
tools,
|
||
hooks: agent_hooks,
|
||
system_prompt: agent_system_prompt,
|
||
memory: committed_memory,
|
||
..
|
||
} = agent;
|
||
let staged_memory = committed_memory
|
||
.as_ref()
|
||
.map(|memory| memory.begin_staged())
|
||
.unwrap_or_else(|| {
|
||
Box::new(VecMemory::new(Vec::new())) as Box<dyn StagedAgentMemory<Message>>
|
||
});
|
||
let cancellation_message = model.build_error_message(&PromptError::CompletionError(
|
||
"prompt future cancelled after tool activity; reconcile before retry".to_string(),
|
||
));
|
||
let mut memory =
|
||
PromptMemoryTransaction::new(committed_memory, staged_memory, cancellation_message);
|
||
// prompt(message) goes here
|
||
memory.append_message(message);
|
||
|
||
let outcome: Result<Vec<PromptOutput>, PromptRunError> = async {
|
||
let mut prompt_result: Vec<PromptOutput> = Vec::new();
|
||
|
||
for _ in 0..max_turns {
|
||
let text = {
|
||
let messages = agent_system_prompt
|
||
.iter()
|
||
.chain(request_system_prompt.iter())
|
||
.chain(memory.get_memory().iter());
|
||
await_with_deadline(model.complete(messages), &mut deadline)
|
||
.await
|
||
.map_err(|error| {
|
||
PromptRunError::new(error, prompt_result.clone())
|
||
})?
|
||
.map_err(|error| {
|
||
PromptRunError::new(error, prompt_result.clone())
|
||
})?
|
||
};
|
||
|
||
// Try to parse the LLM reply as JSON (handle Markdown fences)
|
||
let cleaned = clean_json_response(&text);
|
||
match serde_json::from_str::<LlmJsonResponse>(&cleaned) {
|
||
Ok(json_resp) => {
|
||
let clean_text = json_resp.reply_text.trim().to_string();
|
||
|
||
// Run on_text_reply hooks
|
||
for hook in agent_hooks.iter().chain(extra_hooks.iter()) {
|
||
match hook.on_text_reply(&clean_text) {
|
||
TextFlow::Stop => {
|
||
return Err(PromptRunError::new(
|
||
PromptError::ToolError(
|
||
"text reply rejected by hook".to_string(),
|
||
),
|
||
prompt_result,
|
||
));
|
||
}
|
||
TextFlow::Continue => {}
|
||
}
|
||
}
|
||
memory.append_message(model.build_assistant_message(&clean_text));
|
||
prompt_result.push(Text(clean_text.clone()));
|
||
|
||
let tool_calls: Vec<ToolCall> = json_resp
|
||
.tool_calls
|
||
.into_iter()
|
||
.enumerate()
|
||
.map(|(idx, tc)| ToolCall {
|
||
id: format!("{idx}"),
|
||
name: tc.tool_name,
|
||
args: tc.args,
|
||
})
|
||
.collect();
|
||
|
||
// no tool call, turn terminate.
|
||
if tool_calls.is_empty() {
|
||
return Ok(prompt_result);
|
||
}
|
||
|
||
let mut all_tool_calls_await_user_confirmation = true;
|
||
for (tc_id, tc) in tool_calls.iter().enumerate() {
|
||
// inline run_hooks: before_tool_call hook
|
||
let mut should_skip = false;
|
||
for hook in agent_hooks.iter().chain(extra_hooks.iter()) {
|
||
match hook.before_tool_call(tc) {
|
||
ToolCallFlow::Stop => {
|
||
return Err(PromptRunError::new(
|
||
PromptError::ToolError(
|
||
"tool call rejected by hook".to_string(),
|
||
),
|
||
prompt_result,
|
||
));
|
||
}
|
||
ToolCallFlow::Skip => {
|
||
let msg = model.tool_result_message(
|
||
&tc.name,
|
||
"(skipped by hook)",
|
||
);
|
||
memory.append_message(msg);
|
||
should_skip = true;
|
||
break;
|
||
}
|
||
ToolCallFlow::Continue => {}
|
||
}
|
||
}
|
||
if should_skip {
|
||
all_tool_calls_await_user_confirmation = false;
|
||
continue;
|
||
}
|
||
|
||
ensure_deadline_not_elapsed(&mut deadline)
|
||
.await
|
||
.map_err(|error| {
|
||
PromptRunError::new(error, prompt_result.clone())
|
||
})?;
|
||
|
||
let matching_tool =
|
||
tools.iter().find(|tool| tool.tool_name() == tc.name);
|
||
let requires_user_confirmation = matching_tool
|
||
.is_some_and(|tool| tool.requires_user_confirmation());
|
||
let result = match matching_tool {
|
||
Some(tool) => {
|
||
let arg_json = serde_json::to_string(&tc.args).map_err(
|
||
|error| {
|
||
PromptRunError::new(
|
||
PromptError::InternalError(error.to_string()),
|
||
prompt_result.clone(),
|
||
)
|
||
},
|
||
)?;
|
||
memory.begin_tool(model.tool_result_message(
|
||
&tc.name,
|
||
&format!(
|
||
"[tool_call:{tc_id}] started with args: {arg_json}; result unknown because prompt execution was cancelled"
|
||
),
|
||
));
|
||
let result = tool.call(tc.args.clone()).await;
|
||
memory.finish_tool();
|
||
result
|
||
}
|
||
None => {
|
||
memory.mark_tool_activity();
|
||
ToolExecutionResult::failed(
|
||
Value::Null,
|
||
ToolFailure::invalid_args(format!(
|
||
"unknown tool: {}",
|
||
tc.name
|
||
)),
|
||
)
|
||
}
|
||
};
|
||
if !requires_user_confirmation
|
||
|| !matches!(&result.outcome, ToolOutcome::InternalOk)
|
||
{
|
||
all_tool_calls_await_user_confirmation = false;
|
||
}
|
||
|
||
match result.outcome {
|
||
ToolOutcome::InternalOk => {
|
||
let mut json_output = result.output;
|
||
let mut hook_stop_error = None;
|
||
// Run after_tool_call hooks to allow output modification
|
||
for hook in agent_hooks.iter().chain(extra_hooks.iter()) {
|
||
match hook.after_tool_call(&tc.name, &mut json_output) {
|
||
ToolCallFlow::Stop => {
|
||
hook_stop_error = Some(PromptError::ToolError(
|
||
"tool call output caused this turn to stop by hook"
|
||
.to_string(),
|
||
));
|
||
break;
|
||
}
|
||
ToolCallFlow::Skip => {
|
||
all_tool_calls_await_user_confirmation = false;
|
||
json_output = serde_json::json!({"message":"tool call is ignored by hook"});
|
||
break;
|
||
}
|
||
ToolCallFlow::Continue => {}
|
||
}
|
||
}
|
||
let overall_message =
|
||
format_tool_call_message(tc_id, &tc.args, &json_output)
|
||
.map_err(|error| {
|
||
PromptRunError::new(
|
||
error,
|
||
prompt_result.clone(),
|
||
)
|
||
})?;
|
||
let msg = model.tool_result_message(&tc.name, &overall_message);
|
||
memory.append_message(msg);
|
||
prompt_result.push(Tool(ToolCallOutput {
|
||
tool_call: tc.clone(),
|
||
output: json_output,
|
||
}));
|
||
if let Some(error) = hook_stop_error {
|
||
return Err(PromptRunError::new(error, prompt_result));
|
||
}
|
||
}
|
||
ToolOutcome::InternalError(failure) => {
|
||
let failure_payload = serde_json::json!({
|
||
"status": "failed",
|
||
"failure": &failure,
|
||
"output": &result.output,
|
||
});
|
||
let failure_json = serde_json::to_string(&failure_payload)
|
||
.map_err(|error| {
|
||
PromptRunError::new(
|
||
PromptError::InternalError(error.to_string()),
|
||
prompt_result.clone(),
|
||
)
|
||
})?;
|
||
let overall_message = format!(
|
||
"[tool_call:{tc_id}] failure: {failure_json}"
|
||
);
|
||
let msg = model.tool_result_message(&tc.name, &overall_message);
|
||
memory.append_message(msg);
|
||
let fatal = failure.fatal;
|
||
let error_message = failure.message.clone();
|
||
prompt_result.push(PromptOutput::ToolFailed(
|
||
ToolFailureOutput {
|
||
tool_call: tc.clone(),
|
||
message: overall_message,
|
||
output: result.output,
|
||
failure,
|
||
},
|
||
));
|
||
if fatal {
|
||
return Err(PromptRunError::new(
|
||
PromptError::ToolError(error_message),
|
||
prompt_result,
|
||
));
|
||
}
|
||
}
|
||
}
|
||
|
||
ensure_deadline_not_elapsed(&mut deadline)
|
||
.await
|
||
.map_err(|error| {
|
||
PromptRunError::new(error, prompt_result.clone())
|
||
})?;
|
||
}
|
||
|
||
if all_tool_calls_await_user_confirmation {
|
||
return Ok(prompt_result);
|
||
}
|
||
}
|
||
Err(_) => {
|
||
// TODO replace the whole impl with native tool call
|
||
// append the correction inside this staged turn and retry without
|
||
// putting it into final prompt result
|
||
memory.append_message(model.build_assistant_message(&text));
|
||
memory.append_message(
|
||
model.build_system_message(INVALID_JSON_RESPONSE_REMINDER),
|
||
);
|
||
continue;
|
||
}
|
||
}
|
||
}
|
||
|
||
Err(PromptRunError::new(
|
||
PromptError::MaxTurnsReached { max_turns },
|
||
prompt_result,
|
||
))
|
||
}
|
||
.await;
|
||
|
||
let commit_staged_memory = outcome.is_ok() || memory.has_tool_activity();
|
||
let terminal_message = outcome
|
||
.as_ref()
|
||
.err()
|
||
.map(|error| model.build_error_message(&error.error));
|
||
memory.finalize(commit_staged_memory, terminal_message);
|
||
outcome
|
||
})
|
||
}
|
||
}
|
||
|
||
#[derive(Deserialize)]
|
||
struct LlmJsonResponse {
|
||
reply_text: String,
|
||
#[serde(default)]
|
||
tool_calls: Vec<LlmToolCallRequest>,
|
||
}
|
||
|
||
#[derive(Deserialize)]
|
||
struct LlmToolCallRequest {
|
||
tool_name: String,
|
||
#[serde(default)]
|
||
args: Value,
|
||
}
|
||
|
||
pub fn clean_json_response(text: &str) -> String {
|
||
let text = text.trim();
|
||
if text.starts_with("```") {
|
||
let lines: Vec<&str> = text.lines().collect();
|
||
let mut cleaned = Vec::new();
|
||
let mut in_code = false;
|
||
for line in lines {
|
||
if line.trim().starts_with("```") {
|
||
in_code = !in_code;
|
||
continue;
|
||
}
|
||
if in_code {
|
||
cleaned.push(line);
|
||
}
|
||
}
|
||
if !cleaned.is_empty() {
|
||
return cleaned.join("\n").trim().to_string();
|
||
}
|
||
}
|
||
text.to_string()
|
||
}
|
||
|
||
#[cfg(test)]
|
||
mod tests {
|
||
use super::*;
|
||
use crate::agent::LlmApiAdaptor;
|
||
use crate::hook::Hook;
|
||
use crate::tool::{Tool, ToolDyn, ToolFailureKind};
|
||
use serde_json::json;
|
||
use std::convert::Infallible;
|
||
use std::sync::Arc;
|
||
use std::sync::Mutex;
|
||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||
|
||
struct RepeatingToolCallModel {
|
||
completion_count: Arc<AtomicUsize>,
|
||
}
|
||
|
||
impl LlmApiAdaptor<String> for RepeatingToolCallModel {
|
||
async fn complete<'a>(
|
||
&self,
|
||
_messages: impl Iterator<Item = &'a String> + Send,
|
||
) -> Result<String, PromptError> {
|
||
self.completion_count.fetch_add(1, Ordering::SeqCst);
|
||
Ok(json!({
|
||
"reply_text": "请确认这次生成",
|
||
"tool_calls": [{
|
||
"tool_name": "test-tool",
|
||
"args": { "prompt": "生成一张图" }
|
||
}]
|
||
})
|
||
.to_string())
|
||
}
|
||
|
||
fn build_system_message(&self, text: &str) -> String {
|
||
format!("system: {text}")
|
||
}
|
||
|
||
fn build_assistant_message(&self, text: &str) -> String {
|
||
text.to_string()
|
||
}
|
||
}
|
||
|
||
struct TestTool {
|
||
requires_user_confirmation: bool,
|
||
}
|
||
|
||
struct CapturingModel {
|
||
messages: Arc<Mutex<Vec<String>>>,
|
||
}
|
||
|
||
struct InvalidJsonThenValidModel {
|
||
completion_count: Arc<AtomicUsize>,
|
||
messages_by_attempt: Arc<Mutex<Vec<Vec<String>>>>,
|
||
}
|
||
|
||
struct OrderedBatchModel;
|
||
|
||
struct FailingCompletionModel;
|
||
|
||
struct PendingCompletionModel;
|
||
|
||
struct ToolThenPendingModel {
|
||
completion_count: Arc<AtomicUsize>,
|
||
}
|
||
|
||
struct SlowToolCallModel;
|
||
|
||
struct FailingToolCallModel {
|
||
include_successful_tool: bool,
|
||
}
|
||
|
||
fn is_system_tool_message(message: &str, tool_name: &str, output: &str) -> bool {
|
||
message.starts_with("system: ")
|
||
&& message.contains(&format!("Tool '{tool_name}' returned:"))
|
||
&& message.contains(output)
|
||
}
|
||
|
||
impl LlmApiAdaptor<String> for OrderedBatchModel {
|
||
async fn complete<'a>(
|
||
&self,
|
||
_messages: impl Iterator<Item = &'a String> + Send,
|
||
) -> Result<String, PromptError> {
|
||
Ok(json!({
|
||
"reply_text": "请确认这批操作",
|
||
"tool_calls": [
|
||
{ "tool_name": "ordered-tool", "args": { "order": 2 } },
|
||
{ "tool_name": "ordered-tool", "args": { "order": 1 } }
|
||
]
|
||
})
|
||
.to_string())
|
||
}
|
||
|
||
fn build_system_message(&self, text: &str) -> String {
|
||
format!("system: {text}")
|
||
}
|
||
|
||
fn build_assistant_message(&self, text: &str) -> String {
|
||
text.to_string()
|
||
}
|
||
}
|
||
|
||
impl LlmApiAdaptor<String> for FailingCompletionModel {
|
||
async fn complete<'a>(
|
||
&self,
|
||
_messages: impl Iterator<Item = &'a String> + Send,
|
||
) -> Result<String, PromptError> {
|
||
Err(PromptError::CompletionError(
|
||
"provider unavailable".to_string(),
|
||
))
|
||
}
|
||
|
||
fn build_system_message(&self, text: &str) -> String {
|
||
format!("system: {text}")
|
||
}
|
||
|
||
fn build_assistant_message(&self, text: &str) -> String {
|
||
text.to_string()
|
||
}
|
||
}
|
||
|
||
impl LlmApiAdaptor<String> for PendingCompletionModel {
|
||
async fn complete<'a>(
|
||
&self,
|
||
_messages: impl Iterator<Item = &'a String> + Send,
|
||
) -> Result<String, PromptError> {
|
||
std::future::pending().await
|
||
}
|
||
|
||
fn build_system_message(&self, text: &str) -> String {
|
||
format!("system: {text}")
|
||
}
|
||
|
||
fn build_assistant_message(&self, text: &str) -> String {
|
||
text.to_string()
|
||
}
|
||
}
|
||
|
||
impl LlmApiAdaptor<String> for ToolThenPendingModel {
|
||
async fn complete<'a>(
|
||
&self,
|
||
_messages: impl Iterator<Item = &'a String> + Send,
|
||
) -> Result<String, PromptError> {
|
||
if self.completion_count.fetch_add(1, Ordering::SeqCst) == 0 {
|
||
return Ok(json!({
|
||
"reply_text": "先执行一个工具",
|
||
"tool_calls": [{ "tool_name": "test-tool", "args": {} }]
|
||
})
|
||
.to_string());
|
||
}
|
||
std::future::pending().await
|
||
}
|
||
|
||
fn build_system_message(&self, text: &str) -> String {
|
||
format!("system: {text}")
|
||
}
|
||
|
||
fn build_assistant_message(&self, text: &str) -> String {
|
||
text.to_string()
|
||
}
|
||
}
|
||
|
||
impl LlmApiAdaptor<String> for SlowToolCallModel {
|
||
async fn complete<'a>(
|
||
&self,
|
||
_messages: impl Iterator<Item = &'a String> + Send,
|
||
) -> Result<String, PromptError> {
|
||
Ok(json!({
|
||
"reply_text": "执行慢工具",
|
||
"tool_calls": [{ "tool_name": "slow-effect-tool", "args": {} }]
|
||
})
|
||
.to_string())
|
||
}
|
||
|
||
fn build_system_message(&self, text: &str) -> String {
|
||
format!("system: {text}")
|
||
}
|
||
|
||
fn build_assistant_message(&self, text: &str) -> String {
|
||
text.to_string()
|
||
}
|
||
}
|
||
|
||
impl LlmApiAdaptor<String> for FailingToolCallModel {
|
||
async fn complete<'a>(
|
||
&self,
|
||
_messages: impl Iterator<Item = &'a String> + Send,
|
||
) -> Result<String, PromptError> {
|
||
let mut tool_calls = Vec::new();
|
||
if self.include_successful_tool {
|
||
tool_calls.push(json!({ "tool_name": "test-tool", "args": {} }));
|
||
}
|
||
tool_calls.push(json!({ "tool_name": "failing-tool", "args": {} }));
|
||
Ok(json!({
|
||
"reply_text": "执行工具",
|
||
"tool_calls": tool_calls,
|
||
})
|
||
.to_string())
|
||
}
|
||
|
||
fn build_system_message(&self, text: &str) -> String {
|
||
format!("system: {text}")
|
||
}
|
||
|
||
fn build_assistant_message(&self, text: &str) -> String {
|
||
text.to_string()
|
||
}
|
||
}
|
||
|
||
struct OrderedTool {
|
||
execution_order: Arc<Mutex<Vec<u64>>>,
|
||
}
|
||
|
||
struct FailingToolDyn {
|
||
fatal: bool,
|
||
}
|
||
|
||
struct SlowEffectTool {
|
||
started: Arc<AtomicUsize>,
|
||
duration: std::time::Duration,
|
||
}
|
||
|
||
#[derive(Clone)]
|
||
struct TailMemory {
|
||
messages: Vec<String>,
|
||
max_messages: usize,
|
||
}
|
||
|
||
impl crate::memory::AgentMemoryBuffer<String> for TailMemory {
|
||
fn get_memory(&self) -> &[String] {
|
||
&self.messages
|
||
}
|
||
|
||
fn append_message(&mut self, message: String) {
|
||
self.messages.push(message);
|
||
let overflow = self.messages.len().saturating_sub(self.max_messages);
|
||
if overflow > 0 {
|
||
self.messages.drain(..overflow);
|
||
}
|
||
}
|
||
}
|
||
|
||
impl AgentMemory<String> for TailMemory {
|
||
fn begin_staged(&self) -> Box<dyn StagedAgentMemory<String>> {
|
||
Box::new(self.clone())
|
||
}
|
||
}
|
||
|
||
impl StagedAgentMemory<String> for TailMemory {
|
||
fn commit(self: Box<Self>) -> Box<dyn AgentMemory<String>> {
|
||
self
|
||
}
|
||
}
|
||
|
||
#[derive(Clone)]
|
||
struct CommitTrackingMemory {
|
||
messages: Vec<String>,
|
||
commits: Arc<AtomicUsize>,
|
||
}
|
||
|
||
struct CommitTrackingStagedMemory {
|
||
messages: Vec<String>,
|
||
commits: Arc<AtomicUsize>,
|
||
}
|
||
|
||
impl crate::memory::AgentMemoryBuffer<String> for CommitTrackingMemory {
|
||
fn get_memory(&self) -> &[String] {
|
||
&self.messages
|
||
}
|
||
|
||
fn append_message(&mut self, message: String) {
|
||
self.messages.push(message);
|
||
}
|
||
}
|
||
|
||
impl AgentMemory<String> for CommitTrackingMemory {
|
||
fn begin_staged(&self) -> Box<dyn StagedAgentMemory<String>> {
|
||
Box::new(CommitTrackingStagedMemory {
|
||
messages: self.messages.clone(),
|
||
commits: self.commits.clone(),
|
||
})
|
||
}
|
||
}
|
||
|
||
impl crate::memory::AgentMemoryBuffer<String> for CommitTrackingStagedMemory {
|
||
fn get_memory(&self) -> &[String] {
|
||
&self.messages
|
||
}
|
||
|
||
fn append_message(&mut self, message: String) {
|
||
self.messages.push(message);
|
||
}
|
||
}
|
||
|
||
impl StagedAgentMemory<String> for CommitTrackingStagedMemory {
|
||
fn commit(self: Box<Self>) -> Box<dyn AgentMemory<String>> {
|
||
self.commits.fetch_add(1, Ordering::SeqCst);
|
||
Box::new(CommitTrackingMemory {
|
||
messages: self.messages,
|
||
commits: self.commits,
|
||
})
|
||
}
|
||
}
|
||
|
||
impl LlmApiAdaptor<String> for CapturingModel {
|
||
async fn complete<'a>(
|
||
&self,
|
||
messages: impl Iterator<Item = &'a String> + Send,
|
||
) -> Result<String, PromptError> {
|
||
*self.messages.lock().expect("messages lock should succeed") =
|
||
messages.cloned().collect();
|
||
Ok(json!({
|
||
"reply_text": "完成",
|
||
"tool_calls": []
|
||
})
|
||
.to_string())
|
||
}
|
||
|
||
fn build_system_message(&self, text: &str) -> String {
|
||
format!("system: {text}")
|
||
}
|
||
|
||
fn build_assistant_message(&self, text: &str) -> String {
|
||
text.to_string()
|
||
}
|
||
}
|
||
|
||
impl LlmApiAdaptor<String> for InvalidJsonThenValidModel {
|
||
async fn complete<'a>(
|
||
&self,
|
||
messages: impl Iterator<Item = &'a String> + Send,
|
||
) -> Result<String, PromptError> {
|
||
self.messages_by_attempt
|
||
.lock()
|
||
.expect("messages lock should succeed")
|
||
.push(messages.cloned().collect());
|
||
if self.completion_count.fetch_add(1, Ordering::SeqCst) == 0 {
|
||
return Ok("this is not json".to_string());
|
||
}
|
||
Ok(json!({
|
||
"reply_text": "已按 JSON 格式重试",
|
||
"tool_calls": []
|
||
})
|
||
.to_string())
|
||
}
|
||
|
||
fn build_system_message(&self, text: &str) -> String {
|
||
format!("system: {text}")
|
||
}
|
||
|
||
fn build_assistant_message(&self, text: &str) -> String {
|
||
format!("assistant: {text}")
|
||
}
|
||
}
|
||
|
||
struct SkipAfterToolCallHook;
|
||
|
||
struct StopAfterToolCallHook;
|
||
|
||
impl Hook for SkipAfterToolCallHook {
|
||
fn after_tool_call(&self, _tool_name: &str, _output: &mut Value) -> ToolCallFlow {
|
||
ToolCallFlow::Skip
|
||
}
|
||
}
|
||
|
||
impl Hook for StopAfterToolCallHook {
|
||
fn after_tool_call(&self, _tool_name: &str, _output: &mut Value) -> ToolCallFlow {
|
||
ToolCallFlow::Stop
|
||
}
|
||
}
|
||
|
||
impl Tool for TestTool {
|
||
const NAME: &'static str = "test-tool";
|
||
type Error = Infallible;
|
||
type Args = Value;
|
||
type Output = Value;
|
||
|
||
fn description(&self) -> String {
|
||
"test tool".to_string()
|
||
}
|
||
|
||
fn parameters(&self) -> Value {
|
||
json!({ "type": "object" })
|
||
}
|
||
|
||
fn call(
|
||
&self,
|
||
_args: Self::Args,
|
||
) -> impl Future<Output = Result<Self::Output, Self::Error>> + Send {
|
||
async { Ok(json!({ "message": "pending user confirmation" })) }
|
||
}
|
||
|
||
fn requires_user_confirmation(&self) -> bool {
|
||
self.requires_user_confirmation
|
||
}
|
||
}
|
||
|
||
impl Tool for SlowEffectTool {
|
||
const NAME: &'static str = "slow-effect-tool";
|
||
type Error = Infallible;
|
||
type Args = Value;
|
||
type Output = Value;
|
||
|
||
fn description(&self) -> String {
|
||
"slow effect tool".to_string()
|
||
}
|
||
|
||
fn parameters(&self) -> Value {
|
||
json!({ "type": "object" })
|
||
}
|
||
|
||
fn call(
|
||
&self,
|
||
_args: Self::Args,
|
||
) -> impl Future<Output = Result<Self::Output, Self::Error>> + Send {
|
||
async move {
|
||
self.started.fetch_add(1, Ordering::SeqCst);
|
||
tokio::time::sleep(self.duration).await;
|
||
Ok(json!({ "message": "effect completed" }))
|
||
}
|
||
}
|
||
}
|
||
|
||
impl Tool for OrderedTool {
|
||
const NAME: &'static str = "ordered-tool";
|
||
type Error = Infallible;
|
||
type Args = Value;
|
||
type Output = Value;
|
||
|
||
fn description(&self) -> String {
|
||
"ordered test tool".to_string()
|
||
}
|
||
|
||
fn parameters(&self) -> Value {
|
||
json!({ "type": "object" })
|
||
}
|
||
|
||
fn call(
|
||
&self,
|
||
args: Self::Args,
|
||
) -> impl Future<Output = Result<Self::Output, Self::Error>> + Send {
|
||
async move {
|
||
let order = args["order"]
|
||
.as_u64()
|
||
.expect("ordered tool should receive an order");
|
||
self.execution_order
|
||
.lock()
|
||
.expect("execution order lock should succeed")
|
||
.push(order);
|
||
Ok(json!({ "message": "pending user confirmation" }))
|
||
}
|
||
}
|
||
|
||
fn requires_user_confirmation(&self) -> bool {
|
||
true
|
||
}
|
||
}
|
||
|
||
impl ToolDyn for FailingToolDyn {
|
||
fn tool_name(&self) -> &'static str {
|
||
"failing-tool"
|
||
}
|
||
|
||
fn description(&self) -> String {
|
||
"failing test tool".to_string()
|
||
}
|
||
|
||
fn parameters(&self) -> Value {
|
||
json!({ "type": "object" })
|
||
}
|
||
|
||
fn requires_user_confirmation(&self) -> bool {
|
||
false
|
||
}
|
||
|
||
fn call(
|
||
&self,
|
||
_args: Value,
|
||
) -> Pin<Box<dyn Future<Output = ToolExecutionResult> + Send + '_>> {
|
||
Box::pin(async move {
|
||
ToolExecutionResult::failed(
|
||
json!({ "attempt": 1 }),
|
||
ToolFailure::new(ToolFailureKind::Network, "network failed")
|
||
.with_fatal(self.fatal),
|
||
)
|
||
})
|
||
}
|
||
}
|
||
|
||
#[tokio::test]
|
||
async fn request_system_prompt_is_added_after_the_agent_system_prompt() {
|
||
let captured_messages = Arc::new(Mutex::new(Vec::new()));
|
||
let model = CapturingModel {
|
||
messages: captured_messages.clone(),
|
||
};
|
||
let mut agent = Agent::new(model).system_prompt("Agent system".to_string());
|
||
|
||
let outputs = agent
|
||
.prompt("User message".to_string())
|
||
.system_prompt("Request system".to_string())
|
||
.await
|
||
.expect("request prompt should succeed");
|
||
|
||
assert_eq!(outputs.len(), 1);
|
||
assert_eq!(
|
||
*captured_messages
|
||
.lock()
|
||
.expect("messages lock should succeed"),
|
||
vec![
|
||
"Agent system".to_string(),
|
||
"Request system".to_string(),
|
||
"User message".to_string(),
|
||
]
|
||
);
|
||
}
|
||
|
||
#[tokio::test]
|
||
async fn completion_failure_rolls_back_the_staged_user_message() {
|
||
let mut agent = Agent::new(FailingCompletionModel)
|
||
.memory(VecMemory::new(vec!["prior message".to_string()]));
|
||
|
||
let error = agent
|
||
.prompt("new user message".to_string())
|
||
.await
|
||
.expect_err("completion should fail");
|
||
|
||
assert!(matches!(error.error, PromptError::CompletionError(_)));
|
||
assert!(error.partial_outputs.is_empty());
|
||
assert_eq!(
|
||
agent
|
||
.memory
|
||
.as_ref()
|
||
.expect("existing memory should be restored")
|
||
.get_memory(),
|
||
&["prior message".to_string()]
|
||
);
|
||
}
|
||
|
||
#[tokio::test]
|
||
async fn invalid_json_correction_is_visible_to_the_next_completion() {
|
||
let completion_count = Arc::new(AtomicUsize::new(0));
|
||
let messages_by_attempt = Arc::new(Mutex::new(Vec::new()));
|
||
let mut agent = Agent::new(InvalidJsonThenValidModel {
|
||
completion_count: completion_count.clone(),
|
||
messages_by_attempt: messages_by_attempt.clone(),
|
||
})
|
||
.max_turns(2);
|
||
|
||
let outputs = agent
|
||
.prompt("生成图片".to_string())
|
||
.await
|
||
.expect("the corrected completion should succeed");
|
||
|
||
assert_eq!(completion_count.load(Ordering::SeqCst), 2);
|
||
assert!(matches!(
|
||
outputs.as_slice(),
|
||
[PromptOutput::Text(text)] if text == "已按 JSON 格式重试"
|
||
));
|
||
|
||
let attempts = messages_by_attempt
|
||
.lock()
|
||
.expect("captured attempts lock should succeed");
|
||
assert_eq!(attempts.len(), 2);
|
||
assert!(
|
||
!attempts[0]
|
||
.iter()
|
||
.any(|message| message.contains("this is not json"))
|
||
);
|
||
assert_eq!(
|
||
attempts[1]
|
||
.iter()
|
||
.rev()
|
||
.take(2)
|
||
.cloned()
|
||
.collect::<Vec<_>>(),
|
||
vec![
|
||
format!("system: {INVALID_JSON_RESPONSE_REMINDER}"),
|
||
"assistant: this is not json".to_string(),
|
||
]
|
||
);
|
||
}
|
||
|
||
#[tokio::test]
|
||
async fn dropping_a_pending_prompt_keeps_the_original_memory() {
|
||
let mut agent = Agent::new(PendingCompletionModel)
|
||
.memory(VecMemory::new(vec!["prior message".to_string()]));
|
||
|
||
let timeout = tokio::time::timeout(
|
||
std::time::Duration::from_millis(1),
|
||
agent.prompt("new user message".to_string()),
|
||
)
|
||
.await;
|
||
|
||
assert!(timeout.is_err());
|
||
assert_eq!(
|
||
agent
|
||
.memory
|
||
.as_ref()
|
||
.expect("external cancellation must keep committed memory")
|
||
.get_memory(),
|
||
&["prior message".to_string()]
|
||
);
|
||
}
|
||
|
||
#[tokio::test]
|
||
async fn dropping_after_a_completed_tool_commits_the_fact_and_cancellation_closure() {
|
||
let completion_count = Arc::new(AtomicUsize::new(0));
|
||
let commits = Arc::new(AtomicUsize::new(0));
|
||
let mut agent = Agent::new(ToolThenPendingModel {
|
||
completion_count: completion_count.clone(),
|
||
})
|
||
.tool(TestTool {
|
||
requires_user_confirmation: false,
|
||
})
|
||
.memory(CommitTrackingMemory {
|
||
messages: vec!["prior message".to_string()],
|
||
commits: commits.clone(),
|
||
});
|
||
|
||
let timeout = tokio::time::timeout(
|
||
std::time::Duration::from_millis(10),
|
||
agent.prompt("execute then wait".to_string()),
|
||
)
|
||
.await;
|
||
|
||
assert!(timeout.is_err());
|
||
assert_eq!(completion_count.load(Ordering::SeqCst), 2);
|
||
assert_eq!(commits.load(Ordering::SeqCst), 1);
|
||
let memory = agent
|
||
.memory
|
||
.as_ref()
|
||
.expect("completed tool cancellation should commit memory")
|
||
.get_memory();
|
||
assert!(
|
||
memory
|
||
.iter()
|
||
.any(|message| is_system_tool_message(message, "test-tool", ""))
|
||
);
|
||
assert!(memory.last().is_some_and(|message| {
|
||
message.contains("prompt future cancelled after tool activity")
|
||
}));
|
||
}
|
||
|
||
#[tokio::test]
|
||
async fn dropping_an_in_flight_tool_commits_an_unknown_result_fact() {
|
||
let started = Arc::new(AtomicUsize::new(0));
|
||
let mut agent = Agent::new(SlowToolCallModel)
|
||
.tool(SlowEffectTool {
|
||
started: started.clone(),
|
||
duration: std::time::Duration::from_secs(60),
|
||
})
|
||
.memory(VecMemory::new(vec!["prior message".to_string()]));
|
||
|
||
let timeout = tokio::time::timeout(
|
||
std::time::Duration::from_millis(10),
|
||
agent.prompt("start slow effect".to_string()),
|
||
)
|
||
.await;
|
||
|
||
assert!(timeout.is_err());
|
||
assert_eq!(started.load(Ordering::SeqCst), 1);
|
||
let memory = agent
|
||
.memory
|
||
.as_ref()
|
||
.expect("in-flight cancellation should commit memory")
|
||
.get_memory();
|
||
assert!(memory.iter().any(|message| {
|
||
message.contains("slow-effect-tool") && message.contains("result unknown")
|
||
}));
|
||
assert!(memory.last().is_some_and(|message| {
|
||
message.contains("prompt future cancelled after tool activity")
|
||
}));
|
||
}
|
||
|
||
#[tokio::test]
|
||
async fn internal_deadline_returns_completed_tools_and_closes_memory() {
|
||
let completion_count = Arc::new(AtomicUsize::new(0));
|
||
let mut agent = Agent::new(ToolThenPendingModel {
|
||
completion_count: completion_count.clone(),
|
||
})
|
||
.tool(TestTool {
|
||
requires_user_confirmation: false,
|
||
});
|
||
|
||
let error = agent
|
||
.prompt("执行后等待".to_string())
|
||
.deadline(
|
||
tokio::time::sleep(std::time::Duration::from_millis(1)),
|
||
PromptError::CompletionError("total deadline reached".to_string()),
|
||
)
|
||
.await
|
||
.expect_err("runner deadline should terminate the pending completion");
|
||
|
||
assert_eq!(completion_count.load(Ordering::SeqCst), 2);
|
||
assert!(matches!(error.error, PromptError::CompletionError(_)));
|
||
assert_eq!(error.partial_outputs.len(), 2);
|
||
assert!(matches!(error.partial_outputs[1], PromptOutput::Tool(_)));
|
||
let memory = agent
|
||
.memory
|
||
.as_ref()
|
||
.expect("completed tool activity should commit staged memory")
|
||
.get_memory();
|
||
assert!(memory.last().is_some_and(|message| {
|
||
is_system_tool_message(message, "agent-error", "total deadline reached")
|
||
}));
|
||
}
|
||
|
||
#[tokio::test]
|
||
async fn internal_deadline_does_not_cancel_an_in_flight_tool() {
|
||
let started = Arc::new(AtomicUsize::new(0));
|
||
let mut agent = Agent::new(SlowToolCallModel).tool(SlowEffectTool {
|
||
started: started.clone(),
|
||
duration: std::time::Duration::from_millis(30),
|
||
});
|
||
|
||
let error = agent
|
||
.prompt("run effect safely".to_string())
|
||
.deadline(
|
||
tokio::time::sleep(std::time::Duration::from_millis(10)),
|
||
PromptError::CompletionError("total deadline reached".to_string()),
|
||
)
|
||
.await
|
||
.expect_err("deadline should close after the started tool returns");
|
||
|
||
assert_eq!(started.load(Ordering::SeqCst), 1);
|
||
assert!(matches!(error.error, PromptError::CompletionError(_)));
|
||
assert!(matches!(error.partial_outputs[1], PromptOutput::Tool(_)));
|
||
let memory = agent
|
||
.memory
|
||
.as_ref()
|
||
.expect("completed tool should commit before deadline closure")
|
||
.get_memory();
|
||
assert!(memory.last().is_some_and(|message| {
|
||
is_system_tool_message(message, "agent-error", "total deadline reached")
|
||
}));
|
||
}
|
||
|
||
#[tokio::test]
|
||
async fn elapsed_deadline_wins_before_polling_the_next_operation() {
|
||
let captured_messages = Arc::new(Mutex::new(Vec::new()));
|
||
let model = CapturingModel {
|
||
messages: captured_messages.clone(),
|
||
};
|
||
let mut agent = Agent::new(model).memory(VecMemory::new(vec!["prior message".to_string()]));
|
||
|
||
let error = agent
|
||
.prompt("new user message".to_string())
|
||
.deadline(
|
||
std::future::ready(()),
|
||
PromptError::CompletionError("deadline already elapsed".to_string()),
|
||
)
|
||
.await
|
||
.expect_err("elapsed deadline should win before completion is polled");
|
||
|
||
assert!(matches!(error.error, PromptError::CompletionError(_)));
|
||
assert!(error.partial_outputs.is_empty());
|
||
assert!(
|
||
captured_messages
|
||
.lock()
|
||
.expect("messages lock should succeed")
|
||
.is_empty()
|
||
);
|
||
assert_eq!(
|
||
agent
|
||
.memory
|
||
.as_ref()
|
||
.expect("no-tool deadline should retain committed memory")
|
||
.get_memory(),
|
||
&["prior message".to_string()]
|
||
);
|
||
}
|
||
|
||
#[tokio::test]
|
||
async fn staged_prompt_preserves_custom_memory_append_semantics() {
|
||
let captured_messages = Arc::new(Mutex::new(Vec::new()));
|
||
let model = CapturingModel {
|
||
messages: captured_messages.clone(),
|
||
};
|
||
let mut agent = Agent::new(model).memory(TailMemory {
|
||
messages: vec!["older".to_string(), "latest".to_string()],
|
||
max_messages: 2,
|
||
});
|
||
|
||
agent
|
||
.prompt("current user".to_string())
|
||
.await
|
||
.expect("bounded staged memory should complete");
|
||
|
||
assert_eq!(
|
||
*captured_messages
|
||
.lock()
|
||
.expect("messages lock should succeed"),
|
||
vec!["latest".to_string(), "current user".to_string()]
|
||
);
|
||
assert_eq!(
|
||
agent
|
||
.memory
|
||
.as_ref()
|
||
.expect("successful staged memory should commit")
|
||
.get_memory(),
|
||
&["current user".to_string(), "完成".to_string()]
|
||
);
|
||
}
|
||
|
||
#[tokio::test]
|
||
async fn staged_memory_uses_explicit_commit_and_drop_as_rollback() {
|
||
let commits = Arc::new(AtomicUsize::new(0));
|
||
let mut failing_agent = Agent::new(FailingCompletionModel).memory(CommitTrackingMemory {
|
||
messages: vec!["prior".to_string()],
|
||
commits: commits.clone(),
|
||
});
|
||
|
||
failing_agent
|
||
.prompt("failed turn".to_string())
|
||
.await
|
||
.expect_err("completion failure without tools should roll back");
|
||
assert_eq!(commits.load(Ordering::SeqCst), 0);
|
||
|
||
let captured_messages = Arc::new(Mutex::new(Vec::new()));
|
||
let mut successful_agent = Agent::new(CapturingModel {
|
||
messages: captured_messages,
|
||
})
|
||
.memory(CommitTrackingMemory {
|
||
messages: vec!["prior".to_string()],
|
||
commits: commits.clone(),
|
||
});
|
||
|
||
successful_agent
|
||
.prompt("successful turn".to_string())
|
||
.await
|
||
.expect("successful turn should commit staged memory");
|
||
assert_eq!(commits.load(Ordering::SeqCst), 1);
|
||
}
|
||
|
||
#[tokio::test]
|
||
async fn non_fatal_tool_failure_preserves_structured_failure_and_output() {
|
||
let mut agent = Agent::new(FailingToolCallModel {
|
||
include_successful_tool: false,
|
||
})
|
||
.max_turns(1);
|
||
agent.tools.push(Box::new(FailingToolDyn { fatal: false }));
|
||
|
||
let error = agent
|
||
.prompt("执行失败工具".to_string())
|
||
.await
|
||
.expect_err("non-fatal failure should still respect max turns");
|
||
|
||
assert!(matches!(
|
||
error.error,
|
||
PromptError::MaxTurnsReached { max_turns: 1 }
|
||
));
|
||
let PromptOutput::ToolFailed(failure_output) = &error.partial_outputs[1] else {
|
||
panic!("structured tool failure should reach the caller");
|
||
};
|
||
assert_eq!(failure_output.failure.kind, ToolFailureKind::Network);
|
||
assert!(failure_output.failure.retryable);
|
||
assert!(!failure_output.failure.fatal);
|
||
assert_eq!(failure_output.output, json!({ "attempt": 1 }));
|
||
assert!(failure_output.message.contains("\"retryable\":true"));
|
||
}
|
||
|
||
#[tokio::test]
|
||
async fn fatal_tool_failure_keeps_earlier_success_and_closes_memory() {
|
||
let mut agent = Agent::new(FailingToolCallModel {
|
||
include_successful_tool: true,
|
||
})
|
||
.tool(TestTool {
|
||
requires_user_confirmation: false,
|
||
});
|
||
agent.tools.push(Box::new(FailingToolDyn { fatal: true }));
|
||
|
||
let error = agent
|
||
.prompt("先成功再失败".to_string())
|
||
.await
|
||
.expect_err("fatal tool should terminate the run");
|
||
|
||
assert!(matches!(error.error, PromptError::ToolError(_)));
|
||
assert!(matches!(error.partial_outputs[1], PromptOutput::Tool(_)));
|
||
assert!(matches!(
|
||
error.partial_outputs[2],
|
||
PromptOutput::ToolFailed(_)
|
||
));
|
||
let memory = agent
|
||
.memory
|
||
.as_ref()
|
||
.expect("tool activity should commit memory")
|
||
.get_memory();
|
||
assert!(memory.last().is_some_and(|message| {
|
||
is_system_tool_message(message, "agent-error", "network failed")
|
||
}));
|
||
}
|
||
|
||
#[tokio::test]
|
||
async fn post_tool_hook_failure_keeps_executed_tool_fact_and_closes_memory() {
|
||
let completion_count = Arc::new(AtomicUsize::new(0));
|
||
let model = RepeatingToolCallModel {
|
||
completion_count: completion_count.clone(),
|
||
};
|
||
let mut agent = Agent::new(model)
|
||
.tool(TestTool {
|
||
requires_user_confirmation: false,
|
||
})
|
||
.hook(StopAfterToolCallHook);
|
||
|
||
let error = agent
|
||
.prompt("执行后由 hook 终止".to_string())
|
||
.await
|
||
.expect_err("post-tool hook should terminate the run");
|
||
|
||
assert!(matches!(error.error, PromptError::ToolError(_)));
|
||
assert!(matches!(error.partial_outputs[1], PromptOutput::Tool(_)));
|
||
let memory = agent
|
||
.memory
|
||
.as_ref()
|
||
.expect("executed tool should commit memory")
|
||
.get_memory();
|
||
assert!(memory.last().is_some_and(|message| {
|
||
is_system_tool_message(
|
||
message,
|
||
"agent-error",
|
||
"tool call output caused this turn to stop by hook",
|
||
)
|
||
}));
|
||
}
|
||
|
||
#[tokio::test]
|
||
async fn multiple_tool_calls_execute_sequentially_in_array_order() {
|
||
let execution_order = Arc::new(Mutex::new(Vec::new()));
|
||
let mut agent = Agent::new(OrderedBatchModel).tool(OrderedTool {
|
||
execution_order: execution_order.clone(),
|
||
});
|
||
|
||
let outputs = agent
|
||
.prompt("执行两项操作".to_string())
|
||
.await
|
||
.expect("ordered pending tools should finish the turn");
|
||
|
||
assert_eq!(outputs.len(), 3);
|
||
assert_eq!(
|
||
*execution_order
|
||
.lock()
|
||
.expect("execution order lock should succeed"),
|
||
vec![2, 1]
|
||
);
|
||
}
|
||
|
||
#[tokio::test]
|
||
async fn pending_confirmation_tool_batch_finishes_without_another_completion() {
|
||
let completion_count = Arc::new(AtomicUsize::new(0));
|
||
let model = RepeatingToolCallModel {
|
||
completion_count: completion_count.clone(),
|
||
};
|
||
let mut agent = Agent::new(model)
|
||
.tool(TestTool {
|
||
requires_user_confirmation: true,
|
||
})
|
||
.max_turns(3);
|
||
|
||
let outputs = agent
|
||
.prompt("生成一张图".to_string())
|
||
.await
|
||
.expect("pending confirmation should finish the planning turn");
|
||
|
||
assert_eq!(completion_count.load(Ordering::SeqCst), 1);
|
||
assert_eq!(outputs.len(), 2);
|
||
assert!(matches!(outputs[0], PromptOutput::Text(_)));
|
||
assert!(matches!(outputs[1], PromptOutput::Tool(_)));
|
||
}
|
||
|
||
#[tokio::test]
|
||
async fn non_confirmation_tool_keeps_the_existing_max_turn_guard() {
|
||
let completion_count = Arc::new(AtomicUsize::new(0));
|
||
let model = RepeatingToolCallModel {
|
||
completion_count: completion_count.clone(),
|
||
};
|
||
let mut agent = Agent::new(model)
|
||
.tool(TestTool {
|
||
requires_user_confirmation: false,
|
||
})
|
||
.max_turns(3);
|
||
|
||
let error = agent
|
||
.prompt("生成一张图".to_string())
|
||
.await
|
||
.expect_err("a continuing tool should still hit the max-turn guard");
|
||
|
||
assert_eq!(completion_count.load(Ordering::SeqCst), 3);
|
||
assert!(matches!(
|
||
error.error,
|
||
PromptError::MaxTurnsReached { max_turns: 3 }
|
||
));
|
||
assert_eq!(error.partial_outputs.len(), 6);
|
||
let memory = agent
|
||
.memory
|
||
.as_ref()
|
||
.expect("tool activity should commit memory")
|
||
.get_memory();
|
||
assert!(
|
||
memory
|
||
.last()
|
||
.is_some_and(|message| { is_system_tool_message(message, "agent-error", "3") })
|
||
);
|
||
}
|
||
|
||
#[tokio::test]
|
||
async fn skipped_confirmation_result_keeps_the_existing_max_turn_guard() {
|
||
let completion_count = Arc::new(AtomicUsize::new(0));
|
||
let model = RepeatingToolCallModel {
|
||
completion_count: completion_count.clone(),
|
||
};
|
||
let mut agent = Agent::new(model)
|
||
.tool(TestTool {
|
||
requires_user_confirmation: true,
|
||
})
|
||
.hook(SkipAfterToolCallHook)
|
||
.max_turns(3);
|
||
|
||
let error = agent
|
||
.prompt("生成一张图".to_string())
|
||
.await
|
||
.expect_err("a skipped result must not finish as pending confirmation");
|
||
|
||
assert_eq!(completion_count.load(Ordering::SeqCst), 3);
|
||
assert!(matches!(
|
||
error.error,
|
||
PromptError::MaxTurnsReached { max_turns: 3 }
|
||
));
|
||
assert_eq!(error.partial_outputs.len(), 6);
|
||
assert_eq!(
|
||
error
|
||
.partial_outputs
|
||
.iter()
|
||
.filter(|output| matches!(output, PromptOutput::Text(_)))
|
||
.count(),
|
||
3
|
||
);
|
||
assert_eq!(
|
||
error
|
||
.partial_outputs
|
||
.iter()
|
||
.filter(|output| matches!(output, PromptOutput::Tool(_)))
|
||
.count(),
|
||
3
|
||
);
|
||
let memory = agent
|
||
.memory
|
||
.as_ref()
|
||
.expect("tool activity should commit memory")
|
||
.get_memory();
|
||
assert!(
|
||
memory
|
||
.last()
|
||
.is_some_and(|message| { is_system_tool_message(message, "agent-error", "3") })
|
||
);
|
||
}
|
||
}
|