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,"cursor":{"type":"string","minLength":1,"maxLength":512},"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 if !mcp::active_owner(&app, &grant)? {
622 return Err(Error::new(401, "This connection's account is no longer authorized."));
623 }
624 let id = string(&grant["id"]);
625 let method = request.method().clone();
626 let path = request
627 .uri()
628 .path()
629 .trim_start_matches("/api/v1/")
630 .to_owned();
631 let bytes = axum::body::to_bytes(request.into_body(), 256 * 1024)
632 .await
633 .map_err(|_| Error::new(400, "Send a shorter command."))?;
634 let body: Value = if bytes.is_empty() {
635 Value::Null
636 } else {
637 serde_json::from_slice(&bytes).map_err(|_| Error::new(400, "Use JSON command fields."))?
638 };
639 let value = match path.split('/').collect::<Vec<_>>().as_slice() {
640 ["machines"] if method == Method::GET => json!(view(&app, id, None)?),
641 ["targets"] if method == Method::PUT => json!(view(&app, id, Some(&body["machine_ids"]))?),
642 ["machines", machine, "commands"] if method == Method::POST => {
643 let method = body["method"]
644 .as_str()
645 .ok_or_else(|| Error::new(400, "Choose a session command."))?;
646 json!({"result":app.relay.dispatch(&app.mcp, id, machine, method, body["params"].clone(), Duration::from_secs(30)).await?})
647 }
648 _ => return Err(Error::new(404, "No relay endpoint here.")),
649 };
650 Ok(axum::Json(value).into_response())
651 }.await;
652 match result {
653 Ok(response) => response,
654 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(),
655 }
656}
657#[derive(Clone)]
658struct Agents(Arc<App>);
659impl ServerHandler for Agents {
660 fn get_info(&self) -> ServerConfig {
661 ServerConfig::new(ServerCapabilities::builder().enable_tools().build())
662 }
663 async fn list_tools(
664 &self,
665 _: Option<PaginatedRequestParams>,
666 _: RequestContext<RoleServer>,
667 ) -> std::result::Result<ListToolsResult, ErrorData> {
668 let tools = [
669 ("list_machines", "List granted machines and their connection state."),
670 ("set_target_machines", "Select default machines for this connection."),
671 ("list_threads", "List recent Codex and Claude Code threads on selected machines."),
672 ("read_thread", "Read a thread and its current status."),
673 ("send_message", "Submit a message. Desktop control requires local opt-in and reopens unloaded chats, switching the selected desktop chat. Submission acknowledges delivery, not task completion."),
674 ("interrupt_thread", "Interrupt an agent-owned session or an opted-in Codex desktop turn."),
675 ("start_thread", "Start an agent-owned session under a locally allowed directory."),
676 ].into_iter().map(|(name, description)| {
677 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})});
678 if !["list_machines","set_target_machines"].contains(&name) {schema["properties"]["machine_id"] = json!({"type":"string","format":"uuid"});}
679 let read = ["list_machines","list_threads","read_thread"].contains(&name);
680 Tool::new(name, description, schema.as_object().unwrap().clone()).with_annotations(ToolAnnotations::new().read_only(read).destructive(name == "interrupt_thread").idempotent(read))
681 }).collect();
682 Ok(ListToolsResult {
683 tools,
684 ..Default::default()
685 })
686 }
687 async fn call_tool(
688 &self,
689 request: CallToolRequestParams,
690 context: RequestContext<RoleServer>,
691 ) -> std::result::Result<CallToolResponse, ErrorData> {
692 let result: Result<CallToolResult> = async {
693 let grant = &context
694 .extensions
695 .get::<axum::http::request::Parts>()
696 .and_then(|parts| parts.extensions.get::<mcp::Grant>())
697 .ok_or_else(|| Error::new(401, "This connection expired. Connect again."))?
698 .0;
699 let mut arguments = request.arguments.unwrap_or_default();
700 let id = string(&grant["id"]);
701 match request.name.as_ref() {
702 "list_machines" if arguments.is_empty() => Ok(CallToolResult::structured(
703 json!({"machines":view(&self.0, id, None)?}),
704 )),
705 "set_target_machines"
706 if arguments.len() == 1 && arguments.contains_key("machine_ids") =>
707 {
708 Ok(CallToolResult::structured(
709 json!({"machines":view(&self.0, id, arguments.get("machine_ids"))?}),
710 ))
711 }
712 method => {
713 let explicit = arguments.remove("machine_id");
714 let machines = view(&self.0, id, None)?;
715 let targets: Vec<_> = if let Some(explicit) = explicit {
716 let explicit = explicit
717 .as_str()
718 .ok_or_else(|| Error::new(400, "Choose a machine ID."))?;
719 if !machines.iter().any(|m| m["id"] == explicit) {
720 return Err(Error::new(403, "Choose a granted machine."));
721 }
722 vec![explicit.to_owned()]
723 } else {
724 machines
725 .iter()
726 .filter(|m| m["selected"] == true)
727 .map(|m| string(&m["id"]).to_owned())
728 .collect()
729 };
730 if targets.is_empty() || method != "list_threads" && targets.len() != 1 {
731 return Err(Error::new(
732 400,
733 "Select one machine or supply machine_id for this command.",
734 ));
735 }
736 let params = command(method, json!(arguments))?;
737 let mut replies: futures::stream::FuturesUnordered<_> = targets.into_iter().map(|machine| {
738 let params = params.clone();
739 async move {
740 match self
741 .0
742 .relay
743 .dispatch(
744 &self.0.mcp,
745 id,
746 &machine,
747 method,
748 params,
749 Duration::from_secs(30),
750 )
751 .await
752 {
753 Ok(result) => json!({"machine_id":machine,"result":result}),
754 Err(error) => json!({"machine_id":machine,"error":error.message}),
755 }
756 }
757 }).collect();
758 let mut results = Vec::new();
759 let mut bytes = 0;
760 while let Some(reply) = futures::StreamExt::next(&mut replies).await {
761 bytes += reply.to_string().len();
762 if bytes > 16*1024*1024 {return Err(Error::new(413, "Machine replies exceed 16 MB. Select fewer machines or request fewer messages."));}
763 results.push(reply);
764 }
765 let errors = results.iter().any(|r| r.get("error").is_some());
766 let mut result = CallToolResult::structured(json!({"results":results}));
767 result.is_error = Some(errors);
768 Ok(result)
769 }
770 }
771 }
772 .await;
773 Ok(match result {
774 Ok(result) => result,
775 Err(error) => CallToolResult::error(vec![ContentBlock::text(if error.status >= 500 {
776 "The relay couldn't answer. Check its dashboard and retry.".to_owned()
777 } else {
778 error.message
779 })]),
780 }
781 .into())
782 }
783}
784async fn installer(State(app): State<Arc<App>>, request: Request) -> Response {
785 let origin = app.mcp.origin.origin().ascii_serialization();
786 let source = if request.uri().path().ends_with(".ps1") {
787 include_str!("../agent/install.ps1")
788 .replace("__SERVER__", &format!("'{}'", origin.replace('\'', "''")))
789 } else {
790 include_str!("../agent/install.sh").replace(
791 "__SERVER__",
792 &format!("'{}'", origin.replace('\'', "'\\''")),
793 )
794 };
795 ([("content-type", "text/plain; charset=utf-8")], source).into_response()
796}
797pub fn router(app: Arc<App>) -> Router {
798 let state = app.clone();
799 let expected_host =
800 app.mcp.origin[url::Position::BeforeHost..url::Position::AfterPort].to_owned();
801 let agent = env("STUDIO_AGENT_DIR", "agent");
802 Router::new()
803 .route("/pairing", any(pairing))
804 .route("/agent/connect", any(connect))
805 .route("/agent/install.sh", axum::routing::get(installer))
806 .route("/agent/install.ps1", axum::routing::get(installer))
807 .route_service(
808 "/agent/setup.mjs",
809 ServeFile::new(format!("{agent}/setup.mjs")),
810 )
811 .route_service(
812 "/agent/relay.mjs",
813 ServeFile::new(format!("{agent}/relay.mjs")),
814 )
815 .route("/api/v1/{*path}", any(rest))
816 .with_state(app.clone())
817 .merge(mcp::router(app, "agents", move || {
818 Ok(Agents(state.clone()))
819 }))
820 .layer(axum::middleware::from_fn(
821 move |request: Request, next: axum::middleware::Next| {
822 let host = expected_host.clone();
823 async move {
824 if request.headers().get("host").and_then(|v| v.to_str().ok()) != Some(&host) {
825 return StatusCode::MISDIRECTED_REQUEST.into_response();
826 }
827 next.run(request).await
828 }
829 },
830 ))
831}
832
833#[cfg(test)]
834mod tests {
835 use super::*;
836 struct Fixture {
837 store: Arc<mcp::Store>,
838 broker: Arc<Broker>,
839 path: PathBuf,
840 machine: String,
841 grant: String,
842 }
843 impl Fixture {
844 fn new(write: bool) -> Self {
845 let path = std::env::temp_dir()
846 .canonicalize()
847 .unwrap()
848 .join(format!("studio-relay-test-{}", uuid::Uuid::new_v4()));
849 let store = Arc::new(mcp::Store::new(&path, "https://globe.studio.test").unwrap());
850 let machine = uuid::Uuid::new_v4().to_string();
851 let grant = uuid::Uuid::new_v4().to_string();
852 {
853 let db = store.db.lock().unwrap();
854 put(&db, &format!("machine:{machine}"), &json!({"id":machine,"user":"owner","name":"Fixture","tokenHash":hash("device-secret")}), 0).unwrap();
855 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();
856 }
857 Self {
858 store,
859 broker: Arc::new(Broker::default()),
860 path,
861 machine,
862 grant,
863 }
864 }
865 fn connect(&self) -> (Arc<Connection>, mpsc::Receiver<Message>) {
866 let (messages, receive) = mpsc::channel(16);
867 let (closed, _) = watch::channel(false);
868 let connection = Arc::new(Connection {
869 messages,
870 closed,
871 pending: Mutex::new(HashMap::new()),
872 });
873 self.broker
874 .connections
875 .lock()
876 .unwrap()
877 .insert(self.machine.clone(), connection.clone());
878 (connection, receive)
879 }
880 fn dispatch(&self, duration: Duration) -> tokio::task::JoinHandle<Result<Value>> {
881 let (store, broker, machine, grant) = (
882 self.store.clone(),
883 self.broker.clone(),
884 self.machine.clone(),
885 self.grant.clone(),
886 );
887 tokio::spawn(async move {
888 broker
889 .dispatch(
890 &store,
891 &grant,
892 &machine,
893 "list_threads",
894 json!({}),
895 duration,
896 )
897 .await
898 })
899 }
900 }
901 impl Drop for Fixture {
902 fn drop(&mut self) {
903 std::fs::remove_dir_all(&self.path).unwrap();
904 }
905 }
906 #[test]
907 fn command_fields_defaults_and_provider_boundaries() {
908 assert_eq!(
909 command("list_threads", json!({})).unwrap(),
910 json!({"limit":30})
911 );
912 let id = uuid::Uuid::new_v4().to_string();
913 assert_eq!(
914 command("read_thread", json!({"provider":"claude","thread_id":id})).unwrap()["limit"],
915 20
916 );
917 assert_eq!(
918 command(
919 "read_thread",
920 json!({"provider":"codex","thread_id":id,"cursor":"page"})
921 )
922 .unwrap()["cursor"],
923 "page"
924 );
925 for (method, params) in [
926 (
927 "read_thread",
928 json!({"provider":"codex","thread_id":id,"cursor":""}),
929 ),
930 (
931 "read_thread",
932 json!({"provider":"codex","thread_id":id,"cursor":"x".repeat(513)}),
933 ),
934 (
935 "read_thread",
936 json!({"provider":"codex","thread_id":"not-a-uuid"}),
937 ),
938 (
939 "read_thread",
940 json!({"provider":"native-chat","thread_id":id}),
941 ),
942 ("list_threads", json!({"limit":101})),
943 ("list_threads", json!({"limit":1.5})),
944 ("list_threads", json!({"approval_response":true})),
945 (
946 "send_message",
947 json!({"provider":"codex","thread_id":id,"message":""}),
948 ),
949 ("start_thread", json!({"provider":"claude","cwd":"/tmp"})),
950 (
951 "interrupt_thread",
952 json!({"provider":"codex","thread_id":id,"expected_turn_id":1}),
953 ),
954 ("approve_tool", json!({})),
955 ] {
956 assert!(command(method, params).is_err(), "{method}");
957 }
958 }
959 #[tokio::test]
960 async fn dispatch_uses_current_owner_grant_and_control_scope() {
961 let fixture = Fixture::new(false);
962 let (connection, mut receive) = fixture.connect();
963 let task = fixture.dispatch(Duration::from_secs(2));
964 let Message::Text(frame) = receive.recv().await.unwrap() else {
965 panic!()
966 };
967 let frame: Value = serde_json::from_str(&frame).unwrap();
968 assert_eq!(frame["params"], json!({"limit":30}));
969 connection
970 .pending
971 .lock()
972 .unwrap()
973 .remove(string(&frame["id"]))
974 .unwrap()
975 .send(Ok(json!({"threads":[]})))
976 .unwrap();
977 assert_eq!(task.await.unwrap().unwrap(), json!({"threads":[]}));
978 let refused = fixture
979 .broker
980 .dispatch(
981 &fixture.store,
982 &fixture.grant,
983 &fixture.machine,
984 "start_thread",
985 json!({"provider":"codex","cwd":"/tmp","message":"owned fixture"}),
986 Duration::from_secs(2),
987 )
988 .await
989 .unwrap_err();
990 assert_eq!(refused.status, 403);
991 assert!(receive.try_recv().is_err());
992 let foreign = uuid::Uuid::new_v4().to_string();
993 {
994 let db = fixture.store.db.lock().unwrap();
995 let mut grant = get(&db, &format!("grant:{}", fixture.grant)).unwrap();
996 grant["resources"] = json!([foreign]);
997 put(&db, &format!("grant:{}", fixture.grant), &grant, 0).unwrap();
998 put(
999 &db,
1000 &format!("machine:{foreign}"),
1001 &json!({"id":foreign,"user":"someone-else"}),
1002 0,
1003 )
1004 .unwrap();
1005 }
1006 assert_eq!(
1007 fixture
1008 .broker
1009 .dispatch(
1010 &fixture.store,
1011 &fixture.grant,
1012 &foreign,
1013 "list_threads",
1014 json!({}),
1015 Duration::from_secs(2)
1016 )
1017 .await
1018 .unwrap_err()
1019 .status,
1020 403
1021 );
1022 delete(
1023 &fixture.store.db.lock().unwrap(),
1024 &format!("grant:{}", fixture.grant),
1025 )
1026 .unwrap();
1027 assert_eq!(
1028 fixture
1029 .dispatch(Duration::from_secs(2))
1030 .await
1031 .unwrap()
1032 .unwrap_err()
1033 .status,
1034 403
1035 );
1036 }
1037 #[tokio::test]
1038 async fn unicode_message_cannot_exceed_local_agent_frame_budget() {
1039 let fixture = Fixture::new(true);
1040 let (connection, mut receive) = fixture.connect();
1041 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();
1042 assert_eq!(error.status, 413);
1043 assert!(receive.try_recv().is_err());
1044 assert!(connection.pending.lock().unwrap().is_empty());
1045 assert_eq!(fixture.broker.pending_slots.available_permits(), 128);
1046 }
1047 #[tokio::test]
1048 async fn eight_pending_limit_and_cancellation_release_capacity() {
1049 let fixture = Fixture::new(true);
1050 let (connection, mut receive) = fixture.connect();
1051 let mut tasks = Vec::new();
1052 for _ in 0..8 {
1053 tasks.push(fixture.dispatch(Duration::from_secs(2)));
1054 receive.recv().await.unwrap();
1055 }
1056 assert_eq!(
1057 fixture
1058 .dispatch(Duration::from_secs(2))
1059 .await
1060 .unwrap()
1061 .unwrap_err()
1062 .status,
1063 429
1064 );
1065 tasks.remove(0).abort();
1066 tokio::task::yield_now().await;
1067 assert_eq!(connection.pending.lock().unwrap().len(), 7);
1068 tasks.push(fixture.dispatch(Duration::from_secs(2)));
1069 receive.recv().await.unwrap();
1070 assert_eq!(connection.pending.lock().unwrap().len(), 8);
1071 fixture.broker.disconnect(&fixture.machine);
1072 for task in tasks {
1073 assert!(task.await.unwrap().unwrap_err().message.contains("unknown"));
1074 }
1075 assert!(connection.pending.lock().unwrap().is_empty());
1076 }
1077 #[tokio::test]
1078 async fn timeout_disconnect_and_reconnect_do_not_replay() {
1079 let fixture = Fixture::new(true);
1080 let (old, mut receive) = fixture.connect();
1081 let task = fixture.dispatch(Duration::from_millis(5));
1082 receive.recv().await.unwrap();
1083 assert!(task.await.unwrap().unwrap_err().message.contains("unknown"));
1084 assert!(old.pending.lock().unwrap().is_empty());
1085 assert!(receive.try_recv().is_err());
1086 let task = fixture.dispatch(Duration::from_secs(2));
1087 receive.recv().await.unwrap();
1088 fixture.broker.disconnect(&fixture.machine);
1089 assert!(task.await.unwrap().unwrap_err().message.contains("unknown"));
1090 let (new, mut receive) = fixture.connect();
1091 fixture.broker.remove(&fixture.machine, &old);
1092 assert!(Arc::ptr_eq(
1093 fixture
1094 .broker
1095 .connections
1096 .lock()
1097 .unwrap()
1098 .get(&fixture.machine)
1099 .unwrap(),
1100 &new
1101 ));
1102 assert!(receive.try_recv().is_err());
1103 let task = fixture.dispatch(Duration::from_secs(2));
1104 receive.recv().await.unwrap();
1105 fixture.broker.disconnect(&fixture.machine);
1106 assert!(task.await.unwrap().unwrap_err().message.contains("unknown"));
1107 }
1108}