Skip to main content

mlua_swarm_server/operator_ws/
login.rs

1//! REST-like Operator session resource.
2//!
3//! Provides the `POST/GET/DELETE /v1/operators` + `WS /v1/operators/:sid/ws`
4//! route family — the sole WS Operator session route. `session.rs` /
5//! `protocol.rs` are unchanged by this module.
6//!
7//! ## Login flow
8//!
9//! ```text
10//! POST /v1/operators { roles?: ["main-ai"], capability_manifest?: {...} }
11//!   → 409 if any role already owns a live entry (roles alias exclusivity,
12//!     v1.md §Auth session flow)
13//!   → { sid: "S-<hex>", token: "<10-hex>", roles: [...] }
14//!   The manifest is pinned to this session and later resolved through the
15//!   Core `AgentBindingProvider` interface before any Runner-backed spawn.
16//!
17//! WS /v1/operators/:sid/ws
18//!   Authorization: Bearer <token>   (mandatory — no empty-string default)
19//!   → 401 missing/empty Bearer, 404 unknown sid, 401 token mismatch
20//!   → registers a `WSOperatorSession` into the engine's 3 registries
21//!     (senior_bridge / spawn_hook / operator) + role aliases, same pattern
22//!     as `handler::handle_socket`. Reconnect (same sid, matching token)
23//!     reuses the existing `WSOperatorSession` via `replace_tx`.
24//!
25//! DELETE /v1/operators/:sid   (Bearer required)
26//!   → unregisters the 3 registries + role aliases + `operator_sessions`
27//!     entry + releases `roles_to_sid` ownership.
28//!
29//! GET /v1/operators/:sid   (Bearer required)
30//!   → { sid, roles, connected }
31//! ```
32//!
33//! `OperatorSessionEntry` is the login-flow record (`AppState.operator_sessions`),
34//! distinct from `mlua_swarm::OperatorSession` (the engine-side
35//! `attach`/session-token record) and from `WSOperatorSession` (the 3-trait WS
36//! session, `session.rs`) — this module owns the mapping `sid → (token, roles,
37//! Option<WSOperatorSession>)` that the login flow is built on.
38
39use 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
59/// Login-flow record for a minted Operator session. Held in
60/// `AppState.operator_sessions`, keyed by `sid`. `ws_session` starts `None`
61/// (login only mints sid+token) and is set on first successful WS connect;
62/// on reconnect the same `WSOperatorSession` is reused (`replace_tx`) rather
63/// than re-registered.
64pub struct OperatorSessionEntry {
65    /// Server-minted session id (typed [`SessionId`] since issue #14).
66    pub sid: SessionId,
67    /// Bearer auth token (10-hex-char) required on the WS upgrade and admin routes.
68    pub token: String,
69    /// Role aliases claimed by this session (roles-exclusivity set).
70    pub roles: Vec<String>,
71    /// Provider-owned effective capability manifest submitted at join.
72    pub capability_manifest: Option<AgentProviderManifest>,
73    /// The reusable 3-trait session object once a WS has connected at least
74    /// once; `None` before first connect. Its sender tracks current connectivity.
75    pub ws_session: Mutex<Option<Arc<WSOperatorSession>>>,
76}
77
78// ─── POST /v1/operators (mint) ──────────────────────────────────────────────
79
80/// Body for `POST /v1/operators`.
81#[derive(Debug, Deserialize, Default)]
82pub struct OperatorsCreateReq {
83    /// Role aliases to claim exclusively (empty = no exclusivity claimed).
84    #[serde(default)]
85    pub roles: Vec<String>,
86    /// Effective execution capabilities supplied by the Operator/MainAI.
87    #[serde(default)]
88    pub capability_manifest: Option<AgentProviderManifest>,
89}
90
91/// Response for `POST /v1/operators`.
92#[derive(Debug, Serialize)]
93pub struct OperatorsCreateResp {
94    /// Newly minted session id (typed [`SessionId`]; serializes as the
95    /// plain `S-<hex>` string — the wire shape is unchanged).
96    pub sid: SessionId,
97    /// Bearer auth token required on the WS upgrade and admin routes.
98    pub token: String,
99    /// Echoes the granted role aliases.
100    pub roles: Vec<String>,
101}
102
103/// `POST /v1/operators`. Mints `sid` (`S-<hex>` — the shared `SessionId`
104/// shape; issue #11) + a 10-hex-char token
105/// (`mlua_swarm::types::secure_hex(5)` — OS-RNG hex, unguessable across
106/// calls and restarts, which is the point: this token is the sole bearer
107/// secret on the short-handle path). When `roles` is non-empty, checks
108/// `AppState.roles_to_sid` for conflicts under a single lock (check + insert
109/// atomic w.r.t. concurrent mints) and returns `409 CONFLICT` with the
110/// conflicting role names on collision. Empty `roles` never conflicts (= no
111/// exclusivity is claimed).
112pub 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    // The sid is the operator-session identity, so it mints in the same
119    // `SessionId` shape (`S-<hex>`) as the engine-side session id — one
120    // session-id form across the system (issue #11 observation 2; the old
121    // `op-<uuid>` shape collided with the operator-backend registry prefix).
122    // It is an identifier, not a secret: `token` (secure_hex) is the sole
123    // bearer credential on this path.
124    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
166// ─── WS /v1/operators/:sid/ws (Bearer required) ─────────────────────────────
167
168/// Extracts `Authorization: Bearer <token>`; missing header, wrong scheme, or
169/// an empty token all resolve to a `401` response. `Authorization` is
170/// mandatory on the WS path — there is no empty-string default.
171fn 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
183/// `GET /v1/operators/:sid/ws` (WS upgrade). Bearer mandatory. `404` on
184/// unknown sid, `401` on token mismatch. On successful upgrade, registers (or
185/// reuses, on reconnect) a `WSOperatorSession` under `sid` — same 3-registry
186/// pattern as `handler::handle_socket`, plus role-alias registration for
187/// every role minted alongside this sid.
188pub 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    // A string that doesn't even parse as a SessionId can't be a known sid.
199    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
218/// Bidirectional pump for a single WS connection, bound to an
219/// `OperatorSessionEntry`. Owns the full wire protocol pump (write task /
220/// read task / `ClientMsg` dispatch / disconnect) for this session.
221async 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            // Reconnect: reuse the existing WSOperatorSession on this entry; only swap out `tx`.
232            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            // Role exclusivity was already resolved at login (POST) time. Here
261            // we just bind the same session into the three registries + factory
262            // under its role aliases (same shape as handler::handle_socket's
263            // ?roles= path).
264            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    // write task: mpsc → WebSocket
282    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    // read task: WS message → ClientMsg parse → session.resolve_pending
296    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    // Clear only this socket's sender. A reconnect may already have installed
350    // a replacement while this older socket was unwinding.
351    session.clear_tx_if(&tx).await;
352    write_task.abort();
353    let _ = read_result;
354}
355
356// ─── DELETE /v1/operators/:sid (Bearer required) ────────────────────────────
357
358/// `DELETE /v1/operators/:sid`. Bearer mandatory. `404` on unknown sid, `401`
359/// on token mismatch. Drops the 3 engine registries + role aliases +
360/// `ws_operator_factory` bindings + `operator_sessions` entry, and releases
361/// this sid's ownership in `roles_to_sid` (re-opening the role names for a
362/// future mint).
363pub 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// ─── GET /v1/operators/:sid (Bearer required) ───────────────────────────────
420
421/// Response for `GET /v1/operators/:sid`.
422#[derive(Debug, Serialize)]
423pub struct OperatorsInfoResp {
424    /// Echoes the requested session id.
425    pub sid: SessionId,
426    /// Role aliases held by this session.
427    pub roles: Vec<String>,
428    /// Capability manifest pinned when this session joined.
429    #[serde(skip_serializing_if = "Option::is_none")]
430    pub capability_manifest: Option<AgentProviderManifest>,
431    /// Whether a WS is currently attached (not merely that the session ever connected).
432    pub connected: bool,
433}
434
435/// `GET /v1/operators/:sid`. Bearer mandatory. `404` on unknown sid, `401` on
436/// token mismatch. `connected` reflects whether the reusable session currently
437/// owns a live sender, not merely whether it connected at least once.
438pub 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}