Skip to main content

agent_block_testkit/
server.rs

1//! Axum-backed in-process mock LLM server.
2//!
3//! [`MockLlm`] serves one provider endpoint from a caller-supplied responder
4//! closure. Per request it records facts useful for assertions (call count,
5//! tool declarations, `tool_call_id`s of returned tool results) into a
6//! shared [`MockState`], derives [`RequestFacts`], and returns whatever the
7//! responder builds (typically via [`crate::shapes`]).
8
9use 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/// Which provider wire the mock speaks.
24#[derive(Clone, Copy, Debug, PartialEq, Eq)]
25pub enum Provider {
26    /// Anthropic Messages API, served at `POST /v1/messages`.
27    Anthropic,
28    /// OpenAI Chat Completions, served at `POST /chat/completions`.
29    OpenAi,
30}
31
32/// Facts derived from one incoming request, handed to the responder.
33#[derive(Clone, Debug)]
34pub struct RequestFacts {
35    /// 0-based index of this call (0 = first call the mock received).
36    pub call_index: usize,
37    /// Whether the request carries tool results from a previous turn
38    /// (Anthropic: `tool_result` content blocks; OpenAI: `role: "tool"` messages).
39    pub has_tool_results: bool,
40    /// Absolute paths extracted from a `Files:` section in any string message
41    /// content (the multi-file lazy-load convention). At most 2, in order.
42    pub paths: Vec<String>,
43    /// The full parsed request body.
44    pub body: Value,
45}
46
47/// Shared request-fact recorder. Cheap to clone; all clones share state.
48#[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    /// Total requests served.
58    pub fn call_count(&self) -> usize {
59        self.call_count.load(Ordering::SeqCst)
60    }
61
62    /// Requests that carried a `tools` array (of any content).
63    pub fn tools_declared_count(&self) -> usize {
64        self.tools_declared_count.load(Ordering::SeqCst)
65    }
66
67    /// Every tool name declared across all requests (with repetition).
68    pub fn declared_tool_names(&self) -> Vec<String> {
69        self.declared_tool_names.lock().expect("lock").clone()
70    }
71
72    /// How many times `name` was declared across all requests.
73    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    /// The `tool_call_id` / `tool_use_id` of every tool result seen in
83    /// requests, in order. A missing id is recorded as an empty string so
84    /// callers can assert its absence.
85    pub fn tool_result_ids(&self) -> Vec<String> {
86        self.tool_result_ids.lock().expect("lock").clone()
87    }
88}
89
90/// Handle to a spawned mock server.
91pub struct MockHandle {
92    /// Base URL (`http://127.0.0.1:<port>`); append nothing — the provider
93    /// route is already registered under it.
94    pub base_url: String,
95    /// Shared fact recorder for assertions.
96    pub state: MockState,
97    /// Cancel to shut the server down gracefully.
98    pub ct: CancellationToken,
99}
100
101type Responder = Arc<dyn Fn(&RequestFacts) -> Value + Send + Sync>;
102
103/// Builder for an in-process mock LLM server.
104pub struct MockLlm {
105    provider: Provider,
106    responder: Responder,
107}
108
109impl MockLlm {
110    /// Mock an Anthropic Messages endpoint (`POST /v1/messages`).
111    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    /// Mock an OpenAI Chat Completions endpoint (`POST /chat/completions`).
119    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    /// Bind an ephemeral local port and serve until the handle's token is
127    /// cancelled.
128    ///
129    /// # Panics
130    /// Panics only on OS-level port bind failure (fatal test infra condition).
131    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
210/// Record whether (and which) tools the request declared. Both wire forms are
211/// accepted: Anthropic `{name, ...}` and OpenAI `{function: {name, ...}}`.
212fn 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
230/// Record the id carried by every tool result in the request (empty string
231/// when absent).
232fn 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
266/// Provider-specific detection of "this request carries tool results".
267fn 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
288/// Extract absolute paths from a `Files:` section in any string message
289/// content (the multi-file lazy-load convention: `Files:\n  <abs>\n  <abs>`).
290/// Returns at most 2 paths in order of appearance.
291fn 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}