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 route = ModelRouteContext {
session_id: context.session_id,
env_id: context.env_id,
request_id: context.request_id,
episode_ids: request.episode_ids,
};
validate_predict_route(&route)?;
let num_envs = route.episode_ids.len();
Ok(ModelObservation {
observation: value_leaves(request.observation.as_ref()).map(<[_]>::to_vec),
num_envs,
env_contract: None,
route,
})
}
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.episode_ids.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_id) in route.episode_ids.iter().enumerate() {
if episode_id.is_empty() {
return Err(Error::Internal(format!(
"model predict episode_ids[{index}] 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 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);
}
}