mlua_swarm_server/operator_ws/
login.rs1use 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 ws_session: Mutex<Option<Arc<WSOperatorSession>>>,
76}
77
78#[derive(Debug, Deserialize, Default)]
82pub struct OperatorsCreateReq {
83 #[serde(default)]
85 pub roles: Vec<String>,
86 #[serde(default)]
88 pub capability_manifest: Option<AgentProviderManifest>,
89}
90
91#[derive(Debug, Serialize)]
93pub struct OperatorsCreateResp {
94 pub sid: SessionId,
97 pub token: String,
99 pub roles: Vec<String>,
101}
102
103pub async fn operators_create(
113 State(state): State<AppState>,
114 Json(req): Json<OperatorsCreateReq>,
115) -> Response {
116 let roles = req.roles;
117 let capability_manifest = req.capability_manifest;
118 let sid = SessionId::new();
125 let token = mlua_swarm::types::secure_hex(5);
126
127 {
128 let mut map = state.roles_to_sid.lock().await;
129 let conflicts: Vec<String> = roles
130 .iter()
131 .filter(|r| map.contains_key(r.as_str()))
132 .cloned()
133 .collect();
134 if !conflicts.is_empty() {
135 return (
136 StatusCode::CONFLICT,
137 Json(json!({"error": "roles conflict", "conflicts": conflicts})),
138 )
139 .into_response();
140 }
141 for r in &roles {
142 map.insert(r.clone(), sid.clone());
143 }
144 }
145
146 let entry = Arc::new(OperatorSessionEntry {
147 sid: sid.clone(),
148 token: token.clone(),
149 roles: roles.clone(),
150 capability_manifest,
151 ws_session: Mutex::new(None),
152 });
153 state
154 .operator_sessions
155 .lock()
156 .await
157 .insert(sid.clone(), entry);
158
159 (
160 StatusCode::OK,
161 Json(OperatorsCreateResp { sid, token, roles }),
162 )
163 .into_response()
164}
165
166fn extract_bearer_token_required(headers: &HeaderMap) -> Result<String, Box<Response>> {
172 let token = headers
173 .get(axum::http::header::AUTHORIZATION)
174 .and_then(|v| v.to_str().ok())
175 .and_then(|s| s.strip_prefix("Bearer "))
176 .map(|s| s.trim().to_string())
177 .filter(|s| !s.is_empty());
178 token.ok_or_else(|| {
179 Box::new((StatusCode::UNAUTHORIZED, "missing or empty Bearer token").into_response())
180 })
181}
182
183pub async fn operators_ws_connect(
189 State(state): State<AppState>,
190 Path(sid): Path<String>,
191 headers: HeaderMap,
192 ws: WebSocketUpgrade,
193) -> Response {
194 let bearer = match extract_bearer_token_required(&headers) {
195 Ok(t) => t,
196 Err(resp) => return *resp,
197 };
198 let Ok(sid) = SessionId::parse(sid) else {
200 return (StatusCode::NOT_FOUND, "unknown sid").into_response();
201 };
202
203 let entry = {
204 let map = state.operator_sessions.lock().await;
205 map.get(&sid).cloned()
206 };
207 let entry = match entry {
208 Some(e) => e,
209 None => return (StatusCode::NOT_FOUND, "unknown sid").into_response(),
210 };
211 if !mlua_swarm::types::ct_eq(entry.token.as_bytes(), bearer.as_bytes()) {
212 return (StatusCode::UNAUTHORIZED, "token mismatch").into_response();
213 }
214
215 ws.on_upgrade(move |socket| handle_operator_socket(socket, state, entry))
216}
217
218async fn handle_operator_socket(
222 socket: WebSocket,
223 state: AppState,
224 entry: Arc<OperatorSessionEntry>,
225) {
226 let (tx, mut rx) = mpsc::unbounded_channel::<ServerMsg>();
227
228 let existing_ws = entry.ws_session.lock().await.clone();
229 let session = match existing_ws {
230 Some(ws_session) => {
231 ws_session.replace_tx(tx.clone()).await;
233 ws_session
234 }
235 None => {
236 let ws_session = Arc::new(WSOperatorSession::new_with_base_url(
237 entry.sid.clone(),
238 tx.clone(),
239 state.base_url.clone(),
240 ));
241 state
242 .engine
243 .register_senior_bridge(
244 entry.sid.clone(),
245 ws_session.clone() as Arc<dyn SeniorBridge>,
246 )
247 .await;
248 state
249 .engine
250 .register_spawn_hook(entry.sid.clone(), ws_session.clone() as Arc<dyn SpawnHook>)
251 .await;
252 state
253 .engine
254 .register_operator(entry.sid.clone(), ws_session.clone() as Arc<dyn Operator>)
255 .await;
256 if let Some(factory) = &state.ws_operator_factory {
257 factory
258 .register_operator(entry.sid.clone(), ws_session.clone() as Arc<dyn Operator>);
259 }
260 for role in &entry.roles {
265 if let Some(factory) = &state.ws_operator_factory {
266 factory
267 .register_operator(role.clone(), ws_session.clone() as Arc<dyn Operator>);
268 }
269 state
270 .engine
271 .register_operator(role.clone(), ws_session.clone() as Arc<dyn Operator>)
272 .await;
273 }
274 *entry.ws_session.lock().await = Some(ws_session.clone());
275 ws_session
276 }
277 };
278
279 let (mut ws_sink, mut ws_stream) = socket.split();
280
281 let write_task = tokio::spawn(async move {
283 while let Some(msg) = rx.recv().await {
284 let txt = match serde_json::to_string(&msg) {
285 Ok(s) => s,
286 Err(_) => continue,
287 };
288 if ws_sink.send(Message::Text(txt)).await.is_err() {
289 break;
290 }
291 }
292 let _ = ws_sink.close().await;
293 });
294
295 let session_for_read = session.clone();
297 let read_result: Result<(), String> = async {
298 while let Some(item) = ws_stream.next().await {
299 match item {
300 Ok(Message::Text(t)) => {
301 let parsed: ClientMsg = match serde_json::from_str(&t) {
302 Ok(p) => p,
303 Err(_) => continue,
304 };
305 match parsed {
306 ClientMsg::Answer { req_id, value } => {
307 session_for_read
308 .resolve_pending(&req_id, PendingReply::Answer(value))
309 .await;
310 }
311 ClientMsg::HookAck { req_id, ok, reason } => {
312 session_for_read
313 .resolve_pending(&req_id, PendingReply::HookAck { ok, reason })
314 .await;
315 }
316 ClientMsg::SpawnAck {
317 req_id,
318 value,
319 ok,
320 error,
321 } => {
322 session_for_read
323 .resolve_pending(
324 &req_id,
325 PendingReply::SpawnAck { value, ok, error },
326 )
327 .await;
328 }
329 ClientMsg::SpawnHalt {
330 req_id,
331 value,
332 reason,
333 } => {
334 session_for_read
335 .resolve_pending(&req_id, PendingReply::SpawnHalt { value, reason })
336 .await;
337 }
338 }
339 }
340 Ok(Message::Ping(_)) | Ok(Message::Pong(_)) => {}
341 Ok(Message::Close(_)) | Err(_) => break,
342 _ => {}
343 }
344 }
345 Ok(())
346 }
347 .await;
348
349 session.clear_tx_if(&tx).await;
352 write_task.abort();
353 let _ = read_result;
354}
355
356pub async fn operators_delete(
364 State(state): State<AppState>,
365 Path(sid): Path<String>,
366 headers: HeaderMap,
367) -> Response {
368 let bearer = match extract_bearer_token_required(&headers) {
369 Ok(t) => t,
370 Err(resp) => return *resp,
371 };
372 let Ok(sid) = SessionId::parse(sid) else {
373 return (StatusCode::NOT_FOUND, "unknown sid").into_response();
374 };
375
376 let entry = {
377 let map = state.operator_sessions.lock().await;
378 map.get(&sid).cloned()
379 };
380 let entry = match entry {
381 Some(e) => e,
382 None => return (StatusCode::NOT_FOUND, "unknown sid").into_response(),
383 };
384 if !mlua_swarm::types::ct_eq(entry.token.as_bytes(), bearer.as_bytes()) {
385 return (StatusCode::UNAUTHORIZED, "token mismatch").into_response();
386 }
387
388 state.engine.unregister_senior_bridge(sid.as_str()).await;
389 state.engine.unregister_spawn_hook(sid.as_str()).await;
390 state.engine.unregister_operator(sid.as_str()).await;
391 if let Some(factory) = &state.ws_operator_factory {
392 factory.unregister_operator(sid.as_str());
393 }
394 for role in &entry.roles {
395 state.engine.unregister_operator(role).await;
396 if let Some(factory) = &state.ws_operator_factory {
397 factory.unregister_operator(role);
398 }
399 }
400
401 if let Some(session) = entry.ws_session.lock().await.take() {
402 session.clear_tx().await;
403 }
404
405 state.operator_sessions.lock().await.remove(&sid);
406
407 {
408 let mut map = state.roles_to_sid.lock().await;
409 for role in &entry.roles {
410 if map.get(role) == Some(&sid) {
411 map.remove(role);
412 }
413 }
414 }
415
416 StatusCode::NO_CONTENT.into_response()
417}
418
419#[derive(Debug, Serialize)]
423pub struct OperatorsInfoResp {
424 pub sid: SessionId,
426 pub roles: Vec<String>,
428 #[serde(skip_serializing_if = "Option::is_none")]
430 pub capability_manifest: Option<AgentProviderManifest>,
431 pub connected: bool,
433}
434
435pub async fn operators_info(
439 State(state): State<AppState>,
440 Path(sid): Path<String>,
441 headers: HeaderMap,
442) -> Response {
443 let bearer = match extract_bearer_token_required(&headers) {
444 Ok(t) => t,
445 Err(resp) => return *resp,
446 };
447 let Ok(sid) = SessionId::parse(sid) else {
448 return (StatusCode::NOT_FOUND, "unknown sid").into_response();
449 };
450
451 let entry = {
452 let map = state.operator_sessions.lock().await;
453 map.get(&sid).cloned()
454 };
455 let entry = match entry {
456 Some(e) => e,
457 None => return (StatusCode::NOT_FOUND, "unknown sid").into_response(),
458 };
459 if !mlua_swarm::types::ct_eq(entry.token.as_bytes(), bearer.as_bytes()) {
460 return (StatusCode::UNAUTHORIZED, "token mismatch").into_response();
461 }
462
463 let session = entry.ws_session.lock().await.clone();
464 let connected = match session {
465 Some(session) => session.is_connected().await,
466 None => false,
467 };
468 (
469 StatusCode::OK,
470 Json(OperatorsInfoResp {
471 sid: entry.sid.clone(),
472 roles: entry.roles.clone(),
473 capability_manifest: entry.capability_manifest.clone(),
474 connected,
475 }),
476 )
477 .into_response()
478}
479
480#[cfg(test)]
481mod tests {
482 use super::*;
483 use axum::http::HeaderValue;
484
485 fn headers_with_bearer(token: &str) -> HeaderMap {
486 let mut h = HeaderMap::new();
487 h.insert(
488 axum::http::header::AUTHORIZATION,
489 HeaderValue::from_str(&format!("Bearer {token}")).unwrap(),
490 );
491 h
492 }
493
494 #[test]
495 fn extract_bearer_token_required_accepts_valid() {
496 let h = headers_with_bearer("abc123");
497 assert_eq!(extract_bearer_token_required(&h).unwrap(), "abc123");
498 }
499
500 #[test]
501 fn extract_bearer_token_required_rejects_missing_header() {
502 let h = HeaderMap::new();
503 assert!(extract_bearer_token_required(&h).is_err());
504 }
505
506 #[test]
507 fn extract_bearer_token_required_rejects_empty_token() {
508 let h = headers_with_bearer("");
509 assert!(extract_bearer_token_required(&h).is_err());
510 }
511
512 #[test]
513 fn extract_bearer_token_required_rejects_wrong_scheme() {
514 let mut h = HeaderMap::new();
515 h.insert(
516 axum::http::header::AUTHORIZATION,
517 HeaderValue::from_static("Basic dXNlcjpwYXNz"),
518 );
519 assert!(extract_bearer_token_required(&h).is_err());
520 }
521
522 #[test]
523 fn operators_create_request_accepts_capability_manifest() {
524 let req: OperatorsCreateReq = serde_json::from_value(serde_json::json!({
525 "roles": ["main-ai"],
526 "capability_manifest": {
527 "provider_id": "main-ai-self-report",
528 "capabilities": [{
529 "launch_variant": "mse-coder",
530 "resolved_model": "claude-sonnet-4",
531 "effective_tools": ["Read", "Edit"]
532 }]
533 }
534 }))
535 .unwrap();
536 assert_eq!(req.roles, ["main-ai"]);
537 assert_eq!(
538 req.capability_manifest.unwrap().provider_id,
539 "main-ai-self-report"
540 );
541 }
542
543 #[test]
544 fn operators_create_request_keeps_manifest_optional_on_wire() {
545 let req: OperatorsCreateReq =
546 serde_json::from_value(serde_json::json!({ "roles": [] })).unwrap();
547 assert!(req.capability_manifest.is_none());
548 }
549}