1use std::collections::{HashMap, HashSet, VecDeque};
16use std::future::Future;
17use std::pin::Pin;
18use std::sync::atomic::{AtomicBool, Ordering};
19use std::sync::{Arc, Mutex, MutexGuard, PoisonError};
20use std::time::{Duration, Instant};
21
22use async_trait::async_trait;
23use futures::stream::{FuturesUnordered, StreamExt};
24use prost::{Message, bytes::Bytes};
25use rlmesh_proto::EndpointPhases;
26use rlmesh_proto::core::v1::AutoresetMode;
27use rlmesh_proto::env::v1::{
28 EpisodeMetadata, ResetRequest, ResetResponse, StepRequest, StepResponse,
29};
30use rlmesh_proto::model::v1::{
31 AdapterContext, ObservationHistoryFrame, PredictRequest, PredictResponse,
32 ReleaseAdapterRequest, ResetAdapterRequest,
33};
34use rlmesh_proto::spaces::v1::{MetaMap, SpaceSpec, SpaceValue, TupleSpec, space_spec};
35use rlmesh_spaces::Advisory;
36use tokio::task::JoinSet;
37use tokio_util::sync::CancellationToken;
38
39use crate::hooks::{
40 ActionReceivedEvent, EpisodeCompletedEvent, EpisodeStartedEvent, Leg, LogEvent, LogLevel,
41 ObservationEmittedEvent, PayloadFacts, RefusingRelayPolicy, RelayAdvisoryEvent, RelayDecision,
42 RelayPolicy, RuntimeEnvContext, RuntimeHooks, SessionEndedEvent, SessionFailedEvent,
43 SessionStartedEvent, StepCompletedEvent, TelemetrySnapshotEvent,
44};
45use crate::spec::{ENV_RESET_OPTIONS_KEY, RuntimeReport, RuntimeSessionSpec, reset_options_for};
46use crate::state::{RequestPhase, RouteSnapshot, RouteState, StartedEpisode};
47use crate::telemetry::{Aggregator, Horizon, Sample, Source, metrics};
48
49mod error;
50
51pub use error::RuntimeError;
52
53macro_rules! fan_out_event {
56 ($self:ident, $method:ident, $event:expr) => {
57 if let Err(err) = $self.hooks.$method($event).await {
58 tracing::warn!(
59 concat!("runtime hook ", stringify!($method), " failed: {}"),
60 err
61 );
62 }
63 };
64}
65
66pub struct RuntimeEnvReset {
68 pub response: ResetResponse,
69 pub endpoint_total_ns: Option<u64>,
72 pub phases: EndpointPhases,
74}
75
76pub struct RuntimeEnvStep {
78 pub response: StepResponse,
79 pub endpoint_total_ns: Option<u64>,
81 pub phases: EndpointPhases,
83}
84
85pub struct RuntimeModelPrediction {
87 pub response: PredictResponse,
88 pub endpoint_total_ns: Option<u64>,
90 pub phases: EndpointPhases,
93 pub group_size: Option<u64>,
97}
98
99pub(crate) struct PeerReport {
102 pub(crate) endpoint_total_ns: Option<u64>,
103 pub(crate) phases: EndpointPhases,
104 pub(crate) group_size: Option<u64>,
105}
106
107#[async_trait]
113pub trait RuntimeEnv: Send {
114 async fn reset(&mut self, request: ResetRequest) -> Result<RuntimeEnvReset, RuntimeError>;
116
117 async fn step(&mut self, request: StepRequest) -> Result<RuntimeEnvStep, RuntimeError>;
119
120 async fn close(&mut self, _timeout: Duration) -> Result<(), String> {
123 Ok(())
124 }
125}
126
127#[async_trait]
130pub trait RuntimeModel: Send + Sync {
131 async fn predict(
134 &self,
135 request: PredictRequest,
136 ) -> Result<RuntimeModelPrediction, RuntimeError>;
137
138 async fn predict_group(
144 &self,
145 requests: Vec<PredictRequest>,
146 ) -> Vec<Result<RuntimeModelPrediction, RuntimeError>> {
147 futures::future::join_all(requests.into_iter().map(|request| self.predict(request))).await
148 }
149
150 fn fuses_predicts(&self) -> bool {
156 false
157 }
158
159 fn wants_history(&self) -> bool {
168 false
169 }
170
171 async fn reset_adapter(&self, _request: ResetAdapterRequest) -> Result<(), RuntimeError> {
177 Ok(())
178 }
179
180 async fn release_adapter(
182 &self,
183 _request: ReleaseAdapterRequest,
184 _timeout: Duration,
185 ) -> Result<(), String> {
186 Ok(())
187 }
188}
189
190pub trait PredictScheduler: Send {
198 fn plan(&mut self, waiting: &[usize], busy: usize) -> Vec<usize>;
199}
200
201pub struct EagerScheduler;
204
205impl PredictScheduler for EagerScheduler {
206 fn plan(&mut self, waiting: &[usize], _busy: usize) -> Vec<usize> {
207 waiting.to_vec()
208 }
209}
210
211const DEFAULT_CANCELLATION_REASON: &str = "cancelled by caller";
214
215fn success_from_final_info(
222 keys: &[&str],
223 final_info: Option<&rlmesh_proto::spaces::v1::MetaMap>,
224) -> Option<bool> {
225 use rlmesh_proto::spaces::v1::meta_value::Kind;
226 let entries = &final_info?.entries;
227 keys.iter()
228 .find_map(|key| match entries.get(*key)?.kind.as_ref()? {
229 Kind::Bool(value) => Some(*value),
230 Kind::Integer(value) => Some(*value != 0),
231 Kind::Number(value) => Some(*value != 0.0),
232 _ => None,
233 })
234}
235
236const SRC_PREDICT: Source = Source {
239 op: "model.predict",
240 component: "model",
241};
242const SRC_RESET_ADAPTER: Source = Source {
243 op: "model.reset_adapter",
244 component: "model",
245};
246const SRC_STEP: Source = Source {
247 op: "env.step",
248 component: "env",
249};
250const SRC_RESET: Source = Source {
251 op: "env.reset",
252 component: "env",
253};
254const SRC_TRANSFORM_OBS: Source = Source {
255 op: "runner.transform_observation",
256 component: "runner",
257};
258const SRC_TRANSFORM_ACTION: Source = Source {
259 op: "runner.transform_action",
260 component: "runner",
261};
262const SRC_ROUND: Source = Source {
266 op: "runner.round",
267 component: "runner",
268};
269
270#[must_use = "a RuntimeDriver does nothing until one of its run methods is awaited"]
273pub struct RuntimeDriver<E, M> {
274 spec: RuntimeSessionSpec,
275 env: E,
277 model: Option<Arc<M>>,
282 prefetch_lead: u32,
288 deliver_history: bool,
291 scheduler: Box<dyn PredictScheduler>,
292 hooks: Arc<dyn RuntimeHooks>,
293 relay_policy: Arc<dyn RelayPolicy>,
294 advisories: Mutex<Vec<Advisory>>,
296 cancellation_reason: String,
297 action_space: Arc<rlmesh_proto::spaces::v1::SpaceSpec>,
301 observation_space: Arc<rlmesh_proto::spaces::v1::SpaceSpec>,
302 pending_evictions: Vec<String>,
305 held_evictions: Vec<(usize, String)>,
310 trial_options_warned: AtomicBool,
315 vector_replay_warned: bool,
318}
319
320#[derive(Debug, Clone, Copy, PartialEq, Eq)]
322enum EnvPhase {
323 Ready,
325 Stepping,
327 Resetting,
329 Idle,
331}
332
333const HISTORY_BACKLOG_CAP: usize = 1023;
339
340enum PredictState {
342 None,
344 Wanted(PredictRequest),
346 InFlight { stale: bool },
350 Ready(VecDeque<Vec<Bytes>>),
352}
353
354struct Group<E> {
356 lanes: Vec<u32>,
358 positions: Vec<usize>,
360 lane_group: bool,
362 whole: bool,
364 phase: EnvPhase,
365 predict: PredictState,
366 env: Option<E>,
368 replay: VecDeque<Vec<Bytes>>,
372 obs_msg: Option<PredictRequest>,
375 history: Vec<ObservationHistoryFrame>,
379 steps: i64,
383 pending_roll: HashMap<u32, String>,
387 pending_completed: Vec<EpisodeCompletedEvent>,
391 pending_start: Option<(Vec<String>, Vec<u64>)>,
393 reset_generation: u64,
394 round_started: Instant,
395 pending_predict: Duration,
398}
399
400impl<E> Group<E> {
401 fn width(&self) -> usize {
402 self.lanes.len()
403 }
404
405 fn busy(&self) -> bool {
406 matches!(self.phase, EnvPhase::Stepping | EnvPhase::Resetting)
407 }
408}
409
410enum EnvOutcome<E> {
412 Reset {
413 group: usize,
414 env: E,
415 initial: bool,
416 request_bytes: u64,
417 rpc: Duration,
418 result: Result<RuntimeEnvReset, RuntimeError>,
419 },
420 Step {
421 group: usize,
422 env: E,
423 request_bytes: u64,
424 rpc: Duration,
425 result: Result<RuntimeEnvStep, RuntimeError>,
426 },
427}
428
429type PredictFuture<'m> = Pin<Box<dyn Future<Output = PredictOutcome> + Send + 'm>>;
433
434struct PredictOutcome {
436 requests: Vec<(usize, Option<AdapterContext>, u64)>,
438 rpc: Duration,
439 result: Result<Vec<Result<RuntimeModelPrediction, RuntimeError>>, RuntimeError>,
440}
441
442impl<E, M> RuntimeDriver<E, M>
443where
444 E: RuntimeEnv + Clone + 'static,
445 M: RuntimeModel,
446{
447 pub fn new(spec: RuntimeSessionSpec, env: E, model: M, hooks: Arc<dyn RuntimeHooks>) -> Self {
448 Self {
449 spec,
450 env,
451 model: Some(Arc::new(model)),
452 prefetch_lead: 0,
453 deliver_history: false,
454 scheduler: Box::new(EagerScheduler),
455 hooks,
456 relay_policy: Arc::new(RefusingRelayPolicy),
457 advisories: Mutex::new(Vec::new()),
458 cancellation_reason: DEFAULT_CANCELLATION_REASON.to_string(),
459 action_space: Arc::default(),
461 observation_space: Arc::default(),
462 pending_evictions: Vec::new(),
463 held_evictions: Vec::new(),
464 trial_options_warned: AtomicBool::new(false),
465 vector_replay_warned: false,
466 }
467 }
468
469 fn planned_trial_indices(&self, state: &mut RouteState, lanes: usize) -> Vec<u64> {
477 if self.driver_owns_resets() {
478 state.claim_trial_indices(self.spec.trial_index_base(), lanes)
479 } else {
480 Vec::new()
481 }
482 }
483
484 fn trial_options(&self, trials: &[u64]) -> Option<MetaMap> {
496 let option_key = self.spec.edition_defaults().trial_index_option_key;
497 let options = reset_options_for(&self.spec.env_contract, option_key, trials);
498 if options.is_none()
499 && !trials.is_empty()
500 && self.spec.trial_index_base() != 0
501 && !self.trial_options_warned.swap(true, Ordering::Relaxed)
502 {
503 tracing::warn!(
504 env_id = %self.spec.env_id,
505 key = ENV_RESET_OPTIONS_KEY,
506 option = option_key,
507 "trial_index_base is set but the env contract declares no such reset \
508 option; the ordinal is recorded on the episode events and summaries \
509 but not delivered to the env",
510 );
511 }
512 options
513 }
514
515 pub fn with_prefetch(mut self, lead: u32) -> Self {
521 self.prefetch_lead = lead;
522 self
523 }
524
525 pub fn with_scheduler(mut self, scheduler: Box<dyn PredictScheduler>) -> Self {
527 self.scheduler = scheduler;
528 self
529 }
530
531 pub fn with_relay_policy(mut self, relay_policy: Arc<dyn RelayPolicy>) -> Self {
533 self.relay_policy = relay_policy;
534 self
535 }
536
537 fn autoreset_mode(&self) -> AutoresetMode {
540 AutoresetMode::try_from(self.spec.env_contract.autoreset_mode)
544 .unwrap_or(AutoresetMode::Disabled)
545 }
546
547 fn driver_owns_resets(&self) -> bool {
550 self.spec
551 .edition_defaults()
552 .driver_owned_reset_modes
553 .contains(&self.autoreset_mode())
554 }
555
556 fn groups(&self) -> Vec<Group<E>> {
559 let num_envs = self.spec.num_envs.max(1);
560 let partitions: Vec<Vec<u32>> = if self.spec.subset_step && num_envs > 1 {
561 (0..num_envs as u32).map(|lane| vec![lane]).collect()
562 } else {
563 vec![(0..num_envs as u32).collect()]
564 };
565 let lane_group = self.spec.subset_step && num_envs > 1;
566 partitions
567 .into_iter()
568 .map(|lanes| Group {
569 positions: lanes.iter().map(|&lane| lane as usize).collect(),
570 whole: !lane_group,
571 lane_group,
572 lanes,
573 phase: EnvPhase::Ready,
574 predict: PredictState::None,
575 env: Some(self.env.clone()),
576 replay: VecDeque::new(),
577 obs_msg: None,
578 history: Vec::new(),
579 steps: 0,
580 pending_roll: HashMap::new(),
581 pending_completed: Vec::new(),
582 pending_start: None,
583 reset_generation: 0,
584 round_started: Instant::now(),
585 pending_predict: Duration::ZERO,
586 })
587 .collect()
588 }
589
590 fn seeds_for(&self, group: &Group<E>, slots: &[u64]) -> Vec<i64> {
594 if !self.spec.episode_seeds.is_empty() {
595 let seeds: Vec<Option<i64>> = slots
596 .iter()
597 .map(|&slot| {
598 self.spec
599 .episode_seeds
600 .get(usize::try_from(slot).unwrap_or(usize::MAX))
601 .copied()
602 })
603 .collect();
604 return if seeds.iter().all(Option::is_some) {
607 seeds.into_iter().flatten().collect()
608 } else {
609 if seeds.iter().any(Option::is_some) {
610 tracing::warn!(
611 lanes = group.width(),
612 "episode_seeds cannot cover this reset batch; it runs unseeded"
613 );
614 }
615 Vec::new()
616 };
617 }
618 let Some(base_seed) = self.spec.base_seed else {
619 return Vec::new();
620 };
621 if group.lane_group {
622 slots
625 .iter()
626 .map(|&slot| deterministic_reset_seed(base_seed, &self.spec.session_id, slot, 0))
627 .collect()
628 } else {
629 group
630 .lanes
631 .iter()
632 .map(|&lane| {
633 deterministic_reset_seed(
634 base_seed,
635 &self.spec.session_id,
636 group.reset_generation,
637 lane as usize,
638 )
639 })
640 .collect()
641 }
642 }
643
644 pub async fn run(self) -> Result<RuntimeReport, RuntimeError> {
645 self.run_with_cancellation(CancellationToken::new()).await
646 }
647
648 pub async fn run_with_cancellation(
649 self,
650 cancellation: CancellationToken,
651 ) -> Result<RuntimeReport, RuntimeError> {
652 self.run_with_cancellation_reason(cancellation, DEFAULT_CANCELLATION_REASON)
653 .await
654 }
655
656 pub async fn run_with_cancellation_reason(
664 mut self,
665 cancellation: CancellationToken,
666 reason: impl Into<String>,
667 ) -> Result<RuntimeReport, RuntimeError> {
668 self.cancellation_reason = reason.into();
669 self.spec.validate().map_err(RuntimeError::InvalidSpec)?;
670 self.deliver_history = self
671 .model
672 .as_ref()
673 .is_some_and(|model| model.wants_history());
674 self.action_space = Arc::new(self.spec.action_space_validated().clone());
677 self.observation_space = Arc::new(self.spec.observation_space_validated().clone());
678 let mut state = RouteState::new(&self.spec);
679 let telemetry = Arc::new(Mutex::new(Aggregator::default()));
685 let ticker = (!self.spec.limits.telemetry_window.is_zero()).then(|| {
688 TelemetryTicker::spawn(
689 Arc::clone(&telemetry),
690 Arc::clone(&self.hooks),
691 self.spec.limits.telemetry_window,
692 state.session_id().to_string(),
693 state.env_context(),
694 )
695 });
696 let mut env_ops: JoinSet<EnvOutcome<E>> = JoinSet::new();
697 let mut predicts: FuturesUnordered<PredictFuture<'_>> = FuturesUnordered::new();
698 let result = self
699 .run_loop(
700 &mut state,
701 &cancellation,
702 &telemetry,
703 &mut env_ops,
704 &mut predicts,
705 )
706 .await;
707 env_ops.abort_all();
711 let mut abandoned = false;
712 if !predicts.is_empty() {
713 let drain = async { while predicts.next().await.is_some() {} };
714 if tokio::time::timeout(self.spec.limits.service_close_timeout, drain)
715 .await
716 .is_err()
717 {
718 abandoned = true;
719 tracing::warn!("predict still in flight at route end; abandoned");
720 }
721 }
722 drop(predicts);
723 if abandoned {
734 self.pending_evictions.clear();
735 self.held_evictions.clear();
736 } else {
737 self.pending_evictions
738 .extend(self.held_evictions.drain(..).map(|(_, id)| id));
739 self.pending_evictions.extend(state.end_live_episodes());
740 self.flush_evictions(
741 &mut state,
742 &telemetry,
743 None,
744 self.spec.limits.service_close_timeout,
745 )
746 .await;
747 }
748 drop(ticker);
752 let final_snapshot = lock_agg(&telemetry).snapshot(Horizon::Session);
753 fan_out_event!(
754 self,
755 on_telemetry,
756 TelemetrySnapshotEvent {
757 session_id: state.session_id().to_string(),
758 route: state.env_context(),
759 snapshot: final_snapshot.clone(),
760 }
761 );
762 match result {
763 Ok(reason) => {
764 let release_request = state.release_adapter_request(reason);
765 self.shutdown_terminal_route(&state, reason, release_request)
766 .await;
767 fan_out_event!(
768 self,
769 session_ended,
770 SessionEndedEvent {
771 session_id: state.session_id().to_string(),
772 route: state.env_context(),
773 reason: reason.to_string(),
774 total_steps: state.total_steps(),
775 total_episodes: state.total_episodes(),
776 }
777 );
778 Ok(RuntimeReport {
779 session_id: state.session_id().to_string(),
780 env_id: self.spec.env_id.clone(),
781 total_steps: state.total_steps(),
782 total_episodes: state.total_episodes(),
783 episodes: state.take_episode_summaries(),
784 telemetry: final_snapshot,
785 advisories: std::mem::take(
786 self.advisories
787 .get_mut()
788 .unwrap_or_else(PoisonError::into_inner),
789 ),
790 })
791 }
792 Err(error) => {
793 self.shutdown_after_failure(&mut state, &error).await;
794 Err(error)
795 }
796 }
797 }
798
799 #[tracing::instrument(
804 name = "rlmesh.route",
805 level = "info",
806 skip_all,
807 fields(
808 session_id = %state.session_id(),
809 env_id = %self.spec.env_id,
810 num_envs = self.spec.num_envs,
811 lanes = self.spec.subset_step,
812 ),
813 )]
814 async fn run_loop<'m>(
815 &mut self,
816 state: &mut RouteState,
817 cancellation: &CancellationToken,
818 telemetry: &Arc<Mutex<Aggregator>>,
819 env_ops: &mut JoinSet<EnvOutcome<E>>,
820 predicts: &mut FuturesUnordered<PredictFuture<'m>>,
821 ) -> Result<&'static str, RuntimeError>
822 where
823 M: 'm,
824 {
825 fan_out_event!(
826 self,
827 session_started,
828 SessionStartedEvent {
829 session_id: state.session_id().to_string(),
830 route: state.env_context(),
831 env_id: self.spec.env_id.clone(),
832 }
833 );
834 self.relay_contract(state).await?;
835
836 let mut groups = self.groups();
837 let group_count = groups.len();
838 for gid in 0..group_count {
839 self.begin_reset(gid, &mut groups, state, env_ops, true);
840 }
841
842 loop {
843 if cancellation.is_cancelled() {
844 return Err(self.cancelled_error(state, &groups));
845 }
846 if groups.iter().all(|group| group.phase == EnvPhase::Idle)
847 && env_ops.is_empty()
848 && predicts.is_empty()
849 {
850 return Ok("completed requested episodes");
851 }
852 if predicts.is_empty() {
856 self.flush_evictions(
857 state,
858 telemetry,
859 Some(cancellation),
860 self.spec.limits.model_predict_timeout,
861 )
862 .await;
863 }
864 self.dispatch_steps(&mut groups, state, env_ops, telemetry)
865 .await?;
866 let force = env_ops.is_empty() && predicts.is_empty();
870 self.dispatch_predict(&mut groups, state, predicts, force);
871 if env_ops.is_empty() && predicts.is_empty() {
872 return Err(RuntimeError::Protocol(format!(
876 "route {} stalled with no operation in flight",
877 state.env_id()
878 )));
879 }
880
881 tokio::select! {
882 _ = cancellation.cancelled() => {
883 return Err(self.cancelled_error(state, &groups));
884 }
885 Some(joined) = env_ops.join_next() => {
886 let outcome = joined.map_err(|error| RuntimeError::Protocol(format!(
887 "env operation task failed: {error}"
888 )))?;
889 self.on_env_outcome(outcome, &mut groups, state, env_ops, telemetry)
890 .await?;
891 }
892 Some(outcome) = predicts.next(), if !predicts.is_empty() => {
893 self.on_predict_outcome(outcome, &mut groups, state, telemetry)?;
894 }
895 }
896 }
897 }
898
899 fn begin_reset(
904 &mut self,
905 gid: usize,
906 groups: &mut [Group<E>],
907 state: &mut RouteState,
908 env_ops: &mut JoinSet<EnvOutcome<E>>,
909 initial: bool,
910 ) {
911 let width = groups[gid].width();
912 let slots = if groups[gid].lane_group {
913 state.claim_slots(width, true).unwrap_or_default()
914 } else {
915 state.claim_slots_upto(width)
916 };
917 if slots.is_empty() {
918 groups[gid].phase = EnvPhase::Idle;
919 groups[gid].predict = PredictState::None;
920 return;
921 }
922 if !initial {
923 groups[gid].reset_generation += 1;
924 }
925 let seeds = self.seeds_for(&groups[gid], &slots);
926 let episode_ids = mint_episode_ids(width);
930 for surplus in &episode_ids[slots.len()..] {
931 state.mark_surplus(surplus);
932 }
933 state.note_episode_seeds(&episode_ids, &seeds);
934 let trials = self.planned_trial_indices(state, width);
935 state.note_episode_trials(&episode_ids, &trials);
936 let group = &mut groups[gid];
937 group.pending_start = Some((episode_ids.clone(), slots));
938 group.replay.clear();
939 group.history.clear();
942 group.pending_roll.clear();
943 let request = ResetRequest {
944 seeds,
945 options: self.trial_options(&trials),
946 timeout_ms: self.spec.limits.env_reset_timeout_ms().max(0) as u64,
947 env_indices: if group.whole {
948 Vec::new()
949 } else {
950 group.lanes.clone()
951 },
952 episode_ids,
953 };
954 let request_bytes = request.encoded_len() as u64;
955 let timeout = self.spec.limits.env_reset_timeout;
956 let timeout_error = RuntimeError::operation_timeout(
957 state.env_id(),
958 state.env_component_id(),
959 "env.reset",
960 0,
961 timeout,
962 );
963 let mut env = group
964 .env
965 .take()
966 .expect("group env handle present while ready");
967 group.phase = EnvPhase::Resetting;
968 env_ops.spawn(async move {
969 let started = Instant::now();
970 let result = match tokio::time::timeout(timeout, env.reset(request)).await {
971 Ok(result) => result,
972 Err(_) => Err(timeout_error),
973 };
974 EnvOutcome::Reset {
975 group: gid,
976 env,
977 initial,
978 request_bytes,
979 rpc: started.elapsed(),
980 result,
981 }
982 });
983 }
984
985 #[allow(clippy::needless_range_loop)]
989 async fn dispatch_steps(
990 &mut self,
991 groups: &mut [Group<E>],
992 state: &mut RouteState,
993 env_ops: &mut JoinSet<EnvOutcome<E>>,
994 telemetry: &Arc<Mutex<Aggregator>>,
995 ) -> Result<(), RuntimeError> {
996 for gid in 0..groups.len() {
997 if groups[gid].phase != EnvPhase::Ready {
998 continue;
999 }
1000 if groups[gid].replay.is_empty() {
1001 match std::mem::replace(&mut groups[gid].predict, PredictState::None) {
1002 PredictState::Ready(frames) => groups[gid].replay = frames,
1003 other => {
1004 groups[gid].predict = other;
1005 continue;
1006 }
1007 }
1008 }
1009 let Some(model_action) = groups[gid].replay.pop_front() else {
1010 continue;
1011 };
1012 if self.prefetch_lead > 0
1018 && groups[gid].replay.len() <= self.prefetch_lead as usize
1019 && matches!(groups[gid].predict, PredictState::None)
1020 && let Some(msg) = self.latest_request(&mut groups[gid], telemetry)
1021 {
1022 groups[gid].predict = PredictState::Wanted(msg);
1023 }
1024
1025 let group = &mut groups[gid];
1026 let snapshot = state.snapshot_at(&group.positions);
1027 let context = state.group_context(&group.positions, group.lane_group);
1028 let action_step = snapshot.step + 1;
1029 let mut action_event = ActionReceivedEvent {
1030 session_id: state.session_id().to_string(),
1031 route: context,
1032 episode_id: snapshot.episode_id.clone(),
1033 episode_record_id: snapshot.episode_record_id.clone(),
1034 episode_ids: snapshot.episode_ids.clone(),
1035 episode_record_ids: snapshot.episode_record_ids.clone(),
1036 step: action_step,
1037 env_index: snapshot.env_index,
1038 action_space: Arc::clone(&self.action_space),
1039 action: Some(model_action.clone()),
1040 raw_action: Some(model_action),
1041 };
1042 let (action, relayed_space) = self
1043 .invoke_transform_action(telemetry, &action_event)
1044 .await?;
1045 action_event.action = action;
1046 if let Some(space) = relayed_space {
1047 action_event.action_space = space;
1048 }
1049 fan_out_event!(self, action_received, action_event.clone());
1050
1051 let group = &mut groups[gid];
1052 let episode_ids = episode_ids_with_roll(
1057 state.episode_ids_at(&group.positions),
1058 &group.lanes,
1059 &group.pending_roll,
1060 );
1061 let request = StepRequest {
1062 action: action_event.action.map(leaves_value),
1063 timeout_ms: self.spec.limits.env_step_timeout_ms().max(0) as u64,
1064 env_indices: if group.whole {
1065 Vec::new()
1066 } else {
1067 group.lanes.clone()
1068 },
1069 episode_ids,
1070 };
1071 let request_bytes = request.encoded_len() as u64;
1072 let timeout = self.spec.limits.env_step_timeout;
1073 let timeout_error = RuntimeError::operation_timeout(
1074 state.env_id(),
1075 state.env_component_id(),
1076 "env.step",
1077 action_step,
1078 timeout,
1079 );
1080 let mut env = group
1081 .env
1082 .take()
1083 .expect("group env handle present while ready");
1084 group.phase = EnvPhase::Stepping;
1085 env_ops.spawn(async move {
1086 let started = Instant::now();
1087 let result = match tokio::time::timeout(timeout, env.step(request)).await {
1088 Ok(result) => result,
1089 Err(_) => Err(timeout_error),
1090 };
1091 EnvOutcome::Step {
1092 group: gid,
1093 env,
1094 request_bytes,
1095 rpc: started.elapsed(),
1096 result,
1097 }
1098 });
1099 }
1100 Ok(())
1101 }
1102
1103 fn dispatch_predict<'m>(
1108 &mut self,
1109 groups: &mut [Group<E>],
1110 state: &mut RouteState,
1111 predicts: &mut FuturesUnordered<PredictFuture<'m>>,
1112 force: bool,
1113 ) where
1114 M: 'm,
1115 {
1116 let fuses = self
1117 .model
1118 .as_ref()
1119 .expect("model handle set for the session")
1120 .fuses_predicts();
1121 if fuses && !predicts.is_empty() {
1122 return;
1123 }
1124 let waiting: Vec<usize> = groups
1125 .iter()
1126 .enumerate()
1127 .filter(|(_, group)| matches!(group.predict, PredictState::Wanted(_)))
1128 .map(|(gid, _)| gid)
1129 .collect();
1130 if waiting.is_empty() {
1131 return;
1132 }
1133 let busy = groups.iter().filter(|group| group.busy()).count();
1134 let mut chosen = self.scheduler.plan(&waiting, busy);
1135 chosen.retain(|gid| waiting.contains(gid));
1136 if chosen.is_empty() {
1137 if !force {
1138 return;
1139 }
1140 chosen = waiting;
1141 }
1142 if fuses {
1143 predicts.push(self.predict_for(groups, state, chosen));
1144 } else {
1145 for gid in chosen {
1146 predicts.push(self.predict_for(groups, state, vec![gid]));
1147 }
1148 }
1149 }
1150
1151 fn predict_for<'m>(
1154 &mut self,
1155 groups: &mut [Group<E>],
1156 state: &mut RouteState,
1157 chosen: Vec<usize>,
1158 ) -> PredictFuture<'m>
1159 where
1160 M: 'm,
1161 {
1162 let mut requests = Vec::with_capacity(chosen.len());
1163 let mut metas = Vec::with_capacity(chosen.len());
1164 let mut step = 0;
1165 for gid in chosen {
1166 let PredictState::Wanted(msg) = std::mem::replace(
1167 &mut groups[gid].predict,
1168 PredictState::InFlight { stale: false },
1169 ) else {
1170 unreachable!("only waiting groups are planned")
1171 };
1172 step = step.max(state.snapshot_at(&groups[gid].positions).step);
1173 state.mark_predicted(&groups[gid].positions);
1174 metas.push((gid, msg.context.clone(), msg.encoded_len() as u64));
1175 requests.push(msg);
1176 }
1177 let timeout = self.spec.limits.model_predict_timeout;
1178 let timeout_error = RuntimeError::operation_timeout(
1179 state.env_id(),
1180 state.model_component_id(),
1181 "model.predict",
1182 step,
1183 timeout,
1184 );
1185 let model = Arc::clone(
1186 self.model
1187 .as_ref()
1188 .expect("model handle set for the session"),
1189 );
1190 Box::pin(async move {
1191 let started = Instant::now();
1192 let result = match tokio::time::timeout(timeout, model.predict_group(requests)).await {
1193 Ok(results) => Ok(results),
1194 Err(_) => Err(timeout_error),
1195 };
1196 PredictOutcome {
1197 requests: metas,
1198 rpc: started.elapsed(),
1199 result,
1200 }
1201 })
1202 }
1203
1204 fn on_predict_outcome(
1208 &mut self,
1209 outcome: PredictOutcome,
1210 groups: &mut [Group<E>],
1211 state: &RouteState,
1212 telemetry: &Arc<Mutex<Aggregator>>,
1213 ) -> Result<(), RuntimeError> {
1214 let results = outcome.result?;
1215 if results.len() != outcome.requests.len() {
1216 return Err(RuntimeError::Protocol(format!(
1217 "model endpoint {} answered {} of {} grouped predicts",
1218 state.model_component_id(),
1219 results.len(),
1220 outcome.requests.len()
1221 )));
1222 }
1223 let group_count = outcome.requests.len() as u64;
1224 let rpc = outcome.rpc;
1225 let mut recorded = false;
1226 for ((gid, expected_context, request_bytes), result) in
1227 outcome.requests.into_iter().zip(results)
1228 {
1229 let prediction = result?;
1230 if prediction.response.context != expected_context {
1231 let request_id = expected_context
1232 .as_ref()
1233 .map(|context| context.request_id.clone())
1234 .unwrap_or_default();
1235 return Err(RuntimeError::ModelRouteMismatch {
1236 component_id: state.model_component_id().to_string(),
1237 request_id,
1238 });
1239 }
1240 if !recorded {
1243 recorded = true;
1244 record_op(
1245 telemetry,
1246 SRC_PREDICT,
1247 outcome.rpc,
1248 PeerReport {
1249 endpoint_total_ns: prediction.endpoint_total_ns,
1250 phases: prediction.phases,
1251 group_size: if group_count > 1 {
1252 Some(group_count)
1253 } else {
1254 prediction.group_size
1255 },
1256 },
1257 request_bytes,
1258 prediction.response.encoded_len() as u64,
1259 );
1260 }
1261 if prediction.response.actions.is_empty() {
1262 return Err(RuntimeError::Protocol(format!(
1263 "model endpoint {} returned a predict response with no actions",
1264 state.model_component_id()
1265 )));
1266 }
1267 let mut frames = VecDeque::with_capacity(prediction.response.actions.len());
1269 for frame in &prediction.response.actions {
1270 if let Some(leaves) = value_leaves(Some(frame))? {
1271 frames.push_back(leaves);
1272 }
1273 }
1274 let group = &mut groups[gid];
1275 group.pending_predict += rpc;
1276 group.predict = match std::mem::replace(&mut group.predict, PredictState::None) {
1277 PredictState::InFlight { stale: true } => {
1296 if group.phase == EnvPhase::Ready && group.replay.is_empty() {
1297 self.latest_request(group, telemetry)
1298 .map_or(PredictState::None, PredictState::Wanted)
1299 } else {
1300 PredictState::None
1301 }
1302 }
1303 PredictState::InFlight { stale: false } => {
1304 if group.replay.is_empty() {
1305 group.replay = frames;
1306 PredictState::None
1307 } else {
1308 PredictState::Ready(frames)
1309 }
1310 }
1311 other => other,
1312 };
1313 self.pending_evictions.extend(
1314 self.held_evictions
1315 .extract_if(.., |(held, _)| *held == gid)
1316 .map(|(_, id)| id),
1317 );
1318 }
1319 Ok(())
1320 }
1321
1322 async fn on_env_outcome(
1323 &mut self,
1324 outcome: EnvOutcome<E>,
1325 groups: &mut [Group<E>],
1326 state: &mut RouteState,
1327 env_ops: &mut JoinSet<EnvOutcome<E>>,
1328 telemetry: &Arc<Mutex<Aggregator>>,
1329 ) -> Result<(), RuntimeError> {
1330 match outcome {
1331 EnvOutcome::Reset {
1332 group: gid,
1333 env,
1334 initial,
1335 request_bytes,
1336 rpc,
1337 result,
1338 } => {
1339 groups[gid].env = Some(env);
1340 let reset = result?;
1341 record_op(
1342 telemetry,
1343 SRC_RESET,
1344 rpc,
1345 PeerReport {
1346 endpoint_total_ns: reset.endpoint_total_ns,
1347 phases: reset.phases,
1348 group_size: None,
1349 },
1350 request_bytes,
1351 reset.response.encoded_len() as u64,
1352 );
1353 let group = &mut groups[gid];
1354 let (episode_ids, slots) = group
1355 .pending_start
1356 .take()
1357 .expect("a reset in flight has its episodes staged");
1358 let context = state.group_context(&group.positions, group.lane_group);
1359 if initial {
1360 fan_out_event!(
1361 self,
1362 log,
1363 LogEvent {
1364 session_id: state.session_id().to_string(),
1365 route: context.clone(),
1366 level: LogLevel::Info,
1367 message: format!(
1368 "env reset complete in {:.0}ms ({} episode(s) ready)",
1369 rpc.as_secs_f64() * 1000.0,
1370 episode_ids.len()
1371 ),
1372 source: Some("runtime".to_string()),
1373 }
1374 );
1375 }
1376 let started =
1377 state.start_episodes_at(&groups[gid].positions, episode_ids, !initial, &slots);
1378 self.invoke_started_episodes(state, &context, started).await;
1379 groups[gid].round_started = Instant::now();
1380 let observation = value_leaves(reset.response.observation.as_ref())?;
1381 self.observe(
1382 gid,
1383 groups,
1384 state,
1385 telemetry,
1386 observation,
1387 reset.response.infos,
1388 RequestPhase::ResetObservation,
1389 true,
1390 )
1391 .await
1392 }
1393 EnvOutcome::Step {
1394 group: gid,
1395 env,
1396 request_bytes,
1397 rpc,
1398 result,
1399 } => {
1400 groups[gid].env = Some(env);
1401 let step = result?;
1402 record_op(
1403 telemetry,
1404 SRC_STEP,
1405 rpc,
1406 PeerReport {
1407 endpoint_total_ns: step.endpoint_total_ns,
1408 phases: step.phases,
1409 group_size: None,
1410 },
1411 request_bytes,
1412 step.response.encoded_len() as u64,
1413 );
1414 self.on_step(gid, groups, state, env_ops, telemetry, step.response, rpc)
1415 .await
1416 }
1417 }
1418 }
1419
1420 #[allow(clippy::too_many_arguments)]
1423 async fn on_step(
1424 &mut self,
1425 gid: usize,
1426 groups: &mut [Group<E>],
1427 state: &mut RouteState,
1428 env_ops: &mut JoinSet<EnvOutcome<E>>,
1429 telemetry: &Arc<Mutex<Aggregator>>,
1430 response: StepResponse,
1431 rpc: Duration,
1432 ) -> Result<(), RuntimeError> {
1433 let positions = groups[gid].positions.clone();
1434 let lane_group = groups[gid].lane_group;
1435 let context = state.group_context(&positions, lane_group);
1436 lock_agg(telemetry).record(Sample::dur(
1437 SRC_ROUND,
1438 metrics::RPC_TOTAL,
1439 groups[gid].round_started.elapsed(),
1440 ));
1441 groups[gid].round_started = Instant::now();
1442 groups[gid].phase = EnvPhase::Ready;
1443 groups[gid].steps += 1;
1444
1445 let step_observation = value_leaves(response.observation.as_ref())?;
1446 let predict = std::mem::take(&mut groups[gid].pending_predict);
1447 state.record_step_at(&positions, &response.rewards, rpc, predict);
1448 let snapshot = state.snapshot_at(&positions);
1449 let mut completed_episodes: Vec<EpisodeMetadata> = Vec::new();
1456 for metadata in &response.completed_episodes {
1457 if state.episode_id_at(metadata.env_index) == Some(metadata.episode_id.as_str()) {
1458 completed_episodes.push(metadata.clone());
1459 } else {
1460 tracing::debug!(
1461 env_index = metadata.env_index,
1462 episode_id = %metadata.episode_id,
1463 "dropping a stale env-reported episode completion: the lane has moved past \
1464 this id"
1465 );
1466 }
1467 }
1468
1469 let capped = self.capped_completions(
1475 state,
1476 &positions,
1477 &completed_episodes,
1478 response.infos.as_ref(),
1479 );
1480 completed_episodes.extend(capped);
1481 let lane_ids = state.episode_ids_at(&positions);
1482 let lane_flags = |ended: fn(&EpisodeMetadata) -> bool| -> Vec<bool> {
1483 lane_ids
1484 .iter()
1485 .map(|id| {
1486 completed_episodes
1487 .iter()
1488 .any(|m| m.episode_id == *id && ended(m))
1489 })
1490 .collect()
1491 };
1492 let terminated = lane_flags(|m| m.terminated);
1493 let truncated = lane_flags(|m| m.truncated);
1494 let autoreset_roll: Vec<bool> = groups[gid]
1495 .lanes
1496 .iter()
1497 .map(|lane| groups[gid].pending_roll.contains_key(lane))
1498 .collect();
1499 let rolled = !groups[gid].pending_roll.is_empty();
1503 fan_out_event!(
1504 self,
1505 step_completed,
1506 StepCompletedEvent {
1507 session_id: state.session_id().to_string(),
1508 route: context.clone(),
1509 episode_id: snapshot.episode_id.clone(),
1510 episode_record_id: snapshot.episode_record_id.clone(),
1511 step: snapshot.step,
1512 env_index: snapshot.env_index,
1513 rewards: response.rewards.clone(),
1514 infos: if rolled { None } else { response.infos.clone() },
1515 terminated,
1516 truncated,
1517 autoreset_roll,
1518 predict_ms: predict.as_secs_f64() * 1e3,
1519 step_ms: rpc.as_secs_f64() * 1e3,
1520 }
1521 );
1522
1523 if rolled {
1527 let completions = std::mem::take(&mut groups[gid].pending_completed);
1530 self.fan_out_completions(completions).await;
1531 let pending_roll = std::mem::take(&mut groups[gid].pending_roll);
1532 self.queue_evictions(gid, &groups[gid], state, pending_roll.keys().copied());
1535 let roll_ids = episode_ids_with_roll(
1536 state.episode_ids_at(&positions),
1537 &groups[gid].lanes,
1538 &pending_roll,
1539 );
1540 let rolling: Vec<Option<u64>> = groups[gid]
1541 .lanes
1542 .iter()
1543 .map(|lane| {
1544 pending_roll
1545 .contains_key(lane)
1546 .then(|| state.claim_slots(1, true))
1547 .flatten()
1548 .and_then(|slots| slots.first().copied())
1549 })
1550 .collect();
1551 for ((lane, id), slot) in groups[gid].lanes.iter().zip(&roll_ids).zip(&rolling) {
1553 if pending_roll.contains_key(lane) && slot.is_none() {
1554 state.mark_surplus(id);
1555 }
1556 }
1557 let started = state.observe_episode_ids_at(&positions, roll_ids, &rolling);
1558 self.invoke_started_episodes(state, &context, started).await;
1559 }
1560
1561 let completions = self.complete_episodes(state, &context, &completed_episodes);
1562 if self.driver_owns_resets() {
1568 self.fan_out_completions(completions).await;
1569 self.queue_evictions(
1570 gid,
1571 &groups[gid],
1572 state,
1573 completed_episodes.iter().map(|c| c.env_index),
1574 );
1575 } else {
1576 groups[gid].pending_completed.extend(completions);
1577 }
1578
1579 if !completed_episodes.is_empty() {
1580 let group = &mut groups[gid];
1585 if !group.replay.is_empty() && group.width() > 1 && !self.vector_replay_warned {
1586 self.vector_replay_warned = true;
1590 tracing::warn!(
1591 num_envs = group.width(),
1592 discarded_frames = group.replay.len(),
1593 "a lane's episode ended mid-chunk on a lockstep vector group: chunk replay \
1594 is whole-batch, so every lane's buffered frames were discarded and the \
1595 group re-plans. Serve lanes, or use an execution horizon of 1.",
1596 );
1597 }
1598 group.replay.clear();
1599 group.predict = match std::mem::replace(&mut group.predict, PredictState::None) {
1600 PredictState::InFlight { .. } => PredictState::InFlight { stale: true },
1601 _ => PredictState::None,
1602 };
1603 if !self.driver_owns_resets() {
1607 for completed in &completed_episodes {
1608 group
1609 .pending_roll
1610 .entry(completed.env_index)
1611 .or_insert_with(mint_episode_id);
1612 }
1613 }
1614 }
1615
1616 let budget_spent = self
1624 .spec
1625 .max_episodes
1626 .is_some_and(|limit| state.total_episodes() >= limit as i64);
1627 if ((!lane_group || !self.driver_owns_resets()) && budget_spent)
1628 || state.all_surplus_at(&positions)
1629 {
1630 let completions = std::mem::take(&mut groups[gid].pending_completed);
1635 self.fan_out_completions(completions).await;
1636 let pending_roll = std::mem::take(&mut groups[gid].pending_roll);
1637 self.queue_evictions(gid, &groups[gid], state, pending_roll.keys().copied());
1638 groups[gid].phase = EnvPhase::Idle;
1639 groups[gid].predict = PredictState::None;
1640 return Ok(());
1641 }
1642
1643 if self.driver_owns_resets() {
1647 let done: HashSet<u32> = completed_episodes
1648 .iter()
1649 .map(|metadata| metadata.env_index)
1650 .collect();
1651 if !done.is_empty() {
1652 self.begin_reset(gid, groups, state, env_ops, false);
1653 return Ok(());
1654 }
1655 }
1656
1657 self.observe(
1658 gid,
1659 groups,
1660 state,
1661 telemetry,
1662 step_observation,
1663 if rolled { response.infos } else { None },
1664 RequestPhase::StepObservation,
1665 false,
1666 )
1667 .await
1668 }
1669
1670 #[allow(clippy::too_many_arguments)]
1674 async fn observe(
1675 &mut self,
1676 gid: usize,
1677 groups: &mut [Group<E>],
1678 state: &mut RouteState,
1679 telemetry: &Arc<Mutex<Aggregator>>,
1680 observation: Option<Vec<Bytes>>,
1681 infos: Option<rlmesh_proto::spaces::v1::MetaMap>,
1682 phase: RequestPhase,
1683 is_reset: bool,
1684 ) -> Result<(), RuntimeError> {
1685 let positions = groups[gid].positions.clone();
1686 let context = state.group_context(&positions, groups[gid].lane_group);
1687 let mut msg = state.predict_request_at(&positions, observation.clone(), phase);
1688 let mut event = self.observation_event(
1689 state,
1690 context,
1691 state.snapshot_at(&positions),
1692 is_reset,
1693 observation,
1694 infos,
1695 groups[gid].width(),
1696 );
1697 let (transformed, relayed_space) =
1698 self.invoke_transform_observation(telemetry, &event).await?;
1699 event.observation = transformed.clone();
1700 if let Some(space) = relayed_space {
1701 event.observation_space = space;
1702 }
1703 msg.observation = transformed.map(leaves_value);
1704 fan_out_event!(self, observation_emitted, event);
1708
1709 let group = &mut groups[gid];
1710 let replans = group.replay.is_empty() && matches!(group.predict, PredictState::None);
1715 if self.deliver_history {
1716 msg.step = Some(group.steps);
1724 if replans {
1725 msg.history = std::mem::take(&mut group.history);
1726 record_history_rows(telemetry, &msg);
1727 } else {
1728 if group.history.len() >= HISTORY_BACKLOG_CAP {
1729 return Err(RuntimeError::Protocol(format!(
1730 "route {} buffered {} observation-history rows without a predict: the \
1731 model endpoint returned a chunk longer than any execution horizon \
1732 the runtime supports",
1733 state.env_id(),
1734 group.history.len()
1735 )));
1736 }
1737 group.history.push(ObservationHistoryFrame {
1738 observation: msg.observation.clone(),
1739 episode_info: msg.episode_info.clone(),
1740 step: group.steps,
1741 });
1742 }
1743 }
1744 group.obs_msg = Some(msg.clone());
1745 group.phase = EnvPhase::Ready;
1746 if replans {
1747 group.predict = PredictState::Wanted(msg);
1748 }
1749 Ok(())
1750 }
1751
1752 fn latest_request(
1764 &self,
1765 group: &mut Group<E>,
1766 telemetry: &Arc<Mutex<Aggregator>>,
1767 ) -> Option<PredictRequest> {
1768 let mut msg = group.obs_msg.clone()?;
1769 if self.deliver_history {
1770 if group.history.last().map(|row| row.step) != msg.step {
1771 return None;
1772 }
1773 group.history.pop();
1774 msg.history = std::mem::take(&mut group.history);
1775 record_history_rows(telemetry, &msg);
1776 }
1777 Some(msg)
1778 }
1779
1780 async fn shutdown_after_failure(&mut self, state: &mut RouteState, error: &RuntimeError) {
1781 let reason = error.to_string();
1782 let request = state.release_adapter_request(reason.clone());
1783 self.shutdown_terminal_route(state, &reason, request).await;
1784
1785 if let Err(err) = self
1786 .hooks
1787 .session_failed(SessionFailedEvent {
1788 session_id: state.session_id().to_string(),
1789 route: state.env_context(),
1790 reason,
1791 })
1792 .await
1793 {
1794 tracing::warn!("runtime hook session_failed failed: {err}");
1795 }
1796 }
1797
1798 async fn shutdown_terminal_route(
1799 &mut self,
1800 state: &RouteState,
1801 reason: &str,
1802 request: ReleaseAdapterRequest,
1803 ) {
1804 let timeout = self.spec.limits.service_close_timeout;
1805 let model_close = async {
1810 match self.model.as_ref() {
1811 Some(model) => {
1812 tokio::time::timeout(timeout, model.release_adapter(request, timeout)).await
1813 }
1814 None => Ok(Ok(())),
1815 }
1816 };
1817 if self.spec.close_env_on_end {
1818 let env_close = tokio::time::timeout(timeout, self.env.close(timeout));
1819 let (env_result, model_result) = tokio::join!(env_close, model_close);
1820 match env_result {
1821 Ok(Err(err)) => {
1822 tracing::warn!(error = %err, "environment close failed during route shutdown");
1823 }
1824 Err(_) => {
1825 tracing::warn!(
1826 timeout_ms = timeout.as_millis(),
1827 "environment close timed out during route shutdown; abandoning close"
1828 );
1829 }
1830 Ok(Ok(())) => {}
1831 }
1832 log_model_close_result(model_result, reason, timeout);
1833 return;
1834 }
1835
1836 tracing::debug!(
1837 env_id = %state.env_id(),
1838 reason,
1839 "skipping environment close for adapter; endpoint remains owned by the run"
1840 );
1841 log_model_close_result(model_close.await, reason, timeout);
1842 }
1843
1844 fn cancelled_error(&self, state: &RouteState, groups: &[Group<E>]) -> RuntimeError {
1845 let step = groups
1846 .iter()
1847 .map(|group| state.snapshot_at(&group.positions).step)
1848 .max()
1849 .unwrap_or(0);
1850 RuntimeError::route_cancelled(state.env_id(), step, self.cancellation_reason.as_str())
1851 }
1852
1853 async fn invoke_started_episodes(
1854 &self,
1855 state: &RouteState,
1856 context: &RuntimeEnvContext,
1857 episodes: Vec<StartedEpisode>,
1858 ) {
1859 for episode in episodes {
1860 if state.is_surplus(&episode.episode_id) {
1861 continue;
1862 }
1863 let record = &episode.record;
1864 fan_out_event!(
1865 self,
1866 episode_started,
1867 EpisodeStartedEvent {
1868 session_id: state.session_id().to_string(),
1869 route: context.clone(),
1870 episode_id: episode.episode_id.clone(),
1871 episode_record_id: record.record_id.clone(),
1872 episode_index: record.index,
1873 env_index: record.env_index,
1874 started_from_auto_reset: record.started_from_auto_reset,
1875 seed: state.seed_for_episode(&episode.episode_id),
1876 trial_index: state.trial_for_episode(&episode.episode_id),
1877 }
1878 );
1879 }
1880 }
1881
1882 fn capped_completions(
1890 &self,
1891 state: &RouteState,
1892 positions: &[usize],
1893 env_completed: &[EpisodeMetadata],
1894 infos: Option<&rlmesh_proto::spaces::v1::MetaMap>,
1895 ) -> Vec<EpisodeMetadata> {
1896 let step_cap = self.spec.max_episode_steps.or_else(|| {
1897 self.driver_owns_resets()
1898 .then_some(self.spec.edition_defaults().default_max_episode_steps)
1899 });
1900 let time_cap = self.spec.max_episode_seconds;
1901 if step_cap.is_none() && time_cap.is_none() {
1902 return Vec::new();
1903 }
1904 let env_done: Vec<u32> = env_completed
1905 .iter()
1906 .map(|metadata| metadata.env_index)
1907 .collect();
1908 let now_ns = crate::state::now_unix_ns();
1909 state
1910 .slots_at(positions)
1911 .into_iter()
1912 .enumerate()
1913 .filter_map(|(lane, slot)| {
1914 let episode = slot.episode.as_ref()?;
1915 let env_index = u32::try_from(slot.env_index).ok()?;
1916 if env_done.contains(&env_index) {
1917 return None;
1918 }
1919 let steps_capped = step_cap.is_some_and(|cap| slot.step >= cap);
1920 let elapsed_seconds = (now_ns - slot.started_at_ns).max(0) as f64 / 1e9;
1921 let time_capped = time_cap.is_some_and(|cap| elapsed_seconds >= cap);
1922 (steps_capped || time_capped).then(|| EpisodeMetadata {
1923 episode_id: episode.episode_id.clone(),
1924 seed: None,
1925 env_index,
1926 step_count: slot.step,
1927 cumulative_reward: slot.cumulative_reward,
1928 terminated: false,
1929 truncated: true,
1930 start_timestamp_ns: slot.started_at_ns,
1931 end_timestamp_ns: now_ns,
1932 final_info: rlmesh_proto::lane_final_info(infos, lane, positions.len()),
1933 })
1934 })
1935 .collect()
1936 }
1937
1938 fn complete_episodes(
1946 &self,
1947 state: &mut RouteState,
1948 context: &RuntimeEnvContext,
1949 episodes: &[EpisodeMetadata],
1950 ) -> Vec<EpisodeCompletedEvent> {
1951 let mut events = Vec::with_capacity(episodes.len());
1952 for completed in episodes {
1953 if state.is_surplus(&completed.episode_id) {
1954 continue;
1955 }
1956 let (predict_ms, step_ms) = state.slot_timings_ms(completed.env_index);
1957 let record = state.complete_episode(&completed.episode_id);
1958 let episode_record_id = record
1959 .as_ref()
1960 .map(|record| record.record_id.clone())
1961 .unwrap_or_default();
1962 let env_index = i32::try_from(completed.env_index).unwrap_or(i32::MAX);
1964 let seed = state.seed_for_episode(&completed.episode_id);
1965 let trial_index = state.trial_for_episode(&completed.episode_id);
1966 if self.spec.max_episodes.is_some() {
1967 state.record_episode_summary(crate::spec::EpisodeSummary {
1968 episode_index: record.as_ref().map_or(0, |record| record.index),
1969 env_index,
1970 seed,
1971 trial_index,
1972 step_count: completed.step_count,
1973 cumulative_reward: completed.cumulative_reward,
1974 terminated: completed.terminated,
1975 truncated: completed.truncated,
1976 duration_ms: (completed.end_timestamp_ns - completed.start_timestamp_ns).max(0)
1977 / 1_000_000,
1978 success: success_from_final_info(
1979 self.spec.edition_defaults().success_info_keys,
1980 completed.final_info.as_ref(),
1981 ),
1982 predict_ms,
1983 step_ms,
1984 });
1985 }
1986 events.push(EpisodeCompletedEvent {
1987 session_id: state.session_id().to_string(),
1988 route: context.clone(),
1989 episode_id: completed.episode_id.clone(),
1990 episode_record_id,
1991 episode_index: record.as_ref().map_or(0, |record| record.index),
1992 env_index,
1993 step_count: completed.step_count,
1994 cumulative_reward: completed.cumulative_reward,
1995 terminated: completed.terminated,
1996 truncated: completed.truncated,
1997 duration_ms: (completed.end_timestamp_ns - completed.start_timestamp_ns).max(0)
1998 / 1_000_000,
1999 success: success_from_final_info(
2000 self.spec.edition_defaults().success_info_keys,
2001 completed.final_info.as_ref(),
2002 ),
2003 final_info: completed.final_info.clone(),
2004 seed,
2005 trial_index,
2006 predict_ms,
2007 step_ms,
2008 });
2009 }
2010 events
2011 }
2012
2013 async fn fan_out_completions(&self, events: Vec<EpisodeCompletedEvent>) {
2015 for event in events {
2016 fan_out_event!(self, episode_completed, event);
2017 }
2018 }
2019
2020 fn queue_evictions(
2038 &mut self,
2039 gid: usize,
2040 group: &Group<E>,
2041 state: &mut RouteState,
2042 env_indices: impl Iterator<Item = u32>,
2043 ) {
2044 let ids = env_indices.filter_map(|env_index| state.end_episode_at(env_index));
2045 if matches!(group.predict, PredictState::InFlight { .. }) {
2046 self.held_evictions.extend(ids.map(|id| (gid, id)));
2047 } else {
2048 self.pending_evictions.extend(ids);
2049 }
2050 }
2051
2052 async fn flush_evictions(
2059 &mut self,
2060 state: &mut RouteState,
2061 telemetry: &Mutex<Aggregator>,
2062 cancellation: Option<&CancellationToken>,
2063 deadline: Duration,
2064 ) {
2065 if self.pending_evictions.is_empty() {
2066 return;
2067 }
2068 let Some(model) = self.model.as_ref() else {
2069 return;
2070 };
2071 let episode_ids = std::mem::take(&mut self.pending_evictions);
2072 let episodes = episode_ids.len();
2073 let request = state.reset_adapter_request(episode_ids);
2074 let cancelled = async {
2075 match cancellation {
2076 Some(cancellation) => cancellation.cancelled().await,
2077 None => std::future::pending().await,
2078 }
2079 };
2080 let started = Instant::now();
2086 tokio::select! {
2087 biased;
2088 () = cancelled => {
2089 tracing::warn!(episodes, "model reset_adapter (evict) abandoned: route cancelled");
2090 }
2091 result = tokio::time::timeout(deadline, model.reset_adapter(request)) => match result {
2092 Ok(result) => {
2093 lock_agg(telemetry).record(Sample::dur(
2094 SRC_RESET_ADAPTER,
2095 metrics::RPC_TOTAL,
2096 started.elapsed(),
2097 ));
2098 if let Err(err) = result {
2099 tracing::warn!("model reset_adapter (evict) failed: {err}");
2100 }
2101 }
2102 Err(_) => tracing::warn!(
2103 episodes,
2104 timeout_ms = deadline.as_millis(),
2105 "model reset_adapter (evict) timed out; abandoning"
2106 ),
2107 }
2108 }
2109 }
2110
2111 async fn invoke_transform_action(
2114 &self,
2115 telemetry: &Mutex<Aggregator>,
2116 event: &ActionReceivedEvent,
2117 ) -> Result<Relayed, RuntimeError> {
2118 let started = Instant::now();
2119 let result = self.hooks.transform_action(event.clone()).await;
2120 lock_agg(telemetry).record(Sample::dur(
2121 SRC_TRANSFORM_ACTION,
2122 metrics::RPC_TOTAL,
2123 started.elapsed(),
2124 ));
2125 match result {
2126 Ok(action) => {
2127 self.relay(
2128 Leg::ModelToEnv,
2129 &event.session_id,
2130 &event.route,
2131 &self.action_space,
2132 action,
2133 )
2134 .await
2135 }
2136 Err(err) => {
2137 tracing::warn!("runtime hook transform_action failed: {err}");
2138 Err(RuntimeError::Hook(err))
2139 }
2140 }
2141 }
2142
2143 async fn invoke_transform_observation(
2146 &self,
2147 telemetry: &Mutex<Aggregator>,
2148 event: &ObservationEmittedEvent,
2149 ) -> Result<Relayed, RuntimeError> {
2150 let started = Instant::now();
2151 let result = self.hooks.transform_observation(event.clone()).await;
2152 lock_agg(telemetry).record(Sample::dur(
2153 SRC_TRANSFORM_OBS,
2154 metrics::RPC_TOTAL,
2155 started.elapsed(),
2156 ));
2157 match result {
2158 Ok(observation) => {
2159 self.relay(
2160 Leg::EnvToModel,
2161 &event.session_id,
2162 &event.route,
2163 &self.observation_space,
2164 observation,
2165 )
2166 .await
2167 }
2168 Err(err) => {
2169 tracing::warn!("runtime hook transform_observation failed: {err}");
2170 Err(RuntimeError::Hook(err))
2171 }
2172 }
2173 }
2174
2175 fn ceiling(&self, leg: Leg) -> Option<&crate::spec::PeerCeiling> {
2177 match leg {
2178 Leg::EnvToModel => self.spec.model_ceiling.as_ref(),
2179 Leg::ModelToEnv => self.spec.env_ceiling.as_ref(),
2180 }
2181 }
2182
2183 async fn relay_contract(&self, state: &RouteState) -> Result<(), RuntimeError> {
2186 let Some(ceiling) = self.ceiling(Leg::EnvToModel) else {
2187 return Ok(());
2188 };
2189 let contract_spaces = SpaceSpec {
2190 spec: Some(space_spec::Spec::Tuple(TupleSpec {
2191 spaces: vec![
2192 SpaceSpec::clone(&self.observation_space),
2193 SpaceSpec::clone(&self.action_space),
2194 ],
2195 })),
2196 ..Default::default()
2197 };
2198 let payload = PayloadFacts {
2199 byte_len: self.spec.env_contract.encoded_len(),
2200 ..PayloadFacts::new(Arc::new(contract_spaces), Vec::new())
2201 };
2202 match self
2203 .relay_policy
2204 .reconcile(Leg::EnvToModel, ceiling, &payload)
2205 {
2206 RelayDecision::Forward => Ok(()),
2207 RelayDecision::Convert { advisory, .. } => {
2208 self.raise_advisory(
2209 state.session_id(),
2210 &state.env_context(),
2211 Leg::EnvToModel,
2212 advisory,
2213 )
2214 .await;
2215 Ok(())
2216 }
2217 RelayDecision::Refuse(reason) => {
2218 Err(relay_refused(Leg::EnvToModel, "contract", reason))
2219 }
2220 }
2221 }
2222
2223 async fn relay(
2225 &self,
2226 leg: Leg,
2227 session_id: &str,
2228 route: &RuntimeEnvContext,
2229 space: &Arc<SpaceSpec>,
2230 leaves: Option<Vec<Bytes>>,
2231 ) -> Result<Relayed, RuntimeError> {
2232 let Some(ceiling) = self.ceiling(leg) else {
2233 return Ok((leaves, None));
2234 };
2235 let Some(leaves) = leaves else {
2236 return Ok((None, None));
2237 };
2238 let payload = PayloadFacts::new(Arc::clone(space), leaves);
2239 match self.relay_policy.reconcile(leg, ceiling, &payload) {
2240 RelayDecision::Forward => Ok((Some(payload.leaves), None)),
2241 RelayDecision::Convert {
2242 leaves,
2243 space,
2244 advisory,
2245 } => {
2246 self.raise_advisory(session_id, route, leg, advisory).await;
2247 Ok((Some(leaves), space))
2248 }
2249 RelayDecision::Refuse(reason) => Err(relay_refused(leg, "payload", reason)),
2250 }
2251 }
2252
2253 async fn raise_advisory(
2255 &self,
2256 session_id: &str,
2257 route: &RuntimeEnvContext,
2258 leg: Leg,
2259 advisory: Advisory,
2260 ) {
2261 {
2262 let mut advisories = self
2263 .advisories
2264 .lock()
2265 .unwrap_or_else(PoisonError::into_inner);
2266 if advisories.contains(&advisory) {
2267 return;
2268 }
2269 advisories.push(advisory.clone());
2270 }
2271 fan_out_event!(
2272 self,
2273 relay_advisory,
2274 RelayAdvisoryEvent {
2275 session_id: session_id.to_string(),
2276 route: route.clone(),
2277 leg,
2278 advisory,
2279 }
2280 );
2281 }
2282
2283 #[allow(clippy::too_many_arguments)]
2284 fn observation_event(
2285 &self,
2286 state: &RouteState,
2287 route: RuntimeEnvContext,
2288 snapshot: RouteSnapshot,
2289 is_reset: bool,
2290 observation: Option<Vec<Bytes>>,
2291 infos: Option<rlmesh_proto::spaces::v1::MetaMap>,
2292 width: usize,
2293 ) -> ObservationEmittedEvent {
2294 ObservationEmittedEvent {
2295 session_id: state.session_id().to_string(),
2296 route,
2297 episode_id: snapshot.episode_id,
2298 episode_record_id: snapshot.episode_record_id,
2299 episode_ids: snapshot.episode_ids,
2300 episode_record_ids: snapshot.episode_record_ids,
2301 step: snapshot.step,
2302 env_index: snapshot.env_index,
2303 is_reset,
2304 num_envs: width as u32,
2305 observation_space: Arc::clone(&self.observation_space),
2306 raw_observation: observation.clone(),
2307 observation,
2308 infos,
2309 }
2310 }
2311}
2312
2313type Relayed = (Option<Vec<Bytes>>, Option<Arc<SpaceSpec>>);
2316
2317fn relay_refused(leg: Leg, what: &str, reason: String) -> RuntimeError {
2318 RuntimeError::Protocol(format!(
2319 "the runtime cannot relay this {what} to the {}: {reason}",
2320 leg.target()
2321 ))
2322}
2323
2324fn deterministic_reset_seed(
2330 base_seed: i64,
2331 session_id: &str,
2332 reset_generation: u64,
2333 env_index: usize,
2334) -> i64 {
2335 const FNV_OFFSET: u64 = 0xcbf2_9ce4_8422_2325;
2336 const FNV_PRIME: u64 = 0x0000_0100_0000_01b3;
2337
2338 fn update(mut hash: u64, bytes: &[u8]) -> u64 {
2339 for byte in bytes {
2340 hash ^= u64::from(*byte);
2341 hash = hash.wrapping_mul(FNV_PRIME);
2342 }
2343 hash
2344 }
2345
2346 let mut hash = FNV_OFFSET;
2347 hash = update(hash, &base_seed.to_le_bytes());
2348 hash = update(hash, &[0xff]);
2349 hash = update(hash, session_id.as_bytes());
2350 hash = update(hash, &[0xfd]);
2351 hash = update(hash, &reset_generation.to_le_bytes());
2352 hash = update(hash, &[0xfc]);
2353 hash = update(hash, &(env_index as u64).to_le_bytes());
2354 (hash & i64::MAX as u64) as i64
2355}
2356
2357fn log_model_close_result(
2358 result: Result<Result<(), String>, tokio::time::error::Elapsed>,
2359 reason: &str,
2360 timeout: Duration,
2361) {
2362 match result {
2363 Ok(Err(err)) => {
2364 tracing::warn!(
2365 error = %err,
2366 reason,
2367 "model route close failed during route shutdown; relying on owner shutdown"
2368 );
2369 }
2370 Err(_) => {
2371 tracing::warn!(
2372 timeout_ms = timeout.as_millis(),
2373 reason,
2374 "model route close timed out during route shutdown; relying on owner shutdown"
2375 );
2376 }
2377 Ok(Ok(())) => {}
2378 }
2379}
2380
2381fn leaves_value(leaves: Vec<Bytes>) -> SpaceValue {
2382 SpaceValue { leaves }
2383}
2384
2385fn mint_episode_id() -> String {
2389 uuid::Uuid::now_v7().to_string()
2390}
2391
2392fn mint_episode_ids(count: usize) -> Vec<String> {
2394 (0..count).map(|_| mint_episode_id()).collect()
2395}
2396
2397fn episode_ids_with_roll(
2407 mut ids: Vec<String>,
2408 lanes: &[u32],
2409 pending_roll: &HashMap<u32, String>,
2410) -> Vec<String> {
2411 for (env_index, new_id) in pending_roll {
2412 if let Some(slot) = lanes
2413 .iter()
2414 .position(|lane| lane == env_index)
2415 .and_then(|position| ids.get_mut(position))
2416 {
2417 *slot = new_id.clone();
2418 }
2419 }
2420 ids
2421}
2422
2423fn value_leaves(payload: Option<&SpaceValue>) -> Result<Option<Vec<Bytes>>, RuntimeError> {
2427 Ok(payload.map(|payload| payload.leaves.clone()))
2428}
2429
2430fn lock_agg(telemetry: &Mutex<Aggregator>) -> MutexGuard<'_, Aggregator> {
2434 telemetry
2435 .lock()
2436 .unwrap_or_else(|poisoned| poisoned.into_inner())
2437}
2438
2439fn record_history_rows(telemetry: &Arc<Mutex<Aggregator>>, msg: &PredictRequest) {
2442 if !msg.history.is_empty() {
2443 lock_agg(telemetry).record(Sample::count(
2444 SRC_PREDICT,
2445 metrics::HISTORY_ROWS,
2446 msg.history.len() as u64,
2447 ));
2448 }
2449}
2450
2451fn record_op(
2452 telemetry: &Mutex<Aggregator>,
2453 src: Source,
2454 rpc: Duration,
2455 peer: PeerReport,
2456 request_bytes: u64,
2457 response_bytes: u64,
2458) {
2459 let mut agg = lock_agg(telemetry);
2460 agg.record(Sample::dur(src, metrics::RPC_TOTAL, rpc));
2461 if let Some(ns) = peer.endpoint_total_ns {
2462 agg.record(Sample::dur(
2463 src,
2464 metrics::ENDPOINT_TOTAL,
2465 Duration::from_nanos(ns),
2466 ));
2467 }
2468 for (metric, ns) in [
2469 (metrics::ENDPOINT_DECODE, peer.phases.decode_ns),
2470 (metrics::ENDPOINT_USER, peer.phases.user_ns),
2471 (metrics::ENDPOINT_ENCODE, peer.phases.encode_ns),
2472 (metrics::ENDPOINT_QUEUE, peer.phases.queue_ns),
2473 (metrics::PREDICT_ADAPTER, peer.phases.adapter_ns),
2474 ] {
2475 if ns != 0 {
2476 agg.record(Sample::dur(src, metric, Duration::from_nanos(ns)));
2477 }
2478 }
2479 if let Some(ns) = peer.phases.lane_skew_ns {
2482 agg.record(Sample::dur(
2483 src,
2484 metrics::LANE_SKEW,
2485 Duration::from_nanos(ns),
2486 ));
2487 }
2488 if peer.phases.in_flight != 0 {
2489 agg.record(Sample::count(
2490 src,
2491 metrics::PREDICT_IN_FLIGHT,
2492 u64::from(peer.phases.in_flight),
2493 ));
2494 }
2495 if let Some(episodes) = peer.phases.held_episodes {
2496 agg.record(Sample::count(
2497 src,
2498 metrics::HELD_EPISODES,
2499 u64::from(episodes),
2500 ));
2501 }
2502 if let Some(bytes) = peer.phases.held_state_bytes {
2503 agg.record(Sample::bytes(src, metrics::HELD_BYTES, bytes));
2504 }
2505 agg.record(Sample::bytes(src, metrics::REQUEST_BYTES, request_bytes));
2506 agg.record(Sample::bytes(src, metrics::RESPONSE_BYTES, response_bytes));
2507 if let Some(group) = peer.group_size {
2508 agg.record(Sample::count(src, metrics::GROUP_SIZE, group));
2509 }
2510}
2511
2512struct TelemetryTicker {
2522 handle: tokio::task::JoinHandle<()>,
2523}
2524
2525impl TelemetryTicker {
2526 fn spawn(
2527 telemetry: Arc<Mutex<Aggregator>>,
2528 hooks: Arc<dyn RuntimeHooks>,
2529 window: Duration,
2530 session_id: String,
2531 route: RuntimeEnvContext,
2532 ) -> Self {
2533 let period = window.max(Duration::from_millis(1));
2536 let handle = tokio::spawn(async move {
2537 let mut ticker = tokio::time::interval(period);
2538 ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
2539 ticker.tick().await; loop {
2541 ticker.tick().await;
2542 let window_snap = {
2546 let mut agg = lock_agg(&telemetry);
2547 let snap = agg.snapshot(Horizon::Window);
2548 agg.flush_window();
2549 snap
2550 };
2551 if window_snap.rows.is_empty() {
2554 continue;
2555 }
2556 let window_event = TelemetrySnapshotEvent {
2559 session_id: session_id.clone(),
2560 route: route.clone(),
2561 snapshot: window_snap,
2562 };
2563 if let Err(err) = hooks.on_telemetry(window_event).await {
2564 tracing::warn!("runtime hook on_telemetry (window) failed: {err}");
2565 }
2566 }
2567 });
2568 Self { handle }
2569 }
2570}
2571
2572impl Drop for TelemetryTicker {
2573 fn drop(&mut self) {
2574 self.handle.abort();
2575 }
2576}