| 1 | use crate::*; |
| 2 | use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade}; |
| 3 | use futures::SinkExt; |
| 4 | use mcp::{delete, get, hash, list, put, secret}; |
| 5 | use rmcp::{ |
| 6 | ErrorData, RoleServer, ServerHandler, |
| 7 | model::{ |
| 8 | CallToolRequestParams, CallToolResponse, CallToolResult, ContentBlock, ListToolsResult, |
| 9 | PaginatedRequestParams, ServerCapabilities, ServerConfig, Tool, ToolAnnotations, |
| 10 | }, |
| 11 | service::RequestContext, |
| 12 | }; |
| 13 | use tokio::sync::{mpsc, oneshot}; |
| 14 | |
| 15 | pub struct Broker { |
| 16 | connections: Mutex<HashMap<String, Arc<Connection>>>, |
| 17 | pending_slots: Semaphore, |
| 18 | pub(crate) changes: watch::Sender<()>, |
| 19 | } |
| 20 | impl 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 | } |
| 29 | struct Connection { |
| 30 | messages: mpsc::Sender<Message>, |
| 31 | pending: Mutex<HashMap<String, oneshot::Sender<Result<Value>>>>, |
| 32 | closed: watch::Sender<bool>, |
| 33 | } |
| 34 | struct Pending { |
| 35 | connection: Arc<Connection>, |
| 36 | id: String, |
| 37 | } |
| 38 | impl Drop for Pending { |
| 39 | fn drop(&mut self) { |
| 40 | self.connection.pending.lock().unwrap().remove(&self.id); |
| 41 | } |
| 42 | } |
| 43 | fn unknown() -> Error { |
| 44 | Error::new( |
| 45 | 409, |
| 46 | "Command outcome is unknown. Read the thread before retrying.", |
| 47 | ) |
| 48 | } |
| 49 | impl 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 | } |
| 182 | pub(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 | } |
| 188 | fn 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 | } |
| 207 | fn 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 | } |
| 213 | pub(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 | } |
| 311 | fn 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 | } |
| 328 | fn 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 | } |
| 352 | fn 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 | } |
| 383 | fn 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 | } |
| 431 | async 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 | } |
| 493 | async 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 | } |
| 538 | async 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 | } |
| 611 | async 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)] |
| 658 | struct Agents(Arc<App>); |
| 659 | impl 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 | } |
| 784 | async 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 | } |
| 797 | pub 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)] |
| 834 | mod 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 | } |