agent_block_testkit/
server.rs1use axum::{
10 extract::State,
11 http::{header, StatusCode},
12 response::IntoResponse,
13 routing::post,
14 Router,
15};
16use serde_json::{json, Value};
17use std::sync::{
18 atomic::{AtomicUsize, Ordering},
19 Arc, Mutex,
20};
21use tokio_util::sync::CancellationToken;
22
23#[derive(Clone, Copy, Debug, PartialEq, Eq)]
25pub enum Provider {
26 Anthropic,
28 OpenAi,
30}
31
32#[derive(Clone, Debug)]
34pub struct RequestFacts {
35 pub call_index: usize,
37 pub has_tool_results: bool,
40 pub paths: Vec<String>,
43 pub body: Value,
45}
46
47#[derive(Clone, Default)]
49pub struct MockState {
50 call_count: Arc<AtomicUsize>,
51 tools_declared_count: Arc<AtomicUsize>,
52 declared_tool_names: Arc<Mutex<Vec<String>>>,
53 tool_result_ids: Arc<Mutex<Vec<String>>>,
54}
55
56impl MockState {
57 pub fn call_count(&self) -> usize {
59 self.call_count.load(Ordering::SeqCst)
60 }
61
62 pub fn tools_declared_count(&self) -> usize {
64 self.tools_declared_count.load(Ordering::SeqCst)
65 }
66
67 pub fn declared_tool_names(&self) -> Vec<String> {
69 self.declared_tool_names.lock().expect("lock").clone()
70 }
71
72 pub fn declared_count_of(&self, name: &str) -> usize {
74 self.declared_tool_names
75 .lock()
76 .expect("lock")
77 .iter()
78 .filter(|n| n.as_str() == name)
79 .count()
80 }
81
82 pub fn tool_result_ids(&self) -> Vec<String> {
86 self.tool_result_ids.lock().expect("lock").clone()
87 }
88}
89
90pub struct MockHandle {
92 pub base_url: String,
95 pub state: MockState,
97 pub ct: CancellationToken,
99}
100
101type Responder = Arc<dyn Fn(&RequestFacts) -> Value + Send + Sync>;
102
103pub struct MockLlm {
105 provider: Provider,
106 responder: Responder,
107}
108
109impl MockLlm {
110 pub fn anthropic(responder: impl Fn(&RequestFacts) -> Value + Send + Sync + 'static) -> Self {
112 Self {
113 provider: Provider::Anthropic,
114 responder: Arc::new(responder),
115 }
116 }
117
118 pub fn openai(responder: impl Fn(&RequestFacts) -> Value + Send + Sync + 'static) -> Self {
120 Self {
121 provider: Provider::OpenAi,
122 responder: Arc::new(responder),
123 }
124 }
125
126 pub async fn spawn(self) -> MockHandle {
132 let state = MockState::default();
133 let ct = CancellationToken::new();
134
135 let route = match self.provider {
136 Provider::Anthropic => "/v1/messages",
137 Provider::OpenAi => "/chat/completions",
138 };
139 let shared = HandlerState {
140 provider: self.provider,
141 responder: self.responder,
142 state: state.clone(),
143 };
144 let router = Router::new()
145 .route(route, post(handle_request))
146 .with_state(shared);
147
148 let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
149 .await
150 .expect("bind ephemeral port for testkit mock LLM");
151 let addr = listener.local_addr().expect("local_addr");
152
153 let ct_shutdown = ct.clone();
154 tokio::spawn(async move {
155 let _ = axum::serve(listener, router)
156 .with_graceful_shutdown(async move { ct_shutdown.cancelled_owned().await })
157 .await;
158 });
159
160 MockHandle {
161 base_url: format!("http://{addr}"),
162 state,
163 ct,
164 }
165 }
166}
167
168#[derive(Clone)]
169struct HandlerState {
170 provider: Provider,
171 responder: Responder,
172 state: MockState,
173}
174
175async fn handle_request(
176 State(hs): State<HandlerState>,
177 body: axum::body::Bytes,
178) -> impl IntoResponse {
179 let parsed = match serde_json::from_slice::<Value>(&body) {
180 Ok(v) => v,
181 Err(e) => {
182 let err_body = json!({ "error": format!("bad request: {e}") }).to_string();
183 return (
184 StatusCode::BAD_REQUEST,
185 [(header::CONTENT_TYPE, "application/json")],
186 err_body,
187 );
188 }
189 };
190
191 let call_index = hs.state.call_count.fetch_add(1, Ordering::SeqCst);
192 record_tools(&hs.state, &parsed);
193 record_tool_result_ids(&hs.state, hs.provider, &parsed);
194
195 let facts = RequestFacts {
196 call_index,
197 has_tool_results: has_tool_results(hs.provider, &parsed),
198 paths: extract_paths(&parsed),
199 body: parsed,
200 };
201
202 let response = (hs.responder)(&facts);
203 (
204 StatusCode::OK,
205 [(header::CONTENT_TYPE, "application/json")],
206 response.to_string(),
207 )
208}
209
210fn record_tools(state: &MockState, body: &Value) {
213 let Some(tools) = body.get("tools").and_then(|t| t.as_array()) else {
214 return;
215 };
216 state.tools_declared_count.fetch_add(1, Ordering::SeqCst);
217 let mut names = state.declared_tool_names.lock().expect("lock");
218 for t in tools {
219 let name = t.get("name").and_then(|n| n.as_str()).or_else(|| {
220 t.get("function")
221 .and_then(|f| f.get("name"))
222 .and_then(|n| n.as_str())
223 });
224 if let Some(name) = name {
225 names.push(name.to_string());
226 }
227 }
228}
229
230fn record_tool_result_ids(state: &MockState, provider: Provider, body: &Value) {
233 let Some(messages) = body.get("messages").and_then(|m| m.as_array()) else {
234 return;
235 };
236 let mut ids = state.tool_result_ids.lock().expect("lock");
237 for msg in messages {
238 match provider {
239 Provider::OpenAi => {
240 if msg.get("role").and_then(|r| r.as_str()) == Some("tool") {
241 let id = msg
242 .get("tool_call_id")
243 .and_then(|v| v.as_str())
244 .unwrap_or("");
245 ids.push(id.to_string());
246 }
247 }
248 Provider::Anthropic => {
249 if msg.get("role").and_then(|r| r.as_str()) != Some("user") {
250 continue;
251 }
252 let Some(blocks) = msg.get("content").and_then(|c| c.as_array()) else {
253 continue;
254 };
255 for b in blocks {
256 if b.get("type").and_then(|t| t.as_str()) == Some("tool_result") {
257 let id = b.get("tool_use_id").and_then(|v| v.as_str()).unwrap_or("");
258 ids.push(id.to_string());
259 }
260 }
261 }
262 }
263 }
264}
265
266fn has_tool_results(provider: Provider, body: &Value) -> bool {
268 let Some(messages) = body.get("messages").and_then(|m| m.as_array()) else {
269 return false;
270 };
271 messages.iter().any(|msg| match provider {
272 Provider::OpenAi => msg.get("role").and_then(|r| r.as_str()) == Some("tool"),
273 Provider::Anthropic => {
274 msg.get("role").and_then(|r| r.as_str()) == Some("user")
275 && msg
276 .get("content")
277 .and_then(|c| c.as_array())
278 .map(|blocks| {
279 blocks
280 .iter()
281 .any(|b| b.get("type").and_then(|t| t.as_str()) == Some("tool_result"))
282 })
283 .unwrap_or(false)
284 }
285 })
286}
287
288fn extract_paths(body: &Value) -> Vec<String> {
292 let mut paths = Vec::new();
293 let Some(messages) = body.get("messages").and_then(|m| m.as_array()) else {
294 return paths;
295 };
296 for msg in messages {
297 let Some(content) = msg.get("content").and_then(|c| c.as_str()) else {
298 continue;
299 };
300 let mut in_files_section = false;
301 for line in content.lines() {
302 if line.trim() == "Files:" {
303 in_files_section = true;
304 continue;
305 }
306 if in_files_section {
307 let trimmed = line.trim();
308 if trimmed.starts_with('/') {
309 let p = trimmed.to_string();
310 if !paths.contains(&p) {
311 paths.push(p);
312 }
313 if paths.len() >= 2 {
314 return paths;
315 }
316 } else if !trimmed.is_empty() {
317 in_files_section = false;
318 }
319 }
320 }
321 }
322 paths
323}