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}
71
72#[async_trait]
74pub trait RuntimeEnv: Send {
75 async fn reset(&mut self, request: ResetRequest) -> Result<RuntimeEnvReset, RuntimeError>;
77
78 async fn step(&mut self, request: StepRequest) -> Result<RuntimeEnvStep, RuntimeError>;
80
81 async fn close(&mut self, _timeout: Duration) -> Result<(), String> {
84 Ok(())
85 }
86}
87
88#[async_trait]
91pub trait RuntimeModel: Send {
92 async fn predict(
95 &mut self,
96 request: PredictRequest,
97 ) -> Result<RuntimeModelPrediction, RuntimeError>;
98
99 async fn reset_adapter(&mut self, _request: ResetAdapterRequest) -> Result<(), RuntimeError> {
105 Ok(())
106 }
107
108 async fn release_adapter(
110 &mut self,
111 _request: ReleaseAdapterRequest,
112 _timeout: Duration,
113 ) -> Result<(), String> {
114 Ok(())
115 }
116}
117
118const DEFAULT_CANCELLATION_REASON: &str = "cancelled by caller";
121
122const SRC_PREDICT: Source = Source {
126 op: "model.predict",
127 component: "model",
128};
129const SRC_STEP: Source = Source {
130 op: "env.step",
131 component: "env",
132};
133const SRC_RESET: Source = Source {
134 op: "env.reset",
135 component: "env",
136};
137
138#[must_use = "a RuntimeDriver does nothing until one of its run methods is awaited"]
141pub struct RuntimeDriver<E, M> {
142 spec: RuntimeSessionSpec,
143 env: E,
144 model: M,
145 hooks: Arc<dyn RuntimeHooks>,
146 cancellation_reason: String,
147 action_space: Arc<rlmesh_proto::spaces::v1::SpaceSpec>,
151 observation_space: Arc<rlmesh_proto::spaces::v1::SpaceSpec>,
152}
153
154impl<E, M> RuntimeDriver<E, M>
155where
156 E: RuntimeEnv,
157 M: RuntimeModel,
158{
159 pub fn new(spec: RuntimeSessionSpec, env: E, model: M, hooks: Arc<dyn RuntimeHooks>) -> Self {
160 Self {
161 spec,
162 env,
163 model,
164 hooks,
165 cancellation_reason: DEFAULT_CANCELLATION_REASON.to_string(),
166 action_space: Arc::default(),
168 observation_space: Arc::default(),
169 }
170 }
171
172 fn reset_seeds(&self, reset_generation: u64) -> Vec<i64> {
173 self.seeds_for(reset_generation, 0..self.spec.num_envs)
174 }
175
176 fn reset_subset_seeds(&self, reset_generation: u64, env_indices: &[u32]) -> Vec<i64> {
179 self.seeds_for(
180 reset_generation,
181 env_indices.iter().map(|&index| index as usize),
182 )
183 }
184
185 fn seeds_for(
188 &self,
189 reset_generation: u64,
190 env_indices: impl Iterator<Item = usize>,
191 ) -> Vec<i64> {
192 let Some(base_seed) = self.spec.base_seed else {
193 return Vec::new();
194 };
195 env_indices
196 .map(|env_index| {
197 deterministic_reset_seed(
198 base_seed,
199 &self.spec.session_id,
200 reset_generation,
201 env_index,
202 )
203 })
204 .collect()
205 }
206
207 fn autoreset_mode(&self) -> AutoresetMode {
210 AutoresetMode::try_from(self.spec.env_contract.autoreset_mode)
214 .unwrap_or(AutoresetMode::Disabled)
215 }
216
217 pub async fn run(self) -> Result<RuntimeReport, RuntimeError> {
218 self.run_with_cancellation(CancellationToken::new()).await
219 }
220
221 pub async fn run_with_cancellation(
222 self,
223 cancellation: CancellationToken,
224 ) -> Result<RuntimeReport, RuntimeError> {
225 self.run_with_cancellation_reason(cancellation, DEFAULT_CANCELLATION_REASON)
226 .await
227 }
228
229 pub async fn run_with_cancellation_reason(
237 mut self,
238 cancellation: CancellationToken,
239 reason: impl Into<String>,
240 ) -> Result<RuntimeReport, RuntimeError> {
241 self.cancellation_reason = reason.into();
242 self.spec.validate().map_err(RuntimeError::InvalidSpec)?;
243 self.action_space = Arc::new(self.spec.action_space_validated().clone());
246 self.observation_space = Arc::new(self.spec.observation_space_validated().clone());
247 let mut state = RouteState::new(&self.spec);
248 let telemetry = Arc::new(Mutex::new(Aggregator::default()));
254 let ticker = (!self.spec.limits.telemetry_window.is_zero()).then(|| {
257 TelemetryTicker::spawn(
258 Arc::clone(&telemetry),
259 Arc::clone(&self.hooks),
260 self.spec.limits.telemetry_window,
261 state.session_id().to_string(),
262 state.env_context(),
263 )
264 });
265 let result = self.run_loop(&mut state, &cancellation, &telemetry).await;
266 drop(ticker);
270 let final_snapshot = lock_agg(&telemetry).snapshot(Horizon::Session);
271 fan_out_event!(
272 self,
273 on_telemetry,
274 TelemetrySnapshotEvent {
275 session_id: state.session_id().to_string(),
276 route: state.env_context(),
277 snapshot: final_snapshot,
278 }
279 );
280 if let Err(error) = &result {
281 self.shutdown_after_failure(&mut state, error).await;
282 }
283 result
284 }
285
286 #[tracing::instrument(
291 name = "rlmesh.route",
292 level = "info",
293 skip_all,
294 fields(
295 session_id = %state.session_id(),
296 env_id = %self.spec.env_id,
297 num_envs = self.spec.num_envs,
298 ),
299 )]
300 async fn run_loop(
301 &mut self,
302 state: &mut RouteState,
303 cancellation: &CancellationToken,
304 telemetry: &Arc<Mutex<Aggregator>>,
305 ) -> Result<RuntimeReport, RuntimeError> {
306 fan_out_event!(
307 self,
308 session_started,
309 SessionStartedEvent {
310 session_id: state.session_id().to_string(),
311 route: state.env_context(),
312 env_id: self.spec.env_id.clone(),
313 }
314 );
315
316 let mut reset_generation = 0_u64;
317 let reset_timeout = self.spec.limits.env_reset_timeout;
318 let reset_timeout_ms = self.spec.limits.env_reset_timeout_ms().max(0) as u64;
320 let reset_seeds = self.reset_seeds(reset_generation);
321 let initial_episode_ids = mint_episode_ids(self.spec.num_envs);
325 let reset_request = ResetRequest {
326 seeds: reset_seeds,
327 options: None,
328 timeout_ms: reset_timeout_ms,
329 env_indices: Vec::new(),
330 episode_ids: initial_episode_ids.clone(),
331 };
332 let reset_request_bytes = reset_request.encoded_len() as u64;
333 let reset_started = Instant::now();
336 let reset_ok = await_runtime_operation(
337 cancellation,
338 reset_timeout,
339 RuntimeError::operation_timeout(
340 state.env_id(),
341 state.env_component_id(),
342 "env.reset",
343 0,
344 reset_timeout,
345 ),
346 self.cancelled_error(state, 0),
347 self.env.reset(reset_request),
348 )
349 .await?;
350 let reset_latency = reset_started.elapsed();
351 record_op(
352 telemetry,
353 SRC_RESET,
354 reset_latency,
355 reset_ok.endpoint_total_ns,
356 reset_request_bytes,
357 reset_ok.response.encoded_len() as u64,
358 );
359 fan_out_event!(
360 self,
361 log,
362 LogEvent {
363 session_id: state.session_id().to_string(),
364 route: state.env_context(),
365 level: LogLevel::Info,
366 message: format!(
367 "env reset complete in {:.0}ms ({} episode(s) ready)",
368 reset_latency.as_secs_f64() * 1000.0,
369 initial_episode_ids.len()
370 ),
371 source: Some("runtime".to_string()),
372 }
373 );
374
375 let reset_observation = value_leaves(reset_ok.response.observation.as_ref())?;
376 let started_episodes = state.start_episodes(initial_episode_ids, false);
377 self.invoke_started_episodes(state, started_episodes).await;
378
379 let mut pending_roll: std::collections::HashMap<u32, String> =
385 std::collections::HashMap::new();
386
387 let mut reset_msg =
388 state.predict_request(reset_observation.clone(), RequestPhase::ResetObservation);
389 let mut reset_event =
390 self.observation_event(state, state.snapshot(), true, reset_observation.clone());
391 let transformed_reset_observation = self
392 .invoke_transform_observation(reset_event.clone())
393 .await?;
394 reset_event.observation = transformed_reset_observation.clone();
395 reset_msg.observation = transformed_reset_observation.map(leaves_value);
396 fan_out_event!(self, observation_emitted, reset_event);
397
398 let mut pending_observation_msg = reset_msg;
399
400 let mut replay_buffer: std::collections::VecDeque<Vec<Bytes>> =
408 std::collections::VecDeque::new();
409
410 loop {
411 if cancellation.is_cancelled() {
412 return Err(self.cancelled_error(state, state.snapshot().step));
413 }
414
415 let predict_snapshot = state.snapshot();
416 if replay_buffer.is_empty() {
423 let predict_timeout = self.spec.limits.model_predict_timeout;
424 let expected_context = pending_observation_msg.context.clone();
425 let predict_request_bytes = pending_observation_msg.encoded_len() as u64;
426 let predict_started = Instant::now();
427 let action_msg = await_runtime_operation(
428 cancellation,
429 predict_timeout,
430 RuntimeError::operation_timeout(
431 state.env_id(),
432 state.model_component_id(),
433 "model.predict",
434 predict_snapshot.step,
435 predict_timeout,
436 ),
437 self.cancelled_error(state, predict_snapshot.step),
438 self.model.predict(pending_observation_msg),
439 )
440 .await?;
441 let predict_rpc = predict_started.elapsed();
442 if action_msg.response.context != expected_context {
443 let request_id = expected_context
444 .as_ref()
445 .map(|context| context.request_id.clone())
446 .unwrap_or_default();
447 return Err(RuntimeError::ModelRouteMismatch {
448 component_id: state.model_component_id().to_string(),
449 request_id,
450 });
451 }
452 record_op(
453 telemetry,
454 SRC_PREDICT,
455 predict_rpc,
456 action_msg.endpoint_total_ns,
457 predict_request_bytes,
458 action_msg.response.encoded_len() as u64,
459 );
460 if action_msg.response.actions.is_empty() {
461 return Err(RuntimeError::Protocol(format!(
462 "model endpoint {} returned a predict response with no actions",
463 state.model_component_id()
464 )));
465 }
466 for frame in &action_msg.response.actions {
470 if let Some(leaves) = value_leaves(Some(frame))? {
471 replay_buffer.push_back(leaves);
472 }
473 }
474 }
475 let model_action = replay_buffer
476 .pop_front()
477 .expect("replay buffer is non-empty after a refill");
478
479 let action_step = predict_snapshot.step + 1;
480 let mut action_event = ActionReceivedEvent {
481 session_id: state.session_id().to_string(),
482 route: state.env_context(),
483 episode_id: predict_snapshot.episode_id.clone(),
484 episode_record_id: predict_snapshot.episode_record_id.clone(),
485 episode_ids: predict_snapshot.episode_ids.clone(),
486 episode_record_ids: predict_snapshot.episode_record_ids.clone(),
487 step: action_step,
488 env_index: predict_snapshot.env_index,
489 action_space: Arc::clone(&self.action_space),
490 action: Some(model_action),
491 };
492 action_event.action = self.invoke_transform_action(action_event.clone()).await?;
493 fan_out_event!(self, action_received, action_event.clone());
494
495 let step_timeout = self.spec.limits.env_step_timeout;
496 let step_timeout_ms = self.spec.limits.env_step_timeout_ms().max(0) as u64;
498 let step_episode_ids = episode_ids_with_roll(state.episode_ids(), &pending_roll);
503 let step_request = StepRequest {
504 action: action_event.action.map(leaves_value),
505 timeout_ms: step_timeout_ms,
506 env_indices: Vec::new(),
507 episode_ids: step_episode_ids,
508 };
509 let step_request_bytes = step_request.encoded_len() as u64;
510 let step_started = Instant::now();
511 let step_ok = await_runtime_operation(
512 cancellation,
513 step_timeout,
514 RuntimeError::operation_timeout(
515 state.env_id(),
516 state.env_component_id(),
517 "env.step",
518 action_step,
519 step_timeout,
520 ),
521 self.cancelled_error(state, action_step),
522 self.env.step(step_request),
523 )
524 .await?;
525 let step_rpc = step_started.elapsed();
526 record_op(
527 telemetry,
528 SRC_STEP,
529 step_rpc,
530 step_ok.endpoint_total_ns,
531 step_request_bytes,
532 step_ok.response.encoded_len() as u64,
533 );
534 let step_observation = value_leaves(step_ok.response.observation.as_ref())?;
535
536 state.record_step();
537 let step_snapshot = state.snapshot();
538 fan_out_event!(
539 self,
540 step_completed,
541 StepCompletedEvent {
542 session_id: state.session_id().to_string(),
543 route: state.env_context(),
544 episode_id: step_snapshot.episode_id.clone(),
545 episode_record_id: step_snapshot.episode_record_id.clone(),
546 step: step_snapshot.step,
547 env_index: step_snapshot.env_index,
548 rewards: step_ok.response.rewards.clone(),
549 }
550 );
551
552 if !pending_roll.is_empty() {
558 let roll_ids = episode_ids_with_roll(state.episode_ids(), &pending_roll);
559 pending_roll.clear();
560 let started_episodes = state.observe_episode_ids(roll_ids);
561 self.invoke_started_episodes(state, started_episodes).await;
562 }
563 self.emit_completed_episodes(state, &step_ok.response.completed_episodes)
570 .await;
571 self.emit_reset_adapter(state, &step_ok.response.completed_episodes)
574 .await;
575
576 if !step_ok.response.completed_episodes.is_empty() {
581 replay_buffer.clear();
582 }
583
584 if matches!(
588 self.autoreset_mode(),
589 AutoresetMode::NextStep | AutoresetMode::SameStep
590 ) {
591 for completed in &step_ok.response.completed_episodes {
592 pending_roll
593 .entry(completed.env_index)
594 .or_insert_with(mint_episode_id);
595 }
596 }
597
598 if self
605 .spec
606 .max_episodes
607 .is_some_and(|limit| state.total_episodes() >= limit as i64)
608 {
609 let release_request = state.release_adapter_request("completed requested episodes");
610 self.shutdown_terminal_route(
611 state,
612 "completed requested episodes",
613 release_request,
614 )
615 .await;
616 let telemetry_snapshot = lock_agg(telemetry).snapshot(Horizon::Session);
620 fan_out_event!(
621 self,
622 session_ended,
623 SessionEndedEvent {
624 session_id: state.session_id().to_string(),
625 route: state.env_context(),
626 reason: "completed requested episodes".to_string(),
627 total_steps: state.total_steps(),
628 total_episodes: state.total_episodes(),
629 }
630 );
631 return Ok(RuntimeReport {
632 session_id: state.session_id().to_string(),
633 env_id: self.spec.env_id.clone(),
634 total_steps: state.total_steps(),
635 total_episodes: state.total_episodes(),
636 telemetry: telemetry_snapshot,
637 });
638 }
639
640 let (next_obs, phase, is_reset_msg) = match self.autoreset_mode() {
644 AutoresetMode::NextStep | AutoresetMode::SameStep => (
649 step_observation.clone(),
650 RequestPhase::StepObservation,
651 false,
652 ),
653 AutoresetMode::Unspecified | AutoresetMode::Disabled => {
659 let mut done_lanes: Vec<u32> = step_ok
662 .response
663 .completed_episodes
664 .iter()
665 .map(|metadata| metadata.env_index)
666 .collect();
667 done_lanes.sort_unstable();
675 done_lanes.dedup();
676 if done_lanes.is_empty() {
677 (
678 step_observation.clone(),
679 RequestPhase::StepObservation,
680 false,
681 )
682 } else {
683 reset_generation += 1;
684 let step = state.snapshot().step;
685 let reset_timeout = self.spec.limits.env_reset_timeout;
686 let reset_timeout_ms =
688 self.spec.limits.env_reset_timeout_ms().max(0) as u64;
689 let whole_vector = done_lanes.len() == self.spec.num_envs;
690 let (reset_seeds, env_indices, reset_episode_ids) = if whole_vector {
694 (
695 self.reset_seeds(reset_generation),
696 Vec::new(),
697 mint_episode_ids(self.spec.num_envs),
698 )
699 } else {
700 (
701 self.reset_subset_seeds(reset_generation, &done_lanes),
702 done_lanes.clone(),
703 mint_episode_ids(done_lanes.len()),
704 )
705 };
706 let reset_request = ResetRequest {
707 seeds: reset_seeds,
708 options: None,
709 timeout_ms: reset_timeout_ms,
710 env_indices,
711 episode_ids: reset_episode_ids.clone(),
712 };
713 let reset_request_bytes = reset_request.encoded_len() as u64;
714 let inloop_reset_started = Instant::now();
715 let reset_ok = await_runtime_operation(
716 cancellation,
717 reset_timeout,
718 RuntimeError::operation_timeout(
719 state.env_id(),
720 state.env_component_id(),
721 "env.reset",
722 step,
723 reset_timeout,
724 ),
725 self.cancelled_error(state, step),
726 self.env.reset(reset_request),
727 )
728 .await?;
729 record_op(
730 telemetry,
731 SRC_RESET,
732 inloop_reset_started.elapsed(),
733 reset_ok.endpoint_total_ns,
734 reset_request_bytes,
735 reset_ok.response.encoded_len() as u64,
736 );
737 let next_obs = value_leaves(reset_ok.response.observation.as_ref())?;
738 let started_episodes = if whole_vector {
742 state.start_episodes(reset_episode_ids, true)
743 } else {
744 let mut full = state.episode_ids();
745 for (lane, id) in done_lanes.iter().zip(reset_episode_ids) {
746 if let Some(slot) = full.get_mut(*lane as usize) {
747 *slot = id;
748 }
749 }
750 state.observe_episode_ids(full)
751 };
752 self.invoke_started_episodes(state, started_episodes).await;
753 (next_obs, RequestPhase::ResetObservation, true)
754 }
755 }
756 };
757
758 let mut obs_msg = state.predict_request(next_obs.clone(), phase);
759 let mut outgoing_observation_event =
760 self.observation_event(state, state.snapshot(), is_reset_msg, next_obs);
761 let transformed_observation = self
762 .invoke_transform_observation(outgoing_observation_event.clone())
763 .await?;
764 outgoing_observation_event.observation = transformed_observation.clone();
765 obs_msg.observation = transformed_observation.map(leaves_value);
766 fan_out_event!(self, observation_emitted, outgoing_observation_event);
770
771 pending_observation_msg = obs_msg;
772 }
773 }
774
775 async fn shutdown_after_failure(&mut self, state: &mut RouteState, error: &RuntimeError) {
776 let reason = error.to_string();
777 let request = state.release_adapter_request(reason.clone());
778 self.shutdown_terminal_route(state, &reason, request).await;
779
780 if let Err(err) = self
781 .hooks
782 .session_failed(SessionFailedEvent {
783 session_id: state.session_id().to_string(),
784 route: state.env_context(),
785 reason,
786 })
787 .await
788 {
789 tracing::warn!("runtime hook session_failed failed: {err}");
790 }
791 }
792
793 async fn shutdown_terminal_route(
794 &mut self,
795 state: &RouteState,
796 reason: &str,
797 request: ReleaseAdapterRequest,
798 ) {
799 let timeout = self.spec.limits.service_close_timeout;
800 let model_close =
805 tokio::time::timeout(timeout, self.model.release_adapter(request, timeout));
806 if self.spec.close_env_on_end {
807 let env_close = tokio::time::timeout(timeout, self.env.close(timeout));
808 let (env_result, model_result) = tokio::join!(env_close, model_close);
809 match env_result {
810 Ok(Err(err)) => {
811 tracing::warn!(error = %err, "environment close failed during route shutdown");
812 }
813 Err(_) => {
814 tracing::warn!(
815 timeout_ms = timeout.as_millis(),
816 "environment close timed out during route shutdown; abandoning close"
817 );
818 }
819 Ok(Ok(())) => {}
820 }
821 log_model_close_result(model_result, reason, timeout);
822 return;
823 }
824
825 tracing::debug!(
826 env_id = %state.env_id(),
827 reason,
828 "skipping environment close for adapter; endpoint remains owned by the run"
829 );
830 log_model_close_result(model_close.await, reason, timeout);
831 }
832
833 fn cancelled_error(&self, state: &RouteState, step: i64) -> RuntimeError {
834 RuntimeError::route_cancelled(state.env_id(), step, self.cancellation_reason.as_str())
835 }
836
837 async fn invoke_started_episodes(&self, state: &RouteState, episodes: Vec<StartedEpisode>) {
838 for episode in episodes {
839 let record = &episode.record;
840 fan_out_event!(
841 self,
842 episode_started,
843 EpisodeStartedEvent {
844 session_id: state.session_id().to_string(),
845 route: state.env_context(),
846 episode_id: episode.episode_id.clone(),
847 episode_record_id: record.record_id.clone(),
848 episode_index: record.index,
849 env_index: record.env_index,
850 started_from_auto_reset: record.started_from_auto_reset,
851 }
852 );
853 }
854 }
855
856 async fn emit_completed_episodes(&self, state: &mut RouteState, episodes: &[EpisodeMetadata]) {
857 for completed in episodes {
858 let record = state.complete_episode(&completed.episode_id);
859 let episode_record_id = record
860 .as_ref()
861 .map(|record| record.record_id.clone())
862 .unwrap_or_default();
863 let env_index = i32::try_from(completed.env_index).unwrap_or(i32::MAX);
865 fan_out_event!(
866 self,
867 episode_completed,
868 EpisodeCompletedEvent {
869 session_id: state.session_id().to_string(),
870 route: state.env_context(),
871 episode_id: completed.episode_id.clone(),
872 episode_record_id,
873 episode_index: record.as_ref().map_or(0, |record| record.index),
874 env_index,
875 step_count: completed.step_count,
876 cumulative_reward: completed.cumulative_reward,
877 terminated: completed.terminated,
878 truncated: completed.truncated,
879 duration_ms: i64::try_from(completed.duration_ms).unwrap_or(i64::MAX),
880 final_info: completed.final_info.clone(),
881 }
882 );
883 }
884 }
885
886 async fn emit_reset_adapter(&mut self, state: &mut RouteState, episodes: &[EpisodeMetadata]) {
897 let slot_ids = state.episode_ids();
898 let episode_ids: Vec<String> = episodes
899 .iter()
900 .filter_map(|completed| slot_ids.get(completed.env_index as usize).cloned())
901 .filter(|id| !id.is_empty())
902 .collect();
903 if episode_ids.is_empty() {
904 return;
905 }
906 let request = state.reset_adapter_request(episode_ids);
907 if let Err(err) = self.model.reset_adapter(request).await {
908 tracing::warn!("model reset_adapter (evict) failed: {err}");
909 }
910 }
911
912 async fn invoke_transform_action(
913 &self,
914 event: ActionReceivedEvent,
915 ) -> Result<Option<Vec<Bytes>>, RuntimeError> {
916 match self.hooks.transform_action(event).await {
917 Ok(action) => Ok(action),
918 Err(err) => {
919 tracing::warn!("runtime hook transform_action failed: {err}");
920 Err(RuntimeError::Hook(err))
921 }
922 }
923 }
924
925 async fn invoke_transform_observation(
926 &self,
927 event: ObservationEmittedEvent,
928 ) -> Result<Option<Vec<Bytes>>, RuntimeError> {
929 match self.hooks.transform_observation(event).await {
930 Ok(observation) => Ok(observation),
931 Err(err) => {
932 tracing::warn!("runtime hook transform_observation failed: {err}");
933 Err(RuntimeError::Hook(err))
934 }
935 }
936 }
937
938 fn observation_event(
939 &self,
940 state: &RouteState,
941 snapshot: RouteSnapshot,
942 is_reset: bool,
943 observation: Option<Vec<Bytes>>,
944 ) -> ObservationEmittedEvent {
945 ObservationEmittedEvent {
946 session_id: state.session_id().to_string(),
947 route: state.env_context(),
948 episode_id: snapshot.episode_id,
949 episode_record_id: snapshot.episode_record_id,
950 episode_ids: snapshot.episode_ids,
951 episode_record_ids: snapshot.episode_record_ids,
952 step: snapshot.step,
953 env_index: snapshot.env_index,
954 is_reset,
955 num_envs: self.spec.num_envs as u32,
956 observation_space: Arc::clone(&self.observation_space),
957 observation,
958 }
959 }
960}
961
962fn deterministic_reset_seed(
967 base_seed: i64,
968 session_id: &str,
969 reset_generation: u64,
970 env_index: usize,
971) -> i64 {
972 const FNV_OFFSET: u64 = 0xcbf2_9ce4_8422_2325;
973 const FNV_PRIME: u64 = 0x0000_0100_0000_01b3;
974
975 fn update(mut hash: u64, bytes: &[u8]) -> u64 {
976 for byte in bytes {
977 hash ^= u64::from(*byte);
978 hash = hash.wrapping_mul(FNV_PRIME);
979 }
980 hash
981 }
982
983 let mut hash = FNV_OFFSET;
984 hash = update(hash, &base_seed.to_le_bytes());
985 hash = update(hash, &[0xff]);
986 hash = update(hash, session_id.as_bytes());
987 hash = update(hash, &[0xfd]);
988 hash = update(hash, &reset_generation.to_le_bytes());
989 hash = update(hash, &[0xfc]);
990 hash = update(hash, &(env_index as u64).to_le_bytes());
991 (hash & i64::MAX as u64) as i64
992}
993
994async fn await_runtime_operation<T, F>(
995 cancellation: &CancellationToken,
996 timeout: Duration,
997 timeout_error: RuntimeError,
998 cancelled_error: RuntimeError,
999 operation: F,
1000) -> Result<T, RuntimeError>
1001where
1002 F: Future<Output = Result<T, RuntimeError>>,
1003{
1004 tokio::select! {
1005 _ = cancellation.cancelled() => Err(cancelled_error),
1006 result = tokio::time::timeout(timeout, operation) => match result {
1007 Ok(result) => result,
1008 Err(_) => Err(timeout_error),
1009 },
1010 }
1011}
1012
1013fn log_model_close_result(
1014 result: Result<Result<(), String>, tokio::time::error::Elapsed>,
1015 reason: &str,
1016 timeout: Duration,
1017) {
1018 match result {
1019 Ok(Err(err)) => {
1020 tracing::warn!(
1021 error = %err,
1022 reason,
1023 "model route close failed during route shutdown; relying on owner shutdown"
1024 );
1025 }
1026 Err(_) => {
1027 tracing::warn!(
1028 timeout_ms = timeout.as_millis(),
1029 reason,
1030 "model route close timed out during route shutdown; relying on owner shutdown"
1031 );
1032 }
1033 Ok(Ok(())) => {}
1034 }
1035}
1036
1037fn leaves_value(leaves: Vec<Bytes>) -> SpaceValue {
1038 SpaceValue { leaves }
1039}
1040
1041fn mint_episode_id() -> String {
1045 uuid::Uuid::now_v7().to_string()
1046}
1047
1048fn mint_episode_ids(count: usize) -> Vec<String> {
1050 (0..count).map(|_| mint_episode_id()).collect()
1051}
1052
1053fn episode_ids_with_roll(
1058 mut ids: Vec<String>,
1059 pending_roll: &std::collections::HashMap<u32, String>,
1060) -> Vec<String> {
1061 for (env_index, new_id) in pending_roll {
1064 if let Some(slot) = ids.get_mut(*env_index as usize) {
1065 *slot = new_id.clone();
1066 }
1067 }
1068 ids
1069}
1070
1071fn value_leaves(payload: Option<&SpaceValue>) -> Result<Option<Vec<Bytes>>, RuntimeError> {
1075 Ok(payload.map(|payload| payload.leaves.clone()))
1076}
1077
1078fn lock_agg(telemetry: &Mutex<Aggregator>) -> MutexGuard<'_, Aggregator> {
1082 telemetry
1083 .lock()
1084 .unwrap_or_else(|poisoned| poisoned.into_inner())
1085}
1086
1087fn record_op(
1092 telemetry: &Mutex<Aggregator>,
1093 src: Source,
1094 rpc: Duration,
1095 endpoint_total_ns: Option<u64>,
1096 request_bytes: u64,
1097 response_bytes: u64,
1098) {
1099 let mut agg = lock_agg(telemetry);
1100 agg.record(Sample::dur(src, metrics::RPC_TOTAL, rpc));
1101 if let Some(ns) = endpoint_total_ns {
1102 agg.record(Sample::dur(
1103 src,
1104 metrics::ENDPOINT_TOTAL,
1105 Duration::from_nanos(ns),
1106 ));
1107 }
1108 agg.record(Sample::bytes(src, metrics::REQUEST_BYTES, request_bytes));
1109 agg.record(Sample::bytes(src, metrics::RESPONSE_BYTES, response_bytes));
1110}
1111
1112struct TelemetryTicker {
1122 handle: tokio::task::JoinHandle<()>,
1123}
1124
1125impl TelemetryTicker {
1126 fn spawn(
1127 telemetry: Arc<Mutex<Aggregator>>,
1128 hooks: Arc<dyn RuntimeHooks>,
1129 window: Duration,
1130 session_id: String,
1131 route: RuntimeEnvContext,
1132 ) -> Self {
1133 let period = window.max(Duration::from_millis(1));
1136 let handle = tokio::spawn(async move {
1137 let mut ticker = tokio::time::interval(period);
1138 ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
1139 ticker.tick().await; loop {
1141 ticker.tick().await;
1142 let window_snap = {
1146 let mut agg = lock_agg(&telemetry);
1147 let snap = agg.snapshot(Horizon::Window);
1148 agg.flush_window();
1149 snap
1150 };
1151 if window_snap.rows.is_empty() {
1154 continue;
1155 }
1156 let window_event = TelemetrySnapshotEvent {
1159 session_id: session_id.clone(),
1160 route: route.clone(),
1161 snapshot: window_snap,
1162 };
1163 if let Err(err) = hooks.on_telemetry(window_event).await {
1164 tracing::warn!("runtime hook on_telemetry (window) failed: {err}");
1165 }
1166 }
1167 });
1168 Self { handle }
1169 }
1170}
1171
1172impl Drop for TelemetryTicker {
1173 fn drop(&mut self) {
1174 self.handle.abort();
1175 }
1176}