git.lucas.co / cce-ui
GPU-accelerated UI toolkit (Vulkan)
git clone https://git.lucas.co/cce-ui.git

src/mcp.rs (18.5K)

  1 //! Minimal MCP (Model Context Protocol) server support for cce-ui apps.
  2 //!
  3 //! Implements the tools-only subset of the spec — `initialize`, `tools/list`,
  4 //! `tools/call`, `ping` — over the Streamable HTTP transport (JSON-RPC 2.0 in
  5 //! HTTP POST bodies), so any MCP client can inspect and drive a running app,
  6 //! e.g. `claude mcp add --transport http <name> http://127.0.0.1:<port>/mcp`.
  7 //! The server is stateless: no sessions, no SSE stream, no server-initiated
  8 //! messages (a GET gets 405).
  9 //!
 10 //! An app declares its [`McpTool`]s and calls [`start_mcp_server`] with the
 11 //! engine's calloop sender plus a constructor wrapping [`McpToolCall`] into
 12 //! its `Application::Message`; each `tools/call` is executed on the app's
 13 //! event loop and answered over the carried mpsc channel — the same bridge
 14 //! pattern as cce-designer's HTTP automation API.
 15 
 16 use std::io::{BufRead, BufReader, Read, Write};
 17 use std::net::{TcpListener, TcpStream};
 18 use std::sync::{mpsc, Arc};
 19 use std::time::Duration;
 20 
 21 use serde_json::{json, Value};
 22 
 23 /// Protocol revisions this server accepts; the newest is offered when the
 24 /// client requests anything else. The tools-only subset is identical across
 25 /// all of them.
 26 const PROTOCOL_VERSIONS: [&str; 3] = ["2025-06-18", "2025-03-26", "2024-11-05"];
 27 
 28 /// How long a `tools/call` waits on the app's event loop before failing.
 29 const CALL_TIMEOUT: Duration = Duration::from_secs(30);
 30 
 31 /// How long a connection may take to deliver its request. Each connection
 32 /// holds a thread, so one that never finishes sending must not hold it forever.
 33 const READ_TIMEOUT: Duration = Duration::from_secs(10);
 34 
 35 /// The largest request body read. The body is allocated at the size the
 36 /// client's `Content-Length` claims, and an allocation that fails aborts the
 37 /// whole app — not just this thread — so one request claiming 2^63 bytes
 38 /// would take the window and its unsaved work down with it.
 39 const MAX_BODY: usize = 16 << 20;
 40 /// Room for the request line and headers on top of the body.
 41 const MAX_HEAD: usize = 64 << 10;
 42 
 43 /// Whether a request came from a program on this machine rather than from a
 44 /// web page the user has open.
 45 ///
 46 /// Binding loopback keeps the network out, not the browser: any page can POST
 47 /// to `127.0.0.1:<port>` (a `text/plain` body needs no CORS preflight, and
 48 /// this server parses whatever arrives as JSON), and with DNS rebinding it can
 49 /// read the replies too — which for cce-notes is the whole vault. Two checks,
 50 /// as the MCP transport spec asks:
 51 ///
 52 /// - **`Origin`**, which a browser attaches to every POST, must be absent (an
 53 ///   MCP client is not a browser and sends none) or name loopback. That stops
 54 ///   a page posting from its own origin.
 55 /// - **`Host`** must name loopback. A rebound page posts to ITS hostname, now
 56 ///   resolving to 127.0.0.1, and that hostname is what arrives here — the one
 57 ///   thing a rebind cannot change.
 58 fn request_from_this_machine(host: Option<&str>, origin: Option<&str>) -> bool {
 59     // "localhost:3002", "[::1]:3002", "127.0.0.1" -> the bare host name.
 60     fn hostname(authority: &str) -> &str {
 61         if let Some(rest) = authority.strip_prefix('[') {
 62             return rest.split(']').next().unwrap_or("");
 63         }
 64         authority.split(':').next().unwrap_or("")
 65     }
 66     fn loopback(authority: &str) -> bool {
 67         let h = hostname(authority).to_ascii_lowercase();
 68         h == "localhost" || h == "::1" || h.parse::<std::net::Ipv4Addr>().is_ok_and(|ip| ip.is_loopback())
 69     }
 70     let host_ok = host.is_some_and(loopback);
 71     let origin_ok = match origin {
 72         None => true,
 73         Some(o) => o
 74             .strip_prefix("http://")
 75             .or_else(|| o.strip_prefix("https://"))
 76             .is_some_and(|authority| loopback(authority.trim_end_matches('/'))),
 77     };
 78     host_ok && origin_ok
 79 }
 80 
 81 /// A tool the app exposes over MCP.
 82 #[derive(Debug, Clone)]
 83 pub struct McpTool {
 84     pub name: String,
 85     pub description: String,
 86     /// JSON Schema for the tool's arguments (the spec's `inputSchema`).
 87     pub input_schema: Value,
 88 }
 89 
 90 /// A `tools/call` in flight: delivered to the app's event loop wrapped in its
 91 /// `Application::Message`; the handler sends the outcome back over `reply`.
 92 /// An `Ok` value becomes the result's text content (strings verbatim, other
 93 /// JSON pretty-printed); an `Err` becomes an `isError` tool result.
 94 #[derive(Debug, Clone)]
 95 pub struct McpToolCall {
 96     pub name: String,
 97     pub arguments: Value,
 98     pub reply: mpsc::Sender<Result<Value, String>>,
 99 }
100 
101 /// Spawn the MCP server on `127.0.0.1:port` (one thread per connection,
102 /// mirroring the raw-HTTP style of cce-designer's api.rs). `wrap` lifts a
103 /// tool call into the app's message type for delivery over `sender`.
104 pub fn start_mcp_server<M, F>(
105     server_name: &str,
106     port: u16,
107     tools: Vec<McpTool>,
108     sender: calloop::channel::Sender<M>,
109     wrap: F,
110 ) where
111     M: Send + 'static,
112     F: Fn(McpToolCall) -> M + Send + Sync + 'static,
113 {
114     let server_name = server_name.to_string();
115     std::thread::spawn(move || {
116         let listener = match TcpListener::bind(("127.0.0.1", port)) {
117             Ok(l) => l,
118             Err(e) => {
119                 eprintln!("Failed to bind MCP server to port {port}: {e:?}");
120                 return;
121             }
122         };
123         println!("MCP server '{server_name}' listening on http://127.0.0.1:{port}");
124 
125         let shared = Arc::new((server_name, tools, sender, wrap));
126         for stream in listener.incoming() {
127             let stream = match stream {
128                 Ok(s) => s,
129                 Err(_) => continue,
130             };
131             let shared = Arc::clone(&shared);
132             std::thread::spawn(move || {
133                 let (server_name, tools, sender, wrap) = &*shared;
134                 handle_connection(stream, server_name, tools, &|name, arguments| {
135                     let (tx, rx) = mpsc::channel();
136                     let call = McpToolCall { name: name.to_string(), arguments, reply: tx };
137                     sender
138                         .send(wrap(call))
139                         .map_err(|_| "app event loop is gone".to_string())?;
140                     rx.recv_timeout(CALL_TIMEOUT)
141                         .map_err(|_| "timed out waiting for the app".to_string())?
142                 });
143             });
144         }
145     });
146 }
147 
148 fn handle_connection<F>(stream: TcpStream, server_name: &str, tools: &[McpTool], call_tool: &F)
149 where
150     F: Fn(&str, Value) -> Result<Value, String>,
151 {
152     let _ = stream.set_read_timeout(Some(READ_TIMEOUT));
153     let mut write_stream = match stream.try_clone() {
154         Ok(s) => s,
155         Err(_) => return,
156     };
157     // Bounded like the body: `read_line` would otherwise buffer one endless
158     // header line for as long as the timeout keeps being met.
159     let mut reader = BufReader::new(stream.take((MAX_BODY + MAX_HEAD) as u64));
160     let mut request_line = String::new();
161     if reader.read_line(&mut request_line).is_err() {
162         return;
163     }
164 
165     let mut content_length = 0usize;
166     let mut host = None;
167     let mut origin = None;
168     loop {
169         let mut line = String::new();
170         if reader.read_line(&mut line).is_err() || line == "\r\n" || line == "\n" || line.is_empty()
171         {
172             break;
173         }
174         let Some((name, value)) = line.split_once(':') else { continue };
175         let value = value.trim().to_string();
176         match name.trim().to_ascii_lowercase().as_str() {
177             "content-length" => {
178                 if let Ok(len) = value.parse::<usize>() {
179                     content_length = len;
180                 }
181             }
182             "host" => host = Some(value),
183             "origin" => origin = Some(value),
184             _ => {}
185         }
186     }
187 
188     if !request_from_this_machine(host.as_deref(), origin.as_deref()) {
189         let _ = write_stream.write_all(
190             b"HTTP/1.1 403 Forbidden\r\nContent-Length: 0\r\nConnection: close\r\n\r\n",
191         );
192         return;
193     }
194 
195     if content_length > MAX_BODY {
196         let _ = write_stream.write_all(
197             b"HTTP/1.1 413 Content Too Large\r\nContent-Length: 0\r\nConnection: close\r\n\r\n",
198         );
199         return;
200     }
201 
202     if !request_line.starts_with("POST ") {
203         // Stateless server: no SSE stream (GET) or session teardown (DELETE).
204         let _ = write_stream.write_all(
205             b"HTTP/1.1 405 Method Not Allowed\r\nAllow: POST\r\nContent-Length: 0\r\nConnection: close\r\n\r\n",
206         );
207         return;
208     }
209 
210     let mut body = vec![0; content_length];
211     if reader.read_exact(&mut body).is_err() {
212         return;
213     }
214 
215     let response = match serde_json::from_slice::<Value>(&body) {
216         Ok(req) => handle_jsonrpc(&req, server_name, tools, call_tool),
217         Err(_) => Some(error_response(Value::Null, -32700, "parse error")),
218     };
219     let http = match response {
220         Some(resp) => {
221             let body = resp.to_string();
222             format!(
223                 "HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}",
224                 body.len(),
225                 body
226             )
227         }
228         // Notifications get no JSON-RPC response, just an HTTP ack.
229         None => "HTTP/1.1 202 Accepted\r\nContent-Length: 0\r\nConnection: close\r\n\r\n".to_string(),
230     };
231     let _ = write_stream.write_all(http.as_bytes());
232     let _ = write_stream.flush();
233 }
234 
235 /// Dispatch one JSON-RPC message. Returns `None` for notifications (no id).
236 fn handle_jsonrpc<F>(req: &Value, server_name: &str, tools: &[McpTool], call_tool: &F) -> Option<Value>
237 where
238     F: Fn(&str, Value) -> Result<Value, String>,
239 {
240     let method = req.get("method").and_then(Value::as_str).unwrap_or("");
241     let id = match req.get("id") {
242         Some(id) if !id.is_null() => id.clone(),
243         _ => return None,
244     };
245 
246     let result = match method {
247         "initialize" => {
248             let requested = req
249                 .pointer("/params/protocolVersion")
250                 .and_then(Value::as_str)
251                 .unwrap_or("");
252             let version = if PROTOCOL_VERSIONS.contains(&requested) {
253                 requested
254             } else {
255                 PROTOCOL_VERSIONS[0]
256             };
257             json!({
258                 "protocolVersion": version,
259                 "capabilities": { "tools": {} },
260                 "serverInfo": { "name": server_name, "version": env!("CARGO_PKG_VERSION") },
261             })
262         }
263         "ping" => json!({}),
264         "tools/list" => json!({
265             "tools": tools.iter().map(|t| json!({
266                 "name": t.name,
267                 "description": t.description,
268                 "inputSchema": t.input_schema,
269             })).collect::<Vec<_>>(),
270         }),
271         "tools/call" => {
272             let name = req
273                 .pointer("/params/name")
274                 .and_then(Value::as_str)
275                 .unwrap_or("");
276             if !tools.iter().any(|t| t.name == name) {
277                 return Some(error_response(id, -32602, &format!("unknown tool: {name}")));
278             }
279             let arguments = req
280                 .pointer("/params/arguments")
281                 .cloned()
282                 .unwrap_or_else(|| json!({}));
283             match call_tool(name, arguments) {
284                 Ok(value) => {
285                     let text = match value {
286                         Value::String(s) => s,
287                         other => serde_json::to_string_pretty(&other).unwrap_or_default(),
288                     };
289                     json!({ "content": [{ "type": "text", "text": text }], "isError": false })
290                 }
291                 Err(e) => json!({ "content": [{ "type": "text", "text": e }], "isError": true }),
292             }
293         }
294         _ => return Some(error_response(id, -32601, &format!("method not found: {method}"))),
295     };
296     Some(json!({ "jsonrpc": "2.0", "id": id, "result": result }))
297 }
298 
299 fn error_response(id: Value, code: i64, message: &str) -> Value {
300     json!({ "jsonrpc": "2.0", "id": id, "error": { "code": code, "message": message } })
301 }
302 
303 #[cfg(test)]
304 mod tests {
305     use super::*;
306 
307     fn tools() -> Vec<McpTool> {
308         vec![McpTool {
309             name: "echo".to_string(),
310             description: "Echo the arguments back".to_string(),
311             input_schema: json!({ "type": "object" }),
312         }]
313     }
314 
315     fn no_calls(_: &str, _: Value) -> Result<Value, String> {
316         panic!("no tool call expected");
317     }
318 
319     #[test]
320     fn initialize_negotiates_protocol_version() {
321         let req = json!({
322             "jsonrpc": "2.0", "id": 1, "method": "initialize",
323             "params": { "protocolVersion": "2025-03-26" }
324         });
325         let resp = handle_jsonrpc(&req, "test", &tools(), &no_calls).unwrap();
326         assert_eq!(resp["result"]["protocolVersion"], "2025-03-26");
327         assert!(resp["result"]["capabilities"]["tools"].is_object());
328 
329         // Unknown requested version falls back to our newest.
330         let req = json!({
331             "jsonrpc": "2.0", "id": 2, "method": "initialize",
332             "params": { "protocolVersion": "2099-01-01" }
333         });
334         let resp = handle_jsonrpc(&req, "test", &tools(), &no_calls).unwrap();
335         assert_eq!(resp["result"]["protocolVersion"], PROTOCOL_VERSIONS[0]);
336     }
337 
338     #[test]
339     fn notifications_get_no_response() {
340         let req = json!({ "jsonrpc": "2.0", "method": "notifications/initialized" });
341         assert!(handle_jsonrpc(&req, "test", &tools(), &no_calls).is_none());
342     }
343 
344     #[test]
345     fn tools_list_reports_declared_tools() {
346         let req = json!({ "jsonrpc": "2.0", "id": 3, "method": "tools/list" });
347         let resp = handle_jsonrpc(&req, "test", &tools(), &no_calls).unwrap();
348         assert_eq!(resp["result"]["tools"][0]["name"], "echo");
349         assert!(resp["result"]["tools"][0]["inputSchema"].is_object());
350     }
351 
352     #[test]
353     fn tools_call_wraps_ok_and_err_results() {
354         let req = json!({
355             "jsonrpc": "2.0", "id": 4, "method": "tools/call",
356             "params": { "name": "echo", "arguments": { "x": 1 } }
357         });
358         let resp = handle_jsonrpc(&req, "test", &tools(), &|name, args| {
359             assert_eq!(name, "echo");
360             Ok(args)
361         })
362         .unwrap();
363         assert_eq!(resp["result"]["isError"], false);
364         assert!(resp["result"]["content"][0]["text"].as_str().unwrap().contains("\"x\": 1"));
365 
366         let resp = handle_jsonrpc(&req, "test", &tools(), &|_, _| Err("boom".to_string())).unwrap();
367         assert_eq!(resp["result"]["isError"], true);
368         assert_eq!(resp["result"]["content"][0]["text"], "boom");
369     }
370 
371     #[test]
372     fn unknown_tool_and_method_are_protocol_errors() {
373         let req = json!({
374             "jsonrpc": "2.0", "id": 5, "method": "tools/call",
375             "params": { "name": "nope" }
376         });
377         let resp = handle_jsonrpc(&req, "test", &tools(), &no_calls).unwrap();
378         assert_eq!(resp["error"]["code"], -32602);
379 
380         let req = json!({ "jsonrpc": "2.0", "id": 6, "method": "resources/list" });
381         let resp = handle_jsonrpc(&req, "test", &tools(), &no_calls).unwrap();
382         assert_eq!(resp["error"]["code"], -32601);
383     }
384 
385     #[test]
386     fn only_requests_from_this_machine_are_served() {
387         // An MCP client: no Origin, a loopback Host in any spelling.
388         for host in ["127.0.0.1:3002", "localhost:3002", "LOCALHOST", "[::1]:3002", "127.0.0.2:3002"] {
389             assert!(request_from_this_machine(Some(host), None), "{host}");
390         }
391         // A page on a loopback origin is this machine too.
392         assert!(request_from_this_machine(Some("127.0.0.1:3002"), Some("http://localhost:5173")));
393         assert!(request_from_this_machine(Some("127.0.0.1:3002"), Some("http://[::1]:8080")));
394 
395         // A web page posting cross-origin: the Host is ours, the Origin is not.
396         assert!(!request_from_this_machine(Some("127.0.0.1:3002"), Some("https://evil.example")));
397         assert!(!request_from_this_machine(Some("127.0.0.1:3002"), Some("null")));
398         // Look-alikes of loopback are not loopback.
399         assert!(!request_from_this_machine(Some("127.0.0.1:3002"), Some("http://localhost.evil.example")));
400         assert!(!request_from_this_machine(Some("127.0.0.1:3002"), Some("http://127.0.0.1.evil.example")));
401         // DNS rebinding: the page's own hostname arrives as the Host.
402         assert!(!request_from_this_machine(Some("rebind.evil.example:3002"), Some("http://rebind.evil.example:3002")));
403         assert!(!request_from_this_machine(Some("rebind.evil.example:3002"), None));
404         assert!(!request_from_this_machine(Some("localhost.evil.example"), None));
405         // No Host at all is not HTTP/1.1 from anyone we serve.
406         assert!(!request_from_this_machine(None, None));
407     }
408 
409     /// Run one raw request through `handle_connection` over a real socket.
410     fn exchange(request: &[u8]) -> String {
411         let listener = TcpListener::bind(("127.0.0.1", 0)).unwrap();
412         let addr = listener.local_addr().unwrap();
413         let server = std::thread::spawn(move || {
414             let (stream, _) = listener.accept().unwrap();
415             handle_connection(stream, "test", &tools(), &|_: &str, args: Value| Ok(args));
416         });
417         let mut client = TcpStream::connect(addr).unwrap();
418         client.write_all(request).unwrap();
419         let _ = client.shutdown(std::net::Shutdown::Write);
420         let mut reply = String::new();
421         let _ = client.read_to_string(&mut reply);
422         server.join().unwrap();
423         reply
424     }
425 
426     fn post(host: &str, extra: &str, body: &str) -> Vec<u8> {
427         format!(
428             "POST /mcp HTTP/1.1\r\nHost: {host}\r\n{extra}Content-Length: {}\r\n\r\n{body}",
429             body.len()
430         )
431         .into_bytes()
432     }
433 
434     #[test]
435     fn a_web_page_gets_403_and_its_tool_never_runs() {
436         let call = r#"{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"echo","arguments":{"x":1}}}"#;
437         let ok = exchange(&post("127.0.0.1:3002", "", call));
438         assert!(ok.starts_with("HTTP/1.1 200"), "{ok}");
439         assert!(ok.contains("isError\":false"), "{ok}");
440 
441         let csrf = exchange(&post("127.0.0.1:3002", "Origin: https://evil.example\r\nContent-Type: text/plain\r\n", call));
442         assert!(csrf.starts_with("HTTP/1.1 403"), "{csrf}");
443         let rebind = exchange(&post("rebind.evil.example:3002", "Origin: http://rebind.evil.example:3002\r\n", call));
444         assert!(rebind.starts_with("HTTP/1.1 403"), "{rebind}");
445     }
446 
447     #[test]
448     fn a_huge_content_length_is_refused_not_allocated() {
449         // Before the cap this allocated the claimed size, and a failed
450         // allocation aborts the process: this test would take the runner down.
451         let req = "POST /mcp HTTP/1.1\r\nHost: 127.0.0.1\r\nContent-Length: 9223372036854775807\r\n\r\n{}";
452         let reply = exchange(req.as_bytes());
453         assert!(reply.starts_with("HTTP/1.1 413"), "{reply}");
454     }
455 }