1use crate::*;
2use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade};
3use futures::SinkExt;
4use mcp::{delete, get, hash, list, put, secret};
5use rmcp::{
6 ErrorData, RoleServer, ServerHandler,
7 model::{
8 CallToolRequestParams, CallToolResponse, CallToolResult, ContentBlock, ListToolsResult,
9 PaginatedRequestParams, ServerCapabilities, ServerConfig, Tool, ToolAnnotations,
10 },
11 service::RequestContext,
12};
13use tokio::sync::{mpsc, oneshot};
14
15pub struct Broker {
16 connections: Mutex<HashMap<String, Arc<Connection>>>,
17 pending_slots: Semaphore,
18 pub(crate) changes: watch::Sender<()>,
19}
20impl Default for Broker {
21 fn default() -> Self {
22 Self {
23 connections: Mutex::new(HashMap::new()),
24 pending_slots: Semaphore::new(128),
25 changes: watch::channel(()).0,
26 }
27 }
28}
29struct Connection {
30 messages: mpsc::Sender<Message>,
31 pending: Mutex<HashMap<String, oneshot::Sender<Result<Value>>>>,
32 closed: watch::Sender<bool>,
33}
34struct Pending {
35 connection: Arc<Connection>,
36 id: String,
37}
38impl Drop for Pending {
39 fn drop(&mut self) {
40 self.connection.pending.lock().unwrap().remove(&self.id);
41 }
42}
43fn unknown() -> Error {
44 Error::new(
45 409,
46 "Command outcome is unknown. Read the thread before retrying.",
47 )
48}
49impl Broker {
50 pub fn disconnect(&self, id: &str) {
51 if let Some(connection) = self.connections.lock().unwrap().remove(id) {
52 self.changes.send_replace(());
53 connection.closed.send_replace(true);
54 for (_, response) in connection.pending.lock().unwrap().drain() {
55 let _ = response.send(Err(unknown()));
56 }
57 }
58 }
59 fn remove(&self, id: &str, connection: &Arc<Connection>) {
60 let mut connections = self.connections.lock().unwrap();
61 if connections
62 .get(id)
63 .is_some_and(|current| Arc::ptr_eq(current, connection))
64 {
65 connections.remove(id);
66 self.changes.send_replace(());
67 }
68 connection.closed.send_replace(true);
69 for (_, response) in connection.pending.lock().unwrap().drain() {
70 let _ = response.send(Err(unknown()));
71 }
72 }
73 pub fn view(&self, machines: Vec<Value>, grant: Option<&Value>) -> Vec<Value> {
74 let connections = self.connections.lock().unwrap();
75 machines
76 .into_iter()
77 .filter(|machine| grant.is_none_or(|g| array(&g["resources"]).contains(&machine["id"])))
78 .map(|mut machine| {
79 machine.as_object_mut().unwrap().remove("tokenHash");
80 machine.as_object_mut().unwrap().remove("user");
81 machine["online"] = json!(connections.contains_key(string(&machine["id"])));
82 if let Some(grant) = grant {
83 machine["selected"] = json!(array(&grant["targets"]).contains(&machine["id"]));
84 }
85 machine
86 })
87 .collect()
88 }
89 async fn dispatch(
90 &self,
91 store: &mcp::Store,
92 grant_id: &str,
93 machine_id: &str,
94 method: &str,
95 params: Value,
96 deadline: Duration,
97 ) -> Result<Value> {
98 let params = command(method, params)?;
99 let (pending, receive, slot) = {
100 let db = store.db.lock().unwrap();
101 let grant = get(&db, &format!("grant:{grant_id}"))?;
102 let machine = get(&db, &format!("machine:{machine_id}"))?;
103 if grant.is_null()
104 || grant["resource"] != store.resource("agents")
105 || machine.is_null()
106 || machine["user"] != grant["user"]
107 || !array(&grant["resources"]).iter().any(|id| id == machine_id)
108 {
109 return Err(Error::new(
110 403,
111 "Choose a machine granted to this connection.",
112 ));
113 }
114 let scope = if matches!(method, "send_message" | "interrupt_thread" | "start_thread") {
115 "sessions:write"
116 } else {
117 "sessions:read"
118 };
119 if !array(&grant["scopes"]).iter().any(|s| s == scope) {
120 return Err(Error::new(
121 403,
122 "This connection has read access only. Connect again to request control.",
123 ));
124 }
125 let slot = self.pending_slots.try_acquire().map_err(|_| {
126 Error::new(
127 429,
128 "The relay has too many pending commands. Try again shortly.",
129 )
130 })?;
131 let connection = self
132 .connections
133 .lock()
134 .unwrap()
135 .get(machine_id)
136 .cloned()
137 .ok_or_else(|| Error::new(409, "Machine is offline. Start its local agent."))?;
138 let (send, receive) = oneshot::channel();
139 let id = uuid::Uuid::new_v4().to_string();
140 {
141 let mut pending = connection.pending.lock().unwrap();
142 if *connection.closed.borrow() {
143 return Err(Error::new(
144 409,
145 "Machine is offline. Start its local agent.",
146 ));
147 }
148 if pending.len() >= 8 {
149 return Err(Error::new(
150 429,
151 "Machine has eight pending commands. Wait for one to finish.",
152 ));
153 }
154 pending.insert(id.clone(), send);
155 }
156 let pending = Pending {
157 connection: connection.clone(),
158 id: id.clone(),
159 };
160 let frame = json!({"id":id,"method":method,"params":params}).to_string();
161 if frame.len() > 256 * 1024 {
162 return Err(Error::new(
163 413,
164 "Send a shorter message. Local agent commands are limited to 256 KB.",
165 ));
166 }
167 let frame = Message::Text(frame.into());
168 if connection.messages.try_send(frame).is_err() {
169 return Err(unknown());
170 }
171 (pending, receive, slot)
172 };
173 let result = tokio::time::timeout(deadline, receive)
174 .await
175 .map_err(|_| unknown())?
176 .map_err(|_| unknown())?;
177 drop(pending);
178 drop(slot);
179 result
180 }
181}
182pub(crate) fn machines(db: &rusqlite::Connection, user: &str) -> Result<Vec<Value>> {
183 Ok(list(db, "machine:")?
184 .into_iter()
185 .filter(|m| m["user"] == user)
186 .collect())
187}
188fn device(db: &rusqlite::Connection, headers: &HeaderMap) -> Result<Value> {
189 let token = headers
190 .get("authorization")
191 .and_then(|v| v.to_str().ok())
192 .and_then(|s| s.strip_prefix("Bearer "))
193 .filter(|s| !s.is_empty() && s.len() <= 256)
194 .ok_or_else(|| Error::new(401, "Pair this machine again."))?;
195 let fingerprint = hash(token);
196 list(db, "machine:")?
197 .into_iter()
198 .find(|machine| {
199 string(&machine["tokenHash"])
200 .as_bytes()
201 .ct_eq(fingerprint.as_bytes())
202 .unwrap_u8()
203 == 1
204 })
205 .ok_or_else(|| Error::new(401, "Pair this machine again."))
206}
207fn name(value: &Value) -> Result<&str> {
208 value
209 .as_str()
210 .filter(|s| !s.trim().is_empty() && s.encode_utf16().count() <= 100)
211 .ok_or_else(|| Error::new(400, "Enter a name up to 100 characters."))
212}
213pub(crate) fn manage(
214 app: &App,
215 db: &rusqlite::Connection,
216 parts: &[&str],
217 method: &Method,
218 owner: &str,
219 body: &Value,
220) -> Result<Value> {
221 match parts {
222 ["pair"] if method == Method::POST => {
223 let code = body["code"]
224 .as_str()
225 .filter(|s| !s.is_empty() && s.len() <= 40)
226 .ok_or_else(|| Error::new(400, "Enter the code from your local agent."))?;
227 let code: String = code
228 .chars()
229 .filter(|c| *c != '-' && !c.is_whitespace())
230 .flat_map(char::to_uppercase)
231 .collect();
232 let key = format!("pair:{}", hash(&code));
233 let pairing = get(db, &key)?;
234 if pairing.is_null() {
235 return Err(Error::new(
236 410,
237 "This code expired or was used. Start pairing again.",
238 ));
239 }
240 if machines(db, owner)?.len() >= 128 {
241 return Err(Error::new(
242 409,
243 "Unlink an unused machine before adding another.",
244 ));
245 }
246 let id = uuid::Uuid::new_v4().to_string();
247 let machine = json!({"id":id,"user":owner,"name":pairing["name"],"platform":pairing["platform"],"tokenHash":pairing["tokenHash"]});
248 put(db, &format!("machine:{id}"), &machine, 0)?;
249 delete(db, &key)?;
250 Ok(json!(app.relay.view(vec![machine], None).remove(0)))
251 }
252 ["machines", id] if method == Method::DELETE || method == Method::PATCH => {
253 let mut machine = get(db, &format!("machine:{id}"))?;
254 if machine["user"] != owner {
255 return Err(Error::new(404, "No linked machine with that ID."));
256 }
257 if method == Method::DELETE {
258 delete(db, &format!("machine:{id}"))?;
259 app.relay.disconnect(id);
260 Ok(Value::Null)
261 } else {
262 machine["name"] = json!(name(&body["name"])?);
263 put(db, &format!("machine:{id}"), &machine, 0)?;
264 Ok(json!(app.relay.view(vec![machine], None).remove(0)))
265 }
266 }
267 ["keys"] if method == Method::POST => {
268 let available = machines(db, owner)?;
269 let resources = selection(
270 &body["resources"],
271 &available
272 .iter()
273 .map(|m| m["id"].clone())
274 .collect::<Vec<_>>(),
275 )?;
276 let name = name(&body["name"])?;
277 let write = body
278 .get("write")
279 .map(|v| {
280 v.as_bool()
281 .ok_or_else(|| Error::new(400, "Choose read access or control."))
282 })
283 .transpose()?
284 .unwrap_or(false);
285 if list(db, "grant:")?
286 .iter()
287 .filter(|g| g["user"] == owner)
288 .count()
289 >= 256
290 {
291 return Err(Error::new(
292 409,
293 "Revoke an unused connection before adding another.",
294 ));
295 }
296 let id = uuid::Uuid::new_v4().to_string();
297 let grant = json!({"id":id,"user":owner,"name":name,"client":null,"resource":app.mcp.resource("agents"),"scopes":if write {json!(["sessions:read","sessions:write"])} else {json!(["sessions:read"])},"resources":resources,"targets":resources,"createdAt":now()});
298 let key = format!("ar_{}", secret());
299 put(db, &format!("grant:{id}"), &grant, 0)?;
300 put(
301 db,
302 &format!("access:{}", hash(&key)),
303 &json!({"grant":id,"resource":grant["resource"]}),
304 0,
305 )?;
306 Ok(json!({"key":key,"id":id}))
307 }
308 _ => Err(Error::new(404, "No endpoint here.")),
309 }
310}
311fn selection(value: &Value, allowed: &[Value]) -> Result<Vec<Value>> {
312 let ids = value
313 .as_array()
314 .filter(|ids| !ids.is_empty() && ids.len() <= 128)
315 .ok_or_else(|| Error::new(400, "Choose at least one linked machine."))?;
316 if ids
317 .iter()
318 .enumerate()
319 .any(|(i, id)| !allowed.contains(id) || ids[..i].contains(id))
320 {
321 return Err(Error::new(
322 403,
323 "Choose machines granted to this connection once each.",
324 ));
325 }
326 Ok(ids.clone())
327}
328fn view(app: &App, grant_id: &str, targets: Option<&Value>) -> Result<Vec<Value>> {
329 let mut db = app.mcp.db.lock().unwrap();
330 let tx = db.transaction()?;
331 let mut grant = get(&tx, &format!("grant:{grant_id}"))?;
332 if grant.is_null() || grant["resource"] != app.mcp.resource("agents") {
333 return Err(Error::new(
334 401,
335 "This connection was revoked. Connect again.",
336 ));
337 }
338 let machines = machines(&tx, string(&grant["user"]))?;
339 if let Some(targets) = targets {
340 let allowed: Vec<_> = machines
341 .iter()
342 .filter(|m| array(&grant["resources"]).contains(&m["id"]))
343 .map(|m| m["id"].clone())
344 .collect();
345 grant["targets"] = json!(selection(targets, &allowed)?);
346 put(&tx, &format!("grant:{grant_id}"), &grant, 0)?;
347 }
348 let result = app.relay.view(machines, Some(&grant));
349 tx.commit()?;
350 Ok(result)
351}
352fn schema(method: &str) -> Option<Value> {
353 let provider = json!({"type":"string","enum":["codex","claude"]});
354 let id = json!({"type":"string","format":"uuid"});
355 let message = json!({"type":"string","minLength":1,"maxLength":100000});
356 let (properties, required) = match method {
357 "list_threads" => (
358 json!({"provider":provider,"limit":{"type":"integer","minimum":1,"maximum":100,"default":30}}),
359 json!([]),
360 ),
361 "read_thread" => (
362 json!({"provider":provider,"thread_id":id,"limit":{"type":"integer","minimum":1,"maximum":100,"default":20}}),
363 json!(["provider", "thread_id"]),
364 ),
365 "send_message" => (
366 json!({"provider":provider,"thread_id":id,"message":message,"expected_turn_id":id}),
367 json!(["provider", "thread_id", "message"]),
368 ),
369 "interrupt_thread" => (
370 json!({"provider":provider,"thread_id":id,"expected_turn_id":id}),
371 json!(["provider", "thread_id"]),
372 ),
373 "start_thread" => (
374 json!({"provider":provider,"cwd":{"type":"string","minLength":1,"maxLength":4096},"message":message}),
375 json!(["provider", "cwd", "message"]),
376 ),
377 _ => return None,
378 };
379 Some(
380 json!({"type":"object","properties":properties,"required":required,"additionalProperties":false}),
381 )
382}
383fn command(method: &str, value: Value) -> Result<Value> {
384 let schema = schema(method)
385 .ok_or_else(|| Error::new(404, "Choose a session tool listed by this connector."))?;
386 let mut params = value
387 .as_object()
388 .cloned()
389 .ok_or_else(|| Error::new(400, "Use the fields listed for this tool."))?;
390 let properties = schema["properties"].as_object().unwrap();
391 if params.keys().any(|key| !properties.contains_key(key))
392 || array(&schema["required"])
393 .iter()
394 .any(|key| !params.contains_key(string(key)))
395 {
396 return Err(Error::new(400, "Use the fields listed for this tool."));
397 }
398 for (key, rule) in properties {
399 if !params.contains_key(key) && rule.get("default").is_some() {
400 params.insert(key.clone(), rule["default"].clone());
401 }
402 let Some(value) = params.get(key) else {
403 continue;
404 };
405 let valid = match string(&rule["type"]) {
406 "integer" => value.as_i64().is_some_and(|n| {
407 n >= rule["minimum"].as_i64().unwrap() && n <= rule["maximum"].as_i64().unwrap()
408 }),
409 "string" => value.as_str().is_some_and(|s| {
410 let length = s.encode_utf16().count() as u64;
411 rule["minLength"].as_u64().is_none_or(|n| length >= n)
412 && rule["maxLength"].as_u64().is_none_or(|n| length <= n)
413 && rule
414 .get("enum")
415 .is_none_or(|values| array(values).contains(value))
416 && (rule["format"] != "uuid"
417 || uuid::Uuid::parse_str(s)
418 .is_ok_and(|id| id.to_string().eq_ignore_ascii_case(s)))
419 }),
420 _ => false,
421 };
422 if !valid {
423 return Err(Error::new(
424 400,
425 format!("Check {key} against the tool's fields."),
426 ));
427 }
428 }
429 Ok(json!(params))
430}
431async fn pairing(State(app): State<Arc<App>>, request: Request) -> Result<Response> {
432 let method = request.method().clone();
433 let headers = request.headers().clone();
434 let bytes = axum::body::to_bytes(request.into_body(), 4096)
435 .await
436 .map_err(|_| Error::new(400, "Enter a shorter machine name."))?;
437 let mut db = app.mcp.db.lock().unwrap();
438 let tx = db.transaction()?;
439 let response = if method == Method::POST {
440 if list(&tx, "pair:")?.len() >= 256 {
441 return Err(Error::new(
442 429,
443 "Too many pairing requests. Try again in ten minutes.",
444 ));
445 }
446 let body: Value = serde_json::from_slice(&bytes)
447 .map_err(|_| Error::new(400, "Enter a machine name and platform."))?;
448 let name = name(&body["name"])?;
449 let platform = body["platform"]
450 .as_str()
451 .filter(|s| s.encode_utf16().count() <= 100)
452 .ok_or_else(|| Error::new(400, "Enter a platform up to 100 characters."))?;
453 let token = secret();
454 let code: String = rand::random::<[u8; 5]>()
455 .iter()
456 .map(|b| format!("{b:02X}"))
457 .collect();
458 put(
459 &tx,
460 &format!("pair:{}", hash(&code)),
461 &json!({"name":name,"platform":platform,"tokenHash":hash(&token)}),
462 600,
463 )?;
464 (StatusCode::CREATED, axum::Json(json!({"code":format!("{}-{}", &code[..5], &code[5..]),"token":token,"expires_in":600}))).into_response()
465 } else if method == Method::GET {
466 match device(&tx, &headers) {
467 Ok(machine) => axum::Json(json!({"machine_id":machine["id"]})).into_response(),
468 Err(_) => {
469 let token = headers
470 .get("authorization")
471 .and_then(|v| v.to_str().ok())
472 .and_then(|s| s.strip_prefix("Bearer "))
473 .filter(|s| !s.is_empty() && s.len() <= 256)
474 .ok_or_else(|| Error::new(401, "Start pairing from your local agent."))?;
475 if !list(&tx, "pair:")?
476 .iter()
477 .any(|p| p["tokenHash"] == hash(token))
478 {
479 return Err(Error::new(
480 410,
481 "This pairing expired. Start pairing again.",
482 ));
483 }
484 (StatusCode::ACCEPTED, axum::Json(json!({"pending":true}))).into_response()
485 }
486 }
487 } else {
488 return Err(Error::new(405, "Use GET or POST for pairing."));
489 };
490 tx.commit()?;
491 Ok(response)
492}
493async fn connect(
494 State(app): State<Arc<App>>,
495 headers: HeaderMap,
496 ws: WebSocketUpgrade,
497) -> Result<Response> {
498 if headers.contains_key("origin") {
499 return Err(Error::new(401, "Connect from your local agent."));
500 }
501 let machine = device(&app.mcp.db.lock().unwrap(), &headers)?;
502 let id = string(&machine["id"]).to_owned();
503 let (send, receive) = mpsc::channel(16);
504 let (closed, _) = watch::channel(false);
505 let connection = Arc::new(Connection {
506 messages: send,
507 pending: Mutex::new(HashMap::new()),
508 closed,
509 });
510 {
511 let mut connections = app.relay.connections.lock().unwrap();
512 if connections.contains_key(&id) {
513 return Err(Error::new(
514 409,
515 "Another agent is connected. Stop it before starting a second copy.",
516 ));
517 }
518 if connections.len() >= 256 {
519 return Err(Error::new(
520 503,
521 "Too many agents are connected. Try again later.",
522 ));
523 }
524 connections.insert(id.clone(), connection.clone());
525 app.relay.changes.send_replace(());
526 }
527 let (failure_app, failure_id, failure_connection) =
528 (app.clone(), id.clone(), connection.clone());
529 Ok(ws
530 .max_frame_size(4 * 1024 * 1024)
531 .max_message_size(4 * 1024 * 1024)
532 .read_buffer_size(16 * 1024)
533 .write_buffer_size(0)
534 .max_write_buffer_size(256 * 1024)
535 .on_failed_upgrade(move |_| failure_app.relay.remove(&failure_id, &failure_connection))
536 .on_upgrade(move |socket| session(app, id, connection, receive, socket)))
537}
538async fn session(
539 app: Arc<App>,
540 id: String,
541 connection: Arc<Connection>,
542 mut messages: mpsc::Receiver<Message>,
543 mut socket: WebSocket,
544) {
545 let mut closed = connection.closed.subscribe();
546 let connected = Message::Text(
547 json!({"type":"connected","machine_id":id})
548 .to_string()
549 .into(),
550 );
551 if tokio::time::timeout(Duration::from_secs(5), socket.send(connected))
552 .await
553 .is_ok_and(|r| r.is_ok())
554 {
555 let mut heartbeat = tokio::time::interval_at(
556 tokio::time::Instant::now() + Duration::from_secs(20),
557 Duration::from_secs(20),
558 );
559 heartbeat.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
560 let mut alive = true;
561 loop {
562 if *closed.borrow()
563 || !get(&app.mcp.db.lock().unwrap(), &format!("machine:{id}"))
564 .is_ok_and(|m| !m.is_null())
565 {
566 break;
567 }
568 let outgoing = tokio::select! {
569 _ = closed.changed() => break,
570 outgoing = messages.recv() => match outgoing {Some(frame) => frame, None => break},
571 _ = heartbeat.tick() => {
572 if !alive {break;}
573 alive = false;
574 Message::Ping(Bytes::new())
575 },
576 incoming = socket.recv() => {
577 match incoming {
578 Some(Ok(Message::Pong(_))) => alive = true,
579 Some(Ok(Message::Ping(bytes))) => {
580 if !tokio::time::timeout(Duration::from_secs(5), socket.send(Message::Pong(bytes))).await.is_ok_and(|r| r.is_ok()) {break;}
581 },
582 Some(Ok(Message::Text(text))) => {
583 let reply = serde_json::from_str::<Value>(&text);
584 let Ok(reply) = reply else {break};
585 let valid = reply.as_object().is_some_and(|fields| fields.keys().all(|k| ["id","result","error"].contains(&k.as_str())))
586 && reply["id"].as_str().is_some_and(|s| uuid::Uuid::parse_str(s).is_ok())
587 && (reply.get("result").is_some() != reply.get("error").is_some())
588 && reply.get("error").is_none_or(|error| error.as_str().is_some_and(|s| s.encode_utf16().count() <= 2000));
589 if !valid {break;}
590 if let Some(send) = connection.pending.lock().unwrap().remove(string(&reply["id"])) {
591 let result = if let Some(error) = reply["error"].as_str() {Err(Error::new(400, error))} else {Ok(reply["result"].clone())};
592 let _ = send.send(result);
593 }
594 },
595 _ => break,
596 }
597 continue;
598 }
599 };
600 if !tokio::time::timeout(Duration::from_secs(5), socket.send(outgoing))
601 .await
602 .is_ok_and(|r| r.is_ok())
603 {
604 break;
605 }
606 }
607 }
608 app.relay.remove(&id, &connection);
609 let _ = tokio::time::timeout(Duration::from_secs(1), socket.close()).await;
610}
611async fn rest(State(app): State<Arc<App>>, request: Request) -> Response {
612 let result: Result<Response> = async {
613 if request.headers().get("origin").is_some_and(|v| {
614 v.to_str().ok() != Some(app.mcp.origin.origin().ascii_serialization().as_str())
615 }) {
616 return Err(Error::new(403, "Use the relay from its own origin."));
617 }
618 let grant = app
619 .mcp
620 .authenticate(request.headers(), &app.mcp.resource("agents"))?;
621 let id = string(&grant["id"]);
622 let method = request.method().clone();
623 let path = request
624 .uri()
625 .path()
626 .trim_start_matches("/api/v1/")
627 .to_owned();
628 let bytes = axum::body::to_bytes(request.into_body(), 256 * 1024)
629 .await
630 .map_err(|_| Error::new(400, "Send a shorter command."))?;
631 let body: Value = if bytes.is_empty() {
632 Value::Null
633 } else {
634 serde_json::from_slice(&bytes).map_err(|_| Error::new(400, "Use JSON command fields."))?
635 };
636 let value = match path.split('/').collect::<Vec<_>>().as_slice() {
637 ["machines"] if method == Method::GET => json!(view(&app, id, None)?),
638 ["targets"] if method == Method::PUT => json!(view(&app, id, Some(&body["machine_ids"]))?),
639 ["machines", machine, "commands"] if method == Method::POST => {
640 let method = body["method"]
641 .as_str()
642 .ok_or_else(|| Error::new(400, "Choose a session command."))?;
643 json!({"result":app.relay.dispatch(&app.mcp, id, machine, method, body["params"].clone(), Duration::from_secs(30)).await?})
644 }
645 _ => return Err(Error::new(404, "No relay endpoint here.")),
646 };
647 Ok(axum::Json(value).into_response())
648 }.await;
649 match result {
650 Ok(response) => response,
651 Err(error) => (StatusCode::from_u16(error.status).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR), axum::Json(json!({"error":if error.status >= 500 {"The relay couldn't answer. Check its dashboard and retry."} else {&error.message}}))).into_response(),
652 }
653}
654#[derive(Clone)]
655struct Agents(Arc<App>);
656impl ServerHandler for Agents {
657 fn get_info(&self) -> ServerConfig {
658 ServerConfig::new(ServerCapabilities::builder().enable_tools().build())
659 }
660 async fn list_tools(
661 &self,
662 _: Option<PaginatedRequestParams>,
663 _: RequestContext<RoleServer>,
664 ) -> std::result::Result<ListToolsResult, ErrorData> {
665 let tools = [
666 ("list_machines", "List granted machines and their connection state."),
667 ("set_target_machines", "Select default machines for this connection."),
668 ("list_threads", "List recent Codex and Claude Code threads on selected machines."),
669 ("read_thread", "Read a thread and its current status."),
670 ("send_message", "Submit a message. Submission acknowledges delivery, not task completion. Desktop control requires local opt-in."),
671 ("interrupt_thread", "Interrupt an agent-owned session or an opted-in Codex desktop turn."),
672 ("start_thread", "Start an agent-owned session under a locally allowed directory."),
673 ].into_iter().map(|(name, description)| {
674 let mut schema = schema(name).unwrap_or_else(|| if name == "set_target_machines" {json!({"type":"object","properties":{"machine_ids":{"type":"array","items":{"type":"string","format":"uuid"},"minItems":1,"maxItems":128,"uniqueItems":true}},"required":["machine_ids"],"additionalProperties":false})} else {json!({"type":"object","properties":{},"additionalProperties":false})});
675 if !["list_machines","set_target_machines"].contains(&name) {schema["properties"]["machine_id"] = json!({"type":"string","format":"uuid"});}
676 let read = ["list_machines","list_threads","read_thread"].contains(&name);
677 Tool::new(name, description, schema.as_object().unwrap().clone()).with_annotations(ToolAnnotations::new().read_only(read).destructive(name == "interrupt_thread").idempotent(read))
678 }).collect();
679 Ok(ListToolsResult {
680 tools,
681 ..Default::default()
682 })
683 }
684 async fn call_tool(
685 &self,
686 request: CallToolRequestParams,
687 context: RequestContext<RoleServer>,
688 ) -> std::result::Result<CallToolResponse, ErrorData> {
689 let result: Result<CallToolResult> = async {
690 let grant = &context
691 .extensions
692 .get::<axum::http::request::Parts>()
693 .and_then(|parts| parts.extensions.get::<mcp::Grant>())
694 .ok_or_else(|| Error::new(401, "This connection expired. Connect again."))?
695 .0;
696 let mut arguments = request.arguments.unwrap_or_default();
697 let id = string(&grant["id"]);
698 match request.name.as_ref() {
699 "list_machines" if arguments.is_empty() => Ok(CallToolResult::structured(
700 json!({"machines":view(&self.0, id, None)?}),
701 )),
702 "set_target_machines"
703 if arguments.len() == 1 && arguments.contains_key("machine_ids") =>
704 {
705 Ok(CallToolResult::structured(
706 json!({"machines":view(&self.0, id, arguments.get("machine_ids"))?}),
707 ))
708 }
709 method => {
710 let explicit = arguments.remove("machine_id");
711 let machines = view(&self.0, id, None)?;
712 let targets: Vec<_> = if let Some(explicit) = explicit {
713 let explicit = explicit
714 .as_str()
715 .ok_or_else(|| Error::new(400, "Choose a machine ID."))?;
716 if !machines.iter().any(|m| m["id"] == explicit) {
717 return Err(Error::new(403, "Choose a granted machine."));
718 }
719 vec![explicit.to_owned()]
720 } else {
721 machines
722 .iter()
723 .filter(|m| m["selected"] == true)
724 .map(|m| string(&m["id"]).to_owned())
725 .collect()
726 };
727 if targets.is_empty() || method != "list_threads" && targets.len() != 1 {
728 return Err(Error::new(
729 400,
730 "Select one machine or supply machine_id for this command.",
731 ));
732 }
733 let params = command(method, json!(arguments))?;
734 let mut replies: futures::stream::FuturesUnordered<_> = targets.into_iter().map(|machine| {
735 let params = params.clone();
736 async move {
737 match self
738 .0
739 .relay
740 .dispatch(
741 &self.0.mcp,
742 id,
743 &machine,
744 method,
745 params,
746 Duration::from_secs(30),
747 )
748 .await
749 {
750 Ok(result) => json!({"machine_id":machine,"result":result}),
751 Err(error) => json!({"machine_id":machine,"error":error.message}),
752 }
753 }
754 }).collect();
755 let mut results = Vec::new();
756 let mut bytes = 0;
757 while let Some(reply) = futures::StreamExt::next(&mut replies).await {
758 bytes += reply.to_string().len();
759 if bytes > 16*1024*1024 {return Err(Error::new(413, "Machine replies exceed 16 MB. Select fewer machines or request fewer messages."));}
760 results.push(reply);
761 }
762 let errors = results.iter().any(|r| r.get("error").is_some());
763 let mut result = CallToolResult::structured(json!({"results":results}));
764 result.is_error = Some(errors);
765 Ok(result)
766 }
767 }
768 }
769 .await;
770 Ok(match result {
771 Ok(result) => result,
772 Err(error) => CallToolResult::error(vec![ContentBlock::text(if error.status >= 500 {
773 "The relay couldn't answer. Check its dashboard and retry.".to_owned()
774 } else {
775 error.message
776 })]),
777 }
778 .into())
779 }
780}
781pub fn router(app: Arc<App>) -> Router {
782 let state = app.clone();
783 let expected_host =
784 app.mcp.origin[url::Position::BeforeHost..url::Position::AfterPort].to_owned();
785 Router::new()
786 .route("/pairing", any(pairing))
787 .route("/agent/connect", any(connect))
788 .route("/api/v1/{*path}", any(rest))
789 .with_state(app.clone())
790 .merge(mcp::router(app, "agents", move || {
791 Ok(Agents(state.clone()))
792 }))
793 .layer(axum::middleware::from_fn(
794 move |request: Request, next: axum::middleware::Next| {
795 let host = expected_host.clone();
796 async move {
797 if request.headers().get("host").and_then(|v| v.to_str().ok()) != Some(&host) {
798 return StatusCode::MISDIRECTED_REQUEST.into_response();
799 }
800 next.run(request).await
801 }
802 },
803 ))
804}
805
806#[cfg(test)]
807mod tests {
808 use super::*;
809 struct Fixture {
810 store: Arc<mcp::Store>,
811 broker: Arc<Broker>,
812 path: PathBuf,
813 machine: String,
814 grant: String,
815 }
816 impl Fixture {
817 fn new(write: bool) -> Self {
818 let path = std::env::temp_dir()
819 .canonicalize()
820 .unwrap()
821 .join(format!("studio-relay-test-{}", uuid::Uuid::new_v4()));
822 let store = Arc::new(mcp::Store::new(&path, "https://globe.studio.test").unwrap());
823 let machine = uuid::Uuid::new_v4().to_string();
824 let grant = uuid::Uuid::new_v4().to_string();
825 {
826 let db = store.db.lock().unwrap();
827 put(&db, &format!("machine:{machine}"), &json!({"id":machine,"user":"owner","name":"Fixture","tokenHash":hash("device-secret")}), 0).unwrap();
828 put(&db, &format!("grant:{grant}"), &json!({"id":grant,"user":"owner","resource":store.resource("agents"),"resources":[machine],"scopes":if write {json!(["sessions:read","sessions:write"])} else {json!(["sessions:read"])}}), 0).unwrap();
829 }
830 Self {
831 store,
832 broker: Arc::new(Broker::default()),
833 path,
834 machine,
835 grant,
836 }
837 }
838 fn connect(&self) -> (Arc<Connection>, mpsc::Receiver<Message>) {
839 let (messages, receive) = mpsc::channel(16);
840 let (closed, _) = watch::channel(false);
841 let connection = Arc::new(Connection {
842 messages,
843 closed,
844 pending: Mutex::new(HashMap::new()),
845 });
846 self.broker
847 .connections
848 .lock()
849 .unwrap()
850 .insert(self.machine.clone(), connection.clone());
851 (connection, receive)
852 }
853 fn dispatch(&self, duration: Duration) -> tokio::task::JoinHandle<Result<Value>> {
854 let (store, broker, machine, grant) = (
855 self.store.clone(),
856 self.broker.clone(),
857 self.machine.clone(),
858 self.grant.clone(),
859 );
860 tokio::spawn(async move {
861 broker
862 .dispatch(
863 &store,
864 &grant,
865 &machine,
866 "list_threads",
867 json!({}),
868 duration,
869 )
870 .await
871 })
872 }
873 }
874 impl Drop for Fixture {
875 fn drop(&mut self) {
876 std::fs::remove_dir_all(&self.path).unwrap();
877 }
878 }
879 #[test]
880 fn command_fields_defaults_and_provider_boundaries() {
881 assert_eq!(
882 command("list_threads", json!({})).unwrap(),
883 json!({"limit":30})
884 );
885 let id = uuid::Uuid::new_v4().to_string();
886 assert_eq!(
887 command("read_thread", json!({"provider":"claude","thread_id":id})).unwrap()["limit"],
888 20
889 );
890 for (method, params) in [
891 (
892 "read_thread",
893 json!({"provider":"codex","thread_id":"not-a-uuid"}),
894 ),
895 (
896 "read_thread",
897 json!({"provider":"native-chat","thread_id":id}),
898 ),
899 ("list_threads", json!({"limit":101})),
900 ("list_threads", json!({"limit":1.5})),
901 ("list_threads", json!({"approval_response":true})),
902 (
903 "send_message",
904 json!({"provider":"codex","thread_id":id,"message":""}),
905 ),
906 ("start_thread", json!({"provider":"claude","cwd":"/tmp"})),
907 (
908 "interrupt_thread",
909 json!({"provider":"codex","thread_id":id,"expected_turn_id":1}),
910 ),
911 ("approve_tool", json!({})),
912 ] {
913 assert!(command(method, params).is_err(), "{method}");
914 }
915 }
916 #[tokio::test]
917 async fn dispatch_uses_current_owner_grant_and_control_scope() {
918 let fixture = Fixture::new(false);
919 let (connection, mut receive) = fixture.connect();
920 let task = fixture.dispatch(Duration::from_secs(2));
921 let Message::Text(frame) = receive.recv().await.unwrap() else {
922 panic!()
923 };
924 let frame: Value = serde_json::from_str(&frame).unwrap();
925 assert_eq!(frame["params"], json!({"limit":30}));
926 connection
927 .pending
928 .lock()
929 .unwrap()
930 .remove(string(&frame["id"]))
931 .unwrap()
932 .send(Ok(json!({"threads":[]})))
933 .unwrap();
934 assert_eq!(task.await.unwrap().unwrap(), json!({"threads":[]}));
935 let refused = fixture
936 .broker
937 .dispatch(
938 &fixture.store,
939 &fixture.grant,
940 &fixture.machine,
941 "start_thread",
942 json!({"provider":"codex","cwd":"/tmp","message":"owned fixture"}),
943 Duration::from_secs(2),
944 )
945 .await
946 .unwrap_err();
947 assert_eq!(refused.status, 403);
948 assert!(receive.try_recv().is_err());
949 let foreign = uuid::Uuid::new_v4().to_string();
950 {
951 let db = fixture.store.db.lock().unwrap();
952 let mut grant = get(&db, &format!("grant:{}", fixture.grant)).unwrap();
953 grant["resources"] = json!([foreign]);
954 put(&db, &format!("grant:{}", fixture.grant), &grant, 0).unwrap();
955 put(
956 &db,
957 &format!("machine:{foreign}"),
958 &json!({"id":foreign,"user":"someone-else"}),
959 0,
960 )
961 .unwrap();
962 }
963 assert_eq!(
964 fixture
965 .broker
966 .dispatch(
967 &fixture.store,
968 &fixture.grant,
969 &foreign,
970 "list_threads",
971 json!({}),
972 Duration::from_secs(2)
973 )
974 .await
975 .unwrap_err()
976 .status,
977 403
978 );
979 delete(
980 &fixture.store.db.lock().unwrap(),
981 &format!("grant:{}", fixture.grant),
982 )
983 .unwrap();
984 assert_eq!(
985 fixture
986 .dispatch(Duration::from_secs(2))
987 .await
988 .unwrap()
989 .unwrap_err()
990 .status,
991 403
992 );
993 }
994 #[tokio::test]
995 async fn unicode_message_cannot_exceed_local_agent_frame_budget() {
996 let fixture = Fixture::new(true);
997 let (connection, mut receive) = fixture.connect();
998 let error = fixture.broker.dispatch(&fixture.store, &fixture.grant, &fixture.machine, "send_message", json!({"provider":"codex","thread_id":uuid::Uuid::new_v4().to_string(),"message":"雪".repeat(100000)}), Duration::from_secs(2)).await.unwrap_err();
999 assert_eq!(error.status, 413);
1000 assert!(receive.try_recv().is_err());
1001 assert!(connection.pending.lock().unwrap().is_empty());
1002 assert_eq!(fixture.broker.pending_slots.available_permits(), 128);
1003 }
1004 #[tokio::test]
1005 async fn eight_pending_limit_and_cancellation_release_capacity() {
1006 let fixture = Fixture::new(true);
1007 let (connection, mut receive) = fixture.connect();
1008 let mut tasks = Vec::new();
1009 for _ in 0..8 {
1010 tasks.push(fixture.dispatch(Duration::from_secs(2)));
1011 receive.recv().await.unwrap();
1012 }
1013 assert_eq!(
1014 fixture
1015 .dispatch(Duration::from_secs(2))
1016 .await
1017 .unwrap()
1018 .unwrap_err()
1019 .status,
1020 429
1021 );
1022 tasks.remove(0).abort();
1023 tokio::task::yield_now().await;
1024 assert_eq!(connection.pending.lock().unwrap().len(), 7);
1025 tasks.push(fixture.dispatch(Duration::from_secs(2)));
1026 receive.recv().await.unwrap();
1027 assert_eq!(connection.pending.lock().unwrap().len(), 8);
1028 fixture.broker.disconnect(&fixture.machine);
1029 for task in tasks {
1030 assert!(task.await.unwrap().unwrap_err().message.contains("unknown"));
1031 }
1032 assert!(connection.pending.lock().unwrap().is_empty());
1033 }
1034 #[tokio::test]
1035 async fn timeout_disconnect_and_reconnect_do_not_replay() {
1036 let fixture = Fixture::new(true);
1037 let (old, mut receive) = fixture.connect();
1038 let task = fixture.dispatch(Duration::from_millis(5));
1039 receive.recv().await.unwrap();
1040 assert!(task.await.unwrap().unwrap_err().message.contains("unknown"));
1041 assert!(old.pending.lock().unwrap().is_empty());
1042 assert!(receive.try_recv().is_err());
1043 let task = fixture.dispatch(Duration::from_secs(2));
1044 receive.recv().await.unwrap();
1045 fixture.broker.disconnect(&fixture.machine);
1046 assert!(task.await.unwrap().unwrap_err().message.contains("unknown"));
1047 let (new, mut receive) = fixture.connect();
1048 fixture.broker.remove(&fixture.machine, &old);
1049 assert!(Arc::ptr_eq(
1050 fixture
1051 .broker
1052 .connections
1053 .lock()
1054 .unwrap()
1055 .get(&fixture.machine)
1056 .unwrap(),
1057 &new
1058 ));
1059 assert!(receive.try_recv().is_err());
1060 let task = fixture.dispatch(Duration::from_secs(2));
1061 receive.recv().await.unwrap();
1062 fixture.broker.disconnect(&fixture.machine);
1063 assert!(task.await.unwrap().unwrap_err().message.contains("unknown"));
1064 }
1065}