use std::sync::Arc;
use rlmesh_grpc::wire::{
decode_batched_partial_values, encode_batched_partial_values, env_spec_to_proto,
};
use rlmesh_proto::model::v1::{
AdapterContext, PredictRequest, ReleaseAdapterRequest, ResolveAdapterRequest,
};
use rlmesh_proto::{SessionOffer, supported_workflow_editions};
use uuid::Uuid;
use crate::{ConnectAddress, Error, Result, spaces};
pub struct RemoteModel {
inner: rlmesh_grpc::ModelClient,
observation_space: Arc<spaces::SpaceSpec>,
action_space: Arc<spaces::SpaceSpec>,
env_contract: spaces::EnvContract,
session_id: String,
env_id: String,
request_counter: u64,
configured: bool,
episode_id: Option<String>,
execution_horizon: u32,
replay_buffer: std::collections::VecDeque<spaces::SpaceValue>,
selected_workflow_edition: String,
}
fn new_session_id() -> String {
format!("remote-model-{}", Uuid::new_v4())
}
fn single_lane_contract(mut env_contract: spaces::EnvContract) -> Result<spaces::EnvContract> {
if env_contract.num_envs == 0 {
env_contract.num_envs = 1;
}
if env_contract.num_envs > 1 {
return Err(Error::Internal(format!(
"RemoteModel drives a single env, but the env contract reports num_envs={}; \
use num_envs=1 (this client sends one observation per predict on lane 0)",
env_contract.num_envs
)));
}
Ok(env_contract)
}
impl RemoteModel {
pub async fn connect(address: &str, env_contract: spaces::EnvContract) -> Result<Self> {
Self::connect_with_token(address, "", env_contract).await
}
pub async fn connect_with_token(
address: &str,
token: &str,
env_contract: spaces::EnvContract,
) -> Result<Self> {
let env_offer = SessionOffer {
editions: supported_workflow_editions(),
};
Self::connect_with_env_offer(address, token, env_contract, env_offer).await
}
pub async fn connect_with_env_offer(
address: &str,
token: &str,
env_contract: spaces::EnvContract,
env_offer: SessionOffer,
) -> Result<Self> {
let address = ConnectAddress::parse(address)?;
let observation_space = Arc::new(
env_contract
.observation_space
.clone()
.ok_or_else(|| Error::Internal("env contract missing observation_space".into()))?,
);
let action_space = Arc::new(
env_contract
.action_space
.clone()
.ok_or_else(|| Error::Internal("env contract missing action_space".into()))?,
);
let env_contract = single_lane_contract(env_contract)?;
let mut inner = rlmesh_grpc::ModelClient::connect(&address.to_string(), token)
.await
.map_err(Error::from)?;
inner.handshake().await.map_err(Error::from)?;
let model_offer = inner.model_session_offer();
let selected_workflow_edition = rlmesh_grpc::env_floor(&env_offer, &model_offer)
.map_err(Error::from)?
.selected_workflow_edition;
Ok(Self {
inner,
observation_space,
action_space,
env_contract,
session_id: new_session_id(),
env_id: crate::mint_id(),
request_counter: 0,
configured: false,
episode_id: None,
execution_horizon: 1,
replay_buffer: std::collections::VecDeque::new(),
selected_workflow_edition,
})
}
pub fn set_execution_horizon(&mut self, execution_horizon: u32) {
self.execution_horizon = execution_horizon;
}
pub fn selected_workflow_edition(&self) -> &str {
&self.selected_workflow_edition
}
pub fn env_id(&self) -> &str {
&self.env_id
}
pub fn address(&self) -> &str {
self.inner.address()
}
pub fn reset(&mut self) {
self.episode_id = Some(crate::mint_id());
self.replay_buffer.clear();
}
pub async fn predict(&mut self, observation: spaces::SpaceValue) -> Result<spaces::SpaceValue> {
let episode_id = self
.episode_id
.clone()
.ok_or_else(|| Error::Internal("call reset() before predict()".into()))?;
if !self.configured {
self.resolve_adapter().await?;
self.configured = true;
}
if self.replay_buffer.is_empty() {
let observation_value = encode_batched_partial_values(
std::slice::from_ref(&observation),
&self.observation_space,
)
.map_err(|error| Error::Internal(error.to_string()))?;
let request = PredictRequest {
context: Some(AdapterContext {
session_id: self.session_id.clone(),
env_id: self.env_id.clone(),
request_id: self.next_request_id(),
}),
observation: Some(observation_value),
episode_ids: vec![episode_id],
};
let response = self.inner.predict(request).await.map_err(Error::from)?;
if response.actions.is_empty() {
return Err(Error::Internal(
"predict returned no actions for a single-env route".to_string(),
));
}
for frame in &response.actions {
let mut frames = decode_batched_partial_values(Some(frame), &self.action_space, 1)
.map_err(|error| Error::Internal(error.to_string()))?;
if frames.len() != 1 {
return Err(Error::Internal(format!(
"predict frame decoded to {} actions for a single-env route",
frames.len()
)));
}
self.replay_buffer.push_back(frames.remove(0));
}
}
let action = self
.replay_buffer
.pop_front()
.expect("replay buffer is non-empty after a refill");
Ok(action)
}
pub async fn close(&mut self) -> Result<()> {
if !self.configured {
return Ok(());
}
self.inner
.release_adapter(ReleaseAdapterRequest {
context: Some(AdapterContext {
session_id: self.session_id.clone(),
env_id: self.env_id.clone(),
request_id: format!("{}:release_adapter", self.env_id),
}),
reason: "remote model session complete".to_string(),
})
.await
.map_err(Error::from)
}
async fn resolve_adapter(&mut self) -> Result<()> {
self.inner
.resolve_adapter(ResolveAdapterRequest {
context: Some(AdapterContext {
session_id: self.session_id.clone(),
env_id: self.env_id.clone(),
request_id: format!("{}:resolve_adapter", self.env_id),
}),
env_spec: Some(env_spec_to_proto(&self.env_contract)),
selected_workflow_edition: self.selected_workflow_edition.clone(),
execution_horizon: self.execution_horizon,
})
.await
.map_err(Error::from)
}
fn next_request_id(&mut self) -> String {
self.request_counter += 1;
format!("{}:predict:{}", self.env_id, self.request_counter)
}
}
#[cfg(test)]
mod tests {
use super::{new_session_id, single_lane_contract};
use crate::spaces;
fn single_lane_test_contract(num_envs: u32) -> spaces::EnvContract {
spaces::EnvContract {
id: "remote-model-test".to_string(),
autoreset_mode: Default::default(),
observation_space: None,
action_space: None,
metadata: None,
render_mode: String::new(),
num_envs,
}
}
#[test]
fn session_ids_are_unique_and_globally_namespaced() {
let first = new_session_id();
let second = new_session_id();
assert_ne!(first, second);
assert!(first.starts_with("remote-model-"), "{first}");
assert!(second.starts_with("remote-model-"), "{second}");
}
#[test]
fn single_lane_contract_rejects_vector_env() {
let err = single_lane_contract(single_lane_test_contract(4))
.expect_err("num_envs > 1 must be rejected");
assert!(err.to_string().contains("num_envs=4"), "{err}");
}
#[test]
fn single_lane_contract_clamps_zero_and_accepts_one() {
assert_eq!(
single_lane_contract(single_lane_test_contract(0))
.unwrap()
.num_envs,
1
);
assert_eq!(
single_lane_contract(single_lane_test_contract(1))
.unwrap()
.num_envs,
1
);
}
}