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::collections::HashSet;
use std::time::Instant;

use rlmesh_grpc::wire::value_leaves;
use rlmesh_proto::model::v1::{
    AdapterContext, ModelError, ModelErrorCode, PredictRequest, PredictResponse, join_response,
};
use rlmesh_proto::spaces::v1::SpaceValue;

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

#[derive(Debug, Clone, PartialEq)]
pub(super) struct ModelAction {
    /// Finished ordered action frames built by the worker (codec already ran):
    /// `actions[0]` is THIS step, `actions[1..]` the open-loop chunk replay
    /// frames. Exactly one element when the predict is not chunking.
    pub(super) actions: Vec<SpaceValue>,
    pub(super) route: ModelRouteContext,
}

pub(super) fn model_error(message: impl Into<String>) -> join_response::Kind {
    join_response::Kind::Error(ModelError {
        code: ModelErrorCode::InvalidRequest as i32,
        message: message.into(),
        is_recoverable: false,
        debug_info: String::new(),
    })
}

/// Build the wire [`ModelError`] for a facade [`Error`] returned by a handler,
/// preserving the handler-fault vs internal-fault distinction and the
/// recoverable flag. Used directly by the grouped-predict path (which carries a
/// `ModelError` per group inside `GroupedPredictResult`) and wrapped by
/// [`model_error_from_error`] for the single-predict `JoinResponse.error` arm.
pub(super) fn model_error_value(error: &Error) -> ModelError {
    let (code, is_recoverable) = match error {
        Error::Model(model) => (ModelErrorCode::Internal, model.is_recoverable),
        _ => (ModelErrorCode::Internal, false),
    };
    ModelError {
        code: code as i32,
        message: error.to_string(),
        is_recoverable,
        debug_info: String::new(),
    }
}

/// Map a facade [`Error`] returned by a handler onto the wire model-error,
/// preserving the handler-fault vs internal-fault distinction and the
/// recoverable flag so the caller can react appropriately.
pub(super) fn model_error_from_error(error: &Error) -> join_response::Kind {
    join_response::Kind::Error(model_error_value(error))
}

/// Endpoint-local op duration in nanoseconds for the per-step
/// `JoinResponse.endpoint_total_ns` scalar. Replaces the old nested per-step
/// telemetry message construction (and with it the dead `labels` map and the
/// always-empty `component_id` — the runtime attributes by connection).
/// Saturates at `u64::MAX`.
pub(super) fn model_endpoint_total_ns(started_at: Instant) -> u64 {
    started_at.elapsed().as_nanos().min(u128::from(u64::MAX)) as u64
}

pub(super) fn model_observation_from_endpoint_request(
    request: PredictRequest,
) -> Result<ModelObservation> {
    let context = request
        .context
        .ok_or_else(|| Error::Internal("model request missing adapter context".to_string()))?;
    // The self-describing batch: episode_info rides the PredictRequest, not the
    // context. Build the route context from it.
    let episodes = request
        .episode_info
        .into_iter()
        .map(|info| crate::model::types::EpisodeInfo {
            episode_id: info.episode_id,
            seed: info.seed,
        })
        .collect();
    let route = ModelRouteContext {
        session_id: context.session_id,
        env_id: context.env_id,
        request_id: context.request_id,
        episodes,
    };
    validate_predict_route(&route)?;
    let num_envs = route.episodes.len();
    let history = request
        .history
        .into_iter()
        .map(|frame| crate::model::types::HistoryFrame {
            observation: value_leaves(frame.observation.as_ref()).map(<[_]>::to_vec),
            episodes: frame
                .episode_info
                .into_iter()
                .map(|info| crate::model::types::EpisodeInfo {
                    episode_id: info.episode_id,
                    seed: info.seed,
                })
                .collect(),
            step: frame.step,
        })
        .collect();

    Ok(ModelObservation {
        observation: value_leaves(request.observation.as_ref()).map(<[_]>::to_vec),
        num_envs,
        env_contract: None,
        route,
        history,
        step: request.step,
    })
}

pub(super) fn model_action_to_endpoint_response(action: ModelAction) -> PredictResponse {
    PredictResponse {
        context: Some((&action.route).into()),
        actions: action.actions,
    }
}

/// Encode the per-step chunk replay frames (each a per-lane action batch) into
/// wire `SpaceValue`s — one per future step — validating each against the route
/// action space. These are `PredictResponse.actions[1..]` for a chunked predict;
/// an empty `replay` yields an empty vec (not chunking).
pub(super) fn encode_replay_frames(
    replay: &[Vec<spaces::SpaceValue>],
    num_envs: usize,
    action_space: &spaces::SpaceSpec,
) -> Result<Vec<SpaceValue>> {
    replay
        .iter()
        .map(|frame| {
            if frame.len() != num_envs {
                return Err(Error::model(format!(
                    "chunk replay frame produced {} actions for {num_envs} lanes",
                    frame.len()
                )));
            }
            check_actions_conform(action_space, frame)?;
            rlmesh_grpc::wire::encode_batched_partial_values(frame, action_space)
                .map_err(|err| Error::model(err.to_string()))
        })
        .collect()
}

/// Structurally validate each per-lane action against the route's action space
/// before the codec encodes it. The spec-directed codec would otherwise silently
/// drop or reinterpret a mismatched typed action (extra Dict keys are skipped by
/// the spec-key walk; a wrong-dtype Box leaf is emitted as raw bytes and read
/// back at the spec dtype), and that value/spec mismatch never reaches the env's
/// own validation. Range deviations (Box bounds) pass through — those are the
/// env's validation policy to decide.
pub(super) fn check_actions_conform(
    action_space: &spaces::SpaceSpec,
    actions: &[spaces::SpaceValue],
) -> Result<()> {
    for (lane, action) in actions.iter().enumerate() {
        if let spaces::Conformance::Structural(err) = spaces::conform(action_space, action) {
            return Err(Error::model(format!(
                "model action for lane {lane} does not match the action space: {err}"
            )));
        }
    }
    Ok(())
}

fn validate_predict_route(route: &ModelRouteContext) -> Result<()> {
    if route.env_id.is_empty() {
        return Err(Error::Internal("model env_id is empty".to_string()));
    }
    if route.request_id.is_empty() {
        return Err(Error::Internal("model request_id is empty".to_string()));
    }
    if route.episodes.is_empty() {
        return Err(Error::Internal(
            "model predict must include at least one episode_id".to_string(),
        ));
    }

    let mut seen = HashSet::new();
    for (index, episode) in route.episodes.iter().enumerate() {
        let episode_id = &episode.episode_id;
        if episode_id.is_empty() {
            return Err(Error::Internal(format!(
                "model predict episode_info[{index}].episode_id is empty"
            )));
        }
        if !seen.insert(episode_id.as_str()) {
            return Err(Error::Internal(format!(
                "model predict has duplicate episode_id {episode_id:?}"
            )));
        }
    }

    Ok(())
}

impl From<&ModelRouteContext> for AdapterContext {
    fn from(value: &ModelRouteContext) -> Self {
        Self {
            session_id: value.session_id.clone(),
            env_id: value.env_id.clone(),
            request_id: value.request_id.clone(),
        }
    }
}

#[cfg(test)]
mod tests {
    use rlmesh_proto::model::v1::EpisodeInfo;

    use super::*;

    fn unwrap_error(kind: join_response::Kind) -> ModelError {
        match kind {
            join_response::Kind::Error(error) => error,
            other => panic!("expected model error, got {other:?}"),
        }
    }

    #[test]
    fn check_actions_conform_rejects_structural_mismatch() {
        let space = spaces::spaces::BoxSpaceBuilder::scalar(0.0, 1.0, vec![1])
            .dtype(spaces::DType::Uint8)
            .build()
            .unwrap();
        let boxed = |data: Vec<u8>, shape: Vec<i64>| {
            spaces::SpaceValue::Box(
                spaces::Tensor::from_vec(data, shape, spaces::DType::Uint8).unwrap(),
            )
        };

        // Matching kind/shape/dtype passes (one lane and many).
        assert!(check_actions_conform(&space, &[boxed(vec![0], vec![1])]).is_ok());
        assert!(
            check_actions_conform(&space, &[boxed(vec![0], vec![1]), boxed(vec![1], vec![1])])
                .is_ok()
        );

        // A Discrete value for a Box space is a structural mismatch.
        assert!(check_actions_conform(&space, &[spaces::SpaceValue::Discrete(0)]).is_err());
        // A wrong-shape Box would otherwise mis-encode -> rejected.
        assert!(check_actions_conform(&space, &[boxed(vec![0, 1], vec![2])]).is_err());
    }

    #[test]
    fn handler_model_error_preserves_recoverability_on_the_wire() {
        // A recoverable handler decline must surface as a recoverable model
        // error, not a non-recoverable internal/transport fault.
        let recoverable = unwrap_error(model_error_from_error(&Error::model_recoverable(
            "retry me",
        )));
        assert!(recoverable.is_recoverable);
        assert_eq!(recoverable.code, ModelErrorCode::Internal as i32);
        assert!(recoverable.message.contains("retry me"));

        let permanent = unwrap_error(model_error_from_error(&Error::model("bad observation")));
        assert!(!permanent.is_recoverable);

        // A genuine internal fault is never reported as recoverable.
        let internal = unwrap_error(model_error_from_error(&Error::Internal("boom".to_string())));
        assert!(!internal.is_recoverable);
    }

    #[test]
    fn model_observation_carries_episode_seed_from_episode_info() {
        let request = PredictRequest {
            history: Vec::new(),
            step: None,
            context: Some(AdapterContext {
                session_id: "s".to_string(),
                env_id: "e".to_string(),
                request_id: "r".to_string(),
            }),
            observation: None,
            episode_info: vec![
                EpisodeInfo {
                    episode_id: "ep-1".to_string(),
                    seed: Some(7),
                },
                EpisodeInfo {
                    episode_id: "ep-2".to_string(),
                    seed: None,
                },
            ],
        };
        let observation = model_observation_from_endpoint_request(request).unwrap();
        assert_eq!(observation.route.episode_ids(), ["ep-1", "ep-2"]);
        assert_eq!(
            observation
                .route
                .episodes
                .iter()
                .map(|episode| episode.seed)
                .collect::<Vec<_>>(),
            [Some(7), None]
        );
    }

    #[test]
    fn model_observation_rejects_blank_episode_id_in_episode_info() {
        let request = PredictRequest {
            history: Vec::new(),
            step: None,
            context: Some(AdapterContext {
                session_id: "s".to_string(),
                env_id: "e".to_string(),
                request_id: "r".to_string(),
            }),
            observation: None,
            episode_info: vec![EpisodeInfo {
                episode_id: String::new(),
                seed: None,
            }],
        };
        assert!(model_observation_from_endpoint_request(request).is_err());
    }
}