1use axum::{
40 extract::{
41 ws::{Message, WebSocket, WebSocketUpgrade},
42 Path, State,
43 },
44 http::{HeaderMap, StatusCode},
45 response::{IntoResponse, Response},
46 Json,
47};
48use futures_util::{sink::SinkExt, stream::StreamExt};
49use mlua_swarm::{AgentProviderManifest, Operator, SeniorBridge, SessionId, SpawnHook};
50use serde::{Deserialize, Serialize};
51use serde_json::json;
52use std::sync::Arc;
53use tokio::sync::{mpsc, Mutex};
54
55use super::protocol::{ClientMsg, PendingReply, ServerMsg};
56use super::session::WSOperatorSession;
57use crate::AppState;
58
59pub struct OperatorSessionEntry {
65 pub sid: SessionId,
67 pub token: String,
69 pub roles: Vec<String>,
71 pub capability_manifest: Option<AgentProviderManifest>,
73 pub joined_at_secs: u64,
78 pub ws_session: Mutex<Option<Arc<WSOperatorSession>>>,
81}
82
83#[derive(Debug, Deserialize, Default)]
87pub struct OperatorsCreateReq {
88 #[serde(default)]
90 pub roles: Vec<String>,
91 #[serde(default)]
93 pub capability_manifest: Option<AgentProviderManifest>,
94}
95
96#[derive(Debug, Serialize)]
98pub struct OperatorsCreateResp {
99 pub sid: SessionId,
102 pub token: String,
104 pub roles: Vec<String>,
106}
107
108pub async fn operators_create(
118 State(state): State<AppState>,
119 Json(req): Json<OperatorsCreateReq>,
120) -> Response {
121 let roles = req.roles;
122 let capability_manifest = req.capability_manifest;
123 let sid = SessionId::new();
130 let token = mlua_swarm::types::secure_hex(5);
131
132 {
133 let mut map = state.roles_to_sid.lock().await;
134 let conflicts: Vec<String> = roles
135 .iter()
136 .filter(|r| map.contains_key(r.as_str()))
137 .cloned()
138 .collect();
139 if !conflicts.is_empty() {
140 let conflicts_detail: Vec<serde_json::Value> = conflicts
147 .iter()
148 .map(|r| {
149 let holder = map.get(r.as_str()).map(|sid| sid.to_string());
150 json!({ "role": r, "sid": holder })
151 })
152 .collect();
153 return (
154 StatusCode::CONFLICT,
155 Json(json!({
156 "error": "roles conflict",
157 "conflicts": conflicts,
158 "conflicts_detail": conflicts_detail,
159 })),
160 )
161 .into_response();
162 }
163 for r in &roles {
164 map.insert(r.clone(), sid.clone());
165 }
166 }
167
168 let joined_at_secs = std::time::SystemTime::now()
169 .duration_since(std::time::UNIX_EPOCH)
170 .map(|d| d.as_secs())
171 .unwrap_or(0);
172 let entry = Arc::new(OperatorSessionEntry {
173 sid: sid.clone(),
174 token: token.clone(),
175 roles: roles.clone(),
176 capability_manifest,
177 joined_at_secs,
178 ws_session: Mutex::new(None),
179 });
180 state
181 .operator_sessions
182 .lock()
183 .await
184 .insert(sid.clone(), entry);
185
186 (
187 StatusCode::OK,
188 Json(OperatorsCreateResp { sid, token, roles }),
189 )
190 .into_response()
191}
192
193fn extract_bearer_token_required(headers: &HeaderMap) -> Result<String, Box<Response>> {
199 let token = headers
200 .get(axum::http::header::AUTHORIZATION)
201 .and_then(|v| v.to_str().ok())
202 .and_then(|s| s.strip_prefix("Bearer "))
203 .map(|s| s.trim().to_string())
204 .filter(|s| !s.is_empty());
205 token.ok_or_else(|| {
206 Box::new((StatusCode::UNAUTHORIZED, "missing or empty Bearer token").into_response())
207 })
208}
209
210pub async fn operators_ws_connect(
216 State(state): State<AppState>,
217 Path(sid): Path<String>,
218 headers: HeaderMap,
219 ws: WebSocketUpgrade,
220) -> Response {
221 let bearer = match extract_bearer_token_required(&headers) {
222 Ok(t) => t,
223 Err(resp) => return *resp,
224 };
225 let Ok(sid) = SessionId::parse(sid) else {
227 return (StatusCode::NOT_FOUND, "unknown sid").into_response();
228 };
229
230 let entry = {
231 let map = state.operator_sessions.lock().await;
232 map.get(&sid).cloned()
233 };
234 let entry = match entry {
235 Some(e) => e,
236 None => return (StatusCode::NOT_FOUND, "unknown sid").into_response(),
237 };
238 if !mlua_swarm::types::ct_eq(entry.token.as_bytes(), bearer.as_bytes()) {
239 return (StatusCode::UNAUTHORIZED, "token mismatch").into_response();
240 }
241
242 ws.on_upgrade(move |socket| handle_operator_socket(socket, state, entry))
243}
244
245async fn handle_operator_socket(
249 socket: WebSocket,
250 state: AppState,
251 entry: Arc<OperatorSessionEntry>,
252) {
253 let (tx, mut rx) = mpsc::unbounded_channel::<ServerMsg>();
254
255 let existing_ws = entry.ws_session.lock().await.clone();
256 let session = match existing_ws {
257 Some(ws_session) => {
258 ws_session.replace_tx(tx.clone()).await;
260 ws_session
261 }
262 None => {
263 let ws_session = Arc::new(WSOperatorSession::new_with_base_url(
264 entry.sid.clone(),
265 tx.clone(),
266 state.base_url.clone(),
267 ));
268 state
269 .engine
270 .register_senior_bridge(
271 entry.sid.clone(),
272 ws_session.clone() as Arc<dyn SeniorBridge>,
273 )
274 .await;
275 state
276 .engine
277 .register_spawn_hook(entry.sid.clone(), ws_session.clone() as Arc<dyn SpawnHook>)
278 .await;
279 state
280 .engine
281 .register_operator(entry.sid.clone(), ws_session.clone() as Arc<dyn Operator>)
282 .await;
283 if let Some(factory) = &state.ws_operator_factory {
284 factory
285 .register_operator(entry.sid.clone(), ws_session.clone() as Arc<dyn Operator>);
286 }
287 for role in &entry.roles {
292 if let Some(factory) = &state.ws_operator_factory {
293 factory
294 .register_operator(role.clone(), ws_session.clone() as Arc<dyn Operator>);
295 }
296 state
297 .engine
298 .register_operator(role.clone(), ws_session.clone() as Arc<dyn Operator>)
299 .await;
300 }
301 *entry.ws_session.lock().await = Some(ws_session.clone());
302 ws_session
303 }
304 };
305
306 let (mut ws_sink, mut ws_stream) = socket.split();
307
308 let write_task = tokio::spawn(async move {
310 while let Some(msg) = rx.recv().await {
311 let txt = match serde_json::to_string(&msg) {
312 Ok(s) => s,
313 Err(_) => continue,
314 };
315 if ws_sink.send(Message::Text(txt)).await.is_err() {
316 break;
317 }
318 }
319 let _ = ws_sink.close().await;
320 });
321
322 let session_for_read = session.clone();
324 let read_result: Result<(), String> = async {
325 while let Some(item) = ws_stream.next().await {
326 match item {
327 Ok(Message::Text(t)) => {
328 let parsed: ClientMsg = match serde_json::from_str(&t) {
329 Ok(p) => p,
330 Err(_) => continue,
331 };
332 match parsed {
333 ClientMsg::Answer { req_id, value } => {
334 session_for_read
335 .resolve_pending(&req_id, PendingReply::Answer(value))
336 .await;
337 }
338 ClientMsg::HookAck { req_id, ok, reason } => {
339 session_for_read
340 .resolve_pending(&req_id, PendingReply::HookAck { ok, reason })
341 .await;
342 }
343 ClientMsg::SpawnAck {
344 req_id,
345 value,
346 ok,
347 error,
348 } => {
349 session_for_read
350 .resolve_pending(
351 &req_id,
352 PendingReply::SpawnAck { value, ok, error },
353 )
354 .await;
355 }
356 ClientMsg::SpawnHalt {
357 req_id,
358 value,
359 reason,
360 } => {
361 session_for_read
362 .resolve_pending(&req_id, PendingReply::SpawnHalt { value, reason })
363 .await;
364 }
365 }
366 }
367 Ok(Message::Ping(_)) | Ok(Message::Pong(_)) => {}
368 Ok(Message::Close(_)) | Err(_) => break,
369 _ => {}
370 }
371 }
372 Ok(())
373 }
374 .await;
375
376 session.clear_tx_if(&tx).await;
379 write_task.abort();
380 let _ = read_result;
381}
382
383async fn teardown_operator_session(
393 state: &AppState,
394 sid: &SessionId,
395 entry: &Arc<OperatorSessionEntry>,
396) {
397 state.engine.unregister_senior_bridge(sid.as_str()).await;
398 state.engine.unregister_spawn_hook(sid.as_str()).await;
399 state.engine.unregister_operator(sid.as_str()).await;
400 if let Some(factory) = &state.ws_operator_factory {
401 factory.unregister_operator(sid.as_str());
402 }
403 for role in &entry.roles {
404 state.engine.unregister_operator(role).await;
405 if let Some(factory) = &state.ws_operator_factory {
406 factory.unregister_operator(role);
407 }
408 }
409
410 if let Some(session) = entry.ws_session.lock().await.take() {
411 session.clear_tx().await;
412 }
413
414 state.operator_sessions.lock().await.remove(sid);
415
416 {
417 let mut map = state.roles_to_sid.lock().await;
418 for role in &entry.roles {
419 if map.get(role) == Some(sid) {
420 map.remove(role);
421 }
422 }
423 }
424}
425
426pub async fn operators_delete(
432 State(state): State<AppState>,
433 Path(sid): Path<String>,
434 headers: HeaderMap,
435) -> Response {
436 let bearer = match extract_bearer_token_required(&headers) {
437 Ok(t) => t,
438 Err(resp) => return *resp,
439 };
440 let Ok(sid) = SessionId::parse(sid) else {
441 return (StatusCode::NOT_FOUND, "unknown sid").into_response();
442 };
443
444 let entry = {
445 let map = state.operator_sessions.lock().await;
446 map.get(&sid).cloned()
447 };
448 let entry = match entry {
449 Some(e) => e,
450 None => return (StatusCode::NOT_FOUND, "unknown sid").into_response(),
451 };
452 if !mlua_swarm::types::ct_eq(entry.token.as_bytes(), bearer.as_bytes()) {
453 return (StatusCode::UNAUTHORIZED, "token mismatch").into_response();
454 }
455
456 teardown_operator_session(&state, &sid, &entry).await;
457
458 StatusCode::NO_CONTENT.into_response()
459}
460
461#[derive(Debug, Serialize)]
468pub struct OperatorsListEntry {
469 pub sid: SessionId,
471 pub roles: Vec<String>,
473 pub joined_at_secs: u64,
476 pub connected: bool,
479}
480
481#[derive(Debug, Serialize)]
483pub struct OperatorsListResp {
484 pub operators: Vec<OperatorsListEntry>,
487}
488
489pub async fn operators_list(State(state): State<AppState>) -> Response {
496 let entries: Vec<(SessionId, Arc<OperatorSessionEntry>)> = {
497 let map = state.operator_sessions.lock().await;
498 map.iter().map(|(k, v)| (k.clone(), v.clone())).collect()
499 };
500 let mut operators = Vec::with_capacity(entries.len());
501 for (sid, entry) in entries {
502 let session = entry.ws_session.lock().await.clone();
503 let connected = match session {
504 Some(session) => session.is_connected().await,
505 None => false,
506 };
507 operators.push(OperatorsListEntry {
508 sid,
509 roles: entry.roles.clone(),
510 joined_at_secs: entry.joined_at_secs,
511 connected,
512 });
513 }
514 operators.sort_by(|a, b| a.sid.as_str().cmp(b.sid.as_str()));
515 (StatusCode::OK, Json(OperatorsListResp { operators })).into_response()
516}
517
518pub async fn operators_delete_by_role(
530 State(state): State<AppState>,
531 Path(role): Path<String>,
532) -> Response {
533 let sid = {
534 let map = state.roles_to_sid.lock().await;
535 match map.get(role.as_str()) {
536 Some(sid) => sid.clone(),
537 None => {
538 return (
539 StatusCode::NOT_FOUND,
540 Json(json!({"error": "no session holds this role", "role": role})),
541 )
542 .into_response();
543 }
544 }
545 };
546 let entry = {
547 let map = state.operator_sessions.lock().await;
548 map.get(&sid).cloned()
549 };
550 let entry = match entry {
551 Some(e) => e,
552 None => {
553 let mut map = state.roles_to_sid.lock().await;
559 if map.get(role.as_str()) == Some(&sid) {
560 map.remove(role.as_str());
561 }
562 return (
563 StatusCode::NOT_FOUND,
564 Json(json!({
565 "error": "torn role mapping cleared; role now open",
566 "role": role,
567 })),
568 )
569 .into_response();
570 }
571 };
572 teardown_operator_session(&state, &sid, &entry).await;
573 StatusCode::NO_CONTENT.into_response()
574}
575
576#[derive(Debug, Serialize)]
580pub struct OperatorsInfoResp {
581 pub sid: SessionId,
583 pub roles: Vec<String>,
585 #[serde(skip_serializing_if = "Option::is_none")]
587 pub capability_manifest: Option<AgentProviderManifest>,
588 pub connected: bool,
590}
591
592pub async fn operators_info(
596 State(state): State<AppState>,
597 Path(sid): Path<String>,
598 headers: HeaderMap,
599) -> Response {
600 let bearer = match extract_bearer_token_required(&headers) {
601 Ok(t) => t,
602 Err(resp) => return *resp,
603 };
604 let Ok(sid) = SessionId::parse(sid) else {
605 return (StatusCode::NOT_FOUND, "unknown sid").into_response();
606 };
607
608 let entry = {
609 let map = state.operator_sessions.lock().await;
610 map.get(&sid).cloned()
611 };
612 let entry = match entry {
613 Some(e) => e,
614 None => return (StatusCode::NOT_FOUND, "unknown sid").into_response(),
615 };
616 if !mlua_swarm::types::ct_eq(entry.token.as_bytes(), bearer.as_bytes()) {
617 return (StatusCode::UNAUTHORIZED, "token mismatch").into_response();
618 }
619
620 let session = entry.ws_session.lock().await.clone();
621 let connected = match session {
622 Some(session) => session.is_connected().await,
623 None => false,
624 };
625 (
626 StatusCode::OK,
627 Json(OperatorsInfoResp {
628 sid: entry.sid.clone(),
629 roles: entry.roles.clone(),
630 capability_manifest: entry.capability_manifest.clone(),
631 connected,
632 }),
633 )
634 .into_response()
635}
636
637#[cfg(test)]
638mod tests {
639 use super::*;
640 use axum::http::HeaderValue;
641
642 fn headers_with_bearer(token: &str) -> HeaderMap {
643 let mut h = HeaderMap::new();
644 h.insert(
645 axum::http::header::AUTHORIZATION,
646 HeaderValue::from_str(&format!("Bearer {token}")).unwrap(),
647 );
648 h
649 }
650
651 #[test]
652 fn extract_bearer_token_required_accepts_valid() {
653 let h = headers_with_bearer("abc123");
654 assert_eq!(extract_bearer_token_required(&h).unwrap(), "abc123");
655 }
656
657 #[test]
658 fn extract_bearer_token_required_rejects_missing_header() {
659 let h = HeaderMap::new();
660 assert!(extract_bearer_token_required(&h).is_err());
661 }
662
663 #[test]
664 fn extract_bearer_token_required_rejects_empty_token() {
665 let h = headers_with_bearer("");
666 assert!(extract_bearer_token_required(&h).is_err());
667 }
668
669 #[test]
670 fn extract_bearer_token_required_rejects_wrong_scheme() {
671 let mut h = HeaderMap::new();
672 h.insert(
673 axum::http::header::AUTHORIZATION,
674 HeaderValue::from_static("Basic dXNlcjpwYXNz"),
675 );
676 assert!(extract_bearer_token_required(&h).is_err());
677 }
678
679 #[test]
680 fn operators_create_request_accepts_capability_manifest() {
681 let req: OperatorsCreateReq = serde_json::from_value(serde_json::json!({
682 "roles": ["main-ai"],
683 "capability_manifest": {
684 "provider_id": "main-ai-self-report",
685 "capabilities": [{
686 "launch_variant": "mse-coder",
687 "resolved_model": "claude-sonnet-4",
688 "effective_tools": ["Read", "Edit"]
689 }]
690 }
691 }))
692 .unwrap();
693 assert_eq!(req.roles, ["main-ai"]);
694 assert_eq!(
695 req.capability_manifest.unwrap().provider_id,
696 "main-ai-self-report"
697 );
698 }
699
700 #[test]
701 fn operators_create_request_keeps_manifest_optional_on_wire() {
702 let req: OperatorsCreateReq =
703 serde_json::from_value(serde_json::json!({ "roles": [] })).unwrap();
704 assert!(req.capability_manifest.is_none());
705 }
706}