use axum::{
extract::{
ws::{Message, WebSocket, WebSocketUpgrade},
Path, State,
},
http::{HeaderMap, StatusCode},
response::{IntoResponse, Response},
Json,
};
use futures_util::{sink::SinkExt, stream::StreamExt};
use mlua_swarm::{AgentProviderManifest, Operator, SeniorBridge, SessionId, SpawnHook};
use serde::{Deserialize, Serialize};
use serde_json::json;
use std::sync::Arc;
use tokio::sync::{mpsc, Mutex};
use super::protocol::{ClientMsg, PendingReply, ServerMsg};
use super::session::WSOperatorSession;
use crate::AppState;
pub struct OperatorSessionEntry {
pub sid: SessionId,
pub token: String,
pub roles: Vec<String>,
pub capability_manifest: Option<AgentProviderManifest>,
pub joined_at_secs: u64,
pub ws_session: Mutex<Option<Arc<WSOperatorSession>>>,
}
#[derive(Debug, Deserialize, Default)]
pub struct OperatorsCreateReq {
#[serde(default)]
pub roles: Vec<String>,
#[serde(default)]
pub capability_manifest: Option<AgentProviderManifest>,
}
#[derive(Debug, Serialize)]
pub struct OperatorsCreateResp {
pub sid: SessionId,
pub token: String,
pub roles: Vec<String>,
}
pub async fn operators_create(
State(state): State<AppState>,
Json(req): Json<OperatorsCreateReq>,
) -> Response {
let roles = req.roles;
let capability_manifest = req.capability_manifest;
let sid = SessionId::new();
let token = mlua_swarm::types::secure_hex(5);
{
let mut map = state.roles_to_sid.lock().await;
let conflicts: Vec<String> = roles
.iter()
.filter(|r| map.contains_key(r.as_str()))
.cloned()
.collect();
if !conflicts.is_empty() {
let conflicts_detail: Vec<serde_json::Value> = conflicts
.iter()
.map(|r| {
let holder = map.get(r.as_str()).map(|sid| sid.to_string());
json!({ "role": r, "sid": holder })
})
.collect();
return (
StatusCode::CONFLICT,
Json(json!({
"error": "roles conflict",
"conflicts": conflicts,
"conflicts_detail": conflicts_detail,
})),
)
.into_response();
}
for r in &roles {
map.insert(r.clone(), sid.clone());
}
}
let joined_at_secs = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0);
let entry = Arc::new(OperatorSessionEntry {
sid: sid.clone(),
token: token.clone(),
roles: roles.clone(),
capability_manifest,
joined_at_secs,
ws_session: Mutex::new(None),
});
state
.operator_sessions
.lock()
.await
.insert(sid.clone(), entry);
(
StatusCode::OK,
Json(OperatorsCreateResp { sid, token, roles }),
)
.into_response()
}
fn extract_bearer_token_required(headers: &HeaderMap) -> Result<String, Box<Response>> {
let token = headers
.get(axum::http::header::AUTHORIZATION)
.and_then(|v| v.to_str().ok())
.and_then(|s| s.strip_prefix("Bearer "))
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty());
token.ok_or_else(|| {
Box::new((StatusCode::UNAUTHORIZED, "missing or empty Bearer token").into_response())
})
}
pub async fn operators_ws_connect(
State(state): State<AppState>,
Path(sid): Path<String>,
headers: HeaderMap,
ws: WebSocketUpgrade,
) -> Response {
let bearer = match extract_bearer_token_required(&headers) {
Ok(t) => t,
Err(resp) => return *resp,
};
let Ok(sid) = SessionId::parse(sid) else {
return (StatusCode::NOT_FOUND, "unknown sid").into_response();
};
let entry = {
let map = state.operator_sessions.lock().await;
map.get(&sid).cloned()
};
let entry = match entry {
Some(e) => e,
None => return (StatusCode::NOT_FOUND, "unknown sid").into_response(),
};
if !mlua_swarm::types::ct_eq(entry.token.as_bytes(), bearer.as_bytes()) {
return (StatusCode::UNAUTHORIZED, "token mismatch").into_response();
}
ws.on_upgrade(move |socket| handle_operator_socket(socket, state, entry))
}
async fn handle_operator_socket(
socket: WebSocket,
state: AppState,
entry: Arc<OperatorSessionEntry>,
) {
let (tx, mut rx) = mpsc::unbounded_channel::<ServerMsg>();
let existing_ws = entry.ws_session.lock().await.clone();
let session = match existing_ws {
Some(ws_session) => {
ws_session.replace_tx(tx.clone()).await;
ws_session
}
None => {
let ws_session = Arc::new(WSOperatorSession::new_with_base_url(
entry.sid.clone(),
tx.clone(),
state.base_url.clone(),
));
state
.engine
.register_senior_bridge(
entry.sid.clone(),
ws_session.clone() as Arc<dyn SeniorBridge>,
)
.await;
state
.engine
.register_spawn_hook(entry.sid.clone(), ws_session.clone() as Arc<dyn SpawnHook>)
.await;
state
.engine
.register_operator(entry.sid.clone(), ws_session.clone() as Arc<dyn Operator>)
.await;
if let Some(factory) = &state.ws_operator_factory {
factory
.register_operator(entry.sid.clone(), ws_session.clone() as Arc<dyn Operator>);
}
for role in &entry.roles {
if let Some(factory) = &state.ws_operator_factory {
factory
.register_operator(role.clone(), ws_session.clone() as Arc<dyn Operator>);
}
state
.engine
.register_operator(role.clone(), ws_session.clone() as Arc<dyn Operator>)
.await;
}
*entry.ws_session.lock().await = Some(ws_session.clone());
ws_session
}
};
let (mut ws_sink, mut ws_stream) = socket.split();
let write_task = tokio::spawn(async move {
while let Some(msg) = rx.recv().await {
let txt = match serde_json::to_string(&msg) {
Ok(s) => s,
Err(_) => continue,
};
if ws_sink.send(Message::Text(txt)).await.is_err() {
break;
}
}
let _ = ws_sink.close().await;
});
let session_for_read = session.clone();
let read_result: Result<(), String> = async {
while let Some(item) = ws_stream.next().await {
match item {
Ok(Message::Text(t)) => {
let parsed: ClientMsg = match serde_json::from_str(&t) {
Ok(p) => p,
Err(_) => continue,
};
match parsed {
ClientMsg::Answer { req_id, value } => {
session_for_read
.resolve_pending(&req_id, PendingReply::Answer(value))
.await;
}
ClientMsg::HookAck { req_id, ok, reason } => {
session_for_read
.resolve_pending(&req_id, PendingReply::HookAck { ok, reason })
.await;
}
ClientMsg::SpawnAck {
req_id,
value,
ok,
error,
stats,
} => {
session_for_read
.resolve_pending(
&req_id,
PendingReply::SpawnAck {
value,
ok,
error,
stats,
},
)
.await;
}
ClientMsg::SpawnHalt {
req_id,
value,
reason,
} => {
session_for_read
.resolve_pending(&req_id, PendingReply::SpawnHalt { value, reason })
.await;
}
}
}
Ok(Message::Ping(_)) | Ok(Message::Pong(_)) => {}
Ok(Message::Close(_)) | Err(_) => break,
_ => {}
}
}
Ok(())
}
.await;
session.clear_tx_if(&tx).await;
write_task.abort();
let _ = read_result;
}
async fn teardown_operator_session(
state: &AppState,
sid: &SessionId,
entry: &Arc<OperatorSessionEntry>,
) {
state.engine.unregister_senior_bridge(sid.as_str()).await;
state.engine.unregister_spawn_hook(sid.as_str()).await;
state.engine.unregister_operator(sid.as_str()).await;
if let Some(factory) = &state.ws_operator_factory {
factory.unregister_operator(sid.as_str());
}
for role in &entry.roles {
state.engine.unregister_operator(role).await;
if let Some(factory) = &state.ws_operator_factory {
factory.unregister_operator(role);
}
}
if let Some(session) = entry.ws_session.lock().await.take() {
session.fail_pending("operator session torn down").await;
session.clear_tx().await;
}
state.operator_sessions.lock().await.remove(sid);
{
let mut map = state.roles_to_sid.lock().await;
for role in &entry.roles {
if map.get(role) == Some(sid) {
map.remove(role);
}
}
}
}
pub async fn operators_delete(
State(state): State<AppState>,
Path(sid): Path<String>,
headers: HeaderMap,
) -> Response {
let bearer = match extract_bearer_token_required(&headers) {
Ok(t) => t,
Err(resp) => return *resp,
};
let Ok(sid) = SessionId::parse(sid) else {
return (StatusCode::NOT_FOUND, "unknown sid").into_response();
};
let entry = {
let map = state.operator_sessions.lock().await;
map.get(&sid).cloned()
};
let entry = match entry {
Some(e) => e,
None => return (StatusCode::NOT_FOUND, "unknown sid").into_response(),
};
if !mlua_swarm::types::ct_eq(entry.token.as_bytes(), bearer.as_bytes()) {
return (StatusCode::UNAUTHORIZED, "token mismatch").into_response();
}
teardown_operator_session(&state, &sid, &entry).await;
StatusCode::NO_CONTENT.into_response()
}
#[derive(Debug, Serialize)]
pub struct OperatorsListEntry {
pub sid: SessionId,
pub roles: Vec<String>,
pub joined_at_secs: u64,
pub connected: bool,
}
#[derive(Debug, Serialize)]
pub struct OperatorsListResp {
pub operators: Vec<OperatorsListEntry>,
}
pub async fn operators_list(State(state): State<AppState>) -> Response {
let entries: Vec<(SessionId, Arc<OperatorSessionEntry>)> = {
let map = state.operator_sessions.lock().await;
map.iter().map(|(k, v)| (k.clone(), v.clone())).collect()
};
let mut operators = Vec::with_capacity(entries.len());
for (sid, entry) in entries {
let session = entry.ws_session.lock().await.clone();
let connected = match session {
Some(session) => session.is_connected().await,
None => false,
};
operators.push(OperatorsListEntry {
sid,
roles: entry.roles.clone(),
joined_at_secs: entry.joined_at_secs,
connected,
});
}
operators.sort_by(|a, b| a.sid.as_str().cmp(b.sid.as_str()));
(StatusCode::OK, Json(OperatorsListResp { operators })).into_response()
}
pub async fn operators_delete_by_role(
State(state): State<AppState>,
Path(role): Path<String>,
axum::extract::Query(query): axum::extract::Query<OperatorsDeleteByRoleQuery>,
) -> Response {
let sid = {
let map = state.roles_to_sid.lock().await;
match map.get(role.as_str()) {
Some(sid) => sid.clone(),
None => {
return (
StatusCode::NOT_FOUND,
Json(json!({"error": "no session holds this role", "role": role})),
)
.into_response();
}
}
};
let entry = {
let map = state.operator_sessions.lock().await;
map.get(&sid).cloned()
};
let entry = match entry {
Some(e) => e,
None => {
let mut map = state.roles_to_sid.lock().await;
if map.get(role.as_str()) == Some(&sid) {
map.remove(role.as_str());
}
return (
StatusCode::NOT_FOUND,
Json(json!({
"error": "torn role mapping cleared; role now open",
"role": role,
})),
)
.into_response();
}
};
if !query.force {
match active_runs_for_sid(&state, &sid).await {
Ok(active_runs) if !active_runs.is_empty() => {
return (
StatusCode::CONFLICT,
Json(json!({
"error": "session is driving in-flight runs; \
tearing it down would fail their parked spawns",
"role": role,
"sid": sid,
"active_runs": active_runs,
"hint": "wait for the runs to finish, or repeat with ?force=true \
to tear the session down anyway",
})),
)
.into_response();
}
Ok(_) => {}
Err(error) => {
tracing::warn!(%role, %sid, %error, "operators_delete_by_role: run occupancy check failed");
return (
StatusCode::CONFLICT,
Json(json!({
"error": format!(
"cannot verify whether this session is driving in-flight runs: {error}"
),
"role": role,
"sid": sid,
"hint": "retry, or repeat with ?force=true to tear the session down \
without the check",
})),
)
.into_response();
}
}
}
teardown_operator_session(&state, &sid, &entry).await;
StatusCode::NO_CONTENT.into_response()
}
#[derive(Debug, Deserialize, Default)]
pub struct OperatorsDeleteByRoleQuery {
#[serde(default)]
pub force: bool,
}
async fn active_runs_for_sid(state: &AppState, sid: &SessionId) -> Result<Vec<String>, String> {
let mut running = state
.run_store
.list_running()
.await
.map_err(|e| e.to_string())?;
running.sort_by_key(|record| record.created_at);
Ok(running
.into_iter()
.filter(|record| record.operator_sid.as_deref() == Some(sid.as_str()))
.map(|record| record.id.to_string())
.collect())
}
#[derive(Debug, Serialize)]
pub struct OperatorsInfoResp {
pub sid: SessionId,
pub roles: Vec<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub capability_manifest: Option<AgentProviderManifest>,
pub connected: bool,
}
pub async fn operators_info(
State(state): State<AppState>,
Path(sid): Path<String>,
headers: HeaderMap,
) -> Response {
let bearer = match extract_bearer_token_required(&headers) {
Ok(t) => t,
Err(resp) => return *resp,
};
let Ok(sid) = SessionId::parse(sid) else {
return (StatusCode::NOT_FOUND, "unknown sid").into_response();
};
let entry = {
let map = state.operator_sessions.lock().await;
map.get(&sid).cloned()
};
let entry = match entry {
Some(e) => e,
None => return (StatusCode::NOT_FOUND, "unknown sid").into_response(),
};
if !mlua_swarm::types::ct_eq(entry.token.as_bytes(), bearer.as_bytes()) {
return (StatusCode::UNAUTHORIZED, "token mismatch").into_response();
}
let session = entry.ws_session.lock().await.clone();
let connected = match session {
Some(session) => session.is_connected().await,
None => false,
};
(
StatusCode::OK,
Json(OperatorsInfoResp {
sid: entry.sid.clone(),
roles: entry.roles.clone(),
capability_manifest: entry.capability_manifest.clone(),
connected,
}),
)
.into_response()
}
#[cfg(test)]
mod tests {
use super::*;
use axum::http::HeaderValue;
fn headers_with_bearer(token: &str) -> HeaderMap {
let mut h = HeaderMap::new();
h.insert(
axum::http::header::AUTHORIZATION,
HeaderValue::from_str(&format!("Bearer {token}")).unwrap(),
);
h
}
#[test]
fn extract_bearer_token_required_accepts_valid() {
let h = headers_with_bearer("abc123");
assert_eq!(extract_bearer_token_required(&h).unwrap(), "abc123");
}
#[test]
fn extract_bearer_token_required_rejects_missing_header() {
let h = HeaderMap::new();
assert!(extract_bearer_token_required(&h).is_err());
}
#[test]
fn extract_bearer_token_required_rejects_empty_token() {
let h = headers_with_bearer("");
assert!(extract_bearer_token_required(&h).is_err());
}
#[test]
fn extract_bearer_token_required_rejects_wrong_scheme() {
let mut h = HeaderMap::new();
h.insert(
axum::http::header::AUTHORIZATION,
HeaderValue::from_static("Basic dXNlcjpwYXNz"),
);
assert!(extract_bearer_token_required(&h).is_err());
}
#[test]
fn operators_create_request_accepts_capability_manifest() {
let req: OperatorsCreateReq = serde_json::from_value(serde_json::json!({
"roles": ["main-ai"],
"capability_manifest": {
"provider_id": "main-ai-self-report",
"capabilities": [{
"launch_variant": "mse-coder",
"resolved_model": "claude-sonnet-4",
"effective_tools": ["Read", "Edit"]
}]
}
}))
.unwrap();
assert_eq!(req.roles, ["main-ai"]);
assert_eq!(
req.capability_manifest.unwrap().provider_id,
"main-ai-self-report"
);
}
#[test]
fn operators_create_request_keeps_manifest_optional_on_wire() {
let req: OperatorsCreateReq =
serde_json::from_value(serde_json::json!({ "roles": [] })).unwrap();
assert!(req.capability_manifest.is_none());
}
mod by_role_in_flight {
use super::*;
use mlua_swarm::core::config::EngineCfg;
use mlua_swarm::core::engine::Engine;
use mlua_swarm::store::output::InMemoryOutputStore;
use mlua_swarm::store::run::{InMemoryRunStore, RunRecord, RunStatus};
use mlua_swarm::store::task::InMemoryTaskStore;
use mlua_swarm::RunId;
use mlua_swarm::TaskId;
use std::collections::HashMap;
fn test_state() -> AppState {
let engine =
Engine::new_with_layers(EngineCfg::default(), crate::default_layer_registry());
let compiler = mlua_swarm::Compiler::new(crate::default_registry());
let launch = Arc::new(mlua_swarm::TaskLaunchService::new(engine.clone(), compiler));
AppState {
engine,
sessions: Arc::new(Mutex::new(crate::SessionStore::default())),
task_app: Arc::new(mlua_swarm::TaskApplication::new_inline_only(launch)),
ws_operator_factory: None,
data_store: Arc::new(InMemoryOutputStore::new()),
operator_sessions: Arc::new(Mutex::new(HashMap::new())),
roles_to_sid: Arc::new(Mutex::new(HashMap::new())),
task_store: Arc::new(InMemoryTaskStore::new()),
run_store: Arc::new(InMemoryRunStore::new()),
replay_store: Arc::new(mlua_swarm::store::replay::InMemoryReplayStore::new()),
run_trace_store: Arc::new(mlua_swarm::store::trace::InMemoryRunTraceStore::new()),
base_url: None,
sync_timeout_secs: 300,
}
}
async fn seed_session(state: &AppState, role: &str) -> SessionId {
let sid = SessionId::new();
let entry = Arc::new(OperatorSessionEntry {
sid: sid.clone(),
token: "token".to_string(),
roles: vec![role.to_string()],
capability_manifest: None,
joined_at_secs: 0,
ws_session: Mutex::new(None),
});
state
.operator_sessions
.lock()
.await
.insert(sid.clone(), entry);
state
.roles_to_sid
.lock()
.await
.insert(role.to_string(), sid.clone());
sid
}
async fn seed_run(state: &AppState, sid: Option<&SessionId>, status: RunStatus) -> RunId {
let run_id = RunId::new();
state
.run_store
.create(RunRecord {
id: run_id.clone(),
task_id: TaskId::new(),
status,
step_entries: Vec::new(),
degradations: Vec::new(),
operator_sid: sid.map(|s| s.to_string()),
result_ref: None,
input_json: None,
created_at: 0,
updated_at: 0,
})
.await
.expect("seed run");
run_id
}
async fn body_json(response: Response) -> serde_json::Value {
let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.expect("read body");
serde_json::from_slice(&bytes).expect("json body")
}
async fn session_is_live(state: &AppState, sid: &SessionId) -> bool {
state.operator_sessions.lock().await.contains_key(sid)
}
#[tokio::test]
async fn idle_holder_is_still_torn_down() {
let state = test_state();
let sid = seed_session(&state, "main-ai").await;
let other = SessionId::new();
seed_run(&state, Some(&other), RunStatus::Running).await;
seed_run(&state, Some(&sid), RunStatus::Done).await;
let response = operators_delete_by_role(
State(state.clone()),
Path("main-ai".to_string()),
axum::extract::Query(OperatorsDeleteByRoleQuery::default()),
)
.await;
assert_eq!(response.status(), StatusCode::NO_CONTENT);
assert!(!session_is_live(&state, &sid).await);
}
#[tokio::test]
async fn holder_driving_a_running_run_is_refused_with_its_run_ids() {
let state = test_state();
let sid = seed_session(&state, "main-ai").await;
let run_id = seed_run(&state, Some(&sid), RunStatus::Running).await;
let response = operators_delete_by_role(
State(state.clone()),
Path("main-ai".to_string()),
axum::extract::Query(OperatorsDeleteByRoleQuery::default()),
)
.await;
assert_eq!(response.status(), StatusCode::CONFLICT);
let body = body_json(response).await;
assert_eq!(
body["active_runs"],
serde_json::json!([run_id.to_string()]),
"the 409 must name the in-flight runs: {body}"
);
assert!(
session_is_live(&state, &sid).await,
"a refused teardown must leave the session (and its parked spawns) alone"
);
assert_eq!(
state.roles_to_sid.lock().await.get("main-ai"),
Some(&sid),
"a refused teardown must not release the role"
);
}
#[tokio::test]
async fn force_tears_down_despite_in_flight_runs() {
let state = test_state();
let sid = seed_session(&state, "main-ai").await;
seed_run(&state, Some(&sid), RunStatus::Running).await;
let response = operators_delete_by_role(
State(state.clone()),
Path("main-ai".to_string()),
axum::extract::Query(OperatorsDeleteByRoleQuery { force: true }),
)
.await;
assert_eq!(response.status(), StatusCode::NO_CONTENT);
assert!(!session_is_live(&state, &sid).await);
}
#[tokio::test]
async fn unknown_role_still_404s() {
let state = test_state();
let response = operators_delete_by_role(
State(state),
Path("nobody-holds-this".to_string()),
axum::extract::Query(OperatorsDeleteByRoleQuery::default()),
)
.await;
assert_eq!(response.status(), StatusCode::NOT_FOUND);
}
}
}