use std::sync::Arc;
use rlmesh_grpc::wire::Bytes;
use crate::spaces;
use crate::{Error, Result};
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct EpisodeInfo {
pub episode_id: String,
pub seed: Option<i64>,
}
pub fn predict_seed(episode_seed: i64, predict_index: u64) -> 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, &episode_seed.to_le_bytes());
hash = update(hash, &[0xfe]);
hash = update(hash, &predict_index.to_le_bytes());
(hash & u64::from(u32::MAX)) as i64
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct ModelRouteContext {
pub session_id: String,
pub env_id: String,
pub request_id: String,
pub episodes: Vec<EpisodeInfo>,
}
impl ModelRouteContext {
pub fn primary_episode_id(&self) -> &str {
self.episodes
.first()
.map(|episode| episode.episode_id.as_str())
.unwrap_or("")
}
pub fn episode_ids(&self) -> Vec<String> {
self.episodes
.iter()
.map(|episode| episode.episode_id.clone())
.collect()
}
}
#[derive(Debug, Clone, Default, PartialEq)]
pub struct ModelObservation {
pub observation: Option<Vec<Bytes>>,
pub route: ModelRouteContext,
pub num_envs: usize,
pub env_contract: Option<Arc<spaces::EnvContract>>,
pub history: Vec<HistoryFrame>,
pub step: Option<i64>,
}
#[derive(Debug, Clone, Default, PartialEq)]
pub struct HistoryFrame {
pub observation: Option<Vec<Bytes>>,
pub episodes: Vec<EpisodeInfo>,
pub step: i64,
}
impl ModelObservation {
pub fn episode_id(&self) -> &str {
self.route.primary_episode_id()
}
pub fn episode_ids(&self) -> Vec<String> {
self.route.episode_ids()
}
pub fn ensure_decodable(&self) -> Result<()> {
if self.observation.is_none() {
return Err(Error::model(
"observation absent; a predict request must carry an observation",
));
}
let contract = self
.env_contract
.as_ref()
.ok_or_else(|| Error::model("observation missing env contract; cannot decode"))?;
if contract.observation_space.is_none() {
return Err(Error::model("env contract missing observation space"));
}
Ok(())
}
pub fn decoded_lanes(&self) -> Result<Vec<spaces::SpaceValue>> {
let leaves = self.observation.as_ref().ok_or_else(|| {
Error::model("observation absent; cannot decode lanes (check is_some() first)")
})?;
self.decode_rows(leaves, self.num_envs)
}
pub fn decoded_history_lanes(&self, frame: &HistoryFrame) -> Result<Vec<spaces::SpaceValue>> {
let leaves = frame
.observation
.as_ref()
.ok_or_else(|| Error::model("history frame carries no observation"))?;
self.decode_rows(leaves, frame.episodes.len())
}
fn decode_rows(&self, leaves: &[Bytes], rows: usize) -> Result<Vec<spaces::SpaceValue>> {
let contract = self
.env_contract
.as_ref()
.ok_or_else(|| Error::model("observation missing env contract; cannot decode"))?;
let space = contract
.observation_space
.as_ref()
.ok_or_else(|| Error::model("env contract missing observation space"))?;
let value = rlmesh_grpc::wire::leaves_value(leaves.to_vec());
rlmesh_grpc::wire::decode_batched_partial_values(Some(&value), space, rows)
.map_err(|err| Error::model(err.to_string()))
}
pub fn decoded(&self) -> Result<spaces::SpaceValue> {
let mut lanes = self.decoded_lanes()?;
if lanes.len() != 1 {
return Err(Error::model(format!(
"decoded() requires a single-env observation, got {} lanes",
lanes.len()
)));
}
Ok(lanes.remove(0))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::spaces::{DType, SpaceValue, Tensor};
fn obs_space() -> spaces::SpaceSpec {
spaces::spaces::BoxSpaceBuilder::scalar(0.0, 255.0, vec![1])
.dtype(DType::Uint8)
.build()
.unwrap()
}
fn box_u8(v: u8) -> SpaceValue {
SpaceValue::Box(Tensor::from_vec(vec![v], vec![1], DType::Uint8).unwrap())
}
fn observation(values: &[SpaceValue]) -> ModelObservation {
let space = obs_space();
let wire = rlmesh_grpc::wire::encode_batched_partial_values(values, &space).unwrap();
let contract = spaces::EnvContract {
id: "T".into(),
action_space: None,
observation_space: Some(space),
metadata: None,
render_mode: String::new(),
num_envs: values.len() as u32,
autoreset_mode: Default::default(),
};
ModelObservation {
history: Vec::new(),
step: None,
observation: Some(wire.leaves),
route: ModelRouteContext::default(),
num_envs: values.len(),
env_contract: Some(Arc::new(contract)),
}
}
#[test]
fn predict_seed_law() {
assert_eq!(predict_seed(7, 3), predict_seed(7, 3));
assert_ne!(predict_seed(7, 0), predict_seed(7, 1));
assert_ne!(predict_seed(7, 0), predict_seed(8, 0));
for (seed, index) in [(0, 0), (-1, 0), (i64::MIN, u64::MAX), (i64::MAX, 1)] {
let derived = predict_seed(seed, index);
assert!((0..=i64::from(u32::MAX)).contains(&derived), "{derived}");
}
}
#[test]
fn decoded_lanes_roundtrips_each_lane() {
let lanes = vec![box_u8(5), box_u8(9), box_u8(0)];
assert_eq!(observation(&lanes).decoded_lanes().unwrap(), lanes);
}
#[test]
fn ensure_decodable_catches_malformed_requests_without_decoding() {
assert!(observation(&[box_u8(3)]).ensure_decodable().is_ok());
let mut no_obs = observation(&[box_u8(3)]);
no_obs.observation = None;
assert!(no_obs.ensure_decodable().is_err());
let mut no_contract = observation(&[box_u8(3)]);
no_contract.env_contract = None;
assert!(no_contract.ensure_decodable().is_err());
}
#[test]
fn decoded_requires_exactly_one_lane() {
assert_eq!(observation(&[box_u8(7)]).decoded().unwrap(), box_u8(7));
assert!(observation(&[box_u8(1), box_u8(2)]).decoded().is_err());
}
#[test]
fn absent_observation_errors_not_empty_vec() {
let mut obs = observation(&[box_u8(3)]);
obs.observation = None;
assert!(obs.decoded_lanes().is_err());
}
}