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