agent_base/engine/
turn_facts.rs1use std::sync::Mutex;
2
3use async_trait::async_trait;
4
5use crate::engine::middleware::{Middleware, PostLlmCtx, UserMessageCtx};
6use crate::types::{AgentResult, Language};
7
8pub struct TurnFactMiddleware {
26 pending_facts: Mutex<Vec<String>>,
27 language: Language,
28}
29
30impl TurnFactMiddleware {
31 pub fn new() -> Self {
33 Self {
34 pending_facts: Mutex::new(Vec::new()),
35 language: Language::Zh,
36 }
37 }
38
39 pub fn with_language(language: Language) -> Self {
41 Self {
42 pending_facts: Mutex::new(Vec::new()),
43 language,
44 }
45 }
46}
47
48impl Default for TurnFactMiddleware {
49 fn default() -> Self {
50 Self::new()
51 }
52}
53
54#[async_trait]
55impl Middleware for TurnFactMiddleware {
56 async fn on_user_message(&self, ctx: &mut UserMessageCtx) -> AgentResult<()> {
57 let facts = {
58 let mut guard = self.pending_facts.lock().unwrap();
59 if guard.is_empty() {
60 return Ok(());
61 }
62 std::mem::take(&mut *guard)
63 };
64
65 let prefix = match self.language {
67 Language::Zh => format!(
68 "[本轮工具调用摘要 — 以下为确定性事实,请以此为准]\n{}\n",
69 facts.join("\n")
70 ),
71 Language::En => format!(
72 "[Previous turn tool-call summary — treat these as ground truth]\n{}\n",
73 facts.join("\n")
74 ),
75 };
76 ctx.user_input = format!("{prefix}\n{original}", original = ctx.user_input);
77
78 Ok(())
79 }
80
81 async fn on_post_llm(&self, ctx: &mut PostLlmCtx) -> AgentResult<()> {
82 if ctx.tool_calls.is_empty() {
83 return Ok(());
84 }
85
86 let mut facts = Vec::new();
87 for (_id, name, _args) in &ctx.tool_calls {
88 let fact = match self.language {
91 Language::Zh => format!("- 调用了工具: {name}"),
92 Language::En => format!("- Called tool: {name}"),
93 };
94 facts.push(fact);
95 }
96
97 let mut guard = self.pending_facts.lock().unwrap();
98 guard.extend(facts);
99
100 Ok(())
101 }
102}
103
104#[cfg(test)]
105mod tests {
106 use super::*;
107 use crate::types::{FinishReason, SessionId};
108
109 #[tokio::test]
110 async fn test_no_facts_no_prefix() {
111 let mw = TurnFactMiddleware::new();
112 let mut ctx = UserMessageCtx {
113 session_id: SessionId::new(1),
114 user_input: "hello".to_string(),
115 };
116 mw.on_user_message(&mut ctx).await.unwrap();
117 assert_eq!(ctx.user_input, "hello");
118 }
119
120 #[tokio::test]
121 async fn test_facts_injected_on_next_user_message() {
122 let mw = TurnFactMiddleware::new();
123
124 let mut post_ctx = PostLlmCtx {
126 session_id: SessionId::new(1),
127 full_text: String::new(),
128 is_tool_call: true,
129 tool_calls: vec![
130 ("id1".into(), "execute_ssh_command".into(), "{}".into()),
131 ("id2".into(), "start_interactive_task".into(), "{}".into()),
132 ],
133 available_tools: vec![],
134 turn_count: 1,
135 total_tool_calls: 0,
136 nudge_count: 0,
137 turn_tool_calls: 0,
138 skip_push: false,
139 follow_up_message: None,
140 finish_reason: FinishReason::Stop,
141 };
142 mw.on_post_llm(&mut post_ctx).await.unwrap();
143
144 let mut user_ctx = UserMessageCtx {
146 session_id: SessionId::new(1),
147 user_input: "继续执行".to_string(),
148 };
149 mw.on_user_message(&mut user_ctx).await.unwrap();
150
151 assert!(user_ctx.user_input.contains("本轮工具调用摘要"));
152 assert!(user_ctx.user_input.contains("execute_ssh_command"));
153 assert!(user_ctx.user_input.contains("start_interactive_task"));
154 assert!(user_ctx.user_input.contains("继续执行"));
155 }
156
157 #[tokio::test]
158 async fn test_facts_cleared_after_injection() {
159 let mw = TurnFactMiddleware::new();
160
161 let mut post_ctx = PostLlmCtx {
163 session_id: SessionId::new(1),
164 full_text: String::new(),
165 is_tool_call: true,
166 tool_calls: vec![("id1".into(), "docker".into(), "{}".into())],
167 available_tools: vec![],
168 turn_count: 1,
169 total_tool_calls: 0,
170 nudge_count: 0,
171 turn_tool_calls: 0,
172 skip_push: false,
173 follow_up_message: None,
174 finish_reason: FinishReason::Stop,
175 };
176 mw.on_post_llm(&mut post_ctx).await.unwrap();
177
178 let mut ctx1 = UserMessageCtx {
180 session_id: SessionId::new(1),
181 user_input: "next".into(),
182 };
183 mw.on_user_message(&mut ctx1).await.unwrap();
184 assert!(ctx1.user_input.contains("本轮工具调用摘要"));
185
186 let mut ctx2 = UserMessageCtx {
188 session_id: SessionId::new(1),
189 user_input: "again".into(),
190 };
191 mw.on_user_message(&mut ctx2).await.unwrap();
192 assert_eq!(ctx2.user_input, "again");
193 }
194
195 #[tokio::test]
196 async fn test_english_language_prefix() {
197 let mw = TurnFactMiddleware::with_language(Language::En);
198
199 let mut post_ctx = PostLlmCtx {
200 session_id: SessionId::new(1),
201 full_text: String::new(),
202 is_tool_call: true,
203 tool_calls: vec![("id1".into(), "docker".into(), "{}".into())],
204 available_tools: vec![],
205 turn_count: 1,
206 total_tool_calls: 0,
207 nudge_count: 0,
208 turn_tool_calls: 0,
209 skip_push: false,
210 follow_up_message: None,
211 finish_reason: FinishReason::Stop,
212 };
213 mw.on_post_llm(&mut post_ctx).await.unwrap();
214
215 let mut ctx = UserMessageCtx {
216 session_id: SessionId::new(1),
217 user_input: "continue".into(),
218 };
219 mw.on_user_message(&mut ctx).await.unwrap();
220
221 assert!(ctx.user_input.contains("Previous turn tool-call summary"));
222 assert!(ctx.user_input.contains("ground truth"));
223 assert!(ctx.user_input.contains("continue"));
224 }
225}