use std::future::Future;
use std::sync::{Arc, Mutex, MutexGuard};
use std::time::{Duration, Instant};
use async_trait::async_trait;
use prost::{Message, bytes::Bytes};
use rlmesh_proto::EndpointPhases;
use rlmesh_proto::core::v1::AutoresetMode;
use rlmesh_proto::env::v1::{
EpisodeMetadata, ResetRequest, ResetResponse, StepRequest, StepResponse,
};
use rlmesh_proto::model::v1::{
AdapterContext, PredictRequest, PredictResponse, ReleaseAdapterRequest, ResetAdapterRequest,
};
use rlmesh_proto::spaces::v1::SpaceValue;
use tokio_util::sync::CancellationToken;
use crate::hooks::{
ActionReceivedEvent, EpisodeCompletedEvent, EpisodeStartedEvent, LogEvent, LogLevel,
ObservationEmittedEvent, RuntimeEnvContext, RuntimeHooks, SessionEndedEvent,
SessionFailedEvent, SessionStartedEvent, StepCompletedEvent, TelemetrySnapshotEvent,
};
use crate::spec::{RuntimeReport, RuntimeSessionSpec};
use crate::state::{RequestPhase, RouteSnapshot, RouteState, StartedEpisode};
use crate::telemetry::{Aggregator, Horizon, Sample, Source, metrics};
mod error;
pub use error::RuntimeError;
macro_rules! fan_out_event {
($self:ident, $method:ident, $event:expr) => {
if let Err(err) = $self.hooks.$method($event).await {
tracing::warn!(
concat!("runtime hook ", stringify!($method), " failed: {}"),
err
);
}
};
}
pub struct RuntimeEnvReset {
pub response: ResetResponse,
pub endpoint_total_ns: Option<u64>,
pub phases: EndpointPhases,
}
pub struct RuntimeEnvStep {
pub response: StepResponse,
pub endpoint_total_ns: Option<u64>,
pub phases: EndpointPhases,
}
pub struct RuntimeModelPrediction {
pub response: PredictResponse,
pub endpoint_total_ns: Option<u64>,
pub phases: EndpointPhases,
pub group_size: Option<u64>,
}
pub(crate) struct PeerReport {
pub(crate) endpoint_total_ns: Option<u64>,
pub(crate) phases: EndpointPhases,
pub(crate) group_size: Option<u64>,
}
#[async_trait]
pub trait RuntimeEnv: Send {
async fn reset(&mut self, request: ResetRequest) -> Result<RuntimeEnvReset, RuntimeError>;
async fn step(&mut self, request: StepRequest) -> Result<RuntimeEnvStep, RuntimeError>;
async fn close(&mut self, _timeout: Duration) -> Result<(), String> {
Ok(())
}
}
#[async_trait]
pub trait RuntimeModel: Send {
async fn predict(
&mut self,
request: PredictRequest,
) -> Result<RuntimeModelPrediction, RuntimeError>;
async fn reset_adapter(&mut self, _request: ResetAdapterRequest) -> Result<(), RuntimeError> {
Ok(())
}
async fn release_adapter(
&mut self,
_request: ReleaseAdapterRequest,
_timeout: Duration,
) -> Result<(), String> {
Ok(())
}
}
const DEFAULT_CANCELLATION_REASON: &str = "cancelled by caller";
const DEFAULT_MAX_EPISODE_STEPS: i64 = 100_000;
fn success_from_final_info(final_info: Option<&rlmesh_proto::spaces::v1::MetaMap>) -> Option<bool> {
use rlmesh_proto::spaces::v1::meta_value::Kind;
let entries = &final_info?.entries;
["is_success", "success"]
.iter()
.find_map(|key| match entries.get(*key)?.kind.as_ref()? {
Kind::Bool(value) => Some(*value),
Kind::Integer(value) => Some(*value != 0),
Kind::Number(value) => Some(*value != 0.0),
_ => None,
})
}
const SRC_PREDICT: Source = Source {
op: "model.predict",
component: "model",
};
const SRC_STEP: Source = Source {
op: "env.step",
component: "env",
};
const SRC_RESET: Source = Source {
op: "env.reset",
component: "env",
};
const SRC_TRANSFORM_OBS: Source = Source {
op: "runner.transform_observation",
component: "runner",
};
const SRC_TRANSFORM_ACTION: Source = Source {
op: "runner.transform_action",
component: "runner",
};
const SRC_ROUND: Source = Source {
op: "runner.round",
component: "runner",
};
#[must_use = "a RuntimeDriver does nothing until one of its run methods is awaited"]
pub struct RuntimeDriver<E, M> {
spec: RuntimeSessionSpec,
env: E,
model: M,
prefetch_model: Option<Box<dyn RuntimeModel + Send>>,
prefetch_lead: u32,
hooks: Arc<dyn RuntimeHooks>,
cancellation_reason: String,
action_space: Arc<rlmesh_proto::spaces::v1::SpaceSpec>,
observation_space: Arc<rlmesh_proto::spaces::v1::SpaceSpec>,
}
impl<E, M> RuntimeDriver<E, M>
where
E: RuntimeEnv,
M: RuntimeModel,
{
pub fn new(spec: RuntimeSessionSpec, env: E, model: M, hooks: Arc<dyn RuntimeHooks>) -> Self {
Self {
spec,
env,
model,
prefetch_model: None,
prefetch_lead: 0,
hooks,
cancellation_reason: DEFAULT_CANCELLATION_REASON.to_string(),
action_space: Arc::default(),
observation_space: Arc::default(),
}
}
pub fn with_prefetch(mut self, model: Box<dyn RuntimeModel + Send>, lead: u32) -> Self {
if lead > 0 {
self.prefetch_model = Some(model);
self.prefetch_lead = lead;
}
self
}
fn planned_reset_seeds(
&self,
state: &mut RouteState,
reset_generation: u64,
env_indices: Option<&[u32]>,
) -> Vec<i64> {
if !self.spec.episode_seeds.is_empty() {
let lanes = env_indices.map_or(self.spec.num_envs, <[u32]>::len);
return state.claim_episode_seeds(&self.spec.episode_seeds, lanes);
}
match env_indices {
None => self.reset_seeds(reset_generation),
Some(indices) => self.reset_subset_seeds(reset_generation, indices),
}
}
fn reset_seeds(&self, reset_generation: u64) -> Vec<i64> {
self.seeds_for(reset_generation, 0..self.spec.num_envs)
}
fn reset_subset_seeds(&self, reset_generation: u64, env_indices: &[u32]) -> Vec<i64> {
self.seeds_for(
reset_generation,
env_indices.iter().map(|&index| index as usize),
)
}
fn seeds_for(
&self,
reset_generation: u64,
env_indices: impl Iterator<Item = usize>,
) -> Vec<i64> {
let Some(base_seed) = self.spec.base_seed else {
return Vec::new();
};
env_indices
.map(|env_index| {
deterministic_reset_seed(
base_seed,
&self.spec.session_id,
reset_generation,
env_index,
)
})
.collect()
}
fn autoreset_mode(&self) -> AutoresetMode {
AutoresetMode::try_from(self.spec.env_contract.autoreset_mode)
.unwrap_or(AutoresetMode::Disabled)
}
pub async fn run(self) -> Result<RuntimeReport, RuntimeError> {
self.run_with_cancellation(CancellationToken::new()).await
}
pub async fn run_with_cancellation(
self,
cancellation: CancellationToken,
) -> Result<RuntimeReport, RuntimeError> {
self.run_with_cancellation_reason(cancellation, DEFAULT_CANCELLATION_REASON)
.await
}
pub async fn run_with_cancellation_reason(
mut self,
cancellation: CancellationToken,
reason: impl Into<String>,
) -> Result<RuntimeReport, RuntimeError> {
self.cancellation_reason = reason.into();
self.spec.validate().map_err(RuntimeError::InvalidSpec)?;
self.action_space = Arc::new(self.spec.action_space_validated().clone());
self.observation_space = Arc::new(self.spec.observation_space_validated().clone());
let mut state = RouteState::new(&self.spec);
let telemetry = Arc::new(Mutex::new(Aggregator::default()));
let ticker = (!self.spec.limits.telemetry_window.is_zero()).then(|| {
TelemetryTicker::spawn(
Arc::clone(&telemetry),
Arc::clone(&self.hooks),
self.spec.limits.telemetry_window,
state.session_id().to_string(),
state.env_context(),
)
});
let result = self.run_loop(&mut state, &cancellation, &telemetry).await;
drop(ticker);
let final_snapshot = lock_agg(&telemetry).snapshot(Horizon::Session);
fan_out_event!(
self,
on_telemetry,
TelemetrySnapshotEvent {
session_id: state.session_id().to_string(),
route: state.env_context(),
snapshot: final_snapshot,
}
);
if let Err(error) = &result {
self.shutdown_after_failure(&mut state, error).await;
}
result
}
#[tracing::instrument(
name = "rlmesh.route",
level = "info",
skip_all,
fields(
session_id = %state.session_id(),
env_id = %self.spec.env_id,
num_envs = self.spec.num_envs,
),
)]
async fn run_loop(
&mut self,
state: &mut RouteState,
cancellation: &CancellationToken,
telemetry: &Arc<Mutex<Aggregator>>,
) -> Result<RuntimeReport, RuntimeError> {
fan_out_event!(
self,
session_started,
SessionStartedEvent {
session_id: state.session_id().to_string(),
route: state.env_context(),
env_id: self.spec.env_id.clone(),
}
);
let mut reset_generation = 0_u64;
let reset_timeout = self.spec.limits.env_reset_timeout;
let reset_timeout_ms = self.spec.limits.env_reset_timeout_ms().max(0) as u64;
let reset_seeds = self.planned_reset_seeds(state, reset_generation, None);
let initial_episode_ids = mint_episode_ids(self.spec.num_envs);
state.note_episode_seeds(&initial_episode_ids, &reset_seeds);
let reset_request = ResetRequest {
seeds: reset_seeds,
options: None,
timeout_ms: reset_timeout_ms,
env_indices: Vec::new(),
episode_ids: initial_episode_ids.clone(),
};
let reset_request_bytes = reset_request.encoded_len() as u64;
let reset_started = Instant::now();
let reset_ok = await_runtime_operation(
cancellation,
reset_timeout,
RuntimeError::operation_timeout(
state.env_id(),
state.env_component_id(),
"env.reset",
0,
reset_timeout,
),
self.cancelled_error(state, 0),
self.env.reset(reset_request),
)
.await?;
let reset_latency = reset_started.elapsed();
record_op(
telemetry,
SRC_RESET,
reset_latency,
PeerReport {
endpoint_total_ns: reset_ok.endpoint_total_ns,
phases: reset_ok.phases,
group_size: None,
},
reset_request_bytes,
reset_ok.response.encoded_len() as u64,
);
fan_out_event!(
self,
log,
LogEvent {
session_id: state.session_id().to_string(),
route: state.env_context(),
level: LogLevel::Info,
message: format!(
"env reset complete in {:.0}ms ({} episode(s) ready)",
reset_latency.as_secs_f64() * 1000.0,
initial_episode_ids.len()
),
source: Some("runtime".to_string()),
}
);
let reset_observation = value_leaves(reset_ok.response.observation.as_ref())?;
let started_episodes = state.start_episodes(initial_episode_ids, false);
self.invoke_started_episodes(state, started_episodes).await;
let mut pending_roll: std::collections::HashMap<u32, String> =
std::collections::HashMap::new();
let mut reset_msg =
state.predict_request(reset_observation.clone(), RequestPhase::ResetObservation);
let mut reset_event = self.observation_event(
state,
state.snapshot(),
true,
reset_observation.clone(),
reset_ok.response.infos.clone(),
);
let transformed_reset_observation = self
.invoke_transform_observation(telemetry, reset_event.clone())
.await?;
reset_event.observation = transformed_reset_observation.clone();
reset_msg.observation = transformed_reset_observation.map(leaves_value);
fan_out_event!(self, observation_emitted, reset_event);
let mut pending_observation_msg = reset_msg;
let mut replay_buffer: std::collections::VecDeque<Vec<Bytes>> =
std::collections::VecDeque::new();
let mut prefetch_inflight: Option<PrefetchInflight> = None;
let mut prefetch_stale = false;
loop {
let round_started = Instant::now();
if cancellation.is_cancelled() {
return Err(self.cancelled_error(state, state.snapshot().step));
}
let predict_snapshot = state.snapshot();
if replay_buffer.is_empty() {
let predict_timeout = self.spec.limits.model_predict_timeout;
let mut prefetched: Option<(
RuntimeModelPrediction,
Option<AdapterContext>,
u64,
Duration,
)> = None;
if let Some(inflight) = prefetch_inflight.take() {
let step = predict_snapshot.step;
let join_started = Instant::now();
let (model, result) = await_runtime_operation(
cancellation,
predict_timeout,
RuntimeError::operation_timeout(
state.env_id(),
state.model_component_id(),
"model.predict",
step,
predict_timeout,
),
self.cancelled_error(state, step),
join_prefetch(inflight.join, state.model_component_id()),
)
.await?;
let join_wait = join_started.elapsed();
self.prefetch_model = Some(model);
if prefetch_stale {
match result {
Ok(discarded) => record_op(
telemetry,
SRC_PREDICT,
join_wait,
PeerReport {
endpoint_total_ns: discarded.endpoint_total_ns,
phases: discarded.phases,
group_size: discarded.group_size,
},
inflight.request_bytes,
discarded.response.encoded_len() as u64,
),
Err(error) => tracing::warn!(
error = %error,
"stale prefetch predict failed; re-planning from the current observation"
),
}
} else {
prefetched = Some((
result?,
inflight.expected_context,
inflight.request_bytes,
join_wait,
));
}
}
prefetch_stale = false;
let (action_msg, expected_context, predict_request_bytes, predict_rpc) =
match prefetched {
Some((chunk, context, bytes, join_wait)) => {
(chunk, context, bytes, join_wait)
}
None => {
let expected_context = pending_observation_msg.context.clone();
let predict_request_bytes =
pending_observation_msg.encoded_len() as u64;
let predict_started = Instant::now();
let action_msg = await_runtime_operation(
cancellation,
predict_timeout,
RuntimeError::operation_timeout(
state.env_id(),
state.model_component_id(),
"model.predict",
predict_snapshot.step,
predict_timeout,
),
self.cancelled_error(state, predict_snapshot.step),
self.model.predict(pending_observation_msg),
)
.await?;
(
action_msg,
expected_context,
predict_request_bytes,
predict_started.elapsed(),
)
}
};
if action_msg.response.context != expected_context {
let request_id = expected_context
.as_ref()
.map(|context| context.request_id.clone())
.unwrap_or_default();
return Err(RuntimeError::ModelRouteMismatch {
component_id: state.model_component_id().to_string(),
request_id,
});
}
record_op(
telemetry,
SRC_PREDICT,
predict_rpc,
PeerReport {
endpoint_total_ns: action_msg.endpoint_total_ns,
phases: action_msg.phases,
group_size: action_msg.group_size,
},
predict_request_bytes,
action_msg.response.encoded_len() as u64,
);
if action_msg.response.actions.is_empty() {
return Err(RuntimeError::Protocol(format!(
"model endpoint {} returned a predict response with no actions",
state.model_component_id()
)));
}
for frame in &action_msg.response.actions {
if let Some(leaves) = value_leaves(Some(frame))? {
replay_buffer.push_back(leaves);
}
}
}
let model_action = replay_buffer
.pop_front()
.expect("replay buffer is non-empty after a refill");
let action_step = predict_snapshot.step + 1;
let mut action_event = ActionReceivedEvent {
session_id: state.session_id().to_string(),
route: state.env_context(),
episode_id: predict_snapshot.episode_id.clone(),
episode_record_id: predict_snapshot.episode_record_id.clone(),
episode_ids: predict_snapshot.episode_ids.clone(),
episode_record_ids: predict_snapshot.episode_record_ids.clone(),
step: action_step,
env_index: predict_snapshot.env_index,
action_space: Arc::clone(&self.action_space),
action: Some(model_action.clone()),
raw_action: Some(model_action),
};
action_event.action = self
.invoke_transform_action(telemetry, action_event.clone())
.await?;
fan_out_event!(self, action_received, action_event.clone());
let step_timeout = self.spec.limits.env_step_timeout;
let step_timeout_ms = self.spec.limits.env_step_timeout_ms().max(0) as u64;
let step_episode_ids = episode_ids_with_roll(state.episode_ids(), &pending_roll);
let step_request = StepRequest {
action: action_event.action.map(leaves_value),
timeout_ms: step_timeout_ms,
env_indices: Vec::new(),
episode_ids: step_episode_ids,
};
let step_request_bytes = step_request.encoded_len() as u64;
let step_started = Instant::now();
let step_ok = await_runtime_operation(
cancellation,
step_timeout,
RuntimeError::operation_timeout(
state.env_id(),
state.env_component_id(),
"env.step",
action_step,
step_timeout,
),
self.cancelled_error(state, action_step),
self.env.step(step_request),
)
.await?;
let step_rpc = step_started.elapsed();
record_op(
telemetry,
SRC_STEP,
step_rpc,
PeerReport {
endpoint_total_ns: step_ok.endpoint_total_ns,
phases: step_ok.phases,
group_size: None,
},
step_request_bytes,
step_ok.response.encoded_len() as u64,
);
let step_observation = value_leaves(step_ok.response.observation.as_ref())?;
state.record_step(&step_ok.response.rewards);
let step_snapshot = state.snapshot();
let rolled = !pending_roll.is_empty();
fan_out_event!(
self,
step_completed,
StepCompletedEvent {
session_id: state.session_id().to_string(),
route: state.env_context(),
episode_id: step_snapshot.episode_id.clone(),
episode_record_id: step_snapshot.episode_record_id.clone(),
step: step_snapshot.step,
env_index: step_snapshot.env_index,
rewards: step_ok.response.rewards.clone(),
infos: if rolled {
None
} else {
step_ok.response.infos.clone()
},
}
);
if !pending_roll.is_empty() {
let roll_ids = episode_ids_with_roll(state.episode_ids(), &pending_roll);
pending_roll.clear();
let started_episodes = state.observe_episode_ids(roll_ids);
self.invoke_started_episodes(state, started_episodes).await;
}
let capped = self.capped_lane_completions(state, &step_ok.response.completed_episodes);
let completed_episodes: std::borrow::Cow<'_, [EpisodeMetadata]> = if capped.is_empty() {
std::borrow::Cow::Borrowed(&step_ok.response.completed_episodes)
} else {
let mut all = step_ok.response.completed_episodes.clone();
all.extend(capped);
std::borrow::Cow::Owned(all)
};
self.emit_completed_episodes(state, &completed_episodes)
.await;
self.emit_reset_adapter(state, &completed_episodes).await;
if !completed_episodes.is_empty() {
replay_buffer.clear();
prefetch_stale = true;
}
if matches!(
self.autoreset_mode(),
AutoresetMode::NextStep | AutoresetMode::SameStep
) {
for completed in completed_episodes.iter() {
pending_roll
.entry(completed.env_index)
.or_insert_with(mint_episode_id);
}
}
if self
.spec
.max_episodes
.is_some_and(|limit| state.total_episodes() >= limit as i64)
{
lock_agg(telemetry).record(Sample::dur(
SRC_ROUND,
metrics::RPC_TOTAL,
round_started.elapsed(),
));
let release_request = state.release_adapter_request("completed requested episodes");
self.shutdown_terminal_route(
state,
"completed requested episodes",
release_request,
)
.await;
let telemetry_snapshot = lock_agg(telemetry).snapshot(Horizon::Session);
fan_out_event!(
self,
session_ended,
SessionEndedEvent {
session_id: state.session_id().to_string(),
route: state.env_context(),
reason: "completed requested episodes".to_string(),
total_steps: state.total_steps(),
total_episodes: state.total_episodes(),
}
);
return Ok(RuntimeReport {
session_id: state.session_id().to_string(),
env_id: self.spec.env_id.clone(),
total_steps: state.total_steps(),
total_episodes: state.total_episodes(),
episodes: state.take_episode_summaries(),
telemetry: telemetry_snapshot,
});
}
let (next_obs, phase, is_reset_msg, reset_infos) = match self.autoreset_mode() {
AutoresetMode::NextStep | AutoresetMode::SameStep => (
step_observation.clone(),
RequestPhase::StepObservation,
false,
if rolled {
step_ok.response.infos.clone()
} else {
None
},
),
AutoresetMode::Unspecified | AutoresetMode::Disabled => {
let mut done_lanes: Vec<u32> = completed_episodes
.iter()
.map(|metadata| metadata.env_index)
.collect();
done_lanes.sort_unstable();
done_lanes.dedup();
if done_lanes.is_empty() {
(
step_observation.clone(),
RequestPhase::StepObservation,
false,
None,
)
} else {
reset_generation += 1;
let step = state.snapshot().step;
let reset_timeout = self.spec.limits.env_reset_timeout;
let reset_timeout_ms =
self.spec.limits.env_reset_timeout_ms().max(0) as u64;
let whole_vector = done_lanes.len() == self.spec.num_envs;
let (reset_seeds, env_indices, reset_episode_ids) = if whole_vector {
(
self.planned_reset_seeds(state, reset_generation, None),
Vec::new(),
mint_episode_ids(self.spec.num_envs),
)
} else {
(
self.planned_reset_seeds(
state,
reset_generation,
Some(&done_lanes),
),
done_lanes.clone(),
mint_episode_ids(done_lanes.len()),
)
};
state.note_episode_seeds(&reset_episode_ids, &reset_seeds);
let reset_request = ResetRequest {
seeds: reset_seeds,
options: None,
timeout_ms: reset_timeout_ms,
env_indices,
episode_ids: reset_episode_ids.clone(),
};
let reset_request_bytes = reset_request.encoded_len() as u64;
let inloop_reset_started = Instant::now();
let reset_ok = await_runtime_operation(
cancellation,
reset_timeout,
RuntimeError::operation_timeout(
state.env_id(),
state.env_component_id(),
"env.reset",
step,
reset_timeout,
),
self.cancelled_error(state, step),
self.env.reset(reset_request),
)
.await?;
record_op(
telemetry,
SRC_RESET,
inloop_reset_started.elapsed(),
PeerReport {
endpoint_total_ns: reset_ok.endpoint_total_ns,
phases: reset_ok.phases,
group_size: None,
},
reset_request_bytes,
reset_ok.response.encoded_len() as u64,
);
let next_obs = value_leaves(reset_ok.response.observation.as_ref())?;
let started_episodes = if whole_vector {
state.start_episodes(reset_episode_ids, true)
} else {
let mut full = state.episode_ids();
for (lane, id) in done_lanes.iter().zip(reset_episode_ids) {
if let Some(slot) = full.get_mut(*lane as usize) {
*slot = id;
}
}
state.observe_episode_ids(full)
};
self.invoke_started_episodes(state, started_episodes).await;
(
next_obs,
RequestPhase::ResetObservation,
true,
reset_ok.response.infos.clone(),
)
}
}
};
let mut obs_msg = state.predict_request(next_obs.clone(), phase);
let mut outgoing_observation_event = self.observation_event(
state,
state.snapshot(),
is_reset_msg,
next_obs,
reset_infos,
);
let transformed_observation = self
.invoke_transform_observation(telemetry, outgoing_observation_event.clone())
.await?;
outgoing_observation_event.observation = transformed_observation.clone();
obs_msg.observation = transformed_observation.map(leaves_value);
fan_out_event!(self, observation_emitted, outgoing_observation_event);
pending_observation_msg = obs_msg;
if self.prefetch_lead > 0
&& prefetch_inflight.is_none()
&& replay_buffer.len() <= self.prefetch_lead as usize
&& let Some(mut model) = self.prefetch_model.take()
{
let request = pending_observation_msg.clone();
let expected_context = request.context.clone();
let request_bytes = request.encoded_len() as u64;
let join = tokio::spawn(async move {
let result = model.predict(request).await;
(model, result)
});
prefetch_inflight = Some(PrefetchInflight {
join,
expected_context,
request_bytes,
});
}
lock_agg(telemetry).record(Sample::dur(
SRC_ROUND,
metrics::RPC_TOTAL,
round_started.elapsed(),
));
}
}
async fn shutdown_after_failure(&mut self, state: &mut RouteState, error: &RuntimeError) {
let reason = error.to_string();
let request = state.release_adapter_request(reason.clone());
self.shutdown_terminal_route(state, &reason, request).await;
if let Err(err) = self
.hooks
.session_failed(SessionFailedEvent {
session_id: state.session_id().to_string(),
route: state.env_context(),
reason,
})
.await
{
tracing::warn!("runtime hook session_failed failed: {err}");
}
}
async fn shutdown_terminal_route(
&mut self,
state: &RouteState,
reason: &str,
request: ReleaseAdapterRequest,
) {
let timeout = self.spec.limits.service_close_timeout;
let model_close =
tokio::time::timeout(timeout, self.model.release_adapter(request, timeout));
if self.spec.close_env_on_end {
let env_close = tokio::time::timeout(timeout, self.env.close(timeout));
let (env_result, model_result) = tokio::join!(env_close, model_close);
match env_result {
Ok(Err(err)) => {
tracing::warn!(error = %err, "environment close failed during route shutdown");
}
Err(_) => {
tracing::warn!(
timeout_ms = timeout.as_millis(),
"environment close timed out during route shutdown; abandoning close"
);
}
Ok(Ok(())) => {}
}
log_model_close_result(model_result, reason, timeout);
return;
}
tracing::debug!(
env_id = %state.env_id(),
reason,
"skipping environment close for adapter; endpoint remains owned by the run"
);
log_model_close_result(model_close.await, reason, timeout);
}
fn cancelled_error(&self, state: &RouteState, step: i64) -> RuntimeError {
RuntimeError::route_cancelled(state.env_id(), step, self.cancellation_reason.as_str())
}
async fn invoke_started_episodes(&self, state: &RouteState, episodes: Vec<StartedEpisode>) {
for episode in episodes {
let record = &episode.record;
fan_out_event!(
self,
episode_started,
EpisodeStartedEvent {
session_id: state.session_id().to_string(),
route: state.env_context(),
episode_id: episode.episode_id.clone(),
episode_record_id: record.record_id.clone(),
episode_index: record.index,
env_index: record.env_index,
started_from_auto_reset: record.started_from_auto_reset,
seed: state.seed_for_episode(&episode.episode_id),
}
);
}
}
fn capped_lane_completions(
&self,
state: &RouteState,
env_completed: &[EpisodeMetadata],
) -> Vec<EpisodeMetadata> {
let driver_owns_resets = matches!(
self.autoreset_mode(),
AutoresetMode::Disabled | AutoresetMode::Unspecified
);
let step_cap = self
.spec
.max_episode_steps
.or_else(|| driver_owns_resets.then_some(DEFAULT_MAX_EPISODE_STEPS));
let time_cap = self.spec.max_episode_seconds;
if step_cap.is_none() && time_cap.is_none() {
return Vec::new();
}
let env_done: Vec<u32> = env_completed
.iter()
.map(|metadata| metadata.env_index)
.collect();
let now_ns = crate::state::now_unix_ns();
state
.slots()
.iter()
.filter_map(|slot| {
let episode = slot.episode.as_ref()?;
let env_index = u32::try_from(slot.env_index).ok()?;
if env_done.contains(&env_index) {
return None;
}
let steps_capped = step_cap.is_some_and(|cap| slot.step >= cap);
let elapsed_seconds = (now_ns - slot.started_at_ns).max(0) as f64 / 1e9;
let time_capped = time_cap.is_some_and(|cap| elapsed_seconds >= cap);
(steps_capped || time_capped).then(|| EpisodeMetadata {
episode_id: episode.episode_id.clone(),
seed: None,
env_index,
step_count: slot.step,
cumulative_reward: slot.cumulative_reward,
terminated: false,
truncated: true,
start_timestamp_ns: slot.started_at_ns,
end_timestamp_ns: now_ns,
final_info: None,
})
})
.collect()
}
async fn emit_completed_episodes(&self, state: &mut RouteState, episodes: &[EpisodeMetadata]) {
for completed in episodes {
let record = state.complete_episode(&completed.episode_id);
let episode_record_id = record
.as_ref()
.map(|record| record.record_id.clone())
.unwrap_or_default();
let env_index = i32::try_from(completed.env_index).unwrap_or(i32::MAX);
let seed = state.seed_for_episode(&completed.episode_id);
if self.spec.max_episodes.is_some() {
state.record_episode_summary(crate::spec::EpisodeSummary {
episode_index: record.as_ref().map_or(0, |record| record.index),
env_index,
seed,
step_count: completed.step_count,
cumulative_reward: completed.cumulative_reward,
terminated: completed.terminated,
truncated: completed.truncated,
duration_ms: (completed.end_timestamp_ns - completed.start_timestamp_ns).max(0)
/ 1_000_000,
success: success_from_final_info(completed.final_info.as_ref()),
});
}
fan_out_event!(
self,
episode_completed,
EpisodeCompletedEvent {
session_id: state.session_id().to_string(),
route: state.env_context(),
episode_id: completed.episode_id.clone(),
episode_record_id,
episode_index: record.as_ref().map_or(0, |record| record.index),
env_index,
step_count: completed.step_count,
cumulative_reward: completed.cumulative_reward,
terminated: completed.terminated,
truncated: completed.truncated,
duration_ms: (completed.end_timestamp_ns - completed.start_timestamp_ns).max(0)
/ 1_000_000,
final_info: completed.final_info.clone(),
seed,
}
);
}
}
async fn emit_reset_adapter(&mut self, state: &mut RouteState, episodes: &[EpisodeMetadata]) {
let slot_ids = state.episode_ids();
let episode_ids: Vec<String> = episodes
.iter()
.filter_map(|completed| slot_ids.get(completed.env_index as usize).cloned())
.filter(|id| !id.is_empty())
.collect();
if episode_ids.is_empty() {
return;
}
let request = state.reset_adapter_request(episode_ids);
if let Err(err) = self.model.reset_adapter(request).await {
tracing::warn!("model reset_adapter (evict) failed: {err}");
}
}
async fn invoke_transform_action(
&self,
telemetry: &Mutex<Aggregator>,
event: ActionReceivedEvent,
) -> Result<Option<Vec<Bytes>>, RuntimeError> {
let started = Instant::now();
let result = self.hooks.transform_action(event).await;
lock_agg(telemetry).record(Sample::dur(
SRC_TRANSFORM_ACTION,
metrics::RPC_TOTAL,
started.elapsed(),
));
match result {
Ok(action) => Ok(action),
Err(err) => {
tracing::warn!("runtime hook transform_action failed: {err}");
Err(RuntimeError::Hook(err))
}
}
}
async fn invoke_transform_observation(
&self,
telemetry: &Mutex<Aggregator>,
event: ObservationEmittedEvent,
) -> Result<Option<Vec<Bytes>>, RuntimeError> {
let started = Instant::now();
let result = self.hooks.transform_observation(event).await;
lock_agg(telemetry).record(Sample::dur(
SRC_TRANSFORM_OBS,
metrics::RPC_TOTAL,
started.elapsed(),
));
match result {
Ok(observation) => Ok(observation),
Err(err) => {
tracing::warn!("runtime hook transform_observation failed: {err}");
Err(RuntimeError::Hook(err))
}
}
}
fn observation_event(
&self,
state: &RouteState,
snapshot: RouteSnapshot,
is_reset: bool,
observation: Option<Vec<Bytes>>,
infos: Option<rlmesh_proto::spaces::v1::MetaMap>,
) -> ObservationEmittedEvent {
ObservationEmittedEvent {
session_id: state.session_id().to_string(),
route: state.env_context(),
episode_id: snapshot.episode_id,
episode_record_id: snapshot.episode_record_id,
episode_ids: snapshot.episode_ids,
episode_record_ids: snapshot.episode_record_ids,
step: snapshot.step,
env_index: snapshot.env_index,
is_reset,
num_envs: self.spec.num_envs as u32,
observation_space: Arc::clone(&self.observation_space),
raw_observation: observation.clone(),
observation,
infos,
}
}
}
fn deterministic_reset_seed(
base_seed: i64,
session_id: &str,
reset_generation: u64,
env_index: usize,
) -> i64 {
const FNV_OFFSET: u64 = 0xcbf2_9ce4_8422_2325;
const FNV_PRIME: u64 = 0x0000_0100_0000_01b3;
fn update(mut hash: u64, bytes: &[u8]) -> u64 {
for byte in bytes {
hash ^= u64::from(*byte);
hash = hash.wrapping_mul(FNV_PRIME);
}
hash
}
let mut hash = FNV_OFFSET;
hash = update(hash, &base_seed.to_le_bytes());
hash = update(hash, &[0xff]);
hash = update(hash, session_id.as_bytes());
hash = update(hash, &[0xfd]);
hash = update(hash, &reset_generation.to_le_bytes());
hash = update(hash, &[0xfc]);
hash = update(hash, &(env_index as u64).to_le_bytes());
(hash & i64::MAX as u64) as i64
}
type PrefetchOutcome = (
Box<dyn RuntimeModel + Send>,
Result<RuntimeModelPrediction, RuntimeError>,
);
struct PrefetchInflight {
join: tokio::task::JoinHandle<PrefetchOutcome>,
expected_context: Option<AdapterContext>,
request_bytes: u64,
}
async fn join_prefetch(
join: tokio::task::JoinHandle<PrefetchOutcome>,
model_component_id: &str,
) -> Result<PrefetchOutcome, RuntimeError> {
join.await.map_err(|error| {
RuntimeError::Protocol(format!(
"model endpoint {model_component_id}: background predict task failed: {error}"
))
})
}
async fn await_runtime_operation<T, F>(
cancellation: &CancellationToken,
timeout: Duration,
timeout_error: RuntimeError,
cancelled_error: RuntimeError,
operation: F,
) -> Result<T, RuntimeError>
where
F: Future<Output = Result<T, RuntimeError>>,
{
tokio::select! {
_ = cancellation.cancelled() => Err(cancelled_error),
result = tokio::time::timeout(timeout, operation) => match result {
Ok(result) => result,
Err(_) => Err(timeout_error),
},
}
}
fn log_model_close_result(
result: Result<Result<(), String>, tokio::time::error::Elapsed>,
reason: &str,
timeout: Duration,
) {
match result {
Ok(Err(err)) => {
tracing::warn!(
error = %err,
reason,
"model route close failed during route shutdown; relying on owner shutdown"
);
}
Err(_) => {
tracing::warn!(
timeout_ms = timeout.as_millis(),
reason,
"model route close timed out during route shutdown; relying on owner shutdown"
);
}
Ok(Ok(())) => {}
}
}
fn leaves_value(leaves: Vec<Bytes>) -> SpaceValue {
SpaceValue { leaves }
}
fn mint_episode_id() -> String {
uuid::Uuid::now_v7().to_string()
}
fn mint_episode_ids(count: usize) -> Vec<String> {
(0..count).map(|_| mint_episode_id()).collect()
}
fn episode_ids_with_roll(
mut ids: Vec<String>,
pending_roll: &std::collections::HashMap<u32, String>,
) -> Vec<String> {
for (env_index, new_id) in pending_roll {
if let Some(slot) = ids.get_mut(*env_index as usize) {
*slot = new_id.clone();
}
}
ids
}
fn value_leaves(payload: Option<&SpaceValue>) -> Result<Option<Vec<Bytes>>, RuntimeError> {
Ok(payload.map(|payload| payload.leaves.clone()))
}
fn lock_agg(telemetry: &Mutex<Aggregator>) -> MutexGuard<'_, Aggregator> {
telemetry
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
}
fn record_op(
telemetry: &Mutex<Aggregator>,
src: Source,
rpc: Duration,
peer: PeerReport,
request_bytes: u64,
response_bytes: u64,
) {
let mut agg = lock_agg(telemetry);
agg.record(Sample::dur(src, metrics::RPC_TOTAL, rpc));
if let Some(ns) = peer.endpoint_total_ns {
agg.record(Sample::dur(
src,
metrics::ENDPOINT_TOTAL,
Duration::from_nanos(ns),
));
}
for (metric, ns) in [
(metrics::ENDPOINT_DECODE, peer.phases.decode_ns),
(metrics::ENDPOINT_USER, peer.phases.user_ns),
(metrics::ENDPOINT_ENCODE, peer.phases.encode_ns),
(metrics::ENDPOINT_QUEUE, peer.phases.queue_ns),
(metrics::PREDICT_ADAPTER, peer.phases.adapter_ns),
] {
if ns != 0 {
agg.record(Sample::dur(src, metric, Duration::from_nanos(ns)));
}
}
if let Some(ns) = peer.phases.lane_skew_ns {
agg.record(Sample::dur(
src,
metrics::LANE_SKEW,
Duration::from_nanos(ns),
));
}
if peer.phases.in_flight != 0 {
agg.record(Sample::count(
src,
metrics::PREDICT_IN_FLIGHT,
u64::from(peer.phases.in_flight),
));
}
if let Some(episodes) = peer.phases.held_episodes {
agg.record(Sample::count(
src,
metrics::HELD_EPISODES,
u64::from(episodes),
));
}
if let Some(bytes) = peer.phases.held_state_bytes {
agg.record(Sample::bytes(src, metrics::HELD_BYTES, bytes));
}
agg.record(Sample::bytes(src, metrics::REQUEST_BYTES, request_bytes));
agg.record(Sample::bytes(src, metrics::RESPONSE_BYTES, response_bytes));
if let Some(group) = peer.group_size {
agg.record(Sample::count(src, metrics::GROUP_SIZE, group));
}
}
struct TelemetryTicker {
handle: tokio::task::JoinHandle<()>,
}
impl TelemetryTicker {
fn spawn(
telemetry: Arc<Mutex<Aggregator>>,
hooks: Arc<dyn RuntimeHooks>,
window: Duration,
session_id: String,
route: RuntimeEnvContext,
) -> Self {
let period = window.max(Duration::from_millis(1));
let handle = tokio::spawn(async move {
let mut ticker = tokio::time::interval(period);
ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
ticker.tick().await; loop {
ticker.tick().await;
let window_snap = {
let mut agg = lock_agg(&telemetry);
let snap = agg.snapshot(Horizon::Window);
agg.flush_window();
snap
};
if window_snap.rows.is_empty() {
continue;
}
let window_event = TelemetrySnapshotEvent {
session_id: session_id.clone(),
route: route.clone(),
snapshot: window_snap,
};
if let Err(err) = hooks.on_telemetry(window_event).await {
tracing::warn!("runtime hook on_telemetry (window) failed: {err}");
}
}
});
Self { handle }
}
}
impl Drop for TelemetryTicker {
fn drop(&mut self) {
self.handle.abort();
}
}