#![warn(missing_docs)]
pub mod binding;
pub mod blueprints;
pub mod config;
pub mod data;
pub mod doctor;
pub mod enhance_log;
pub mod enhance_settings;
pub mod issues;
pub mod operator_ws;
pub mod projection;
pub mod tasks;
pub mod worker;
pub use blueprints::{
build_blueprints_router, build_blueprints_router_with_refs, BindingRequirementsResponse,
};
pub use enhance_log::build_enhance_log_router;
pub use enhance_settings::build_enhance_settings_router;
pub use issues::{build_issues_router, GetIssueResponse, PostIssueRequest, PostIssueResponse};
pub use operator_ws::{
operators_create, operators_delete, operators_info, operators_ws_connect, ClientMsg,
OperatorSessionEntry, ServerMsg, WSOperatorSession,
};
pub use projection::{McpQueryAdapter, ProjectionSource, StepList, StepPathQuery, StepSummary};
pub use tasks::{
RunBindingDifference, RunBindingExplainEntry, RunBindingStatus, RunBindingsExplainResponse,
RunKickRequest, RunKickResponse, RunResumeResponse, TaskDetailResponse,
};
pub use worker::{
worker_artifact, worker_prompt, worker_result, ArtifactQuery, PromptQuery, WorkerResultReq,
};
use axum::{
extract::{DefaultBodyLimit, State},
http::{header::AUTHORIZATION, HeaderMap, StatusCode},
response::{IntoResponse, Response},
routing::{get, post},
Json, Router,
};
use mlua_swarm::application::{BlueprintRef, TaskApplication, TaskApplicationError};
use mlua_swarm::blueprint::store::BlueprintStore;
use mlua_swarm::core::config::CheckPolicy;
use mlua_swarm::service::{TaskLaunchError, TaskLaunchService};
use mlua_swarm::store::replay::{InMemoryReplayStore, ReplayStore};
use mlua_swarm::store::run::{RunContext, RunRecord, RunStatus, RunStore};
use mlua_swarm::store::task::{TaskRecord, TaskRecordStatus, TaskStore};
use mlua_swarm::{
CapToken, Compiler, Engine, LayerRegistry, LuaInProcessSpawnerFactory, MainAIMiddleware,
OperatorDelegateMiddleware, OperatorSpawnerFactory, Role, RunId, RustFnInProcessSpawnerFactory,
SeniorEscalationMiddleware, SessionId, SpawnerRegistry, SubprocessProcessSpawnerFactory,
TaskId,
};
use serde::{Deserialize, Serialize};
use serde_json::{json, Value};
use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::Mutex;
#[derive(Default)]
pub struct SessionStore {
pub map: HashMap<String, CapToken>,
}
#[derive(Clone)]
pub struct AppState {
pub engine: Engine,
pub sessions: Arc<Mutex<SessionStore>>,
pub task_app: Arc<TaskApplication>,
pub ws_operator_factory: Option<Arc<OperatorSpawnerFactory>>,
pub data_store: Arc<dyn mlua_swarm::store::output::OutputStore>,
pub operator_sessions:
Arc<Mutex<HashMap<SessionId, Arc<crate::operator_ws::login::OperatorSessionEntry>>>>,
pub roles_to_sid: Arc<Mutex<HashMap<String, SessionId>>>,
pub task_store: Arc<dyn TaskStore>,
pub run_store: Arc<dyn RunStore>,
pub replay_store: Arc<dyn ReplayStore>,
pub base_url: Option<Arc<str>>,
pub sync_timeout_secs: u64,
}
pub fn build_router(engine: Engine) -> Router {
build_router_with(engine, default_registry(), None)
}
pub fn default_layer_registry() -> LayerRegistry {
LayerRegistry::new()
.with_hint("main_ai", |_engine| Arc::new(MainAIMiddleware::new()))
.with_hint("senior_escalation", |_engine| {
Arc::new(SeniorEscalationMiddleware::new())
})
.with_hint("operator_delegate", |_engine| {
Arc::new(OperatorDelegateMiddleware::new())
})
}
pub fn build_router_with(
engine: Engine,
registry: SpawnerRegistry,
store: Option<Arc<dyn BlueprintStore>>,
) -> Router {
build_router_with_ws_factory(engine, registry, store, None)
}
pub fn build_router_with_ws_factory(
engine: Engine,
registry: SpawnerRegistry,
store: Option<Arc<dyn BlueprintStore>>,
ws_operator_factory: Option<Arc<OperatorSpawnerFactory>>,
) -> Router {
build_router_with_ws_factory_and_output(engine, registry, store, ws_operator_factory, None)
}
pub fn build_router_with_ws_factory_and_output(
engine: Engine,
registry: SpawnerRegistry,
store: Option<Arc<dyn BlueprintStore>>,
ws_operator_factory: Option<Arc<OperatorSpawnerFactory>>,
output_store: Option<Arc<dyn mlua_swarm::store::output::OutputStore>>,
) -> Router {
build_router_full(
engine,
registry,
store,
ws_operator_factory,
output_store,
None,
None,
None,
None,
crate::config::default_sync_timeout_secs(),
)
}
#[allow(clippy::too_many_arguments)]
pub fn build_router_full(
engine: Engine,
registry: SpawnerRegistry,
store: Option<Arc<dyn BlueprintStore>>,
ws_operator_factory: Option<Arc<OperatorSpawnerFactory>>,
output_store: Option<Arc<dyn mlua_swarm::store::output::OutputStore>>,
base_url: Option<Arc<str>>,
task_store: Option<Arc<dyn TaskStore>>,
run_store: Option<Arc<dyn RunStore>>,
replay_store: Option<Arc<dyn ReplayStore>>,
sync_timeout_secs: u64,
) -> Router {
build_router_full_with_legacy_worker_binding_policy(
engine,
registry,
store,
ws_operator_factory,
output_store,
base_url,
task_store,
run_store,
replay_store,
sync_timeout_secs,
mlua_swarm::LegacyWorkerBindingPolicy::Allow,
)
}
#[allow(clippy::too_many_arguments)]
pub fn build_router_full_with_legacy_worker_binding_policy(
engine: Engine,
registry: SpawnerRegistry,
store: Option<Arc<dyn BlueprintStore>>,
ws_operator_factory: Option<Arc<OperatorSpawnerFactory>>,
output_store: Option<Arc<dyn mlua_swarm::store::output::OutputStore>>,
base_url: Option<Arc<str>>,
task_store: Option<Arc<dyn TaskStore>>,
run_store: Option<Arc<dyn RunStore>>,
replay_store: Option<Arc<dyn ReplayStore>>,
sync_timeout_secs: u64,
legacy_worker_binding_policy: mlua_swarm::LegacyWorkerBindingPolicy,
) -> Router {
let operator_sessions = Arc::new(Mutex::new(HashMap::new()));
let roles_to_sid = Arc::new(Mutex::new(HashMap::new()));
let compiler = Compiler::new(registry);
let binding_provider = Arc::new(binding::OperatorSessionBindingProvider::new(
operator_sessions.clone(),
roles_to_sid.clone(),
));
let launch = Arc::new(
TaskLaunchService::new(engine.clone(), compiler)
.with_binding_provider(binding_provider)
.with_legacy_worker_binding_policy(legacy_worker_binding_policy),
);
let task_app = Arc::new(match store {
Some(s) => TaskApplication::new(launch, s),
None => TaskApplication::new_inline_only(launch),
});
let data_store: Arc<dyn mlua_swarm::store::output::OutputStore> = match output_store {
Some(s) => s,
None => Arc::new(mlua_swarm::store::output::InMemoryOutputStore::new()),
};
engine.set_output_store(data_store.clone());
let task_store: Arc<dyn TaskStore> = match task_store {
Some(s) => s,
None => Arc::new(mlua_swarm::store::task::InMemoryTaskStore::new()),
};
let run_store: Arc<dyn RunStore> = match run_store {
Some(s) => s,
None => Arc::new(mlua_swarm::store::run::InMemoryRunStore::new()),
};
let replay_store: Arc<dyn ReplayStore> = match replay_store {
Some(s) => s,
None => Arc::new(InMemoryReplayStore::new()),
};
let state = AppState {
engine,
sessions: Arc::new(Mutex::new(SessionStore::default())),
task_app,
ws_operator_factory,
data_store,
operator_sessions,
roles_to_sid,
task_store,
run_store,
replay_store,
base_url,
sync_timeout_secs,
};
Router::new()
.route("/v1/healthz", get(healthz))
.route("/v1/status", get(status_get))
.route(
"/v1/sessions",
post(sessions_attach).delete(sessions_detach),
)
.route("/v1/tasks", post(tasks_start).get(tasks::tasks_list))
.route("/v1/tasks/:id", get(tasks::task_get))
.route("/v1/tasks/:id/runs", post(tasks::task_rekick))
.route("/v1/tasks/:id/runs/:run/steps", get(projection::steps_list))
.route(
"/v1/tasks/:id/runs/:run/steps/:step",
get(projection::step_get),
)
.route(
"/v1/tasks/:id/runs/:run/steps/:step/content",
get(projection::step_content),
)
.route("/v1/runs/:id", get(tasks::run_get))
.route("/v1/runs/:id/bindings", get(tasks::run_bindings_explain))
.route("/v1/runs/:id/resume", post(tasks::run_resume))
.route("/v1/runs/:id/rerun-from", post(tasks::run_rerun_from))
.route("/v1/operators", post(operators_create))
.route("/v1/operators/:sid/ws", get(operators_ws_connect))
.route(
"/v1/operators/:sid",
get(operators_info).delete(operators_delete),
)
.route("/v1/worker/prompt", get(worker::worker_prompt))
.route("/v1/worker/result", post(worker::worker_result))
.route(
"/v1/worker/submit",
post(worker::worker_submit).layer(DefaultBodyLimit::max(2 * 1024 * 1024)),
)
.route(
"/v1/worker/artifact",
post(worker::worker_artifact).layer(DefaultBodyLimit::max(2 * 1024 * 1024)),
)
.route(
"/v1/worker/prompt/system",
get(worker::worker_prompt_system),
)
.route(
"/v1/agents/:name/render-size",
get(worker::agent_render_size),
)
.route("/v1/worker/degradation", post(worker::worker_degradation))
.route("/v1/data/emit", post(data::data_emit))
.route(
"/v1/data/:key",
get(data::data_get).post(data::data_emit_named),
)
.with_state(state)
}
pub fn default_registry() -> SpawnerRegistry {
let rustfn_factory =
mlua_swarm::worker::baseline::extend_with_baseline(RustFnInProcessSpawnerFactory::new());
let mut reg = SpawnerRegistry::new();
reg.register::<SubprocessProcessSpawnerFactory>(Arc::new(SubprocessProcessSpawnerFactory));
reg.register::<RustFnInProcessSpawnerFactory>(Arc::new(rustfn_factory));
reg.register::<LuaInProcessSpawnerFactory>(Arc::new(LuaInProcessSpawnerFactory::new()));
reg.register::<OperatorSpawnerFactory>(Arc::new(OperatorSpawnerFactory::new()));
reg
}
pub fn default_registry_with_enhance_flow() -> SpawnerRegistry {
let lua_factory =
mlua_swarm::enhance::blueprint::extend_factory(LuaInProcessSpawnerFactory::new());
let agent_block_factory =
mlua_swarm::worker::agent_block::AgentBlockInProcessSpawnerFactory::new();
let rustfn_factory =
mlua_swarm::worker::baseline::extend_with_baseline(RustFnInProcessSpawnerFactory::new());
let mut reg = SpawnerRegistry::new();
reg.register::<SubprocessProcessSpawnerFactory>(Arc::new(SubprocessProcessSpawnerFactory));
reg.register::<RustFnInProcessSpawnerFactory>(Arc::new(rustfn_factory));
reg.register::<LuaInProcessSpawnerFactory>(Arc::new(lua_factory));
reg.register::<mlua_swarm::worker::agent_block::AgentBlockInProcessSpawnerFactory>(Arc::new(
agent_block_factory,
));
reg.register::<OperatorSpawnerFactory>(Arc::new(OperatorSpawnerFactory::new()));
reg
}
async fn healthz() -> &'static str {
"ok"
}
#[derive(Debug, Clone, Serialize, schemars::JsonSchema)]
pub struct StatusResponse {
pub running_runs: usize,
pub attached_operators: usize,
}
async fn status_get(State(state): State<AppState>) -> Json<StatusResponse> {
let running_runs = state
.run_store
.list_running()
.await
.map(|v| v.len())
.unwrap_or_else(|e| {
tracing::warn!(error = %e, "status_get: list_running failed");
0
});
let attached_operators = state.engine.list_operator_ids().await.len();
Json(StatusResponse {
running_runs,
attached_operators,
})
}
#[derive(Deserialize)]
struct AttachReq {
agent_id: String,
role: String,
ttl_secs: u64,
}
#[derive(Serialize)]
struct AttachResp {
session_id: String,
role: String,
}
async fn sessions_attach(
State(state): State<AppState>,
Json(req): Json<AttachReq>,
) -> Result<Json<AttachResp>, ApiError> {
let role = parse_role(&req.role)?;
let token = state
.engine
.attach(req.agent_id, role, Duration::from_secs(req.ttl_secs))
.await
.map_err(ApiError::engine)?;
let sid = token.nonce.clone();
let key = token.fingerprint();
state.sessions.lock().await.map.insert(key, token);
Ok(Json(AttachResp {
session_id: sid,
role: req.role,
}))
}
async fn sessions_detach(
State(state): State<AppState>,
headers: HeaderMap,
) -> Result<StatusCode, ApiError> {
let sid = extract_bearer(&headers)?;
let token = take_session_token(&state, &sid).await?;
state
.engine
.detach(&token)
.await
.map_err(ApiError::engine)?;
Ok(StatusCode::NO_CONTENT)
}
#[derive(Deserialize, schemars::JsonSchema)]
pub struct TaskLaunchRequest {
#[schemars(with = "Value")]
blueprint: BlueprintRef,
#[schemars(with = "Value")]
init_ctx: Value,
#[serde(default)]
project_root: Option<String>,
#[serde(default)]
work_dir: Option<String>,
#[serde(default)]
#[schemars(with = "Option<Value>")]
task_metadata: Option<Value>,
#[serde(default)]
ttl_secs: Option<u64>,
#[serde(default)]
operator: Option<OperatorReq>,
#[serde(default)]
operator_sid: Option<String>,
#[serde(default)]
timeout_secs: Option<u64>,
#[serde(default)]
goal: Option<String>,
#[serde(default)]
check_policy: Option<CheckPolicy>,
#[serde(default)]
detach: bool,
}
#[derive(Deserialize, Default, schemars::JsonSchema)]
pub struct OperatorReq {
#[serde(default)]
kind: Option<String>,
#[serde(default)]
id: Option<String>,
#[serde(default)]
spawn_hook_id: Option<String>,
#[serde(default)]
senior_bridge_id: Option<String>,
#[serde(default)]
operator_backend_id: Option<String>,
#[serde(default)]
per_agent_kinds: Option<HashMap<String, String>>,
}
fn parse_operator_kind_str(s: &str) -> Result<mlua_swarm::OperatorKind, ApiError> {
use mlua_swarm::OperatorKind;
match s {
"main_ai" => Ok(OperatorKind::MainAi),
"composite" => Ok(OperatorKind::Composite),
"automate" => Ok(OperatorKind::Automate),
other => Err(ApiError::bad_request(format!(
"operator kind: unknown value '{other}' (expected main_ai|automate|composite)"
))),
}
}
#[derive(Serialize, schemars::JsonSchema)]
pub struct TaskLaunchResponse {
#[schemars(with = "Value")]
final_ctx: Value,
bound_version: Option<String>,
effective_ttl_secs: u64,
ttl_source: TtlSource,
#[schemars(with = "String")]
task_id: TaskId,
#[schemars(with = "String")]
run_id: RunId,
status: RunStatus,
}
pub struct TaskLaunchReply(pub TaskLaunchResponse, pub StatusCode);
impl IntoResponse for TaskLaunchReply {
fn into_response(self) -> Response {
(self.1, Json(self.0)).into_response()
}
}
#[derive(Serialize, Clone, Copy, Debug, PartialEq, Eq, schemars::JsonSchema)]
#[serde(rename_all = "snake_case")]
pub enum TtlSource {
RequestBody,
BpMetadata,
ServerDefault,
}
async fn tasks_start(
State(state): State<AppState>,
Json(req): Json<TaskLaunchRequest>,
) -> Result<TaskLaunchReply, ApiError> {
run_flow_form(&state, req).await
}
async fn run_flow_form(
state: &AppState,
req: TaskLaunchRequest,
) -> Result<TaskLaunchReply, ApiError> {
use mlua_swarm::application::{BlueprintRef as AppBlueprintRef, TaskApplicationInput};
use mlua_swarm::OperatorKind;
let blueprint_ref_json = serde_json::to_value(&req.blueprint)
.map_err(|e| ApiError::bad_request(format!("blueprint snapshot: {e}")))?;
let input_ctx_snapshot = req.init_ctx.clone();
let goal = req.goal.clone().unwrap_or_default();
let task_input_spec = build_task_input_spec_from_request(&req);
let task_input_spec_snapshot = task_input_spec
.clone()
.map(|spec| serde_json::to_value(&spec))
.transpose()
.map_err(|e| ApiError::bad_request(format!("task_input_spec snapshot: {e}")))?;
let init_ctx = req.init_ctx.clone();
let mut op_req = req.operator.unwrap_or_default();
if let Some(sid) = &req.operator_sid {
let known_ids = state.engine.list_operator_ids().await;
if !known_ids.iter().any(|id| id == sid) {
return Err(ApiError::bad_request(format!(
"operator_sid: no such registered operator session '{sid}'"
)));
}
op_req.operator_backend_id = Some(sid.clone());
}
let detach = req.detach;
let sync_timeout_secs = match (detach, req.timeout_secs) {
(true, Some(_)) => {
return Err(ApiError::bad_request(
"timeout_secs is the synchronous launch ceiling and does not apply to a \
detached launch (detach: true), whose lifetime bound is ttl_secs — omit \
timeout_secs"
.into(),
));
}
(false, Some(0)) => {
return Err(ApiError::bad_request(
"timeout_secs: 0 is invalid; omit the field to use the server default".into(),
));
}
(false, Some(v)) => v,
(_, None) => state.sync_timeout_secs,
};
if let Some(backend_id) = op_req.operator_backend_id.as_deref() {
let attached = state.engine.list_operator_ids().await;
if attached.is_empty() {
return Err(ApiError::unavailable(format!(
"no operator attached to serve this launch (operator backend '{backend_id}' \
requested): attach an operator via POST /v1/operators + WS, or use the \
poll-style flow (GET /v1/worker/prompt + POST /v1/worker/submit)"
)));
}
}
let operator_kind = op_req
.kind
.as_deref()
.map(parse_operator_kind_str)
.transpose()?;
let operator_id = op_req.id.unwrap_or_else(|| "http-run".to_string());
let mut operator_kind_overrides: HashMap<String, OperatorKind> = HashMap::new();
for (agent, kind_str) in op_req.per_agent_kinds.take().unwrap_or_default() {
operator_kind_overrides.insert(agent, parse_operator_kind_str(&kind_str)?);
}
let blueprint: AppBlueprintRef = match req.blueprint {
AppBlueprintRef::Inline { value } => AppBlueprintRef::Inline { value },
AppBlueprintRef::Id { id, version } => AppBlueprintRef::Id { id, version },
};
let (ttl_secs, ttl_source) = match req.ttl_secs {
Some(v) => (v, TtlSource::RequestBody),
None => {
let (resolved_bp, _ver) = state
.task_app
.resolve(&blueprint)
.await
.map_err(|e| ApiError::bad_request(format!("bp resolve: {e}")))?;
match resolved_bp.metadata.default_run_ttl_secs {
Some(v) => (v, TtlSource::BpMetadata),
None => (default_run_ttl(), TtlSource::ServerDefault),
}
}
};
let input = TaskApplicationInput {
blueprint,
operator_id: operator_id.clone(),
role: Role::Operator,
ttl: Duration::from_secs(ttl_secs),
init_ctx,
operator_kind,
bridge_id: op_req.senior_bridge_id,
hook_id: op_req.spawn_hook_id,
operator_backend_id: op_req.operator_backend_id,
operator_kind_overrides,
task_input: task_input_spec,
check_policy: req.check_policy,
};
let input_json = Some(tasks::snapshot_launch_input(&input)?);
let task_id = TaskId::new();
let run_id = RunId::new();
let now = tasks::now_secs();
state
.task_store
.create(TaskRecord {
id: task_id.clone(),
goal,
blueprint_ref: blueprint_ref_json,
input_ctx: input_ctx_snapshot,
task_input_spec: task_input_spec_snapshot,
status: TaskRecordStatus::Running,
created_at: now,
updated_at: now,
})
.await
.map_err(ApiError::engine)?;
state
.run_store
.create(RunRecord {
id: run_id.clone(),
task_id: task_id.clone(),
status: RunStatus::Running,
step_entries: Vec::new(),
degradations: Vec::new(),
operator_sid: req.operator_sid.clone(),
result_ref: None,
input_json,
created_at: now,
updated_at: now,
})
.await
.map_err(ApiError::engine)?;
let run_ctx = RunContext::new(run_id.clone(), state.run_store.clone())
.with_replay_store(state.replay_store.clone());
if detach {
let bg_state = state.clone();
let bg_task_id = task_id.clone();
let bg_run_id = run_id.clone();
tokio::spawn(async move {
let outcome = match tokio::time::timeout(
Duration::from_secs(ttl_secs),
bg_state.task_app.handle_with_run(input, Some(run_ctx)),
)
.await
{
Ok(outcome) => outcome,
Err(_elapsed) => {
let reason = json!({
"error": format!("detached run exceeded {ttl_secs}s ttl ceiling"),
});
if let Err(e) = bg_state.run_store.set_result(&bg_run_id, reason).await {
tracing::warn!(%bg_run_id, error = %e, "run_flow_form: detached ttl set_result failed");
}
if let Err(e) = bg_state
.run_store
.update_status(&bg_run_id, RunStatus::Failed)
.await
{
tracing::warn!(%bg_run_id, error = %e, "run_flow_form: detached ttl run update_status(Failed) failed");
}
if let Err(e) = bg_state
.task_store
.update_status(&bg_task_id, TaskRecordStatus::Failed)
.await
{
tracing::warn!(%bg_task_id, error = %e, "run_flow_form: detached ttl task update_status(Failed) failed");
}
return;
}
};
let _ = tasks::finalize_run(&bg_state, &bg_task_id, &bg_run_id, outcome).await;
});
return Ok(TaskLaunchReply(
TaskLaunchResponse {
final_ctx: Value::Null,
bound_version: None,
effective_ttl_secs: ttl_secs,
ttl_source,
task_id,
run_id,
status: RunStatus::Running,
},
StatusCode::ACCEPTED,
));
}
let outcome = match tokio::time::timeout(
Duration::from_secs(sync_timeout_secs),
state.task_app.handle_with_run(input, Some(run_ctx)),
)
.await
{
Ok(outcome) => outcome,
Err(_elapsed) => {
let reason = json!({
"error": format!("sync launch exceeded {sync_timeout_secs}s timeout ceiling"),
});
if let Err(e) = state.run_store.set_result(&run_id, reason).await {
tracing::warn!(%run_id, error = %e, "run_flow_form: timeout run set_result failed");
}
if let Err(e) = state
.run_store
.update_status(&run_id, RunStatus::Failed)
.await
{
tracing::warn!(%run_id, error = %e, "run_flow_form: timeout run update_status(Failed) failed");
}
if let Err(e) = state
.task_store
.update_status(&task_id, TaskRecordStatus::Failed)
.await
{
tracing::warn!(%task_id, error = %e, "run_flow_form: timeout task update_status(Failed) failed");
}
return Err(ApiError::timeout(format!(
"sync launch exceeded {sync_timeout_secs}s timeout ceiling: the in-process flow \
eval was abandoned (dropping the future cancels it); attach an operator that \
acks promptly (POST /v1/operators + WS), or raise timeout_secs / sync_timeout_secs"
)));
}
};
let out = tasks::finalize_run(state, &task_id, &run_id, outcome)
.await
.map_err(flow_eval_error_to_api_error)?;
Ok(TaskLaunchReply(
TaskLaunchResponse {
final_ctx: out.final_ctx,
bound_version: out.bound_version.map(|v| format!("{:?}", v)),
effective_ttl_secs: ttl_secs,
ttl_source,
task_id,
run_id,
status: RunStatus::Done,
},
StatusCode::OK,
))
}
fn build_task_input_spec_from_request(
req: &TaskLaunchRequest,
) -> Option<mlua_swarm::service::TaskInputSpec> {
let project_root = req.project_root.clone().or_else(|| {
req.init_ctx
.get("project_root")
.and_then(Value::as_str)
.map(String::from)
});
let work_dir = req.work_dir.clone().or_else(|| {
req.init_ctx
.get("work_dir")
.and_then(Value::as_str)
.map(String::from)
});
let task_metadata = req.task_metadata.clone().or_else(|| {
req.init_ctx
.get("task_metadata")
.filter(|v| v.is_object())
.cloned()
});
if project_root.is_none() && work_dir.is_none() && task_metadata.is_none() {
None
} else {
Some(mlua_swarm::service::TaskInputSpec {
project_root,
work_dir,
task_metadata,
})
}
}
async fn take_session_token(state: &AppState, sid: &str) -> Result<CapToken, ApiError> {
let key = mlua_swarm::types::token_fingerprint(sid);
state
.sessions
.lock()
.await
.map
.remove(&key)
.ok_or_else(|| ApiError::not_found(format!("session: fp={key}")))
}
fn extract_bearer(headers: &HeaderMap) -> Result<String, ApiError> {
let v = headers
.get(AUTHORIZATION)
.ok_or_else(|| ApiError::bad_request("missing Authorization header".into()))?
.to_str()
.map_err(|_| ApiError::bad_request("invalid Authorization header encoding".into()))?;
let sid = v
.strip_prefix("Bearer ")
.ok_or_else(|| ApiError::bad_request("Authorization must be 'Bearer <sid>'".into()))?
.trim();
if sid.is_empty() {
return Err(ApiError::bad_request("Bearer sid is empty".into()));
}
Ok(sid.to_string())
}
fn parse_role(s: &str) -> Result<Role, ApiError> {
match s.to_ascii_lowercase().as_str() {
"operator" => Ok(Role::Operator),
"worker" => Ok(Role::Worker),
"observer" => Ok(Role::Observer),
"senior" => Ok(Role::Senior),
other => Err(ApiError::bad_request(format!("unknown role: {other}"))),
}
}
fn flow_eval_error_to_api_error(e: TaskApplicationError) -> ApiError {
if let TaskApplicationError::Launch(TaskLaunchError::FlowEval {
message,
failed_step,
verdict_value,
partial_ctx,
}) = &e
{
let details = json!({
"failed_step": failed_step,
"verdict_value": verdict_value,
"partial_ctx": partial_ctx,
});
return ApiError::bad_request(format!("run: flow eval: {message}")).with_details(details);
}
ApiError::bad_request(format!("run: {e}"))
}
#[derive(Debug)]
pub struct ApiError {
status: StatusCode,
message: String,
details: Option<Value>,
}
impl ApiError {
pub fn engine(e: impl std::fmt::Display) -> Self {
Self {
status: StatusCode::INTERNAL_SERVER_ERROR,
message: format!("engine: {e}"),
details: None,
}
}
pub fn not_found(m: String) -> Self {
Self {
status: StatusCode::NOT_FOUND,
message: m,
details: None,
}
}
pub fn bad_request(m: String) -> Self {
Self {
status: StatusCode::BAD_REQUEST,
message: m,
details: None,
}
}
pub fn conflict(m: String) -> Self {
Self {
status: StatusCode::CONFLICT,
message: m,
details: None,
}
}
pub fn unavailable(m: String) -> Self {
Self {
status: StatusCode::SERVICE_UNAVAILABLE,
message: m,
details: None,
}
}
pub fn timeout(m: String) -> Self {
Self {
status: StatusCode::GATEWAY_TIMEOUT,
message: m,
details: None,
}
}
pub fn gone(m: String) -> Self {
Self {
status: StatusCode::GONE,
message: m,
details: None,
}
}
pub fn payload_too_large(m: String) -> Self {
Self {
status: StatusCode::PAYLOAD_TOO_LARGE,
message: m,
details: None,
}
}
pub fn unprocessable(m: impl Into<String>) -> Self {
Self {
status: StatusCode::UNPROCESSABLE_ENTITY,
message: m.into(),
details: None,
}
}
pub fn with_details(mut self, details: Value) -> Self {
self.details = Some(details);
self
}
}
impl IntoResponse for ApiError {
fn into_response(self) -> Response {
let body = match self.details {
Some(details) => json!({"error": self.message, "details": details}),
None => json!({"error": self.message}),
};
(self.status, Json(body)).into_response()
}
}
fn default_run_ttl() -> u64 {
1800
}
#[cfg(test)]
fn resolve_ttl_from_metadata(metadata_ttl: Option<u64>) -> u64 {
metadata_ttl.unwrap_or_else(default_run_ttl)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn ttl_cascade_request_body_wins_over_metadata() {
let req_ttl: Option<u64> = Some(100);
let metadata_ttl: Option<u64> = Some(3600);
let effective = match req_ttl {
Some(v) => v,
None => resolve_ttl_from_metadata(metadata_ttl),
};
assert_eq!(
effective, 100,
"request body ttl_secs=100 must win over metadata=3600 (cascade priority (1) > (2))"
);
}
#[test]
fn ttl_cascade_metadata_used_when_body_missing() {
let req_ttl: Option<u64> = None;
let metadata_ttl: Option<u64> = Some(3600);
let effective = match req_ttl {
Some(v) => v,
None => resolve_ttl_from_metadata(metadata_ttl),
};
assert_eq!(
effective, 3600,
"body None + metadata=3600 must resolve to 3600 (cascade (2))"
);
}
#[test]
fn ttl_cascade_server_default_when_both_missing() {
let req_ttl: Option<u64> = None;
let metadata_ttl: Option<u64> = None;
let effective = match req_ttl {
Some(v) => v,
None => resolve_ttl_from_metadata(metadata_ttl),
};
assert_eq!(
effective,
default_run_ttl(),
"body None + metadata None must fall back to default_run_ttl() = 1800s"
);
assert_eq!(effective, 1800, "default_run_ttl() literal = 1800s");
}
#[test]
fn resolve_ttl_from_metadata_none_returns_server_default() {
assert_eq!(resolve_ttl_from_metadata(None), 1800);
}
#[test]
fn resolve_ttl_from_metadata_some_returns_value() {
assert_eq!(resolve_ttl_from_metadata(Some(7200)), 7200);
assert_eq!(resolve_ttl_from_metadata(Some(60)), 60);
}
#[test]
fn task_launch_request_parses_check_policy_wire_field() {
let body = json!({
"blueprint": { "kind": "id", "id": "some-bp" },
"init_ctx": {},
"check_policy": "silent",
});
let req: TaskLaunchRequest =
serde_json::from_value(body).expect("request must deserialize");
assert_eq!(req.check_policy, Some(CheckPolicy::Silent));
}
#[test]
fn task_launch_request_check_policy_defaults_to_none_when_omitted() {
let body = json!({
"blueprint": { "kind": "id", "id": "some-bp" },
"init_ctx": {},
});
let req: TaskLaunchRequest =
serde_json::from_value(body).expect("request must deserialize");
assert_eq!(req.check_policy, None);
}
fn task_req(
init_ctx: Value,
project_root: Option<&str>,
work_dir: Option<&str>,
task_metadata: Option<Value>,
) -> TaskLaunchRequest {
TaskLaunchRequest {
blueprint: BlueprintRef::Id {
id: mlua_swarm::blueprint::store::BlueprintId::new("ut"),
version: Default::default(),
},
init_ctx,
project_root: project_root.map(String::from),
work_dir: work_dir.map(String::from),
task_metadata,
ttl_secs: None,
operator: None,
operator_sid: None,
timeout_secs: None,
goal: None,
detach: false,
check_policy: None,
}
}
#[test]
fn build_task_input_spec_from_request_returns_sibling_fields_when_present() {
let req = task_req(
json!({"free": "form"}),
Some("/repo/sibling"),
Some("/repo/sibling/work"),
Some(json!({"issue": 19})),
);
let spec = build_task_input_spec_from_request(&req).expect("spec must be Some");
assert_eq!(spec.project_root.as_deref(), Some("/repo/sibling"));
assert_eq!(spec.work_dir.as_deref(), Some("/repo/sibling/work"));
assert_eq!(spec.task_metadata, Some(json!({"issue": 19})));
}
#[test]
fn build_task_input_spec_from_request_falls_back_to_legacy_init_ctx_shape() {
let req = task_req(
json!({
"project_root": "/repo/legacy",
"work_dir": "/repo/legacy/work",
"task_metadata": {"issue": 17},
}),
None,
None,
None,
);
let spec = build_task_input_spec_from_request(&req).expect("spec must be Some");
assert_eq!(spec.project_root.as_deref(), Some("/repo/legacy"));
assert_eq!(spec.work_dir.as_deref(), Some("/repo/legacy/work"));
assert_eq!(spec.task_metadata, Some(json!({"issue": 17})));
}
#[test]
fn build_task_input_spec_from_request_sibling_wins_over_legacy_shape() {
let req = task_req(
json!({
"project_root": "/repo/legacy",
"work_dir": "/repo/legacy/work",
"task_metadata": {"issue": 17},
}),
Some("/repo/sibling"),
Some("/repo/sibling/work"),
Some(json!({"issue": 19})),
);
let spec = build_task_input_spec_from_request(&req).expect("spec must be Some");
assert_eq!(
spec.project_root.as_deref(),
Some("/repo/sibling"),
"sibling field must win over the legacy init_ctx-nested value"
);
assert_eq!(spec.work_dir.as_deref(), Some("/repo/sibling/work"));
assert_eq!(spec.task_metadata, Some(json!({"issue": 19})));
}
#[test]
fn build_task_input_spec_from_request_returns_none_when_no_fields_present() {
let req = task_req(json!({"unrelated": "value"}), None, None, None);
assert!(build_task_input_spec_from_request(&req).is_none());
}
fn status_test_state() -> AppState {
let engine = Engine::new(mlua_swarm::EngineCfg::default());
let compiler = mlua_swarm::Compiler::new(default_registry());
let launch = Arc::new(mlua_swarm::TaskLaunchService::new(engine.clone(), compiler));
AppState {
engine,
sessions: Arc::new(Mutex::new(SessionStore::default())),
task_app: Arc::new(mlua_swarm::TaskApplication::new_inline_only(launch)),
ws_operator_factory: None,
data_store: Arc::new(mlua_swarm::store::output::InMemoryOutputStore::new()),
operator_sessions: Arc::new(Mutex::new(HashMap::new())),
roles_to_sid: Arc::new(Mutex::new(HashMap::new())),
task_store: Arc::new(mlua_swarm::store::task::InMemoryTaskStore::new()),
run_store: Arc::new(mlua_swarm::store::run::InMemoryRunStore::new()),
replay_store: Arc::new(mlua_swarm::store::replay::InMemoryReplayStore::new()),
base_url: None,
sync_timeout_secs: 300,
}
}
#[tokio::test]
async fn status_get_reports_running_runs_and_operators() {
let state = status_test_state();
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0);
state
.run_store
.create(RunRecord {
id: RunId::new(),
task_id: TaskId::new(),
status: RunStatus::Running,
step_entries: Vec::new(),
degradations: Vec::new(),
operator_sid: None,
result_ref: None,
input_json: None,
created_at: now,
updated_at: now,
})
.await
.expect("seed running RunRecord");
struct NoopOperator;
#[async_trait::async_trait]
impl mlua_swarm::Operator for NoopOperator {
async fn execute(
&self,
_ctx: &mlua_swarm::Ctx,
_system: Option<String>,
_prompt: Value,
_worker: Option<mlua_swarm::WorkerBinding>,
_worker_token: mlua_swarm::CapToken,
) -> Result<mlua_swarm::WorkerResult, mlua_swarm::WorkerError> {
unimplemented!("not exercised by this test — only registration/list matters")
}
}
state
.engine
.register_operator("test-op", Arc::new(NoopOperator))
.await;
let Json(resp) = status_get(State(state)).await;
assert_eq!(resp.running_runs, 1);
assert_eq!(resp.attached_operators, 1);
}
#[tokio::test]
async fn api_error_details_render_into_response_body() {
use axum::body::to_bytes;
use axum::response::IntoResponse;
let bare = ApiError::bad_request("something".to_string()).into_response();
let (parts, body) = bare.into_parts();
assert_eq!(parts.status, StatusCode::BAD_REQUEST);
let bytes = to_bytes(body, 1024).await.expect("bare body");
let parsed: serde_json::Value = serde_json::from_slice(&bytes).expect("parse bare");
assert_eq!(parsed, json!({"error": "something"}));
let with_details = ApiError::bad_request("run: flow eval: blocked".to_string())
.with_details(json!({
"failed_step": "gate",
"verdict_value": {"verdict": "BLOCKED"},
"partial_ctx": {"steps": {}},
}));
let resp = with_details.into_response();
let (parts, body) = resp.into_parts();
assert_eq!(parts.status, StatusCode::BAD_REQUEST);
let bytes = to_bytes(body, 4096).await.expect("details body");
let parsed: serde_json::Value = serde_json::from_slice(&bytes).expect("parse details");
assert_eq!(parsed["error"], "run: flow eval: blocked");
assert_eq!(parsed["details"]["failed_step"], "gate");
assert_eq!(parsed["details"]["verdict_value"]["verdict"], "BLOCKED");
assert!(parsed["details"]["partial_ctx"].is_object());
}
#[test]
fn flow_eval_error_to_api_error_lifts_structural_fields_into_details() {
let err = TaskApplicationError::Launch(TaskLaunchError::FlowEval {
message: "blocked: {\"verdict\":\"BLOCKED\"}".to_string(),
failed_step: Some("gate".to_string()),
verdict_value: Some(json!({"verdict": "BLOCKED"})),
partial_ctx: Some(json!({"steps": {}})),
});
let api_err = flow_eval_error_to_api_error(err);
assert_eq!(api_err.status, StatusCode::BAD_REQUEST);
assert!(
api_err.message.starts_with("run: flow eval: "),
"message must preserve pre-#76 `run: flow eval: <msg>` prefix, got: {}",
api_err.message
);
let details = api_err.details.expect("details must be Some for FlowEval");
assert_eq!(details["failed_step"], "gate");
assert_eq!(details["verdict_value"]["verdict"], "BLOCKED");
assert!(details["partial_ctx"].is_object());
}
#[test]
fn flow_eval_error_to_api_error_non_flow_eval_falls_back_to_message_only() {
let api_err = flow_eval_error_to_api_error(TaskApplicationError::NoStore);
assert_eq!(api_err.status, StatusCode::BAD_REQUEST);
assert!(api_err.message.starts_with("run: "));
assert!(
api_err.details.is_none(),
"non-FlowEval errors must not carry a details field (pre-#76 shape)"
);
}
#[test]
fn flow_eval_error_to_api_error_with_all_none_still_populates_details_with_nulls() {
let err = TaskApplicationError::Launch(TaskLaunchError::FlowEval {
message: "unresolved extern".to_string(),
failed_step: None,
verdict_value: None,
partial_ctx: None,
});
let api_err = flow_eval_error_to_api_error(err);
let details = api_err
.details
.expect("details Some even when fields are None");
assert_eq!(details["failed_step"], Value::Null);
assert_eq!(details["verdict_value"], Value::Null);
assert_eq!(details["partial_ctx"], Value::Null);
}
}