rlmesh 0.1.0

Internal RLMesh crate (unstable Rust API): Rust bindings for model-environment evaluation; build on the rlmesh Python package.
Documentation
use std::sync::Arc;

use async_trait::async_trait;

use super::types::ModelObservation;
use crate::{Result, spaces};

/// Resolves per-env adapter state (e.g. an env→model adapter) when an adapter is
/// resolved.
///
/// Obtained once from [`ModelHandler::route_setup`] when serving begins and
/// shared (`Arc`) across every env, so the server can run it at `ResolveAdapter`
/// **without** taking the predict-serialization lock: resolving one env's adapter
/// never blocks on an in-flight `predict` on another. Implementations must
/// therefore synchronize their own state. Resolution happens before any `predict`
/// on the env, and per-env ordering guarantees an adapter is fully resolved
/// before that env's first predict — so an adapter is never re-resolved while its
/// own predict is in flight.
#[async_trait]
pub trait ModelRouteSetup: Send + Sync {
    /// Resolve and cache the adapter for `env_id` from its `env_contract`, and
    /// answer what the resolved route needs back.
    /// Returning an error fails adapter resolution, so the client never predicts
    /// against an unresolved adapter. Idempotent upsert: a later call updates it.
    async fn resolve_adapter(
        &self,
        env_id: &str,
        env_contract: &spaces::EnvContract,
        options: ResolveOptions,
    ) -> Result<RouteNeeds>;

    /// Drop the route-local per-episode adapter state (an engine's episode-keyed
    /// frame windows) for `env_id`'s ended `episode_ids` at `ResetAdapter` —
    /// empty means all of the env's. Runs **off** the predict lock, before
    /// [`ModelHandler::reset_adapter`] fires the model's own hook under it, so
    /// a forward in flight on another env does not gate the cleanup. Per-env
    /// ordering still holds: this never overlaps a predict on the same env.
    /// Defaults to a no-op.
    async fn reset_adapter(&self, _env_id: &str, _episode_ids: &[String]) -> Result<()> {
        Ok(())
    }

    /// Tear down the adapter cached for `env_id` at `ReleaseAdapter`, so a
    /// long-lived server does not retain per-env state for every session it ever
    /// served. Defaults to a no-op.
    async fn release_adapter(&self, _env_id: &str) -> Result<()> {
        Ok(())
    }
}

/// The runtime-pinned knobs a route resolves against, carried on
/// `ResolveAdapter`.
///
/// Runtime-local scheduling decisions, not part of the model spec: the model
/// reads them, it does not choose them.
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct ResolveOptions {
    /// How many actions of each predicted chunk the runtime executes before
    /// re-planning (1 = no chunking). The model returns its NATIVE chunk and the
    /// runtime executes a prefix of it, so an autoregressive head can read this
    /// to decode exactly that many. Bounded by
    /// [`MAX_EXECUTION_HORIZON`](rlmesh_adapters::v1::MAX_EXECUTION_HORIZON).
    pub execution_horizon: u32,
    /// Whether the runtime will deliver observation-history frames on `Predict`
    /// (every env step it executed from a replayed chunk without predicting).
    /// A runtime that does not is refused a stacked adapter at a horizon above
    /// 1, since the window would hold only decision-point frames.
    pub delivers_history: bool,
}

/// What a resolved route needs from the runtime, answered on
/// `ResolveAdapterResponse`.
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct RouteNeeds {
    /// The model's native chunk length K: how many per-step actions one chunk
    /// corner call returns. `None` = undeclared (the elastic contract — the
    /// runtime takes the `min(len, horizon)` prefix of whatever comes back).
    pub native_chunk: Option<u32>,
    /// Set when the runtime offered observation history
    /// ([`ResolveOptions::delivers_history`]) and this route's adapter keeps a
    /// frame window: the runtime must then deliver every replayed step as a
    /// history row on the next predict, stamped with steps.
    pub history: Option<HistoryNeeds>,
}

/// The frame windows a route keeps, answered on `ResolveAdapterResponse.history`.
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct HistoryNeeds {
    /// Canonical placement keys of the stacked inputs (informational: the
    /// runtime sends whole observations per history row).
    pub keys: Vec<String>,
    /// Whether the runtime may send pruned rows (only the keys' leaves). Always
    /// `false` today.
    pub prunable: bool,
}

/// One predict's per-lane action plus any open-loop chunk replay frames.
///
/// `actions` is frame 0 — one action per lane (`len == num_envs`), applied this
/// step. `replay` is the future-step frames the runtime buffers and applies
/// WITHOUT re-calling the model (action chunking): `replay[j]` is the per-lane
/// actions for future step `j + 1` (each `len == num_envs`). An empty `replay`
/// means the model is not chunking — one action this step, re-plan next step.
#[derive(Debug)]
pub struct PredictFrames {
    /// Frame 0: one action per lane, applied this step.
    pub actions: Vec<spaces::SpaceValue>,
    /// Future-step frames (`replay[step][lane]`); empty when not chunking.
    pub replay: Vec<Vec<spaces::SpaceValue>>,
}

/// Engine-held per-episode adapter state at a model endpoint, as stamped on
/// predict responses: episodes with live frame-stack windows and the bytes
/// those windows hold.
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct HeldState {
    /// Episodes with live frame-stack windows across all routes.
    pub episodes: u64,
    /// Bytes those windows hold.
    pub bytes: u64,
}

/// User policy plus episode lifecycle hooks.
///
/// Implement [`predict`](ModelHandler::predict) to map an observation to encoded
/// action bytes. The default hooks let stateful policies track resets, episode
/// ends, and shutdown. Drive the handler with
/// [`ModelWorker::run_local`](crate::ModelWorker::run_local) or host it with
/// [`ModelWorker::serve`](crate::ModelWorker::serve).
#[async_trait]
pub trait ModelHandler: Send {
    /// Produce an action for `observation`.
    ///
    /// Read the observation with
    /// [`decoded_lanes`](ModelObservation::decoded_lanes) (or
    /// [`decoded`](ModelObservation::decoded) for a single-env route) and return
    /// **one typed action per row** — `Vec` length `== observation.num_envs`
    /// (`== route.episode_ids.len()`); a single-env route returns a 1-element
    /// `Vec`. The codec turns the typed values into wire leaves; policy code
    /// never touches bytes.
    ///
    /// Return [`Error::model`](crate::Error::model) or
    /// [`Error::model_recoverable`](crate::Error::model_recoverable) when the
    /// policy declines a request.
    ///
    /// # Concurrency contract (pipelined predict)
    ///
    /// The served model endpoint pipelines Join-stream requests, so responses
    /// may complete out of arrival order. The handler itself is never invoked
    /// concurrently: `predict` and the lifecycle hooks take `&mut self`, and
    /// the server holds a per-handler mutex across each call.
    ///
    /// Per-route lifecycle order is preserved. Calls for different routes may
    /// interleave, still one at a time. The `model.concurrent_predict.v1`
    /// handshake capability advertises this pipelining to clients.
    async fn predict(&mut self, observation: ModelObservation) -> Result<Vec<spaces::SpaceValue>>;

    /// Produce an action for `observation` plus any open-loop chunk replay frames
    /// (action chunking).
    ///
    /// The default emits no replay frames — behaviorally identical to
    /// [`predict`](ModelHandler::predict) — so a non-chunking handler needs no
    /// change. The stateful engine overrides this to split a chunked policy's
    /// output into per-step frames. The runtime driver calls this and replays the
    /// frames itself; the served endpoint packs them into the ordered
    /// `PredictResponse.actions` list (frame 0 first, replay frames after).
    /// `PredictFrames::actions` keeps the same `== num_envs` length contract as
    /// `predict`.
    async fn predict_chunked(&mut self, observation: ModelObservation) -> Result<PredictFrames> {
        Ok(PredictFrames {
            actions: self.predict(observation).await?,
            replay: Vec::new(),
        })
    }

    /// Produce action frames for a batch of routed observations in one call.
    ///
    /// The server calls this for a `GroupedPredictRequest` — a control-plane-
    /// grouped batch where each observation belongs to a *different* configured
    /// route (and so a different env spec/adapter). The default fans out to
    /// [`predict_chunked`](ModelHandler::predict_chunked) per group,
    /// sequentially, which is behaviorally identical to handling each group as
    /// its own predict — including action chunking: a group whose route pinned
    /// an `execution_horizon > 1` returns frame 0 plus replay frames, and each
    /// group's horizon stays its own (routes pin independently at
    /// `ResolveAdapter`). A handler overrides this to fuse the groups into ONE
    /// forward pass (e.g. a single batched GPU inference across env types) —
    /// this is the only seam a fusing model must implement.
    ///
    /// The returned `Vec` aligns 1:1 and in order with `observations`; each
    /// element is that group's own `Result`, so one group's failure is reported
    /// per-group and never sinks the others. An override MUST preserve that
    /// length and order. A non-chunking group returns an empty
    /// [`PredictFrames::replay`], keeping `PredictFrames::actions`'s
    /// `== num_envs` lane contract per group.
    async fn predict_grouped(
        &mut self,
        observations: Vec<ModelObservation>,
    ) -> Vec<Result<PredictFrames>> {
        let mut results = Vec::with_capacity(observations.len());
        for observation in observations {
            results.push(self.predict_chunked(observation).await);
        }
        results
    }

    /// Engine-held per-episode adapter state (frame-stack windows) at this
    /// endpoint, summed across routes — the state that grows with concurrent
    /// episodes and shrinks on episode-end GC. `None` for a handler that keeps
    /// no such accounting; `Some` with zeros when it does and holds nothing.
    fn held_state(&self) -> Option<HeldState> {
        None
    }

    /// Adapter time (obs assembly + action apply) inside the last
    /// predict-family call, in nanoseconds — the share of the handler's own
    /// work that was RLMesh adapter transform rather than the model's forward.
    /// Read-and-clear: the caller drains it once per request. Defaults to `0`
    /// for a handler that does not measure it.
    fn take_adapter_ns(&mut self) -> u64 {
        0
    }

    /// Per-route setup invoked at `ResolveAdapter`, before any `predict` on the
    /// route. Returns a cheaply-cloned, independently-synchronized handle (or
    /// `None` for no per-route setup), obtained once when serving begins so the
    /// server runs it **off** the predict-serialization lock — see
    /// [`ModelRouteSetup`].
    ///
    /// A spec'd model returns a setup that resolves and caches its env→model
    /// adapter per route (from the contract's spaces and adapter tags) for
    /// `predict` to apply. Defaults to `None`. The served path runs it at
    /// `ResolveAdapter`; the in-process
    /// [`run_local`](crate::ModelWorker::run_local) path runs it once at
    /// connect, before driving.
    fn route_setup(&self) -> Option<Arc<dyn ModelRouteSetup>> {
        None
    }

    /// Drop per-episode policy/adapter state for the given episodes (the explicit
    /// `ResetAdapter` op). The runtime fires this when episodes end, keyed by
    /// `env_id`; with `episode_ids` empty it means "evict ALL of this env's
    /// episode state". The adapter itself stays resolved (see
    /// [`ModelRouteSetup::release_adapter`] for full teardown).
    ///
    /// Defaults to a no-op. Because episode ids never repeat (UUIDv7), a missed
    /// `reset_adapter` only leaks memory — it can never alias a new episode — so
    /// a stateful policy lazy-seeds per-episode state on first `predict` and
    /// evicts it here, with no position-diffing.
    ///
    /// This runs under the predict lock, so it is the place for a hook that
    /// must not overlap a forward. State a [`ModelRouteSetup`] owns per route
    /// is dropped through [`ModelRouteSetup::reset_adapter`] instead, off the
    /// lock and before this fires.
    async fn reset_adapter(&mut self, _env_id: &str, _episode_ids: Vec<String>) -> Result<()> {
        Ok(())
    }

    /// Called once when the worker/session shuts down. Defaults to a no-op.
    async fn on_close(&mut self) -> Result<()> {
        Ok(())
    }
}