use std::sync::Arc;
use rlmesh_grpc::wire::Bytes;
use crate::spaces;
use crate::{Error, Result};
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct ModelRouteContext {
pub session_id: String,
pub env_id: String,
pub request_id: String,
pub episode_ids: Vec<String>,
}
impl ModelRouteContext {
pub fn primary_episode_id(&self) -> &str {
self.episode_ids.first().map(String::as_str).unwrap_or("")
}
pub fn episode_ids(&self) -> Vec<String> {
self.episode_ids.clone()
}
}
#[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>>,
}
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)")
})?;
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.clone());
rlmesh_grpc::wire::decode_batched_partial_values(Some(&value), space, self.num_envs)
.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 {
observation: Some(wire.leaves),
route: ModelRouteContext::default(),
num_envs: values.len(),
env_contract: Some(Arc::new(contract)),
}
}
#[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());
}
}