use std::collections::{HashMap, HashSet, VecDeque};
use std::future::Future;
use std::pin::Pin;
use std::sync::atomic::{AtomicBool, Ordering};
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::meta_value::Kind as MetaKind;
use rlmesh_proto::spaces::v1::{MetaList, MetaMap, MetaValue, SpaceValue};
use tokio::task::JoinSet;
use tokio_util::sync::CancellationToken;
use crate::hooks::{
ActionReceivedEvent, EpisodeCompletedEvent, EpisodeStartedEvent, LogEvent, LogLevel,
ObservationEmittedEvent, RuntimeEnvContext, RuntimeHooks, SessionEndedEvent,
SessionFailedEvent, SessionStartedEvent, StepCompletedEvent, TelemetrySnapshotEvent,
};
use crate::spec::{ENV_RESET_OPTIONS_KEY, 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 + Sync {
async fn predict(
&self,
request: PredictRequest,
) -> Result<RuntimeModelPrediction, RuntimeError>;
async fn predict_group(
&self,
requests: Vec<PredictRequest>,
) -> Vec<Result<RuntimeModelPrediction, RuntimeError>> {
futures::future::join_all(requests.into_iter().map(|request| self.predict(request))).await
}
async fn reset_adapter(&self, _request: ResetAdapterRequest) -> Result<(), RuntimeError> {
Ok(())
}
async fn release_adapter(
&self,
_request: ReleaseAdapterRequest,
_timeout: Duration,
) -> Result<(), String> {
Ok(())
}
}
pub trait PredictScheduler: Send {
fn plan(&mut self, waiting: &[usize], busy: usize) -> Vec<usize>;
}
pub struct EagerScheduler;
impl PredictScheduler for EagerScheduler {
fn plan(&mut self, waiting: &[usize], _busy: usize) -> Vec<usize> {
waiting.to_vec()
}
}
const DEFAULT_CANCELLATION_REASON: &str = "cancelled by caller";
const DEFAULT_MAX_EPISODE_STEPS: i64 = 100_000;
const TRIAL_INDEX_OPTION: &str = "trial_index";
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: Option<M>,
prefetch_lead: u32,
scheduler: Box<dyn PredictScheduler>,
hooks: Arc<dyn RuntimeHooks>,
cancellation_reason: String,
action_space: Arc<rlmesh_proto::spaces::v1::SpaceSpec>,
observation_space: Arc<rlmesh_proto::spaces::v1::SpaceSpec>,
pending_evictions: Vec<String>,
trial_options_warned: AtomicBool,
vector_replay_warned: bool,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum EnvPhase {
Ready,
Stepping,
Resetting,
Idle,
}
enum PredictState {
None,
Wanted(PredictRequest),
InFlight { stale: bool },
Ready(VecDeque<Vec<Bytes>>),
}
struct Group<E> {
lanes: Vec<u32>,
positions: Vec<usize>,
lane_group: bool,
whole: bool,
phase: EnvPhase,
predict: PredictState,
env: Option<E>,
replay: VecDeque<Vec<Bytes>>,
obs_msg: Option<PredictRequest>,
pending_roll: HashMap<u32, String>,
pending_start: Option<(Vec<String>, Vec<u64>)>,
reset_generation: u64,
round_started: Instant,
}
impl<E> Group<E> {
fn width(&self) -> usize {
self.lanes.len()
}
fn busy(&self) -> bool {
matches!(self.phase, EnvPhase::Stepping | EnvPhase::Resetting)
}
}
enum EnvOutcome<E> {
Reset {
group: usize,
env: E,
initial: bool,
request_bytes: u64,
rpc: Duration,
result: Result<RuntimeEnvReset, RuntimeError>,
},
Step {
group: usize,
env: E,
request_bytes: u64,
rpc: Duration,
result: Result<RuntimeEnvStep, RuntimeError>,
},
}
type PredictFuture<'m, M> = Pin<Box<dyn Future<Output = PredictOutcome<M>> + Send + 'm>>;
struct PredictOutcome<M> {
model: M,
requests: Vec<(usize, Option<AdapterContext>, u64)>,
rpc: Duration,
result: Result<Vec<Result<RuntimeModelPrediction, RuntimeError>>, RuntimeError>,
}
impl<E, M> RuntimeDriver<E, M>
where
E: RuntimeEnv + Clone + 'static,
M: RuntimeModel,
{
pub fn new(spec: RuntimeSessionSpec, env: E, model: M, hooks: Arc<dyn RuntimeHooks>) -> Self {
Self {
spec,
env,
model: Some(model),
prefetch_lead: 0,
scheduler: Box::new(EagerScheduler),
hooks,
cancellation_reason: DEFAULT_CANCELLATION_REASON.to_string(),
action_space: Arc::default(),
observation_space: Arc::default(),
pending_evictions: Vec::new(),
trial_options_warned: AtomicBool::new(false),
vector_replay_warned: false,
}
}
fn planned_trial_indices(&self, state: &mut RouteState, lanes: usize) -> Vec<u64> {
match self.spec.trial_index_base {
Some(base) => state.claim_trial_indices(base, lanes),
None => Vec::new(),
}
}
fn trial_options(&self, trials: &[u64]) -> Option<MetaMap> {
if trials.is_empty() {
return None;
}
if !self.env_declares_trial_index() {
if !self.trial_options_warned.swap(true, Ordering::Relaxed) {
tracing::warn!(
env_id = %self.spec.env_id,
key = ENV_RESET_OPTIONS_KEY,
option = TRIAL_INDEX_OPTION,
"trial_index_base is set but the env contract declares no such reset \
option; the ordinal is recorded on the episode events and summaries \
but not delivered to the env",
);
}
return None;
}
let value = if trials.len() == 1 {
MetaValue {
kind: Some(MetaKind::Integer(trials[0] as i64)),
}
} else {
MetaValue {
kind: Some(MetaKind::List(MetaList {
items: trials
.iter()
.map(|trial| MetaValue {
kind: Some(MetaKind::Integer(*trial as i64)),
})
.collect(),
})),
}
};
Some(MetaMap {
entries: [(TRIAL_INDEX_OPTION.to_string(), value)]
.into_iter()
.collect(),
})
}
fn env_declares_trial_index(&self) -> bool {
let Some(declared) = self
.spec
.env_contract
.spec
.as_ref()
.and_then(|spec| spec.metadata.as_ref())
.and_then(|metadata| metadata.entries.get(ENV_RESET_OPTIONS_KEY))
.and_then(|declared| declared.kind.as_ref())
else {
return false;
};
match declared {
MetaKind::List(list) => list.items.iter().any(|item| {
matches!(item.kind.as_ref(), Some(MetaKind::Text(key)) if key == TRIAL_INDEX_OPTION)
}),
MetaKind::Text(key) => key == TRIAL_INDEX_OPTION,
_ => false,
}
}
pub fn with_prefetch(mut self, lead: u32) -> Self {
self.prefetch_lead = lead;
self
}
pub fn with_scheduler(mut self, scheduler: Box<dyn PredictScheduler>) -> Self {
self.scheduler = scheduler;
self
}
fn autoreset_mode(&self) -> AutoresetMode {
AutoresetMode::try_from(self.spec.env_contract.autoreset_mode)
.unwrap_or(AutoresetMode::Disabled)
}
fn driver_owns_resets(&self) -> bool {
matches!(
self.autoreset_mode(),
AutoresetMode::Disabled | AutoresetMode::Unspecified
)
}
fn groups(&self) -> Vec<Group<E>> {
let num_envs = self.spec.num_envs.max(1);
let partitions: Vec<Vec<u32>> = if self.spec.subset_step && num_envs > 1 {
(0..num_envs as u32).map(|lane| vec![lane]).collect()
} else {
vec![(0..num_envs as u32).collect()]
};
let lane_group = self.spec.subset_step && num_envs > 1;
partitions
.into_iter()
.map(|lanes| Group {
positions: lanes.iter().map(|&lane| lane as usize).collect(),
whole: !lane_group,
lane_group,
lanes,
phase: EnvPhase::Ready,
predict: PredictState::None,
env: Some(self.env.clone()),
replay: VecDeque::new(),
obs_msg: None,
pending_roll: HashMap::new(),
pending_start: None,
reset_generation: 0,
round_started: Instant::now(),
})
.collect()
}
fn seeds_for(&self, group: &Group<E>, slots: &[u64]) -> Vec<i64> {
if !self.spec.episode_seeds.is_empty() {
let seeds: Vec<Option<i64>> = slots
.iter()
.map(|&slot| {
self.spec
.episode_seeds
.get(usize::try_from(slot).unwrap_or(usize::MAX))
.copied()
})
.collect();
return if seeds.iter().all(Option::is_some) {
seeds.into_iter().flatten().collect()
} else {
if seeds.iter().any(Option::is_some) {
tracing::warn!(
lanes = group.width(),
"episode_seeds cannot cover this reset batch; it runs unseeded"
);
}
Vec::new()
};
}
let Some(base_seed) = self.spec.base_seed else {
return Vec::new();
};
if group.lane_group {
slots
.iter()
.map(|&slot| deterministic_reset_seed(base_seed, &self.spec.session_id, slot, 0))
.collect()
} else {
group
.lanes
.iter()
.map(|&lane| {
deterministic_reset_seed(
base_seed,
&self.spec.session_id,
group.reset_generation,
lane as usize,
)
})
.collect()
}
}
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 mut env_ops: JoinSet<EnvOutcome<E>> = JoinSet::new();
let mut predict: Option<PredictFuture<'_, M>> = None;
let result = self
.run_loop(
&mut state,
&cancellation,
&telemetry,
&mut env_ops,
&mut predict,
)
.await;
env_ops.abort_all();
if let Some(inflight) = predict.take() {
match tokio::time::timeout(self.spec.limits.service_close_timeout, inflight).await {
Ok(outcome) => self.model = Some(outcome.model),
Err(_) => tracing::warn!(
"grouped predict still in flight at route end; model release skipped"
),
}
}
self.flush_evictions(&mut state).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.clone(),
}
);
match result {
Ok(reason) => {
let release_request = state.release_adapter_request(reason);
self.shutdown_terminal_route(&state, reason, release_request)
.await;
fan_out_event!(
self,
session_ended,
SessionEndedEvent {
session_id: state.session_id().to_string(),
route: state.env_context(),
reason: reason.to_string(),
total_steps: state.total_steps(),
total_episodes: state.total_episodes(),
}
);
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: final_snapshot,
})
}
Err(error) => {
self.shutdown_after_failure(&mut state, &error).await;
Err(error)
}
}
}
#[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,
lanes = self.spec.subset_step,
),
)]
async fn run_loop<'m>(
&mut self,
state: &mut RouteState,
cancellation: &CancellationToken,
telemetry: &Arc<Mutex<Aggregator>>,
env_ops: &mut JoinSet<EnvOutcome<E>>,
predict: &mut Option<PredictFuture<'m, M>>,
) -> Result<&'static str, RuntimeError>
where
M: 'm,
{
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 groups = self.groups();
let group_count = groups.len();
for gid in 0..group_count {
self.begin_reset(gid, &mut groups, state, env_ops, true);
}
loop {
if cancellation.is_cancelled() {
return Err(self.cancelled_error(state, &groups));
}
if groups.iter().all(|group| group.phase == EnvPhase::Idle)
&& env_ops.is_empty()
&& predict.is_none()
{
return Ok("completed requested episodes");
}
self.flush_evictions(state).await;
self.dispatch_steps(&mut groups, state, env_ops, telemetry)
.await?;
let force = env_ops.is_empty() && predict.is_none();
self.dispatch_predict(&mut groups, state, predict, force);
if env_ops.is_empty() && predict.is_none() {
return Err(RuntimeError::Protocol(format!(
"route {} stalled with no operation in flight",
state.env_id()
)));
}
tokio::select! {
_ = cancellation.cancelled() => {
return Err(self.cancelled_error(state, &groups));
}
Some(joined) = env_ops.join_next() => {
let outcome = joined.map_err(|error| RuntimeError::Protocol(format!(
"env operation task failed: {error}"
)))?;
self.on_env_outcome(outcome, &mut groups, state, env_ops, telemetry)
.await?;
}
outcome = async {
match predict.as_mut() {
Some(inflight) => inflight.as_mut().await,
None => std::future::pending().await,
}
}, if predict.is_some() => {
*predict = None;
self.on_predict_outcome(outcome, &mut groups, state, telemetry)?;
}
}
}
}
fn begin_reset(
&mut self,
gid: usize,
groups: &mut [Group<E>],
state: &mut RouteState,
env_ops: &mut JoinSet<EnvOutcome<E>>,
initial: bool,
) {
let width = groups[gid].width();
let bounded = groups[gid].lane_group;
let Some(slots) = state.claim_slots(width, bounded) else {
groups[gid].phase = EnvPhase::Idle;
groups[gid].predict = PredictState::None;
return;
};
if !initial {
groups[gid].reset_generation += 1;
}
let seeds = self.seeds_for(&groups[gid], &slots);
let episode_ids = mint_episode_ids(width);
state.note_episode_seeds(&episode_ids, &seeds);
let trials = self.planned_trial_indices(state, width);
state.note_episode_trials(&episode_ids, &trials);
let group = &mut groups[gid];
group.pending_start = Some((episode_ids.clone(), slots));
group.replay.clear();
group.pending_roll.clear();
let request = ResetRequest {
seeds,
options: self.trial_options(&trials),
timeout_ms: self.spec.limits.env_reset_timeout_ms().max(0) as u64,
env_indices: if group.whole {
Vec::new()
} else {
group.lanes.clone()
},
episode_ids,
};
let request_bytes = request.encoded_len() as u64;
let timeout = self.spec.limits.env_reset_timeout;
let timeout_error = RuntimeError::operation_timeout(
state.env_id(),
state.env_component_id(),
"env.reset",
0,
timeout,
);
let mut env = group
.env
.take()
.expect("group env handle present while ready");
group.phase = EnvPhase::Resetting;
env_ops.spawn(async move {
let started = Instant::now();
let result = match tokio::time::timeout(timeout, env.reset(request)).await {
Ok(result) => result,
Err(_) => Err(timeout_error),
};
EnvOutcome::Reset {
group: gid,
env,
initial,
request_bytes,
rpc: started.elapsed(),
result,
}
});
}
#[allow(clippy::needless_range_loop)]
async fn dispatch_steps(
&mut self,
groups: &mut [Group<E>],
state: &mut RouteState,
env_ops: &mut JoinSet<EnvOutcome<E>>,
telemetry: &Arc<Mutex<Aggregator>>,
) -> Result<(), RuntimeError> {
for gid in 0..groups.len() {
if groups[gid].phase != EnvPhase::Ready {
continue;
}
if groups[gid].replay.is_empty() {
match std::mem::replace(&mut groups[gid].predict, PredictState::None) {
PredictState::Ready(frames) => groups[gid].replay = frames,
other => {
groups[gid].predict = other;
continue;
}
}
}
let Some(model_action) = groups[gid].replay.pop_front() else {
continue;
};
if self.prefetch_lead > 0
&& groups[gid].replay.len() <= self.prefetch_lead as usize
&& matches!(groups[gid].predict, PredictState::None)
&& let Some(msg) = groups[gid].obs_msg.clone()
{
groups[gid].predict = PredictState::Wanted(msg);
}
let group = &mut groups[gid];
let snapshot = state.snapshot_at(&group.positions);
let context = state.group_context(&group.positions, group.lane_group);
let action_step = snapshot.step + 1;
let mut action_event = ActionReceivedEvent {
session_id: state.session_id().to_string(),
route: context,
episode_id: snapshot.episode_id.clone(),
episode_record_id: snapshot.episode_record_id.clone(),
episode_ids: snapshot.episode_ids.clone(),
episode_record_ids: snapshot.episode_record_ids.clone(),
step: action_step,
env_index: 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 group = &mut groups[gid];
let episode_ids =
episode_ids_with_roll(state.episode_ids_at(&group.positions), &group.pending_roll);
let request = StepRequest {
action: action_event.action.map(leaves_value),
timeout_ms: self.spec.limits.env_step_timeout_ms().max(0) as u64,
env_indices: if group.whole {
Vec::new()
} else {
group.lanes.clone()
},
episode_ids,
};
let request_bytes = request.encoded_len() as u64;
let timeout = self.spec.limits.env_step_timeout;
let timeout_error = RuntimeError::operation_timeout(
state.env_id(),
state.env_component_id(),
"env.step",
action_step,
timeout,
);
let mut env = group
.env
.take()
.expect("group env handle present while ready");
group.phase = EnvPhase::Stepping;
env_ops.spawn(async move {
let started = Instant::now();
let result = match tokio::time::timeout(timeout, env.step(request)).await {
Ok(result) => result,
Err(_) => Err(timeout_error),
};
EnvOutcome::Step {
group: gid,
env,
request_bytes,
rpc: started.elapsed(),
result,
}
});
}
Ok(())
}
fn dispatch_predict<'m>(
&mut self,
groups: &mut [Group<E>],
state: &RouteState,
predict: &mut Option<PredictFuture<'m, M>>,
force: bool,
) where
M: 'm,
{
if predict.is_some() || self.model.is_none() {
return;
}
let waiting: Vec<usize> = groups
.iter()
.enumerate()
.filter(|(_, group)| matches!(group.predict, PredictState::Wanted(_)))
.map(|(gid, _)| gid)
.collect();
if waiting.is_empty() {
return;
}
let busy = groups.iter().filter(|group| group.busy()).count();
let mut chosen = self.scheduler.plan(&waiting, busy);
chosen.retain(|gid| waiting.contains(gid));
if chosen.is_empty() {
if !force {
return;
}
chosen = waiting;
}
let mut requests = Vec::with_capacity(chosen.len());
let mut metas = Vec::with_capacity(chosen.len());
let mut step = 0;
for gid in chosen {
let PredictState::Wanted(msg) = std::mem::replace(
&mut groups[gid].predict,
PredictState::InFlight { stale: false },
) else {
unreachable!("only waiting groups are planned")
};
step = step.max(state.snapshot_at(&groups[gid].positions).step);
metas.push((gid, msg.context.clone(), msg.encoded_len() as u64));
requests.push(msg);
}
let timeout = self.spec.limits.model_predict_timeout;
let timeout_error = RuntimeError::operation_timeout(
state.env_id(),
state.model_component_id(),
"model.predict",
step,
timeout,
);
let model = self.model.take().expect("model handle checked above");
*predict = Some(Box::pin(async move {
let started = Instant::now();
let result = match tokio::time::timeout(timeout, model.predict_group(requests)).await {
Ok(results) => Ok(results),
Err(_) => Err(timeout_error),
};
PredictOutcome {
model,
requests: metas,
rpc: started.elapsed(),
result,
}
}));
}
fn on_predict_outcome(
&mut self,
outcome: PredictOutcome<M>,
groups: &mut [Group<E>],
state: &RouteState,
telemetry: &Arc<Mutex<Aggregator>>,
) -> Result<(), RuntimeError> {
self.model = Some(outcome.model);
let results = outcome.result?;
if results.len() != outcome.requests.len() {
return Err(RuntimeError::Protocol(format!(
"model endpoint {} answered {} of {} grouped predicts",
state.model_component_id(),
results.len(),
outcome.requests.len()
)));
}
let group_count = outcome.requests.len() as u64;
let mut recorded = false;
for ((gid, expected_context, request_bytes), result) in
outcome.requests.into_iter().zip(results)
{
let prediction = result?;
if prediction.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,
});
}
if !recorded {
recorded = true;
record_op(
telemetry,
SRC_PREDICT,
outcome.rpc,
PeerReport {
endpoint_total_ns: prediction.endpoint_total_ns,
phases: prediction.phases,
group_size: if group_count > 1 {
Some(group_count)
} else {
prediction.group_size
},
},
request_bytes,
prediction.response.encoded_len() as u64,
);
}
if prediction.response.actions.is_empty() {
return Err(RuntimeError::Protocol(format!(
"model endpoint {} returned a predict response with no actions",
state.model_component_id()
)));
}
let mut frames = VecDeque::with_capacity(prediction.response.actions.len());
for frame in &prediction.response.actions {
if let Some(leaves) = value_leaves(Some(frame))? {
frames.push_back(leaves);
}
}
let group = &mut groups[gid];
group.predict = match std::mem::replace(&mut group.predict, PredictState::None) {
PredictState::InFlight { stale: true } => PredictState::None,
PredictState::InFlight { stale: false } => {
if group.replay.is_empty() {
group.replay = frames;
PredictState::None
} else {
PredictState::Ready(frames)
}
}
other => other,
};
}
Ok(())
}
async fn on_env_outcome(
&mut self,
outcome: EnvOutcome<E>,
groups: &mut [Group<E>],
state: &mut RouteState,
env_ops: &mut JoinSet<EnvOutcome<E>>,
telemetry: &Arc<Mutex<Aggregator>>,
) -> Result<(), RuntimeError> {
match outcome {
EnvOutcome::Reset {
group: gid,
env,
initial,
request_bytes,
rpc,
result,
} => {
groups[gid].env = Some(env);
let reset = result?;
record_op(
telemetry,
SRC_RESET,
rpc,
PeerReport {
endpoint_total_ns: reset.endpoint_total_ns,
phases: reset.phases,
group_size: None,
},
request_bytes,
reset.response.encoded_len() as u64,
);
let group = &mut groups[gid];
let (episode_ids, slots) = group
.pending_start
.take()
.expect("a reset in flight has its episodes staged");
let context = state.group_context(&group.positions, group.lane_group);
if initial {
fan_out_event!(
self,
log,
LogEvent {
session_id: state.session_id().to_string(),
route: context.clone(),
level: LogLevel::Info,
message: format!(
"env reset complete in {:.0}ms ({} episode(s) ready)",
rpc.as_secs_f64() * 1000.0,
episode_ids.len()
),
source: Some("runtime".to_string()),
}
);
}
let started =
state.start_episodes_at(&groups[gid].positions, episode_ids, !initial, &slots);
self.invoke_started_episodes(state, &context, started).await;
groups[gid].round_started = Instant::now();
let observation = value_leaves(reset.response.observation.as_ref())?;
self.observe(
gid,
groups,
state,
telemetry,
observation,
reset.response.infos,
RequestPhase::ResetObservation,
true,
)
.await
}
EnvOutcome::Step {
group: gid,
env,
request_bytes,
rpc,
result,
} => {
groups[gid].env = Some(env);
let step = result?;
record_op(
telemetry,
SRC_STEP,
rpc,
PeerReport {
endpoint_total_ns: step.endpoint_total_ns,
phases: step.phases,
group_size: None,
},
request_bytes,
step.response.encoded_len() as u64,
);
self.on_step(gid, groups, state, env_ops, telemetry, step.response)
.await
}
}
}
#[allow(clippy::too_many_arguments)]
async fn on_step(
&mut self,
gid: usize,
groups: &mut [Group<E>],
state: &mut RouteState,
env_ops: &mut JoinSet<EnvOutcome<E>>,
telemetry: &Arc<Mutex<Aggregator>>,
response: StepResponse,
) -> Result<(), RuntimeError> {
let positions = groups[gid].positions.clone();
let lane_group = groups[gid].lane_group;
let context = state.group_context(&positions, lane_group);
lock_agg(telemetry).record(Sample::dur(
SRC_ROUND,
metrics::RPC_TOTAL,
groups[gid].round_started.elapsed(),
));
groups[gid].round_started = Instant::now();
groups[gid].phase = EnvPhase::Ready;
let step_observation = value_leaves(response.observation.as_ref())?;
state.record_step_at(&positions, &response.rewards);
let snapshot = state.snapshot_at(&positions);
let rolled = !groups[gid].pending_roll.is_empty();
fan_out_event!(
self,
step_completed,
StepCompletedEvent {
session_id: state.session_id().to_string(),
route: context.clone(),
episode_id: snapshot.episode_id.clone(),
episode_record_id: snapshot.episode_record_id.clone(),
step: snapshot.step,
env_index: snapshot.env_index,
rewards: response.rewards.clone(),
infos: if rolled { None } else { response.infos.clone() },
}
);
if rolled {
let pending_roll = std::mem::take(&mut groups[gid].pending_roll);
let roll_ids = episode_ids_with_roll(state.episode_ids_at(&positions), &pending_roll);
let rolling: Vec<Option<u64>> = groups[gid]
.lanes
.iter()
.map(|lane| {
pending_roll
.contains_key(lane)
.then(|| state.claim_slots(1, false))
.flatten()
.and_then(|slots| slots.first().copied())
})
.collect();
let started = state.observe_episode_ids_at(&positions, roll_ids, &rolling);
self.invoke_started_episodes(state, &context, started).await;
}
let capped = self.capped_completions(state, &positions, &response.completed_episodes);
let mut completed_episodes = response.completed_episodes.clone();
completed_episodes.extend(capped);
self.emit_completed_episodes(state, &context, &completed_episodes)
.await;
self.queue_evictions(state, &completed_episodes);
if !completed_episodes.is_empty() {
let group = &mut groups[gid];
if !group.replay.is_empty() && group.width() > 1 && !self.vector_replay_warned {
self.vector_replay_warned = true;
tracing::warn!(
num_envs = group.width(),
discarded_frames = group.replay.len(),
"a lane's episode ended mid-chunk on a lockstep vector group: chunk replay \
is whole-batch, so every lane's buffered frames were discarded and the \
group re-plans. Serve lanes, or use an execution horizon of 1.",
);
}
group.replay.clear();
group.predict = match std::mem::replace(&mut group.predict, PredictState::None) {
PredictState::InFlight { .. } => PredictState::InFlight { stale: true },
_ => PredictState::None,
};
if !self.driver_owns_resets() {
for completed in &completed_episodes {
group
.pending_roll
.entry(completed.env_index)
.or_insert_with(mint_episode_id);
}
}
}
if !lane_group
&& self
.spec
.max_episodes
.is_some_and(|limit| state.total_episodes() >= limit as i64)
{
groups[gid].phase = EnvPhase::Idle;
groups[gid].predict = PredictState::None;
return Ok(());
}
if self.driver_owns_resets() {
let done: HashSet<u32> = completed_episodes
.iter()
.map(|metadata| metadata.env_index)
.collect();
if !done.is_empty() {
self.begin_reset(gid, groups, state, env_ops, false);
return Ok(());
}
}
self.observe(
gid,
groups,
state,
telemetry,
step_observation,
if rolled { response.infos } else { None },
RequestPhase::StepObservation,
false,
)
.await
}
#[allow(clippy::too_many_arguments)]
async fn observe(
&mut self,
gid: usize,
groups: &mut [Group<E>],
state: &mut RouteState,
telemetry: &Arc<Mutex<Aggregator>>,
observation: Option<Vec<Bytes>>,
infos: Option<rlmesh_proto::spaces::v1::MetaMap>,
phase: RequestPhase,
is_reset: bool,
) -> Result<(), RuntimeError> {
let positions = groups[gid].positions.clone();
let context = state.group_context(&positions, groups[gid].lane_group);
let mut msg = state.predict_request_at(&positions, observation.clone(), phase);
let mut event = self.observation_event(
state,
context,
state.snapshot_at(&positions),
is_reset,
observation,
infos,
groups[gid].width(),
);
let transformed = self
.invoke_transform_observation(telemetry, event.clone())
.await?;
event.observation = transformed.clone();
msg.observation = transformed.map(leaves_value);
fan_out_event!(self, observation_emitted, event);
let group = &mut groups[gid];
group.obs_msg = Some(msg.clone());
group.phase = EnvPhase::Ready;
if group.replay.is_empty() && matches!(group.predict, PredictState::None) {
group.predict = PredictState::Wanted(msg);
}
Ok(())
}
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 = async {
match self.model.as_ref() {
Some(model) => {
tokio::time::timeout(timeout, model.release_adapter(request, timeout)).await
}
None => Ok(Ok(())),
}
};
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, groups: &[Group<E>]) -> RuntimeError {
let step = groups
.iter()
.map(|group| state.snapshot_at(&group.positions).step)
.max()
.unwrap_or(0);
RuntimeError::route_cancelled(state.env_id(), step, self.cancellation_reason.as_str())
}
async fn invoke_started_episodes(
&self,
state: &RouteState,
context: &RuntimeEnvContext,
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: context.clone(),
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),
trial_index: state.trial_for_episode(&episode.episode_id),
}
);
}
}
fn capped_completions(
&self,
state: &RouteState,
positions: &[usize],
env_completed: &[EpisodeMetadata],
) -> Vec<EpisodeMetadata> {
let step_cap = self.spec.max_episode_steps.or_else(|| {
self.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_at(positions)
.into_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,
context: &RuntimeEnvContext,
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);
let trial_index = state.trial_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,
trial_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,
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: context.clone(),
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,
trial_index,
}
);
}
}
fn queue_evictions(&mut self, state: &RouteState, episodes: &[EpisodeMetadata]) {
let all_positions: Vec<usize> = (0..self.spec.num_envs.max(1)).collect();
let slot_ids = state.episode_ids_at(&all_positions);
self.pending_evictions.extend(
episodes
.iter()
.filter_map(|completed| {
state
.slot_position(completed.env_index)
.and_then(|position| slot_ids.get(position))
.cloned()
})
.filter(|id| !id.is_empty()),
);
}
async fn flush_evictions(&mut self, state: &mut RouteState) {
if self.pending_evictions.is_empty() {
return;
}
let Some(model) = self.model.as_ref() else {
return;
};
let episode_ids = std::mem::take(&mut self.pending_evictions);
let request = state.reset_adapter_request(episode_ids);
if let Err(err) = 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))
}
}
}
#[allow(clippy::too_many_arguments)]
fn observation_event(
&self,
state: &RouteState,
route: RuntimeEnvContext,
snapshot: RouteSnapshot,
is_reset: bool,
observation: Option<Vec<Bytes>>,
infos: Option<rlmesh_proto::spaces::v1::MetaMap>,
width: usize,
) -> ObservationEmittedEvent {
ObservationEmittedEvent {
session_id: state.session_id().to_string(),
route,
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: width 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
}
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: &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();
}
}