返回 CodeWhale
server.rs
根目录 / crates / tui / src / conformance / mcp / server.rs
1 //! The loopback Streamable HTTP transcript server for the `mcp` family.
2 //!
3 //! One server per case, scripted by the case's `server` object. It speaks
4 //! just enough HTTP for any MCP client to connect to it by URL, so a child
5 //! Node process can be pointed at it exactly like at a user's server.
6 //!
7 //! Per-method answers (`initialize`, `tools/list`, `resources/list`,
8 //! `resources/templates/list`, `prompts/list`, `tools/call`,
9 //! `resources/read`, `prompts/get`): either one entry, or an array of
10 //! candidate entries. For an array the first candidate whose `match` fields
11 //! all equal the request params wins (no `match` matches anything), and a
12 //! candidate marked `"once": true` is skipped after it has answered once —
13 //! that is how a case scripts "the first `tools/list` differs from the
14 //! second". Unknown methods get JSON-RPC `-32601`.
15 //!
16 //! An entry answers with exactly one of:
17 //!
18 //! - `result` / `error`: a JSON-RPC reply. `progress` (progress
19 //! notifications) and `notifications` (arbitrary notification objects) are
20 //! sent as SSE events ahead of the reply, which then becomes SSE too.
21 //! - `hold`: never answer until the client goes away (or [`HOLD_LIMIT`]).
22 //! - `http`: a raw HTTP response, `{status, headers, body}`, instead of a
23 //! JSON-RPC one. `{{authorization}}` in `body` is replaced by the request's
24 //! `Authorization` header value (the "server echoes your credential" case),
25 //! and `{{sink_url}}` anywhere in the spec by the case's sink server URL.
26 //! - `expire_session`: the server forgets the current session and answers
27 //! `404`, without running the request.
28 //! - `generate_tools` `{count, prefix}` / `generate_pages` `{count, prefix}`:
29 //! a `tools/list` result built at request time, so a cap-sized catalog does
30 //! not have to live in a fixture. `generate_pages` serves one tool per page
31 //! and links pages with `nextCursor` `page-N`.
32 //!
33 //! [`ServerOptions`] adds server-wide behaviour: numbered sessions (the
34 //! `initialize` reply issues `session-1`, `session-2`, … and every later
35 //! request must carry the latest one, else `404`) and a required bearer token
36 //! (anything else is `401`).
37 //!
38 //! What the server records, for the golden:
39 //!
40 //! - `received`: side-effecting requests it actually dispatched
41 //! (`tools/call`, `resources/read`, `prompts/get`). A request it refused at
42 //! the HTTP layer (stale session) is not here: the server never ran it.
43 //! - `requests`: every POST, as `{rpc, session, protocol_version,
44 //! authorization, status}`. The `Authorization` value is recorded only as a
45 //! shape (`<bearer>` / `<other>`), never as bytes, so a recorded frame
46 //! cannot carry a credential.
47
48 use std::collections::{HashMap, HashSet};
49 use std::sync::atomic::{AtomicUsize, Ordering};
50 use std::sync::{Arc, Mutex};
51 use std::time::Duration;
52
53 use serde_json::{Value, json};
54 use tokio::io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader};
55 use tokio::sync::Notify;
56
57 /// A held `tools/call` is released after this long even if nobody cancels,
58 /// so a dispatch that ignores cancellation fails the case instead of hanging.
59 pub(super) const HOLD_LIMIT: Duration = Duration::from_secs(20);
60
61 /// Server-wide behaviour a case opts into (`server_options`).
62 #[derive(Clone, Default)]
63 pub(super) struct ServerOptions {
64 /// `initialize` issues `session-N`; every later request must carry the
65 /// latest one or is refused with `404`.
66 pub(super) numbered_sessions: bool,
67 /// Every request must carry `Authorization: Bearer <token>`, else `401`.
68 pub(super) require_bearer: Option<String>,
69 }
70
71 pub(super) struct State {
72 spec: Value,
73 options: ServerOptions,
74 received: Mutex<Vec<Value>>,
75 requests: Mutex<Vec<Value>>,
76 /// Signalled when a `hold` entry starts holding a request.
77 pub(super) held: Notify,
78 sessions_issued: AtomicUsize,
79 current_session: Mutex<Option<String>>,
80 used: Mutex<HashSet<(String, usize)>>,
81 }
82
83 pub(super) struct TranscriptServer {
84 pub(super) url: String,
85 pub(super) addr: String,
86 pub(super) state: Arc<State>,
87 task: tokio::task::JoinHandle<()>,
88 }
89
90 impl Drop for TranscriptServer {
91 fn drop(&mut self) {
92 self.task.abort();
93 }
94 }
95
96 impl TranscriptServer {
97 pub(super) async fn start(spec: Value, options: ServerOptions) -> Self {
98 let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
99 .await
100 .expect("bind transcript server");
101 let addr = listener.local_addr().expect("addr").to_string();
102 let url = format!("http://{addr}/mcp");
103 let state = Arc::new(State {
104 spec,
105 options,
106 received: Mutex::new(Vec::new()),
107 requests: Mutex::new(Vec::new()),
108 held: Notify::new(),
109 sessions_issued: AtomicUsize::new(0),
110 current_session: Mutex::new(None),
111 used: Mutex::new(HashSet::new()),
112 });
113 let served = Arc::clone(&state);
114 let task = tokio::spawn(async move {
115 let mut connections = tokio::task::JoinSet::new();
116 loop {
117 let accepted = tokio::select! {
118 accepted = listener.accept() => accepted,
119 _ = connections.join_next(), if !connections.is_empty() => continue,
120 };
121 let Ok((socket, _)) = accepted else {
122 break;
123 };
124 let state = Arc::clone(&served);
125 connections.spawn(async move {
126 answer(socket, &state).await;
127 });
128 }
129 });
130 Self {
131 url,
132 addr,
133 state,
134 task,
135 }
136 }
137
138 pub(super) fn received(&self) -> Vec<Value> {
139 self.state.received.lock().expect("log").clone()
140 }
141
142 pub(super) fn requests(&self) -> Vec<Value> {
143 self.state.requests.lock().expect("log").clone()
144 }
145 }
146
147 struct Request {
148 method: String,
149 headers: HashMap<String, String>,
150 body: Vec<u8>,
151 }
152
153 async fn read_request(socket: tokio::net::TcpStream) -> Option<(Request, tokio::net::TcpStream)> {
154 let mut reader = BufReader::new(socket);
155 let mut line = String::new();
156 if reader.read_line(&mut line).await.unwrap_or(0) == 0 {
157 return None;
158 }
159 let method = line
160 .split_whitespace()
161 .next()
162 .unwrap_or_default()
163 .to_string();
164 let mut headers = HashMap::new();
165 loop {
166 line.clear();
167 if reader.read_line(&mut line).await.unwrap_or(0) == 0 {
168 return None;
169 }
170 if line == "\r\n" || line == "\n" {
171 break;
172 }
173 if let Some((name, value)) = line.split_once(':') {
174 headers.insert(name.trim().to_ascii_lowercase(), value.trim().to_string());
175 }
176 }
177 let content_length = headers
178 .get("content-length")
179 .and_then(|value| value.parse().ok())
180 .unwrap_or(0usize);
181 let mut body = vec![0u8; content_length];
182 reader.read_exact(&mut body).await.ok()?;
183 Some((
184 Request {
185 method,
186 headers,
187 body,
188 },
189 reader.into_inner(),
190 ))
191 }
192
193 fn reason(status: u16) -> &'static str {
194 match status {
195 200 => "OK",
196 202 => "Accepted",
197 307 => "Temporary Redirect",
198 400 => "Bad Request",
199 401 => "Unauthorized",
200 404 => "Not Found",
201 405 => "Method Not Allowed",
202 500 => "Internal Server Error",
203 _ => "Status",
204 }
205 }
206
207 struct Reply {
208 status: u16,
209 headers: Vec<(String, String)>,
210 body: String,
211 }
212
213 impl Reply {
214 fn new(status: u16, body: impl Into<String>) -> Self {
215 Self {
216 status,
217 headers: Vec::new(),
218 body: body.into(),
219 }
220 }
221
222 fn header(mut self, name: &str, value: impl Into<String>) -> Self {
223 self.headers.push((name.to_string(), value.into()));
224 self
225 }
226
227 fn bytes(&self) -> Vec<u8> {
228 let mut head = format!("HTTP/1.1 {} {}\r\n", self.status, reason(self.status));
229 for (name, value) in &self.headers {
230 head.push_str(&format!("{name}: {value}\r\n"));
231 }
232 head.push_str(&format!(
233 "Content-Length: {}\r\nConnection: close\r\n\r\n{}",
234 self.body.len(),
235 self.body
236 ));
237 head.into_bytes()
238 }
239 }
240
241 /// The shape of an `Authorization` header, never its bytes.
242 fn authorization_shape(value: Option<&String>) -> Value {
243 match value {
244 None => Value::Null,
245 Some(value) if value.to_ascii_lowercase().starts_with("bearer ") => json!("<bearer>"),
246 Some(_) => json!("<other>"),
247 }
248 }
249
250 async fn answer(socket: tokio::net::TcpStream, state: &State) {
251 let Some((request, mut socket)) = read_request(socket).await else {
252 return;
253 };
254 if request.method != "POST" {
255 // No server-initiated stream: the spec's answer for a GET.
256 let _ = socket
257 .write_all(b"HTTP/1.1 405 Method Not Allowed\r\nAllow: POST\r\nContent-Length: 0\r\nConnection: close\r\n\r\n")
258 .await;
259 return;
260 }
261 let Ok(message) = serde_json::from_slice::<Value>(&request.body) else {
262 let _ = socket
263 .write_all(
264 b"HTTP/1.1 400 Bad Request\r\nContent-Length: 0\r\nConnection: close\r\n\r\n",
265 )
266 .await;
267 return;
268 };
269 let rpc_method = message["method"].as_str().unwrap_or_default().to_string();
270 let params = message.get("params").cloned().unwrap_or(Value::Null);
271 let session = request.headers.get("mcp-session-id").cloned();
272 let authorization = request.headers.get("authorization").cloned();
273 let record = |status: u16| {
274 state.requests.lock().expect("log").push(json!({
275 "rpc": rpc_method,
276 "session": session,
277 "protocol_version": request.headers.get("mcp-protocol-version"),
278 "authorization": authorization_shape(authorization.as_ref()),
279 "status": status,
280 }));
281 };
282 let finish = async |socket: &mut tokio::net::TcpStream, reply: Reply| {
283 record(reply.status);
284 let _ = socket.write_all(&reply.bytes()).await;
285 let _ = socket.shutdown().await;
286 };
287
288 if let Some(token) = &state.options.require_bearer
289 && authorization.as_deref() != Some(format!("Bearer {token}").as_str())
290 {
291 let reply = Reply::new(401, r#"{"error":"unauthorized"}"#)
292 .header("WWW-Authenticate", "Bearer realm=\"conformance\"")
293 .header("Content-Type", "application/json");
294 finish(&mut socket, reply).await;
295 return;
296 }
297 let is_initialize = rpc_method == "initialize";
298 if state.options.numbered_sessions && !is_initialize {
299 let current = state.current_session.lock().expect("session").clone();
300 if current.is_none() || current != session {
301 finish(&mut socket, Reply::new(404, "session not found")).await;
302 return;
303 }
304 }
305 let Some(id) = message.get("id").cloned() else {
306 // Notifications (initialized, cancelled, progress) are accepted.
307 finish(&mut socket, Reply::new(202, "")).await;
308 return;
309 };
310
311 let entry = scripted_entry(state, &rpc_method, &params);
312 if entry.get("expire_session").and_then(Value::as_bool) == Some(true) {
313 *state.current_session.lock().expect("session") = None;
314 finish(&mut socket, Reply::new(404, "session not found")).await;
315 return;
316 }
317 match rpc_method.as_str() {
318 "tools/call" => state.received.lock().expect("log").push(json!({
319 "method": "tools/call",
320 "name": params["name"],
321 "arguments": params.get("arguments").cloned().unwrap_or(Value::Null),
322 })),
323 "resources/read" => state
324 .received
325 .lock()
326 .expect("log")
327 .push(json!({ "method": "resources/read", "uri": params["uri"] })),
328 "prompts/get" => state
329 .received
330 .lock()
331 .expect("log")
332 .push(json!({ "method": "prompts/get", "name": params["name"] })),
333 _ => {}
334 }
335 if entry.get("hold").and_then(Value::as_bool) == Some(true) {
336 state.held.notify_one();
337 // Answer nothing until the client goes away (or the hold limit).
338 let mut byte = [0u8; 1];
339 let _ = tokio::time::timeout(HOLD_LIMIT, socket.read(&mut byte)).await;
340 return;
341 }
342 if let Some(http) = entry.get("http") {
343 let status = http["status"].as_u64().unwrap_or(500) as u16;
344 let body = http["body"].as_str().unwrap_or_default().replace(
345 "{{authorization}}",
346 authorization.as_deref().unwrap_or("<none>"),
347 );
348 let mut reply = Reply::new(status, body);
349 if let Some(headers) = http["headers"].as_object() {
350 for (name, value) in headers {
351 reply = reply.header(name, value.as_str().unwrap_or_default());
352 }
353 }
354 finish(&mut socket, reply).await;
355 return;
356 }
357 let body = match entry.get("error") {
358 Some(error) => json!({ "jsonrpc": "2.0", "id": id, "error": error }),
359 None => json!({ "jsonrpc": "2.0", "id": id, "result": entry["result"] }),
360 };
361 let mut reply = if entry.get("progress").is_some() || entry.get("notifications").is_some() {
362 let token = params
363 .pointer("/_meta/progressToken")
364 .cloned()
365 .unwrap_or_else(|| json!("conformance"));
366 let mut sse = String::new();
367 for step in entry
368 .get("progress")
369 .and_then(Value::as_array)
370 .into_iter()
371 .flatten()
372 {
373 let notification = json!({
374 "jsonrpc": "2.0",
375 "method": "notifications/progress",
376 "params": { "progressToken": token, "progress": step["progress"], "total": step["total"] },
377 });
378 sse.push_str(&format!("event: message\ndata: {notification}\n\n"));
379 }
380 for notification in entry
381 .get("notifications")
382 .and_then(Value::as_array)
383 .into_iter()
384 .flatten()
385 {
386 sse.push_str(&format!("event: message\ndata: {notification}\n\n"));
387 }
388 sse.push_str(&format!("event: message\ndata: {body}\n\n"));
389 Reply::new(200, sse).header("Content-Type", "text/event-stream")
390 } else {
391 Reply::new(200, body.to_string()).header("Content-Type", "application/json")
392 };
393 if is_initialize {
394 let session = if state.options.numbered_sessions {
395 let number = state.sessions_issued.fetch_add(1, Ordering::SeqCst) + 1;
396 let issued = format!("session-{number}");
397 *state.current_session.lock().expect("session") = Some(issued.clone());
398 issued
399 } else {
400 "conformance-session".to_string()
401 };
402 reply = reply.header("Mcp-Session-Id", session);
403 }
404 finish(&mut socket, reply).await;
405 }
406
407 /// The transcript's answer for one request: one entry, or the first usable
408 /// candidate of an array (see the module doc). Unknown methods get -32601.
409 fn scripted_entry(state: &State, method: &str, params: &Value) -> Value {
410 let not_found =
411 || json!({ "error": { "code": -32601, "message": format!("method not found: {method}") } });
412 let Some(entry) = state.spec.get(method) else {
413 return not_found();
414 };
415 let Some(candidates) = entry.as_array() else {
416 return expand_generated(entry.clone(), params);
417 };
418 for (index, candidate) in candidates.iter().enumerate() {
419 let matches = candidate["match"].as_object().is_none_or(|fields| {
420 fields
421 .iter()
422 .all(|(key, value)| params.get(key) == Some(value))
423 });
424 if !matches {
425 continue;
426 }
427 if candidate.get("once").and_then(Value::as_bool) == Some(true)
428 && !state
429 .used
430 .lock()
431 .expect("used")
432 .insert((method.to_string(), index))
433 {
434 continue;
435 }
436 return expand_generated(candidate.clone(), params);
437 }
438 json!({ "error": { "code": -32602, "message": format!("no scripted {method} answer for {params}") } })
439 }
440
441 fn generated_tool(prefix: &str, index: usize) -> Value {
442 json!({ "name": format!("{prefix}-{index:05}"), "inputSchema": { "type": "object" } })
443 }
444
445 fn expand_generated(mut entry: Value, params: &Value) -> Value {
446 if let Some(spec) = entry.get("generate_tools") {
447 let count = spec["count"].as_u64().unwrap_or(0) as usize;
448 let prefix = spec["prefix"].as_str().unwrap_or("tool");
449 let tools: Vec<Value> = (0..count)
450 .map(|index| generated_tool(prefix, index))
451 .collect();
452 entry["result"] = json!({ "tools": tools });
453 } else if let Some(spec) = entry.get("generate_pages") {
454 let count = spec["count"].as_u64().unwrap_or(1) as usize;
455 let prefix = spec["prefix"].as_str().unwrap_or("tool");
456 let page = params["cursor"]
457 .as_str()
458 .and_then(|cursor| cursor.strip_prefix("page-"))
459 .and_then(|page| page.parse::<usize>().ok())
460 .unwrap_or(1);
461 let mut result = json!({ "tools": [generated_tool(prefix, page)] });
462 if page < count {
463 result["nextCursor"] = json!(format!("page-{}", page + 1));
464 }
465 entry["result"] = result;
466 }
467 entry
468 }
469
469 lines RUST