mlua_swarm/middleware/
worker_binding.rs1use crate::core::ctx::Ctx;
23use crate::core::engine::Engine;
24use crate::middleware::SpawnerLayer;
25use crate::operator::WorkerBinding;
26use crate::types::{CapToken, StepId};
27use crate::worker::adapter::{SpawnError, SpawnerAdapter};
28use crate::worker::Worker;
29use async_trait::async_trait;
30use std::collections::HashMap;
31use std::sync::Arc;
32
33pub const WORKER_BINDING_KEY: &str = "worker_binding";
36
37pub struct WorkerBindingMiddleware {
40 bindings: Arc<HashMap<String, WorkerBinding>>,
41}
42
43impl WorkerBindingMiddleware {
44 pub fn new(bindings: HashMap<String, WorkerBinding>) -> Self {
47 Self {
48 bindings: Arc::new(bindings),
49 }
50 }
51}
52
53impl SpawnerLayer for WorkerBindingMiddleware {
54 fn wrap(&self, inner: Arc<dyn SpawnerAdapter>) -> Arc<dyn SpawnerAdapter> {
55 Arc::new(WorkerBindingWrapped {
56 inner,
57 bindings: self.bindings.clone(),
58 })
59 }
60}
61
62struct WorkerBindingWrapped {
63 inner: Arc<dyn SpawnerAdapter>,
64 bindings: Arc<HashMap<String, WorkerBinding>>,
65}
66
67#[async_trait]
68impl SpawnerAdapter for WorkerBindingWrapped {
69 async fn spawn(
70 &self,
71 engine: &Engine,
72 ctx: &Ctx,
73 task_id: StepId,
74 attempt: u32,
75 token: CapToken,
76 ) -> Result<Box<dyn Worker>, SpawnError> {
77 let Some(binding) = self.bindings.get(&ctx.agent) else {
78 return self.inner.spawn(engine, ctx, task_id, attempt, token).await;
81 };
82 let value = serde_json::to_value(binding).map_err(|e| {
83 SpawnError::Internal(format!(
84 "worker_binding for agent '{}' failed to serialize: {e}",
85 ctx.agent
86 ))
87 })?;
88 let mut new_ctx = ctx.clone();
89 new_ctx
90 .meta
91 .runtime
92 .insert(WORKER_BINDING_KEY.to_string(), value);
93 self.inner
94 .spawn(engine, &new_ctx, task_id, attempt, token)
95 .await
96 }
97}
98
99#[cfg(test)]
100mod tests {
101 use super::*;
102 use crate::core::config::EngineCfg;
103 use crate::types::Role;
104 use std::sync::Mutex;
105 use std::time::Duration;
106
107 struct CtxProbe {
110 seen: Arc<Mutex<Option<Ctx>>>,
111 }
112
113 #[async_trait]
114 impl SpawnerAdapter for CtxProbe {
115 async fn spawn(
116 &self,
117 _engine: &Engine,
118 ctx: &Ctx,
119 _task_id: StepId,
120 _attempt: u32,
121 _token: CapToken,
122 ) -> Result<Box<dyn Worker>, SpawnError> {
123 *self.seen.lock().unwrap() = Some(ctx.clone());
124 Err(SpawnError::Internal("probe stop".into()))
125 }
126 }
127
128 fn probe_stack(
129 bindings: HashMap<String, WorkerBinding>,
130 ) -> (Arc<dyn SpawnerAdapter>, Arc<Mutex<Option<Ctx>>>) {
131 let seen = Arc::new(Mutex::new(None));
132 let inner = Arc::new(CtxProbe { seen: seen.clone() });
133 let wrapped = WorkerBindingMiddleware::new(bindings).wrap(inner);
134 (wrapped, seen)
135 }
136
137 #[tokio::test]
138 async fn injects_binding_into_ctx_meta_runtime_on_hit() {
139 let mut map = HashMap::new();
140 map.insert(
141 "planner".to_string(),
142 WorkerBinding {
143 variant: "knowledge-worker".to_string(),
144 tools: vec!["Read".to_string()],
145 request_digest: Some(
146 "sha256:1111111111111111111111111111111111111111111111111111111111111111"
147 .parse()
148 .unwrap(),
149 ),
150 requested_model: Some("claude-sonnet".to_string()),
151 },
152 );
153 let (stack, seen) = probe_stack(map);
154 let engine = Engine::new(EngineCfg::default());
155 let task_id = StepId::parse("ST-1").unwrap();
156 let ctx = Ctx::new(task_id.clone(), 1, "planner");
157 let token = engine
158 .attach("ut-op", Role::Operator, Duration::from_secs(30))
159 .await
160 .expect("attach");
161 let _ = stack.spawn(&engine, &ctx, task_id, 1, token).await;
162
163 let observed = seen.lock().unwrap().clone().expect("inner ctx captured");
164 let v = observed
165 .meta
166 .runtime
167 .get(WORKER_BINDING_KEY)
168 .expect("worker_binding key present");
169 let wb: WorkerBinding = serde_json::from_value(v.clone()).expect("round-trip");
170 assert_eq!(wb.variant, "knowledge-worker");
171 assert_eq!(wb.tools, vec!["Read".to_string()]);
172 assert!(wb
175 .request_digest
176 .as_ref()
177 .expect("request_digest present")
178 .as_str()
179 .starts_with("sha256:"));
180 assert_eq!(wb.requested_model.as_deref(), Some("claude-sonnet"));
181 }
182
183 #[test]
186 fn omits_self_check_fields_when_none() {
187 let json = serde_json::to_value(WorkerBinding {
188 variant: "v".to_string(),
189 tools: vec![],
190 request_digest: None,
191 requested_model: None,
192 })
193 .expect("serialize");
194 let obj = json.as_object().expect("object");
195 assert!(!obj.contains_key("request_digest"));
196 assert!(!obj.contains_key("requested_model"));
197 }
198
199 #[tokio::test]
200 async fn passes_through_untouched_on_miss() {
201 let (stack, seen) = probe_stack(HashMap::new());
202 let engine = Engine::new(EngineCfg::default());
203 let task_id = StepId::parse("ST-2").unwrap();
204 let ctx = Ctx::new(task_id.clone(), 1, "unbound-agent");
205 let token = engine
206 .attach("ut-op", Role::Operator, Duration::from_secs(30))
207 .await
208 .expect("attach");
209 let _ = stack.spawn(&engine, &ctx, task_id, 1, token).await;
210
211 let observed = seen.lock().unwrap().clone().expect("inner ctx captured");
212 assert!(
213 !observed.meta.runtime.contains_key(WORKER_BINDING_KEY),
214 "no binding entry must be injected on miss"
215 );
216 }
217}