1use std::future::Future;
9use std::sync::{Arc, Mutex, MutexGuard};
10use std::time::{Duration, Instant};
11
12use async_trait::async_trait;
13use prost::{Message, bytes::Bytes};
14use rlmesh_proto::core::v1::AutoresetMode;
15use rlmesh_proto::env::v1::{
16 EpisodeMetadata, ResetRequest, ResetResponse, StepRequest, StepResponse,
17};
18use rlmesh_proto::model::v1::{
19 PredictRequest, PredictResponse, ReleaseAdapterRequest, ResetAdapterRequest,
20};
21use rlmesh_proto::spaces::v1::SpaceValue;
22use tokio_util::sync::CancellationToken;
23
24use crate::hooks::{
25 ActionReceivedEvent, EpisodeCompletedEvent, EpisodeStartedEvent, LogEvent, LogLevel,
26 ObservationEmittedEvent, RuntimeEnvContext, RuntimeHooks, SessionEndedEvent,
27 SessionFailedEvent, SessionStartedEvent, StepCompletedEvent, TelemetrySnapshotEvent,
28};
29use crate::spec::{RuntimeReport, RuntimeSessionSpec};
30use crate::state::{RequestPhase, RouteSnapshot, RouteState, StartedEpisode};
31use crate::telemetry::{Aggregator, Horizon, Sample, Source, metrics};
32
33mod error;
34
35pub use error::RuntimeError;
36
37macro_rules! fan_out_event {
40 ($self:ident, $method:ident, $event:expr) => {
41 if let Err(err) = $self.hooks.$method($event).await {
42 tracing::warn!(
43 concat!("runtime hook ", stringify!($method), " failed: {}"),
44 err
45 );
46 }
47 };
48}
49
50pub struct RuntimeEnvReset {
52 pub response: ResetResponse,
53 pub endpoint_total_ns: Option<u64>,
56}
57
58pub struct RuntimeEnvStep {
60 pub response: StepResponse,
61 pub endpoint_total_ns: Option<u64>,
63}
64
65pub struct RuntimeModelPrediction {
67 pub response: PredictResponse,
68 pub endpoint_total_ns: Option<u64>,
70 pub group_size: Option<u64>,
74}
75
76#[async_trait]
78pub trait RuntimeEnv: Send {
79 async fn reset(&mut self, request: ResetRequest) -> Result<RuntimeEnvReset, RuntimeError>;
81
82 async fn step(&mut self, request: StepRequest) -> Result<RuntimeEnvStep, RuntimeError>;
84
85 async fn close(&mut self, _timeout: Duration) -> Result<(), String> {
88 Ok(())
89 }
90}
91
92#[async_trait]
95pub trait RuntimeModel: Send {
96 async fn predict(
99 &mut self,
100 request: PredictRequest,
101 ) -> Result<RuntimeModelPrediction, RuntimeError>;
102
103 async fn reset_adapter(&mut self, _request: ResetAdapterRequest) -> Result<(), RuntimeError> {
109 Ok(())
110 }
111
112 async fn release_adapter(
114 &mut self,
115 _request: ReleaseAdapterRequest,
116 _timeout: Duration,
117 ) -> Result<(), String> {
118 Ok(())
119 }
120}
121
122const DEFAULT_CANCELLATION_REASON: &str = "cancelled by caller";
125
126const DEFAULT_MAX_EPISODE_STEPS: i64 = 100_000;
133
134fn success_from_final_info(final_info: Option<&rlmesh_proto::spaces::v1::MetaMap>) -> Option<bool> {
139 use rlmesh_proto::spaces::v1::meta_value::Kind;
140 let entries = &final_info?.entries;
141 ["is_success", "success"]
142 .iter()
143 .find_map(|key| match entries.get(*key)?.kind.as_ref()? {
144 Kind::Bool(value) => Some(*value),
145 Kind::Integer(value) => Some(*value != 0),
146 Kind::Number(value) => Some(*value != 0.0),
147 _ => None,
148 })
149}
150
151const SRC_PREDICT: Source = Source {
155 op: "model.predict",
156 component: "model",
157};
158const SRC_STEP: Source = Source {
159 op: "env.step",
160 component: "env",
161};
162const SRC_RESET: Source = Source {
163 op: "env.reset",
164 component: "env",
165};
166
167#[must_use = "a RuntimeDriver does nothing until one of its run methods is awaited"]
170pub struct RuntimeDriver<E, M> {
171 spec: RuntimeSessionSpec,
172 env: E,
173 model: M,
174 hooks: Arc<dyn RuntimeHooks>,
175 cancellation_reason: String,
176 action_space: Arc<rlmesh_proto::spaces::v1::SpaceSpec>,
180 observation_space: Arc<rlmesh_proto::spaces::v1::SpaceSpec>,
181}
182
183impl<E, M> RuntimeDriver<E, M>
184where
185 E: RuntimeEnv,
186 M: RuntimeModel,
187{
188 pub fn new(spec: RuntimeSessionSpec, env: E, model: M, hooks: Arc<dyn RuntimeHooks>) -> Self {
189 Self {
190 spec,
191 env,
192 model,
193 hooks,
194 cancellation_reason: DEFAULT_CANCELLATION_REASON.to_string(),
195 action_space: Arc::default(),
197 observation_space: Arc::default(),
198 }
199 }
200
201 fn planned_reset_seeds(
206 &self,
207 state: &mut RouteState,
208 reset_generation: u64,
209 env_indices: Option<&[u32]>,
210 ) -> Vec<i64> {
211 if !self.spec.episode_seeds.is_empty() {
212 let lanes = env_indices.map_or(self.spec.num_envs, <[u32]>::len);
213 return state.claim_episode_seeds(&self.spec.episode_seeds, lanes);
214 }
215 match env_indices {
216 None => self.reset_seeds(reset_generation),
217 Some(indices) => self.reset_subset_seeds(reset_generation, indices),
218 }
219 }
220
221 fn reset_seeds(&self, reset_generation: u64) -> Vec<i64> {
222 self.seeds_for(reset_generation, 0..self.spec.num_envs)
223 }
224
225 fn reset_subset_seeds(&self, reset_generation: u64, env_indices: &[u32]) -> Vec<i64> {
228 self.seeds_for(
229 reset_generation,
230 env_indices.iter().map(|&index| index as usize),
231 )
232 }
233
234 fn seeds_for(
237 &self,
238 reset_generation: u64,
239 env_indices: impl Iterator<Item = usize>,
240 ) -> Vec<i64> {
241 let Some(base_seed) = self.spec.base_seed else {
242 return Vec::new();
243 };
244 env_indices
245 .map(|env_index| {
246 deterministic_reset_seed(
247 base_seed,
248 &self.spec.session_id,
249 reset_generation,
250 env_index,
251 )
252 })
253 .collect()
254 }
255
256 fn autoreset_mode(&self) -> AutoresetMode {
259 AutoresetMode::try_from(self.spec.env_contract.autoreset_mode)
263 .unwrap_or(AutoresetMode::Disabled)
264 }
265
266 pub async fn run(self) -> Result<RuntimeReport, RuntimeError> {
267 self.run_with_cancellation(CancellationToken::new()).await
268 }
269
270 pub async fn run_with_cancellation(
271 self,
272 cancellation: CancellationToken,
273 ) -> Result<RuntimeReport, RuntimeError> {
274 self.run_with_cancellation_reason(cancellation, DEFAULT_CANCELLATION_REASON)
275 .await
276 }
277
278 pub async fn run_with_cancellation_reason(
286 mut self,
287 cancellation: CancellationToken,
288 reason: impl Into<String>,
289 ) -> Result<RuntimeReport, RuntimeError> {
290 self.cancellation_reason = reason.into();
291 self.spec.validate().map_err(RuntimeError::InvalidSpec)?;
292 self.action_space = Arc::new(self.spec.action_space_validated().clone());
295 self.observation_space = Arc::new(self.spec.observation_space_validated().clone());
296 let mut state = RouteState::new(&self.spec);
297 let telemetry = Arc::new(Mutex::new(Aggregator::default()));
303 let ticker = (!self.spec.limits.telemetry_window.is_zero()).then(|| {
306 TelemetryTicker::spawn(
307 Arc::clone(&telemetry),
308 Arc::clone(&self.hooks),
309 self.spec.limits.telemetry_window,
310 state.session_id().to_string(),
311 state.env_context(),
312 )
313 });
314 let result = self.run_loop(&mut state, &cancellation, &telemetry).await;
315 drop(ticker);
319 let final_snapshot = lock_agg(&telemetry).snapshot(Horizon::Session);
320 fan_out_event!(
321 self,
322 on_telemetry,
323 TelemetrySnapshotEvent {
324 session_id: state.session_id().to_string(),
325 route: state.env_context(),
326 snapshot: final_snapshot,
327 }
328 );
329 if let Err(error) = &result {
330 self.shutdown_after_failure(&mut state, error).await;
331 }
332 result
333 }
334
335 #[tracing::instrument(
340 name = "rlmesh.route",
341 level = "info",
342 skip_all,
343 fields(
344 session_id = %state.session_id(),
345 env_id = %self.spec.env_id,
346 num_envs = self.spec.num_envs,
347 ),
348 )]
349 async fn run_loop(
350 &mut self,
351 state: &mut RouteState,
352 cancellation: &CancellationToken,
353 telemetry: &Arc<Mutex<Aggregator>>,
354 ) -> Result<RuntimeReport, RuntimeError> {
355 fan_out_event!(
356 self,
357 session_started,
358 SessionStartedEvent {
359 session_id: state.session_id().to_string(),
360 route: state.env_context(),
361 env_id: self.spec.env_id.clone(),
362 }
363 );
364
365 let mut reset_generation = 0_u64;
366 let reset_timeout = self.spec.limits.env_reset_timeout;
367 let reset_timeout_ms = self.spec.limits.env_reset_timeout_ms().max(0) as u64;
369 let reset_seeds = self.planned_reset_seeds(state, reset_generation, None);
370 let initial_episode_ids = mint_episode_ids(self.spec.num_envs);
374 state.note_episode_seeds(&initial_episode_ids, &reset_seeds);
375 let reset_request = ResetRequest {
376 seeds: reset_seeds,
377 options: None,
378 timeout_ms: reset_timeout_ms,
379 env_indices: Vec::new(),
380 episode_ids: initial_episode_ids.clone(),
381 };
382 let reset_request_bytes = reset_request.encoded_len() as u64;
383 let reset_started = Instant::now();
386 let reset_ok = await_runtime_operation(
387 cancellation,
388 reset_timeout,
389 RuntimeError::operation_timeout(
390 state.env_id(),
391 state.env_component_id(),
392 "env.reset",
393 0,
394 reset_timeout,
395 ),
396 self.cancelled_error(state, 0),
397 self.env.reset(reset_request),
398 )
399 .await?;
400 let reset_latency = reset_started.elapsed();
401 record_op(
402 telemetry,
403 SRC_RESET,
404 reset_latency,
405 reset_ok.endpoint_total_ns,
406 reset_request_bytes,
407 reset_ok.response.encoded_len() as u64,
408 None,
409 );
410 fan_out_event!(
411 self,
412 log,
413 LogEvent {
414 session_id: state.session_id().to_string(),
415 route: state.env_context(),
416 level: LogLevel::Info,
417 message: format!(
418 "env reset complete in {:.0}ms ({} episode(s) ready)",
419 reset_latency.as_secs_f64() * 1000.0,
420 initial_episode_ids.len()
421 ),
422 source: Some("runtime".to_string()),
423 }
424 );
425
426 let reset_observation = value_leaves(reset_ok.response.observation.as_ref())?;
427 let started_episodes = state.start_episodes(initial_episode_ids, false);
428 self.invoke_started_episodes(state, started_episodes).await;
429
430 let mut pending_roll: std::collections::HashMap<u32, String> =
436 std::collections::HashMap::new();
437
438 let mut reset_msg =
439 state.predict_request(reset_observation.clone(), RequestPhase::ResetObservation);
440 let mut reset_event =
441 self.observation_event(state, state.snapshot(), true, reset_observation.clone());
442 let transformed_reset_observation = self
443 .invoke_transform_observation(reset_event.clone())
444 .await?;
445 reset_event.observation = transformed_reset_observation.clone();
446 reset_msg.observation = transformed_reset_observation.map(leaves_value);
447 fan_out_event!(self, observation_emitted, reset_event);
448
449 let mut pending_observation_msg = reset_msg;
450
451 let mut replay_buffer: std::collections::VecDeque<Vec<Bytes>> =
459 std::collections::VecDeque::new();
460
461 loop {
462 if cancellation.is_cancelled() {
463 return Err(self.cancelled_error(state, state.snapshot().step));
464 }
465
466 let predict_snapshot = state.snapshot();
467 if replay_buffer.is_empty() {
474 let predict_timeout = self.spec.limits.model_predict_timeout;
475 let expected_context = pending_observation_msg.context.clone();
476 let predict_request_bytes = pending_observation_msg.encoded_len() as u64;
477 let predict_started = Instant::now();
478 let action_msg = await_runtime_operation(
479 cancellation,
480 predict_timeout,
481 RuntimeError::operation_timeout(
482 state.env_id(),
483 state.model_component_id(),
484 "model.predict",
485 predict_snapshot.step,
486 predict_timeout,
487 ),
488 self.cancelled_error(state, predict_snapshot.step),
489 self.model.predict(pending_observation_msg),
490 )
491 .await?;
492 let predict_rpc = predict_started.elapsed();
493 if action_msg.response.context != expected_context {
494 let request_id = expected_context
495 .as_ref()
496 .map(|context| context.request_id.clone())
497 .unwrap_or_default();
498 return Err(RuntimeError::ModelRouteMismatch {
499 component_id: state.model_component_id().to_string(),
500 request_id,
501 });
502 }
503 record_op(
504 telemetry,
505 SRC_PREDICT,
506 predict_rpc,
507 action_msg.endpoint_total_ns,
508 predict_request_bytes,
509 action_msg.response.encoded_len() as u64,
510 action_msg.group_size,
511 );
512 if action_msg.response.actions.is_empty() {
513 return Err(RuntimeError::Protocol(format!(
514 "model endpoint {} returned a predict response with no actions",
515 state.model_component_id()
516 )));
517 }
518 for frame in &action_msg.response.actions {
522 if let Some(leaves) = value_leaves(Some(frame))? {
523 replay_buffer.push_back(leaves);
524 }
525 }
526 }
527 let model_action = replay_buffer
528 .pop_front()
529 .expect("replay buffer is non-empty after a refill");
530
531 let action_step = predict_snapshot.step + 1;
532 let mut action_event = ActionReceivedEvent {
533 session_id: state.session_id().to_string(),
534 route: state.env_context(),
535 episode_id: predict_snapshot.episode_id.clone(),
536 episode_record_id: predict_snapshot.episode_record_id.clone(),
537 episode_ids: predict_snapshot.episode_ids.clone(),
538 episode_record_ids: predict_snapshot.episode_record_ids.clone(),
539 step: action_step,
540 env_index: predict_snapshot.env_index,
541 action_space: Arc::clone(&self.action_space),
542 action: Some(model_action),
543 };
544 action_event.action = self.invoke_transform_action(action_event.clone()).await?;
545 fan_out_event!(self, action_received, action_event.clone());
546
547 let step_timeout = self.spec.limits.env_step_timeout;
548 let step_timeout_ms = self.spec.limits.env_step_timeout_ms().max(0) as u64;
550 let step_episode_ids = episode_ids_with_roll(state.episode_ids(), &pending_roll);
555 let step_request = StepRequest {
556 action: action_event.action.map(leaves_value),
557 timeout_ms: step_timeout_ms,
558 env_indices: Vec::new(),
559 episode_ids: step_episode_ids,
560 };
561 let step_request_bytes = step_request.encoded_len() as u64;
562 let step_started = Instant::now();
563 let step_ok = await_runtime_operation(
564 cancellation,
565 step_timeout,
566 RuntimeError::operation_timeout(
567 state.env_id(),
568 state.env_component_id(),
569 "env.step",
570 action_step,
571 step_timeout,
572 ),
573 self.cancelled_error(state, action_step),
574 self.env.step(step_request),
575 )
576 .await?;
577 let step_rpc = step_started.elapsed();
578 record_op(
579 telemetry,
580 SRC_STEP,
581 step_rpc,
582 step_ok.endpoint_total_ns,
583 step_request_bytes,
584 step_ok.response.encoded_len() as u64,
585 None,
586 );
587 let step_observation = value_leaves(step_ok.response.observation.as_ref())?;
588
589 state.record_step(&step_ok.response.rewards);
590 let step_snapshot = state.snapshot();
591 fan_out_event!(
592 self,
593 step_completed,
594 StepCompletedEvent {
595 session_id: state.session_id().to_string(),
596 route: state.env_context(),
597 episode_id: step_snapshot.episode_id.clone(),
598 episode_record_id: step_snapshot.episode_record_id.clone(),
599 step: step_snapshot.step,
600 env_index: step_snapshot.env_index,
601 rewards: step_ok.response.rewards.clone(),
602 }
603 );
604
605 if !pending_roll.is_empty() {
611 let roll_ids = episode_ids_with_roll(state.episode_ids(), &pending_roll);
612 pending_roll.clear();
613 let started_episodes = state.observe_episode_ids(roll_ids);
614 self.invoke_started_episodes(state, started_episodes).await;
615 }
616 let capped = self.capped_lane_completions(state, &step_ok.response.completed_episodes);
623 let completed_episodes: std::borrow::Cow<'_, [EpisodeMetadata]> = if capped.is_empty() {
624 std::borrow::Cow::Borrowed(&step_ok.response.completed_episodes)
625 } else {
626 let mut all = step_ok.response.completed_episodes.clone();
627 all.extend(capped);
628 std::borrow::Cow::Owned(all)
629 };
630
631 self.emit_completed_episodes(state, &completed_episodes)
632 .await;
633 self.emit_reset_adapter(state, &completed_episodes).await;
636
637 if !completed_episodes.is_empty() {
642 replay_buffer.clear();
643 }
644
645 if matches!(
649 self.autoreset_mode(),
650 AutoresetMode::NextStep | AutoresetMode::SameStep
651 ) {
652 for completed in completed_episodes.iter() {
653 pending_roll
654 .entry(completed.env_index)
655 .or_insert_with(mint_episode_id);
656 }
657 }
658
659 if self
666 .spec
667 .max_episodes
668 .is_some_and(|limit| state.total_episodes() >= limit as i64)
669 {
670 let release_request = state.release_adapter_request("completed requested episodes");
671 self.shutdown_terminal_route(
672 state,
673 "completed requested episodes",
674 release_request,
675 )
676 .await;
677 let telemetry_snapshot = lock_agg(telemetry).snapshot(Horizon::Session);
681 fan_out_event!(
682 self,
683 session_ended,
684 SessionEndedEvent {
685 session_id: state.session_id().to_string(),
686 route: state.env_context(),
687 reason: "completed requested episodes".to_string(),
688 total_steps: state.total_steps(),
689 total_episodes: state.total_episodes(),
690 }
691 );
692 return Ok(RuntimeReport {
693 session_id: state.session_id().to_string(),
694 env_id: self.spec.env_id.clone(),
695 total_steps: state.total_steps(),
696 total_episodes: state.total_episodes(),
697 episodes: state.take_episode_summaries(),
698 telemetry: telemetry_snapshot,
699 });
700 }
701
702 let (next_obs, phase, is_reset_msg) = match self.autoreset_mode() {
706 AutoresetMode::NextStep | AutoresetMode::SameStep => (
711 step_observation.clone(),
712 RequestPhase::StepObservation,
713 false,
714 ),
715 AutoresetMode::Unspecified | AutoresetMode::Disabled => {
721 let mut done_lanes: Vec<u32> = completed_episodes
724 .iter()
725 .map(|metadata| metadata.env_index)
726 .collect();
727 done_lanes.sort_unstable();
735 done_lanes.dedup();
736 if done_lanes.is_empty() {
737 (
738 step_observation.clone(),
739 RequestPhase::StepObservation,
740 false,
741 )
742 } else {
743 reset_generation += 1;
744 let step = state.snapshot().step;
745 let reset_timeout = self.spec.limits.env_reset_timeout;
746 let reset_timeout_ms =
748 self.spec.limits.env_reset_timeout_ms().max(0) as u64;
749 let whole_vector = done_lanes.len() == self.spec.num_envs;
750 let (reset_seeds, env_indices, reset_episode_ids) = if whole_vector {
754 (
755 self.planned_reset_seeds(state, reset_generation, None),
756 Vec::new(),
757 mint_episode_ids(self.spec.num_envs),
758 )
759 } else {
760 (
761 self.planned_reset_seeds(
762 state,
763 reset_generation,
764 Some(&done_lanes),
765 ),
766 done_lanes.clone(),
767 mint_episode_ids(done_lanes.len()),
768 )
769 };
770 state.note_episode_seeds(&reset_episode_ids, &reset_seeds);
771 let reset_request = ResetRequest {
772 seeds: reset_seeds,
773 options: None,
774 timeout_ms: reset_timeout_ms,
775 env_indices,
776 episode_ids: reset_episode_ids.clone(),
777 };
778 let reset_request_bytes = reset_request.encoded_len() as u64;
779 let inloop_reset_started = Instant::now();
780 let reset_ok = await_runtime_operation(
781 cancellation,
782 reset_timeout,
783 RuntimeError::operation_timeout(
784 state.env_id(),
785 state.env_component_id(),
786 "env.reset",
787 step,
788 reset_timeout,
789 ),
790 self.cancelled_error(state, step),
791 self.env.reset(reset_request),
792 )
793 .await?;
794 record_op(
795 telemetry,
796 SRC_RESET,
797 inloop_reset_started.elapsed(),
798 reset_ok.endpoint_total_ns,
799 reset_request_bytes,
800 reset_ok.response.encoded_len() as u64,
801 None,
802 );
803 let next_obs = value_leaves(reset_ok.response.observation.as_ref())?;
804 let started_episodes = if whole_vector {
808 state.start_episodes(reset_episode_ids, true)
809 } else {
810 let mut full = state.episode_ids();
811 for (lane, id) in done_lanes.iter().zip(reset_episode_ids) {
812 if let Some(slot) = full.get_mut(*lane as usize) {
813 *slot = id;
814 }
815 }
816 state.observe_episode_ids(full)
817 };
818 self.invoke_started_episodes(state, started_episodes).await;
819 (next_obs, RequestPhase::ResetObservation, true)
820 }
821 }
822 };
823
824 let mut obs_msg = state.predict_request(next_obs.clone(), phase);
825 let mut outgoing_observation_event =
826 self.observation_event(state, state.snapshot(), is_reset_msg, next_obs);
827 let transformed_observation = self
828 .invoke_transform_observation(outgoing_observation_event.clone())
829 .await?;
830 outgoing_observation_event.observation = transformed_observation.clone();
831 obs_msg.observation = transformed_observation.map(leaves_value);
832 fan_out_event!(self, observation_emitted, outgoing_observation_event);
836
837 pending_observation_msg = obs_msg;
838 }
839 }
840
841 async fn shutdown_after_failure(&mut self, state: &mut RouteState, error: &RuntimeError) {
842 let reason = error.to_string();
843 let request = state.release_adapter_request(reason.clone());
844 self.shutdown_terminal_route(state, &reason, request).await;
845
846 if let Err(err) = self
847 .hooks
848 .session_failed(SessionFailedEvent {
849 session_id: state.session_id().to_string(),
850 route: state.env_context(),
851 reason,
852 })
853 .await
854 {
855 tracing::warn!("runtime hook session_failed failed: {err}");
856 }
857 }
858
859 async fn shutdown_terminal_route(
860 &mut self,
861 state: &RouteState,
862 reason: &str,
863 request: ReleaseAdapterRequest,
864 ) {
865 let timeout = self.spec.limits.service_close_timeout;
866 let model_close =
871 tokio::time::timeout(timeout, self.model.release_adapter(request, timeout));
872 if self.spec.close_env_on_end {
873 let env_close = tokio::time::timeout(timeout, self.env.close(timeout));
874 let (env_result, model_result) = tokio::join!(env_close, model_close);
875 match env_result {
876 Ok(Err(err)) => {
877 tracing::warn!(error = %err, "environment close failed during route shutdown");
878 }
879 Err(_) => {
880 tracing::warn!(
881 timeout_ms = timeout.as_millis(),
882 "environment close timed out during route shutdown; abandoning close"
883 );
884 }
885 Ok(Ok(())) => {}
886 }
887 log_model_close_result(model_result, reason, timeout);
888 return;
889 }
890
891 tracing::debug!(
892 env_id = %state.env_id(),
893 reason,
894 "skipping environment close for adapter; endpoint remains owned by the run"
895 );
896 log_model_close_result(model_close.await, reason, timeout);
897 }
898
899 fn cancelled_error(&self, state: &RouteState, step: i64) -> RuntimeError {
900 RuntimeError::route_cancelled(state.env_id(), step, self.cancellation_reason.as_str())
901 }
902
903 async fn invoke_started_episodes(&self, state: &RouteState, episodes: Vec<StartedEpisode>) {
904 for episode in episodes {
905 let record = &episode.record;
906 fan_out_event!(
907 self,
908 episode_started,
909 EpisodeStartedEvent {
910 session_id: state.session_id().to_string(),
911 route: state.env_context(),
912 episode_id: episode.episode_id.clone(),
913 episode_record_id: record.record_id.clone(),
914 episode_index: record.index,
915 env_index: record.env_index,
916 started_from_auto_reset: record.started_from_auto_reset,
917 }
918 );
919 }
920 }
921
922 fn capped_lane_completions(
929 &self,
930 state: &RouteState,
931 env_completed: &[EpisodeMetadata],
932 ) -> Vec<EpisodeMetadata> {
933 let driver_owns_resets = matches!(
934 self.autoreset_mode(),
935 AutoresetMode::Disabled | AutoresetMode::Unspecified
936 );
937 let step_cap = self
938 .spec
939 .max_episode_steps
940 .or_else(|| driver_owns_resets.then_some(DEFAULT_MAX_EPISODE_STEPS));
941 let time_cap = self.spec.max_episode_seconds;
942 if step_cap.is_none() && time_cap.is_none() {
943 return Vec::new();
944 }
945 let env_done: Vec<u32> = env_completed
946 .iter()
947 .map(|metadata| metadata.env_index)
948 .collect();
949 let now_ns = crate::state::now_unix_ns();
950 state
951 .slots()
952 .iter()
953 .filter_map(|slot| {
954 let episode = slot.episode.as_ref()?;
955 let env_index = u32::try_from(slot.env_index).ok()?;
956 if env_done.contains(&env_index) {
957 return None;
958 }
959 let steps_capped = step_cap.is_some_and(|cap| slot.step >= cap);
960 let elapsed_seconds = (now_ns - slot.started_at_ns).max(0) as f64 / 1e9;
961 let time_capped = time_cap.is_some_and(|cap| elapsed_seconds >= cap);
962 (steps_capped || time_capped).then(|| EpisodeMetadata {
963 episode_id: episode.episode_id.clone(),
964 seed: None,
965 env_index,
966 step_count: slot.step,
967 cumulative_reward: slot.cumulative_reward,
968 terminated: false,
969 truncated: true,
970 start_timestamp_ns: slot.started_at_ns,
971 end_timestamp_ns: now_ns,
972 final_info: None,
973 })
974 })
975 .collect()
976 }
977
978 async fn emit_completed_episodes(&self, state: &mut RouteState, episodes: &[EpisodeMetadata]) {
984 for completed in episodes {
985 let record = state.complete_episode(&completed.episode_id);
986 let episode_record_id = record
987 .as_ref()
988 .map(|record| record.record_id.clone())
989 .unwrap_or_default();
990 let env_index = i32::try_from(completed.env_index).unwrap_or(i32::MAX);
992 let seed = state.take_episode_seed(&completed.episode_id);
993 if self.spec.max_episodes.is_some() {
994 state.record_episode_summary(crate::spec::EpisodeSummary {
995 episode_index: record.as_ref().map_or(0, |record| record.index),
996 env_index,
997 seed,
998 step_count: completed.step_count,
999 cumulative_reward: completed.cumulative_reward,
1000 terminated: completed.terminated,
1001 truncated: completed.truncated,
1002 duration_ms: (completed.end_timestamp_ns - completed.start_timestamp_ns).max(0)
1003 / 1_000_000,
1004 success: success_from_final_info(completed.final_info.as_ref()),
1005 });
1006 }
1007 fan_out_event!(
1008 self,
1009 episode_completed,
1010 EpisodeCompletedEvent {
1011 session_id: state.session_id().to_string(),
1012 route: state.env_context(),
1013 episode_id: completed.episode_id.clone(),
1014 episode_record_id,
1015 episode_index: record.as_ref().map_or(0, |record| record.index),
1016 env_index,
1017 step_count: completed.step_count,
1018 cumulative_reward: completed.cumulative_reward,
1019 terminated: completed.terminated,
1020 truncated: completed.truncated,
1021 duration_ms: (completed.end_timestamp_ns - completed.start_timestamp_ns).max(0)
1022 / 1_000_000,
1023 final_info: completed.final_info.clone(),
1024 }
1025 );
1026 }
1027 }
1028
1029 async fn emit_reset_adapter(&mut self, state: &mut RouteState, episodes: &[EpisodeMetadata]) {
1040 let slot_ids = state.episode_ids();
1041 let episode_ids: Vec<String> = episodes
1042 .iter()
1043 .filter_map(|completed| slot_ids.get(completed.env_index as usize).cloned())
1044 .filter(|id| !id.is_empty())
1045 .collect();
1046 if episode_ids.is_empty() {
1047 return;
1048 }
1049 let request = state.reset_adapter_request(episode_ids);
1050 if let Err(err) = self.model.reset_adapter(request).await {
1051 tracing::warn!("model reset_adapter (evict) failed: {err}");
1052 }
1053 }
1054
1055 async fn invoke_transform_action(
1056 &self,
1057 event: ActionReceivedEvent,
1058 ) -> Result<Option<Vec<Bytes>>, RuntimeError> {
1059 match self.hooks.transform_action(event).await {
1060 Ok(action) => Ok(action),
1061 Err(err) => {
1062 tracing::warn!("runtime hook transform_action failed: {err}");
1063 Err(RuntimeError::Hook(err))
1064 }
1065 }
1066 }
1067
1068 async fn invoke_transform_observation(
1069 &self,
1070 event: ObservationEmittedEvent,
1071 ) -> Result<Option<Vec<Bytes>>, RuntimeError> {
1072 match self.hooks.transform_observation(event).await {
1073 Ok(observation) => Ok(observation),
1074 Err(err) => {
1075 tracing::warn!("runtime hook transform_observation failed: {err}");
1076 Err(RuntimeError::Hook(err))
1077 }
1078 }
1079 }
1080
1081 fn observation_event(
1082 &self,
1083 state: &RouteState,
1084 snapshot: RouteSnapshot,
1085 is_reset: bool,
1086 observation: Option<Vec<Bytes>>,
1087 ) -> ObservationEmittedEvent {
1088 ObservationEmittedEvent {
1089 session_id: state.session_id().to_string(),
1090 route: state.env_context(),
1091 episode_id: snapshot.episode_id,
1092 episode_record_id: snapshot.episode_record_id,
1093 episode_ids: snapshot.episode_ids,
1094 episode_record_ids: snapshot.episode_record_ids,
1095 step: snapshot.step,
1096 env_index: snapshot.env_index,
1097 is_reset,
1098 num_envs: self.spec.num_envs as u32,
1099 observation_space: Arc::clone(&self.observation_space),
1100 observation,
1101 }
1102 }
1103}
1104
1105fn deterministic_reset_seed(
1110 base_seed: i64,
1111 session_id: &str,
1112 reset_generation: u64,
1113 env_index: usize,
1114) -> i64 {
1115 const FNV_OFFSET: u64 = 0xcbf2_9ce4_8422_2325;
1116 const FNV_PRIME: u64 = 0x0000_0100_0000_01b3;
1117
1118 fn update(mut hash: u64, bytes: &[u8]) -> u64 {
1119 for byte in bytes {
1120 hash ^= u64::from(*byte);
1121 hash = hash.wrapping_mul(FNV_PRIME);
1122 }
1123 hash
1124 }
1125
1126 let mut hash = FNV_OFFSET;
1127 hash = update(hash, &base_seed.to_le_bytes());
1128 hash = update(hash, &[0xff]);
1129 hash = update(hash, session_id.as_bytes());
1130 hash = update(hash, &[0xfd]);
1131 hash = update(hash, &reset_generation.to_le_bytes());
1132 hash = update(hash, &[0xfc]);
1133 hash = update(hash, &(env_index as u64).to_le_bytes());
1134 (hash & i64::MAX as u64) as i64
1135}
1136
1137async fn await_runtime_operation<T, F>(
1138 cancellation: &CancellationToken,
1139 timeout: Duration,
1140 timeout_error: RuntimeError,
1141 cancelled_error: RuntimeError,
1142 operation: F,
1143) -> Result<T, RuntimeError>
1144where
1145 F: Future<Output = Result<T, RuntimeError>>,
1146{
1147 tokio::select! {
1148 _ = cancellation.cancelled() => Err(cancelled_error),
1149 result = tokio::time::timeout(timeout, operation) => match result {
1150 Ok(result) => result,
1151 Err(_) => Err(timeout_error),
1152 },
1153 }
1154}
1155
1156fn log_model_close_result(
1157 result: Result<Result<(), String>, tokio::time::error::Elapsed>,
1158 reason: &str,
1159 timeout: Duration,
1160) {
1161 match result {
1162 Ok(Err(err)) => {
1163 tracing::warn!(
1164 error = %err,
1165 reason,
1166 "model route close failed during route shutdown; relying on owner shutdown"
1167 );
1168 }
1169 Err(_) => {
1170 tracing::warn!(
1171 timeout_ms = timeout.as_millis(),
1172 reason,
1173 "model route close timed out during route shutdown; relying on owner shutdown"
1174 );
1175 }
1176 Ok(Ok(())) => {}
1177 }
1178}
1179
1180fn leaves_value(leaves: Vec<Bytes>) -> SpaceValue {
1181 SpaceValue { leaves }
1182}
1183
1184fn mint_episode_id() -> String {
1188 uuid::Uuid::now_v7().to_string()
1189}
1190
1191fn mint_episode_ids(count: usize) -> Vec<String> {
1193 (0..count).map(|_| mint_episode_id()).collect()
1194}
1195
1196fn episode_ids_with_roll(
1201 mut ids: Vec<String>,
1202 pending_roll: &std::collections::HashMap<u32, String>,
1203) -> Vec<String> {
1204 for (env_index, new_id) in pending_roll {
1207 if let Some(slot) = ids.get_mut(*env_index as usize) {
1208 *slot = new_id.clone();
1209 }
1210 }
1211 ids
1212}
1213
1214fn value_leaves(payload: Option<&SpaceValue>) -> Result<Option<Vec<Bytes>>, RuntimeError> {
1218 Ok(payload.map(|payload| payload.leaves.clone()))
1219}
1220
1221fn lock_agg(telemetry: &Mutex<Aggregator>) -> MutexGuard<'_, Aggregator> {
1225 telemetry
1226 .lock()
1227 .unwrap_or_else(|poisoned| poisoned.into_inner())
1228}
1229
1230fn record_op(
1236 telemetry: &Mutex<Aggregator>,
1237 src: Source,
1238 rpc: Duration,
1239 endpoint_total_ns: Option<u64>,
1240 request_bytes: u64,
1241 response_bytes: u64,
1242 group_size: Option<u64>,
1243) {
1244 let mut agg = lock_agg(telemetry);
1245 agg.record(Sample::dur(src, metrics::RPC_TOTAL, rpc));
1246 if let Some(ns) = endpoint_total_ns {
1247 agg.record(Sample::dur(
1248 src,
1249 metrics::ENDPOINT_TOTAL,
1250 Duration::from_nanos(ns),
1251 ));
1252 }
1253 agg.record(Sample::bytes(src, metrics::REQUEST_BYTES, request_bytes));
1254 agg.record(Sample::bytes(src, metrics::RESPONSE_BYTES, response_bytes));
1255 if let Some(group) = group_size {
1256 agg.record(Sample::count(src, metrics::GROUP_SIZE, group));
1257 }
1258}
1259
1260struct TelemetryTicker {
1270 handle: tokio::task::JoinHandle<()>,
1271}
1272
1273impl TelemetryTicker {
1274 fn spawn(
1275 telemetry: Arc<Mutex<Aggregator>>,
1276 hooks: Arc<dyn RuntimeHooks>,
1277 window: Duration,
1278 session_id: String,
1279 route: RuntimeEnvContext,
1280 ) -> Self {
1281 let period = window.max(Duration::from_millis(1));
1284 let handle = tokio::spawn(async move {
1285 let mut ticker = tokio::time::interval(period);
1286 ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
1287 ticker.tick().await; loop {
1289 ticker.tick().await;
1290 let window_snap = {
1294 let mut agg = lock_agg(&telemetry);
1295 let snap = agg.snapshot(Horizon::Window);
1296 agg.flush_window();
1297 snap
1298 };
1299 if window_snap.rows.is_empty() {
1302 continue;
1303 }
1304 let window_event = TelemetrySnapshotEvent {
1307 session_id: session_id.clone(),
1308 route: route.clone(),
1309 snapshot: window_snap,
1310 };
1311 if let Err(err) = hooks.on_telemetry(window_event).await {
1312 tracing::warn!("runtime hook on_telemetry (window) failed: {err}");
1313 }
1314 }
1315 });
1316 Self { handle }
1317 }
1318}
1319
1320impl Drop for TelemetryTicker {
1321 fn drop(&mut self) {
1322 self.handle.abort();
1323 }
1324}