Files
Genarrative/server-rs/crates/platform-agent-harness/src/run.rs
T
k88936 b09a98db48 重构画布代理提示词与工具调用循环
统一系统消息、工具结果与异常响应的消息构造。
补充无效响应纠正、批量工具调用和待确认状态处理。
精简图片上下文提示并完善越界引用回退规则。
补齐运行循环和提示词行为测试。
2026-07-31 11:30:48 +08:00

1666 lines
58 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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") })
);
}
}