1use crate::pb::contracts::{
7 robonix_system_executor_control_plan_client::RobonixSystemExecutorControlPlanClient,
8 robonix_system_executor_execute_client::RobonixSystemExecutorExecuteClient,
9 robonix_system_executor_list_active_plans_client::RobonixSystemExecutorListActivePlansClient,
10 robonix_system_pilot_get_health_server::RobonixSystemPilotGetHealth,
11 robonix_system_pilot_server::RobonixSystemPilot,
12};
13use crate::pb::executor::ControlPlanRequest;
14use crate::pb::module_health::{
15 GetModuleHealthRequest, GetModuleHealthResponse, ModuleHealth, ModuleHealthReport,
16};
17use crate::pb::pilot::{
18 BatchResult, PilotEvent, Plan, RtdlNodeState, SessionStatusEvent, Task, TaskStateEvent,
19};
20use crate::planner::{self, ExecutorConn, TaskState};
21use crate::vlm::{Message, VlmClient};
22use anyhow::Context;
23use robonix_atlas::client::{self as atlas_client, AtlasClient};
24use robonix_scribe::{debug, error};
25use std::collections::{HashMap, VecDeque};
26use std::sync::Arc;
27use std::sync::atomic::{AtomicU64, Ordering};
28use tokio::sync::{Mutex, broadcast, mpsc, watch};
29use tokio_stream::wrappers::ReceiverStream;
30use tonic::{Request, Response, Status};
31use uuid::Uuid;
32
33#[derive(Clone, Copy)]
34#[repr(u32)]
35#[allow(dead_code)]
36pub enum SessionState {
37 Active = 0, Completed = 1, Failed = 2, }
41
42pub const EVT_TEXT_CHUNK: u32 = 0;
46pub const EVT_PLAN: u32 = 1;
47pub const EVT_BATCH_RESULT: u32 = 2;
48pub const EVT_STATUS: u32 = 3;
49pub const EVT_FINAL_TEXT: u32 = 4;
50pub const EVT_NODE_STATE: u32 = 5;
51pub const EVT_TASK_STATE: u32 = 6;
52const MODULE_HEALTH_SCHEMA_VERSION: u32 = 1;
53const MODULE_HEALTH_OK: u32 = 0;
54const MODULE_HEALTH_TTL_MS: u32 = 5000;
55
56#[allow(dead_code)]
57pub enum PilotStreamBody {
58 TextChunk(String),
59 FinalText(String),
60 Plan(Plan),
61 BatchResult(BatchResult),
62 Status(SessionStatusEvent),
63 NodeState(RtdlNodeState),
64 TaskState(TaskStateEvent),
65}
66
67pub fn pack(session_id: &str, body: PilotStreamBody) -> PilotEvent {
68 let mut e = PilotEvent {
69 session_id: session_id.to_string(),
70 ..Default::default()
71 };
72 match body {
73 PilotStreamBody::TextChunk(s) => {
74 e.event_kind = EVT_TEXT_CHUNK;
75 e.text_chunk = s;
76 }
77 PilotStreamBody::Plan(g) => {
78 e.event_kind = EVT_PLAN;
79 e.plan = Some(g);
80 }
81 PilotStreamBody::BatchResult(b) => {
82 e.event_kind = EVT_BATCH_RESULT;
83 e.batch_result = Some(b);
84 }
85 PilotStreamBody::Status(s) => {
86 e.event_kind = EVT_STATUS;
87 e.status = Some(s);
88 }
89 PilotStreamBody::FinalText(s) => {
90 e.event_kind = EVT_FINAL_TEXT;
91 e.final_text = s;
92 }
93 PilotStreamBody::NodeState(ns) => {
94 e.event_kind = EVT_NODE_STATE;
95 e.node_state = Some(ns);
96 }
97 PilotStreamBody::TaskState(ts) => {
98 e.event_kind = EVT_TASK_STATE;
99 e.task_state = Some(ts);
100 }
101 }
102 e
103}
104
105type Histories = Arc<Mutex<HashMap<String, Arc<Mutex<Vec<Message>>>>>>;
108type TaskStates = Arc<Mutex<HashMap<String, Arc<Mutex<Option<TaskState>>>>>>;
109
110#[derive(Clone)]
111struct ActiveTurnInput {
112 turn_id: String,
113 tx: mpsc::Sender<Task>,
114 events: broadcast::Sender<Result<PilotEvent, String>>,
115 reply_generation: Arc<AtomicU64>,
116}
117
118fn subscribe_turn_events(
122 events: &broadcast::Sender<Result<PilotEvent, String>>,
123 reply_generation: &Arc<AtomicU64>,
124) -> ReceiverStream<Result<PilotEvent, Status>> {
125 let generation = reply_generation.fetch_add(1, Ordering::AcqRel) + 1;
126 let reply_generation = Arc::clone(reply_generation);
127 let mut subscriber = events.subscribe();
128 let (tx, rx) = mpsc::channel(64);
129 tokio::spawn(async move {
130 loop {
131 match subscriber.recv().await {
132 Ok(Ok(event)) => {
133 if reply_generation.load(Ordering::Acquire) != generation {
138 break;
139 }
140 let complete = event.event_kind == EVT_FINAL_TEXT;
141 if tx.send(Ok(event)).await.is_err() || complete {
142 break;
143 }
144 }
145 Ok(Err(error)) => {
146 let _ = tx.send(Err(Status::internal(error))).await;
147 break;
148 }
149 Err(broadcast::error::RecvError::Lagged(skipped)) => {
150 let _ = tx
151 .send(Err(Status::resource_exhausted(format!(
152 "Pilot event subscriber lagged by {skipped} event(s)"
153 ))))
154 .await;
155 break;
156 }
157 Err(broadcast::error::RecvError::Closed) => break,
158 }
159 }
160 });
161 ReceiverStream::new(rx)
162}
163
164const SEEN_TASK_IDS_PER_SESSION: usize = 256;
165
166#[derive(Clone)]
167pub struct PilotServiceImpl {
168 atlas: AtlasClient,
172 provider_id: String,
176 vlm: VlmClient,
177 soma_prompt_block: Arc<String>,
178 histories: Histories,
179 task_states: TaskStates,
183 cancels: Arc<Mutex<HashMap<String, watch::Sender<bool>>>>,
186 steers: Arc<Mutex<HashMap<String, ActiveTurnInput>>>,
190 seen_task_ids: Arc<Mutex<HashMap<String, VecDeque<String>>>>,
194 plan_seq: Arc<AtomicU64>,
197}
198
199impl PilotServiceImpl {
200 pub fn new(
201 atlas: AtlasClient,
202 provider_id: String,
203 vlm: VlmClient,
204 soma_prompt_block: String,
205 ) -> Self {
206 Self {
207 atlas,
208 provider_id,
209 vlm,
210 soma_prompt_block: Arc::new(soma_prompt_block),
211 histories: Arc::new(Mutex::new(HashMap::new())),
212 task_states: Arc::new(Mutex::new(HashMap::new())),
213 cancels: Arc::new(Mutex::new(HashMap::new())),
214 steers: Arc::new(Mutex::new(HashMap::new())),
215 seen_task_ids: Arc::new(Mutex::new(HashMap::new())),
216 plan_seq: Arc::new(AtomicU64::new(0)),
217 }
218 }
219
220 async fn get_or_create_history(&self, session_id: &str) -> Arc<Mutex<Vec<Message>>> {
221 let mut map = self.histories.lock().await;
222 map.entry(session_id.to_string())
223 .or_insert_with(|| Arc::new(Mutex::new(Vec::new())))
224 .clone()
225 }
226
227 async fn get_or_create_task_state(&self, session_id: &str) -> Arc<Mutex<Option<TaskState>>> {
228 let mut map = self.task_states.lock().await;
229 map.entry(session_id.to_string())
230 .or_insert_with(|| Arc::new(Mutex::new(None)))
231 .clone()
232 }
233
234 async fn accept_task_id_once(&self, session_id: &str, task_id: &str) -> bool {
235 if task_id.is_empty() {
236 return true;
237 }
238 let mut sessions = self.seen_task_ids.lock().await;
239 let ids = sessions.entry(session_id.to_string()).or_default();
240 if ids.iter().any(|seen| seen == task_id) {
241 return false;
242 }
243 ids.push_back(task_id.to_string());
244 while ids.len() > SEEN_TASK_IDS_PER_SESSION {
245 ids.pop_front();
246 }
247 true
248 }
249}
250
251fn task_context(task: &Task) -> Option<serde_json::Value> {
252 let raw = task.context_json.trim();
253 if raw.is_empty() {
254 return None;
255 }
256 serde_json::from_str(raw).ok()
257}
258
259fn task_is_abort_turn(task: &Task) -> bool {
260 task_context(task)
261 .and_then(|v| v.get("abort_turn").and_then(|x| x.as_bool()))
262 .unwrap_or(false)
263}
264
265fn task_is_steer(task: &Task) -> bool {
266 task_context(task).is_some_and(|v| {
267 v.get("steer").and_then(|x| x.as_bool()).unwrap_or(false)
268 || v.get("interaction_mode").and_then(|x| x.as_str()) == Some("steer")
269 })
270}
271
272fn expected_turn_id(task: &Task) -> Option<String> {
273 task_context(task).and_then(|v| {
274 v.get("expected_turn_id")
275 .and_then(|x| x.as_str())
276 .filter(|id| !id.is_empty())
277 .map(str::to_string)
278 })
279}
280
281fn strict_expected_turn(task: &Task) -> bool {
282 task_context(task)
283 .and_then(|value| {
284 value
285 .get("strict_expected_turn")
286 .and_then(|field| field.as_bool())
287 })
288 .unwrap_or(false)
289}
290
291#[tonic::async_trait]
292impl RobonixSystemPilot for PilotServiceImpl {
293 type SubmitTaskStream = ReceiverStream<Result<PilotEvent, Status>>;
294
295 async fn submit_task(
296 &self,
297 request: Request<Task>,
298 ) -> Result<Response<Self::SubmitTaskStream>, Status> {
299 let mut task = request.into_inner();
300
301 if task.session_id.is_empty() {
302 task.session_id = Uuid::new_v4().to_string();
303 }
304 if task.task_id.is_empty() {
305 task.task_id = Uuid::new_v4().to_string();
306 }
307
308 if !self
309 .accept_task_id_once(&task.session_id, &task.task_id)
310 .await
311 {
312 debug!(
313 "[pilot] duplicate task ignored session={} task_id={}",
314 task.session_id, task.task_id
315 );
316 let (_tx, rx) = tokio::sync::mpsc::channel::<Result<PilotEvent, Status>>(1);
317 return Ok(Response::new(ReceiverStream::new(rx)));
318 }
319
320 if task_is_abort_turn(&task) {
321 let id = task.session_id.clone();
322 let turn_signaled = if let Some(tx) = self.cancels.lock().await.get(&id) {
323 tx.send_if_modified(|interrupted| {
324 if *interrupted {
325 false
326 } else {
327 *interrupted = true;
328 true
329 }
330 })
331 } else {
332 false
333 };
334 let executor_cancelled =
335 cancel_all_executor_plans(self.atlas.clone(), &self.provider_id).await;
336 debug!(
337 "[pilot] deterministic stop session {id} (turn_signaled={turn_signaled}, executor={executor_cancelled:?})"
338 );
339 let (tx, rx) = tokio::sync::mpsc::channel::<Result<PilotEvent, Status>>(1);
340 let message = match executor_cancelled {
341 Ok(true) => "stop completed",
342 Ok(false) => "stop reached Executor but cancellation was not accepted",
343 Err(ref error) => error.as_str(),
344 };
345 let state = if executor_cancelled == Ok(true) {
346 SessionState::Completed
347 } else {
348 SessionState::Failed
349 };
350 let _ = tx
351 .send(Ok(pack(
352 &id,
353 PilotStreamBody::Status(SessionStatusEvent {
354 session_id: id.clone(),
355 state: state as u32,
356 message: message.to_string(),
357 }),
358 )))
359 .await;
360 return Ok(Response::new(ReceiverStream::new(rx)));
361 }
362
363 let (steer_tx, steer_rx) = mpsc::channel::<Task>(32);
368 let (candidate_events, _) = broadcast::channel(128);
369 let candidate_reply_generation = Arc::new(AtomicU64::new(0));
370 let explicit_steer = task_is_steer(&task);
371 let expected_turn = expected_turn_id(&task);
372 let existing_turn = {
373 let mut steers = self.steers.lock().await;
374 match steers.get(&task.session_id) {
375 Some(existing) => Some(existing.clone()),
376 None => {
377 if explicit_steer {
378 debug!(
379 "[pilot] steer for session {} has no active turn; starting a new turn",
380 task.session_id
381 );
382 }
383 steers.insert(
384 task.session_id.clone(),
385 ActiveTurnInput {
386 turn_id: task.task_id.clone(),
387 tx: steer_tx.clone(),
388 events: candidate_events.clone(),
389 reply_generation: Arc::clone(&candidate_reply_generation),
390 },
391 );
392 None
393 }
394 }
395 };
396 if let Some(existing) = existing_turn {
397 if let Some(expected) = expected_turn
398 && expected != existing.turn_id
399 {
400 if strict_expected_turn(&task) {
401 return Err(Status::failed_precondition(format!(
402 "steer expected turn {expected}, but active turn is {}",
403 existing.turn_id
404 )));
405 }
406 debug!(
407 "[pilot] accepting same-session steer with stale expected turn {} (active={})",
408 expected, existing.turn_id
409 );
410 }
411 let id = task.session_id.clone();
416 let rx = subscribe_turn_events(&existing.events, &existing.reply_generation);
417 let ok = existing.tx.send(task).await.is_ok();
418 debug!("[pilot] steer task for session {id} (queued={ok})");
419 if !ok {
420 return Err(Status::unavailable(
421 "active Pilot turn stopped before steer was queued",
422 ));
423 }
424 return Ok(Response::new(rx));
425 }
426
427 let history_arc = self.get_or_create_history(&task.session_id).await;
428 let task_state_arc = self.get_or_create_task_state(&task.session_id).await;
429 let plan_seq = Arc::clone(&self.plan_seq);
430 let (tx, mut internal_rx) = tokio::sync::mpsc::channel::<Result<PilotEvent, Status>>(64);
435 let rx = subscribe_turn_events(&candidate_events, &candidate_reply_generation);
436 let relay_events = candidate_events.clone();
437 tokio::spawn(async move {
438 while let Some(item) = internal_rx.recv().await {
439 let shared = item.map_err(|status| status.message().to_string());
440 let _ = relay_events.send(shared);
441 }
442 });
443 let atlas = self.atlas.clone();
444 let provider_id = self.provider_id.clone();
445 let vlm = self.vlm.clone();
446 let soma_prompt_block = Arc::clone(&self.soma_prompt_block);
447 let session_id = task.session_id.clone();
448 let cancels = Arc::clone(&self.cancels);
449 let steers = Arc::clone(&self.steers);
450
451 let (cancel_tx, cancel_rx) = watch::channel(false);
452 cancels.lock().await.insert(session_id.clone(), cancel_tx);
453 tokio::spawn(async move {
458 let _ = tx
459 .send(Ok(pack(
460 &session_id,
461 PilotStreamBody::Status(SessionStatusEvent {
462 session_id: session_id.clone(),
463 state: SessionState::Active as u32,
464 message: format!("turn_id={}", task.task_id),
465 }),
466 )))
467 .await;
468
469 let mut atlas_for_turn = atlas.clone();
470 let mut executor = match build_executor_conn(atlas, &provider_id).await {
471 Ok(e) => e,
472 Err(e) => {
473 let _ = tx
474 .send(Err(Status::unavailable(format!(
475 "cannot reach Executor via atlas: {e:#}"
476 ))))
477 .await;
478 cancels.lock().await.remove(&session_id);
479 steers.lock().await.remove(&session_id);
480 return;
481 }
482 };
483
484 let mut history = history_arc.lock().await;
485 let mut standing_task = task_state_arc.lock().await;
486 if let Err(e) = planner::run_turn(
487 &task,
488 &mut history,
489 &mut standing_task,
490 &vlm,
491 &mut executor,
492 &mut atlas_for_turn,
493 &provider_id,
494 &tx,
495 cancel_rx,
496 steer_rx,
497 plan_seq,
498 soma_prompt_block.as_str(),
499 )
500 .await
501 {
502 error!("[pilot] turn error for session '{session_id}': {e:#}");
503 let _ = tx.send(Err(Status::internal(e.to_string()))).await;
504 }
505
506 cancels.lock().await.remove(&session_id);
507 steers.lock().await.remove(&session_id);
508 });
509
510 Ok(Response::new(rx))
511 }
512}
513
514#[tonic::async_trait]
515impl RobonixSystemPilotGetHealth for PilotServiceImpl {
516 async fn get_module_health(
517 &self,
518 _request: Request<GetModuleHealthRequest>,
519 ) -> Result<Response<GetModuleHealthResponse>, Status> {
520 Ok(Response::new(GetModuleHealthResponse {
521 report: Some(pilot_health_report(&self.provider_id)),
522 }))
523 }
524}
525
526fn pilot_health_report(provider_id: &str) -> ModuleHealthReport {
527 ModuleHealthReport {
528 schema_version: MODULE_HEALTH_SCHEMA_VERSION,
529 module: Some(ModuleHealth {
530 module_key: String::new(),
531 module_id: "pilot".to_string(),
532 provider_id: provider_id.to_string(),
533 health: MODULE_HEALTH_OK,
534 state: "active".to_string(),
535 reason_code: "OK".to_string(),
536 detail: "pilot serving".to_string(),
537 source: String::new(),
538 received_ts_ns: 0,
539 ttl_ms: MODULE_HEALTH_TTL_MS,
540 }),
541 }
542}
543
544async fn cancel_all_executor_plans(
545 mut atlas: AtlasClient,
546 consumer_id: &str,
547) -> Result<bool, String> {
548 let (_, _, channel) = atlas_client::connect_to_capability(
549 &mut atlas,
550 consumer_id,
551 "robonix/system/executor/control_plan",
552 )
553 .await
554 .map_err(|error| format!("stop could not reach Executor: {error:#}"))?;
555 let mut client = RobonixSystemExecutorControlPlanClient::new(channel);
556 client
557 .control_plan(ControlPlanRequest {
558 action: "cancel_all".into(),
559 wait_ms: 5_000,
560 ..Default::default()
561 })
562 .await
563 .map(|response| response.into_inner().success)
564 .map_err(|error| format!("Executor cancellation failed: {error}"))
565}
566
567async fn build_executor_conn(
571 mut atlas: AtlasClient,
572 consumer_id: &str,
573) -> anyhow::Result<ExecutorConn> {
574 let (_, executor_provider_id, exec_ch) = atlas_client::connect_to_capability(
575 &mut atlas,
576 consumer_id,
577 "robonix/system/executor/execute",
578 )
579 .await
580 .context("connect_to_capability robonix/system/executor/execute")?;
581 let (_, control_provider_id, control_ch) = atlas_client::connect_to_capability(
582 &mut atlas,
583 consumer_id,
584 "robonix/system/executor/control_plan",
585 )
586 .await
587 .context("connect_to_capability robonix/system/executor/control_plan")?;
588 let (_, active_provider_id, active_ch) = atlas_client::connect_to_capability(
589 &mut atlas,
590 consumer_id,
591 "robonix/system/executor/list_active_plans",
592 )
593 .await
594 .context("connect_to_capability robonix/system/executor/list_active_plans")?;
595 if control_provider_id != executor_provider_id || active_provider_id != executor_provider_id {
596 anyhow::bail!(
597 "Executor execute/control/active capabilities resolved to different providers: {executor_provider_id} vs {control_provider_id} vs {active_provider_id}"
598 );
599 }
600 Ok(ExecutorConn {
601 graph: RobonixSystemExecutorExecuteClient::new(exec_ch),
602 control: RobonixSystemExecutorControlPlanClient::new(control_ch),
603 active: RobonixSystemExecutorListActivePlansClient::new(active_ch),
604 })
605}
606
607#[cfg(test)]
608mod tests {
609 use super::{
610 EVT_FINAL_TEXT, EVT_STATUS, MODULE_HEALTH_OK, MODULE_HEALTH_SCHEMA_VERSION,
611 MODULE_HEALTH_TTL_MS, expected_turn_id, pilot_health_report, strict_expected_turn,
612 subscribe_turn_events, task_is_abort_turn, task_is_steer,
613 };
614 use crate::pb::pilot::{PilotEvent, Task};
615 use tokio_stream::StreamExt;
616
617 fn task(ctx: &str) -> Task {
618 Task {
619 task_id: "t".into(),
620 session_id: "s".into(),
621 source: 0,
622 text: String::new(),
623 audio_data: Vec::new(),
624 context_json: ctx.into(),
625 timestamp_ms: 0,
626 }
627 }
628
629 #[test]
630 fn pilot_health_report_uses_minimal_module_health_v1_fields() {
631 let report = pilot_health_report("pilot");
632 assert_eq!(report.schema_version, MODULE_HEALTH_SCHEMA_VERSION);
633 let module = report.module.expect("module health");
634 assert_eq!(module.module_id, "pilot");
635 assert_eq!(module.provider_id, "pilot");
636 assert_eq!(module.health, MODULE_HEALTH_OK);
637 assert_eq!(module.state, "active");
638 assert_eq!(module.reason_code, "OK");
639 assert_eq!(module.detail, "pilot serving");
640 assert_eq!(module.ttl_ms, MODULE_HEALTH_TTL_MS);
641 assert!(module.module_key.is_empty());
642 assert!(module.source.is_empty());
643 assert_eq!(module.received_ts_ns, 0);
644 }
645
646 #[test]
647 fn abort_turn_detected() {
648 assert!(task_is_abort_turn(&task(r#"{"abort_turn":true}"#)));
649 assert!(!task_is_abort_turn(&task(r#"{"abort_turn":false}"#)));
650 assert!(!task_is_abort_turn(&task(r#"{"foo":1}"#)));
651 assert!(!task_is_abort_turn(&task("")));
652 assert!(!task_is_abort_turn(&task("not json")));
653 }
654
655 #[test]
656 fn explicit_steer_and_expected_turn_are_parsed() {
657 let value = task(r#"{"interaction_mode":"steer","expected_turn_id":"turn-7"}"#);
658 assert!(task_is_steer(&value));
659 assert_eq!(expected_turn_id(&value).as_deref(), Some("turn-7"));
660 assert!(!strict_expected_turn(&value));
661 assert!(strict_expected_turn(&task(
662 r#"{"expected_turn_id":"turn-7","strict_expected_turn":true}"#
663 )));
664 assert!(task_is_steer(&task(r#"{"steer":true}"#)));
665 assert!(!task_is_steer(&task(r#"{"interaction_mode":"task"}"#)));
666 }
667
668 #[tokio::test]
669 async fn submit_subscriber_closes_at_its_final_text_boundary() {
670 let (events, _) = tokio::sync::broadcast::channel(8);
671 let generation = std::sync::Arc::new(std::sync::atomic::AtomicU64::new(0));
672 let mut stream = subscribe_turn_events(&events, &generation);
673 events
674 .send(Ok(PilotEvent {
675 event_kind: EVT_STATUS,
676 ..Default::default()
677 }))
678 .unwrap();
679 events
680 .send(Ok(PilotEvent {
681 event_kind: EVT_FINAL_TEXT,
682 final_text: "still running".into(),
683 ..Default::default()
684 }))
685 .unwrap();
686 events
687 .send(Ok(PilotEvent {
688 event_kind: EVT_STATUS,
689 ..Default::default()
690 }))
691 .unwrap();
692
693 assert_eq!(stream.next().await.unwrap().unwrap().event_kind, EVT_STATUS);
694 let final_event = stream.next().await.unwrap().unwrap();
695 assert_eq!(final_event.event_kind, EVT_FINAL_TEXT);
696 assert_eq!(final_event.final_text, "still running");
697 assert!(stream.next().await.is_none());
698 }
699
700 #[tokio::test]
701 async fn newest_submit_subscriber_exclusively_owns_future_replies() {
702 let (events, _) = tokio::sync::broadcast::channel(8);
703 let generation = std::sync::Arc::new(std::sync::atomic::AtomicU64::new(0));
704 let mut stale = subscribe_turn_events(&events, &generation);
705 let mut current = subscribe_turn_events(&events, &generation);
706
707 events
708 .send(Ok(PilotEvent {
709 event_kind: EVT_FINAL_TEXT,
710 final_text: "one reply".into(),
711 ..Default::default()
712 }))
713 .unwrap();
714
715 assert!(stale.next().await.is_none());
716 let final_event = current.next().await.unwrap().unwrap();
717 assert_eq!(final_event.event_kind, EVT_FINAL_TEXT);
718 assert_eq!(final_event.final_text, "one reply");
719 assert!(current.next().await.is_none());
720 }
721}