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 {
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(),
})
}
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(),
}
}
pub(super) fn model_error_from_error(error: &Error) -> join_response::Kind {
join_response::Kind::Error(model_error_value(error))
}
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()))?;
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,
}
}
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()
}
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(),
)
};
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()
);
assert!(check_actions_conform(&space, &[spaces::SpaceValue::Discrete(0)]).is_err());
assert!(check_actions_conform(&space, &[boxed(vec![0, 1], vec![2])]).is_err());
}
#[test]
fn handler_model_error_preserves_recoverability_on_the_wire() {
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);
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());
}
}