Skip to main content

rlmesh_runtime/
driver.rs

1//! The single-route run loop.
2//!
3//! Drives one ready model/env session: reset, then predict/step until a step or
4//! episode limit, cancellation, or failure. Records per-op telemetry and fans
5//! every state change out to the session's
6//! [`RuntimeHooks`](crate::hooks::RuntimeHooks).
7
8use 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
37/// Sends `$event` to the best-effort hook `$method`, logging any failure and
38/// keeping the route moving.
39macro_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
50/// The env's reset reply plus the endpoint-local op duration the peer stamped.
51pub struct RuntimeEnvReset {
52    pub response: ResetResponse,
53    /// Endpoint-local op duration (ns) from `JoinResponse.endpoint_total_ns`
54    /// (replaces the old nested per-step telemetry message).
55    pub endpoint_total_ns: Option<u64>,
56}
57
58/// The env's step reply plus the endpoint-local op duration the peer stamped.
59pub struct RuntimeEnvStep {
60    pub response: StepResponse,
61    /// Endpoint-local op duration (ns) from `JoinResponse.endpoint_total_ns`.
62    pub endpoint_total_ns: Option<u64>,
63}
64
65/// The model's predict reply plus the endpoint-local op duration the peer stamped.
66pub struct RuntimeModelPrediction {
67    pub response: PredictResponse,
68    /// Endpoint-local op duration (ns) from `JoinResponse.endpoint_total_ns`.
69    pub endpoint_total_ns: Option<u64>,
70}
71
72/// The environment side of a route: reset and step over the wire, plus close.
73#[async_trait]
74pub trait RuntimeEnv: Send {
75    /// Reset the requested lanes and return their initial observation.
76    async fn reset(&mut self, request: ResetRequest) -> Result<RuntimeEnvReset, RuntimeError>;
77
78    /// Advance the lanes one step under `request.action`.
79    async fn step(&mut self, request: StepRequest) -> Result<RuntimeEnvStep, RuntimeError>;
80
81    /// Release the env endpoint within `timeout`. Default no-op for an endpoint
82    /// the run does not own.
83    async fn close(&mut self, _timeout: Duration) -> Result<(), String> {
84        Ok(())
85    }
86}
87
88/// The model side of a route: predict, plus the per-episode adapter-state
89/// lifecycle (evict on episode end, release at session end).
90#[async_trait]
91pub trait RuntimeModel: Send {
92    /// Predict the ordered action frames for the batched observation (frame 0 is
93    /// this step; any further frames replay open-loop before the next call).
94    async fn predict(
95        &mut self,
96        request: PredictRequest,
97    ) -> Result<RuntimeModelPrediction, RuntimeError>;
98
99    /// Evict the model's per-episode adapter state (frame-stack buffers) for the
100    /// ended episodes. Best-effort GC, not a correctness gate: because episode
101    /// ids never repeat (UUIDv7), a dropped ResetAdapter only leaks memory and
102    /// can never alias a new episode. Default no-op for impls that hold no
103    /// per-episode state.
104    async fn reset_adapter(&mut self, _request: ResetAdapterRequest) -> Result<(), RuntimeError> {
105        Ok(())
106    }
107
108    /// Release the model endpoint within `timeout`. Default no-op.
109    async fn release_adapter(
110        &mut self,
111        _request: ReleaseAdapterRequest,
112        _timeout: Duration,
113    ) -> Result<(), String> {
114        Ok(())
115    }
116}
117
118/// Default reason attributed to a cancellation when the caller does not supply
119/// one via [`RuntimeDriver::run_with_cancellation_reason`].
120const DEFAULT_CANCELLATION_REASON: &str = "cancelled by caller";
121
122// Telemetry sources for the three driver ops. `component` is a coarse class
123// label — the serial single-route driver has one model + one env, and `op`
124// already distinguishes them (see telemetry::Source).
125const 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/// Drives one ready model/env session through its `reset -> predict -> step`
139/// loop. Inert until a `run*` method is awaited.
140#[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/observation space specs shared into every per-step hook event.
148    /// Populated once after [`validate`](RuntimeSessionSpec::validate) so the
149    /// hot path clones an `Arc` instead of deep-copying the spec each step.
150    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            // Filled from the validated spec at run time; default until then.
167            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    /// Deterministic seeds for a partial (`reset_subset`) reset, positionally
177    /// aligned to `env_indices`. Empty when no base seed is configured.
178    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    /// Deterministic per-lane reset seeds for `env_indices`, positionally
186    /// aligned. Empty when no base seed is configured.
187    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    /// Per-lane autoreset convention declared by the served env's contract.
208    /// `UNSPECIFIED` is treated as `DISABLED` (explicit reset only).
209    fn autoreset_mode(&self) -> AutoresetMode {
210        // Unknown modes are rejected at RuntimeSessionSpec::validate() (run before
211        // the loop); a value that still fails to decode falls back to the safe
212        // explicit-reset DISABLED rather than silently aliasing a newer mode.
213        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    /// Runs the session, attributing any cancellation of `cancellation` to
230    /// `reason`.
231    ///
232    /// The reason is carried into [`RuntimeError::RouteCancelled`], the
233    /// `session_failed` hook event, and the `CloseRoute` reason, so callers
234    /// (e.g. an owner that cancels for Ctrl+C, a deadline, or a sibling-route
235    /// failure) can supply an accurate cause instead of a hardcoded one.
236    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        // validate() confirmed both spaces are present; cache them as shared
244        // Arcs so per-step hook events clone a pointer, not the whole spec.
245        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        // Telemetry lives here, not in run_loop, so the final Session snapshot is
249        // delivered on EVERY exit (including aborts). The background ticker only
250        // ever pushes Window snapshots (the live tier); the cumulative Session
251        // total is pushed once below and returned on the report (the durable
252        // tier), so a late ticker tick cannot race or supersede it.
253        let telemetry = Arc::new(Mutex::new(Aggregator::default()));
254        // A zero window disables live streaming (it would otherwise be a 1ms hot
255        // loop); the final session push below still fires.
256        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        // Stop the ticker (it only emits Window snapshots, so it cannot contend
267        // this Session push), then deliver the durable session total exactly once
268        // on every exit path.
269        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    /// Session/route-level span (enabling-only): lets a closed-side OTel
287    /// subscriber attach to `rlmesh.route` later; inert under the default
288    /// subscriber. Created once per session, not per step; `skip_all` records only
289    /// the cheap ids. Any future per-step span MUST be trace-level + target-gated.
290    #[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        // Spec timeout getter returns a clamped-non-negative i64; proto field is uint64.
319        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        // The runtime is the sole id authority (R1): mint a fresh UUIDv7 per lane
322        // and push them DOWN so the env tags its episodes with our ids; we never
323        // read ids back from the env.
324        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        // Time only the RPC (after building the request), matching the predict /
334        // step / in-loop-reset sites so rpc.total is consistent across ops.
335        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        // Runtime-side mirror of the env's NEXT_STEP autoreset expectation: maps a
380        // lane that completed this step to the fresh UUIDv7 we minted for its next
381        // episode. Set on completion (step t), consumed on the autoreset roll
382        // (step t+1) — both the down-push to the env and our own slot roll read
383        // from here, so the env tags the rolled episode with exactly our id.
384        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        // Runtime-owned action-chunk replay buffer. A predict returns its ordered
401        // action frames in `PredictResponse.actions` (frame 0 = this step, frames
402        // 1.. = open-loop replay); the driver pushes every frame here (whole-batch
403        // frames, one `SpaceValue`'s leaves per step, covering every lane), pops one
404        // per step, and re-calls the model only when the buffer drains. A
405        // non-chunking predict returns exactly one frame, so the predict-every-step
406        // path is unchanged.
407        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            // Re-call model.predict only when the replay buffer is empty; otherwise
417            // replay a buffered chunk frame (no RPC, no obs re-encode). A predict
418            // returns its ordered frames in `actions` (frame 0 = this step, frames
419            // 1.. = open-loop replay): push every frame, then pop one. env.step +
420            // the observation transform/emit below run every iteration regardless,
421            // so a replay step still feeds the observation ledger.
422            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                // Push every ordered frame (frame 0 = this step, frames 1.. =
467                // open-loop replay); the pop below applies frame 0 now and the rest
468                // on the following steps without re-calling the model.
469                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            // Spec timeout getter returns a clamped-non-negative i64; proto field is uint64.
497            let step_timeout_ms = self.spec.limits.env_step_timeout_ms().max(0) as u64;
498            // Down-push the authoritative per-lane ids. For a NEXT_STEP autoreset
499            // roll (a lane in `pending_roll`), substitute the freshly minted id so
500            // the env tags its rolled episode with our id; the slot itself rolls to
501            // the same id below, after the step.
502            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            // Apply any NEXT_STEP autoreset roll the env just performed. The env
553            // rolled the lanes in `pending_roll` (from the previous step's
554            // completions) using the ids we pushed above; mirror that here by
555            // rolling our own slots to the same ids. The env no longer mints or
556            // returns ids — the runtime is authoritative.
557            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            // The observation_emitted hook fires once per observation actually
564            // sent to the model, post-transform, below (or at the initial
565            // reset). Emitting the raw step observation here would expose
566            // pre-transform bytes and, when the episode completes, an
567            // observation the model never sees.
568
569            self.emit_completed_episodes(state, &step_ok.response.completed_episodes)
570                .await;
571            // Tell the model to evict the ended episodes' frame-stack buffers
572            // (best-effort GC; ids never repeat so a miss only leaks memory).
573            self.emit_reset_adapter(state, &step_ok.response.completed_episodes)
574                .await;
575
576            // A lane that completed this step gets a fresh episode, so its buffered
577            // future actions are stale. The replay buffer holds whole-batch frames
578            // and cannot be partially invalidated, so flush it and re-plan on the
579            // next step (receding horizon on reset). No-op when not chunking.
580            if !step_ok.response.completed_episodes.is_empty() {
581                replay_buffer.clear();
582            }
583
584            // Under NEXT_STEP, a lane that completed this step (t) autoresets at
585            // t+1. Mint its next id now and stash it; the next step's down-push and
586            // slot roll both consume it. Mirrors the env's `expect_autoreset`.
587            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            // Under NEXT_STEP, the final episode completes at its done step `t`
599            // and this early-return fires before the `t+1` roll. So the
600            // model-side `on_episode_end` for the final episode is not delivered
601            // at `t+1`; instead it fires via the close-time `finish_lifecycle`
602            // sweep during shutdown. Asymmetric versus mid-run episodes, but the
603            // callback is not lost.
604            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                // The single final session push is delivered by the epilogue in
617                // run_with_cancellation_reason (on every exit path); here we only
618                // capture the durable pull snapshot for the returned report.
619                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            // Mode-aware next observation. The reflexive "any lane completed =>
641            // reset the whole vector" trigger is gone. That was the category
642            // error that cut healthy lanes short.
643            let (next_obs, phase, is_reset_msg) = match self.autoreset_mode() {
644                // NEXT_STEP (and the unreachable SAME_STEP): the env auto-resets a
645                // done lane itself and the rolled episode ids already arrived via
646                // observe_episode_ids above. The driver is purely observational;
647                // it never resets on the hot path.
648                AutoresetMode::NextStep | AutoresetMode::SameStep => (
649                    step_observation.clone(),
650                    RequestPhase::StepObservation,
651                    false,
652                ),
653                // DISABLED (and the single-env default): the env does not
654                // autoreset, so restart the lanes that just completed. When every
655                // lane completed this is a whole-vector reset (also the num_envs==1
656                // path); a strict subset uses a per-lane seeded reset_subset, the
657                // controlled / reproducible path.
658                AutoresetMode::Unspecified | AutoresetMode::Disabled => {
659                    // Proto env_index is uint32; thread it straight into the
660                    // uint32 ResetRequest.env_indices without a round-trip.
661                    let mut done_lanes: Vec<u32> = step_ok
662                        .response
663                        .completed_episodes
664                        .iter()
665                        .map(|metadata| metadata.env_index)
666                        .collect();
667                    // completed_episodes can carry duplicate env_index entries
668                    // (e.g. drained interrupted episodes), which would inflate
669                    // the lane count and misfire the whole_vector decision below.
670                    // Dedupe so the count reflects distinct lanes. Sorting is
671                    // safe: reset_subset_seeds is derived FROM done_lanes (so
672                    // seeds stay positionally aligned) and env_indices is
673                    // done_lanes.clone().
674                    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                        // Spec timeout getter returns a clamped-non-negative i64; proto field is uint64.
687                        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                        // Mint the authoritative ids for the lanes this reset
691                        // restarts (full-width for a whole reset, aligned to
692                        // done_lanes for a partial one) and push them DOWN.
693                        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                        // Whole-vector reset starts every lane; a partial reset
739                        // rolls only the lanes it restarted (build a full-width id
740                        // vector: current ids with the reset lanes replaced).
741                        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            // Emit the transformed observation actually sent to the model, for
767            // both step and reset observations, so hooks always see the same
768            // payload model.predict receives.
769            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        // The timeout is also forwarded to the impls, but the driver enforces
801        // it independently: a close impl that blocks (e.g. an RPC on a hung
802        // connection) without honoring the deadline must not be able to hang
803        // run()/run_with_cancellation() forever during shutdown.
804        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            // Proto env_index/duration_ms are uint32/uint64; events are i32/i64.
864            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    /// Tell the model to evict the ended episodes' per-episode adapter state.
887    /// Best-effort GC (R2): a failure is logged and the route keeps moving — a
888    /// missed evict only leaks model memory, never corrupts state (ids never
889    /// repeat). Skips the call when nothing completed.
890    ///
891    /// The id to evict is resolved POSITIONALLY from the runtime's own slot by
892    /// `env_index` (decision A) — the env's `completed_episodes[].episode_id`
893    /// echo is never trusted as the authority. At the completion step the slot
894    /// still holds the completing id (the autoreset roll lands at t+1), so this
895    /// is the id the model lazily seeded and must drop.
896    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
962/// The per-lane reset seed, derived purely from reproducible inputs: the user's
963/// base_seed, the session id, the reset generation, and the lane index. The
964/// container env_id is deliberately NOT mixed in — it is a per-attach random
965/// UUIDv7, so including it would make a base_seed non-reproducible across runs.
966fn 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
1041/// Mint one authoritative episode id. UUIDv7 is time-ordered (sortable by
1042/// creation) and never repeats, so a missed ResetAdapter can only leak memory —
1043/// never alias a fresh episode.
1044fn mint_episode_id() -> String {
1045    uuid::Uuid::now_v7().to_string()
1046}
1047
1048/// Mint `count` fresh episode ids (one per lane being started).
1049fn mint_episode_ids(count: usize) -> Vec<String> {
1050    (0..count).map(|_| mint_episode_id()).collect()
1051}
1052
1053/// Current per-lane ids with the pending NEXT_STEP autoreset rolls substituted
1054/// in. Lanes not rolling keep their current id; a rolling lane takes its freshly
1055/// minted next id. Used for both the env down-push and our own slot roll so they
1056/// stay byte-identical.
1057fn episode_ids_with_roll(
1058    mut ids: Vec<String>,
1059    pending_roll: &std::collections::HashMap<u32, String>,
1060) -> Vec<String> {
1061    // Consume the already-allocated id vector and roll in place; an empty
1062    // `pending_roll` (the steady-state common case) is then zero-copy.
1063    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
1071/// The relay is content-blind: it carries the peer's leaf vector through
1072/// unchanged (structure/dtype live in the route spec, never inline). Kept
1073/// returning `Result` so the existing `?` call sites are untouched.
1074fn value_leaves(payload: Option<&SpaceValue>) -> Result<Option<Vec<Bytes>>, RuntimeError> {
1075    Ok(payload.map(|payload| payload.leaves.clone()))
1076}
1077
1078/// Locks the telemetry aggregator, recovering from a poisoned mutex instead of
1079/// panicking. Telemetry is best-effort and must never take down the route, so a
1080/// panic under the guard degrades telemetry rather than killing the session.
1081fn lock_agg(telemetry: &Mutex<Aggregator>) -> MutexGuard<'_, Aggregator> {
1082    telemetry
1083        .lock()
1084        .unwrap_or_else(|poisoned| poisoned.into_inner())
1085}
1086
1087/// Record the four per-op telemetry samples — RPC latency, the optional
1088/// endpoint-local duration the peer stamped, and request + response wire bytes —
1089/// under a single lock. Every driver op (predict, step, reset) records this same
1090/// shape, so they all route through here.
1091fn 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
1112/// Background wall-clock telemetry emitter. On a fixed real-time cadence it
1113/// snapshots the aggregator's Window horizon and pushes it to the hooks — so live
1114/// Window deltas keep arriving even while the run loop is parked in a stalled
1115/// predict/step/reset (which a step-gated path cannot see). It does NOT push
1116/// Session snapshots: the cumulative session total is the durable tier, delivered
1117/// once by the run epilogue and on `RuntimeReport.telemetry`. Empty windows (no
1118/// samples since the last flush) are skipped. Aborts when the returned handle is
1119/// dropped; because it only ever emits Window snapshots, a late tick can never
1120/// race the epilogue's authoritative Session push.
1121struct 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        // The caller skips spawning for a zero window (disabled live streaming).
1134        // Defensive floor for any sub-ms value: interval panics on a zero period.
1135        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; // the first tick is immediate; skip it
1140            loop {
1141                ticker.tick().await;
1142                // Snapshot + clear the Window horizon under a scoped lock; the
1143                // guard is NEVER held across the await below (keeps the std Mutex
1144                // sound + the task future Send).
1145                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                // Nothing recorded this window — skip the push rather than emit an
1152                // empty snapshot to consumers.
1153                if window_snap.rows.is_empty() {
1154                    continue;
1155                }
1156                // Tag the snapshot with the route/session it belongs to (one
1157                // shared hooks instance serves all concurrent routes).
1158                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}