use std::collections::HashMap;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Duration;
use async_trait::async_trait;
use rlmesh_proto::model::v1::{
AdapterContext, CloseParticipantRequest, EpisodeInfo, GroupedPredictRequest,
GroupedPredictResponse, GroupedPredictResult, JoinRequest, PredictRequest,
ReleaseAdapterRequest, ResetAdapterRequest, ResolveAdapterRequest, grouped_predict_result,
join_request, join_response,
};
use tokio::net::TcpListener;
use tokio::sync::{Mutex, mpsc};
use tokio_stream::wrappers::TcpListenerStream;
use super::server::{Admission, ModelRouteConfig, handle_model_request};
use super::*;
use crate::{BindAddress, ConnectAddress, Result, ServeOptions, spaces};
struct SmokeEnv {
obs_space: spaces::SpaceSpec,
action_space: spaces::SpaceSpec,
env_contract: spaces::EnvContract,
reset_seeds: Option<Arc<Mutex<Vec<Option<i64>>>>>,
}
impl SmokeEnv {
fn recording(sink: Arc<Mutex<Vec<Option<i64>>>>) -> Self {
Self {
reset_seeds: Some(sink),
..Self::new()
}
}
fn new() -> Self {
let obs_space = spaces::spaces::BoxSpaceBuilder::scalar(0.0, 255.0, vec![1])
.dtype(spaces::DType::Uint8)
.build()
.unwrap();
let action_space = spaces::spaces::BoxSpaceBuilder::scalar(0.0, 1.0, vec![1])
.dtype(spaces::DType::Uint8)
.build()
.unwrap();
let env_contract = spaces::EnvContract {
id: "SmokeEnv-v0".to_string(),
autoreset_mode: Default::default(),
observation_space: Some(obs_space.clone()),
action_space: Some(action_space.clone()),
metadata: None,
render_mode: String::new(),
num_envs: 1,
};
Self {
obs_space,
action_space,
env_contract,
reset_seeds: None,
}
}
}
#[async_trait]
impl crate::Env for SmokeEnv {
fn observation_space(&self) -> &spaces::SpaceSpec {
&self.obs_space
}
fn action_space(&self) -> &spaces::SpaceSpec {
&self.action_space
}
fn env_contract(&self) -> &spaces::EnvContract {
&self.env_contract
}
async fn reset(
&mut self,
req: spaces::request::ResetRequest,
) -> std::result::Result<spaces::request::ResetResult, spaces::EnvRuntimeError> {
if let Some(seeds) = &self.reset_seeds {
seeds.lock().await.push(req.seed);
}
Ok(spaces::request::ResetResult {
observation: Some(spaces::SpaceValue::Box(
spaces::Tensor::from_vec(vec![0], vec![1], spaces::DType::Uint8).unwrap(),
)),
info: None,
episode_id: Some("ep-smoke".to_string()),
})
}
async fn step(
&mut self,
_req: spaces::request::StepRequest,
) -> std::result::Result<spaces::request::StepResult, spaces::EnvRuntimeError> {
Ok(spaces::request::StepResult {
observation: Some(spaces::SpaceValue::Box(
spaces::Tensor::from_vec(vec![1], vec![1], spaces::DType::Uint8).unwrap(),
)),
reward: 1.0,
terminated: true,
truncated: false,
info: None,
})
}
async fn render(
&mut self,
_req: spaces::RenderRequest,
) -> std::result::Result<spaces::RenderResult, spaces::EnvRuntimeError> {
Ok(spaces::RenderResult::default())
}
async fn close(
&mut self,
_req: spaces::CloseRequest,
) -> std::result::Result<spaces::request::CloseResult, spaces::EnvRuntimeError> {
Ok(spaces::request::CloseResult)
}
}
struct SmokeModel {
predicts: Arc<AtomicUsize>,
closes: Arc<AtomicUsize>,
}
#[async_trait]
impl ModelHandler for SmokeModel {
async fn predict(&mut self, observation: ModelObservation) -> Result<Vec<spaces::SpaceValue>> {
self.predicts.fetch_add(1, Ordering::SeqCst);
Ok((0..observation.num_envs)
.map(|_| {
spaces::SpaceValue::Box(
spaces::Tensor::from_vec(vec![0u8], vec![1], spaces::DType::Uint8).unwrap(),
)
})
.collect())
}
async fn on_close(&mut self) -> Result<()> {
self.closes.fetch_add(1, Ordering::SeqCst);
Ok(())
}
}
fn spawn_bound_server(bound: BoundModelServer) -> (u16, tokio::task::JoinHandle<Result<()>>) {
let port = match bound.local_addr().clone() {
BindAddress::Tcp { port, .. } => port,
other => panic!("expected tcp bind address, got {other:?}"),
};
assert_ne!(port, 0, "port 0 must resolve to a real port");
let server = tokio::spawn(async move { bound.serve().await });
(port, server)
}
async fn shutdown_and_join(server: tokio::task::JoinHandle<Result<()>>) {
tokio::time::timeout(Duration::from_secs(2), server)
.await
.unwrap()
.unwrap()
.unwrap();
}
#[tokio::test]
async fn user_set_base_seed_reaches_the_env_reset_seeds() {
async fn run_with_base_seed(base_seed: Option<i64>) -> Vec<Option<i64>> {
let reset_seeds = Arc::new(Mutex::new(Vec::new()));
let env = SmokeEnv::recording(Arc::clone(&reset_seeds));
let bound = crate::EnvServer::new(env)
.bind(BindAddress::Tcp {
host: "127.0.0.1".to_string(),
port: 0,
})
.await
.unwrap();
let address = bound.local_addr().to_string();
let server = tokio::spawn(async move { bound.serve().await });
let mut options = RunLocalOptions::parse(&address).unwrap().for_episodes(1);
options.base_seed = base_seed;
ModelWorker::new(SmokeModel {
predicts: Arc::new(AtomicUsize::new(0)),
closes: Arc::new(AtomicUsize::new(0)),
})
.run_local_async(options)
.await
.unwrap();
server.abort();
reset_seeds.lock().await.clone()
}
let seeded = run_with_base_seed(Some(4242)).await;
assert!(
seeded.first().map(Option::is_some).unwrap_or(false),
"expected a concrete reset seed, got {seeded:?}"
);
let seeded_again = run_with_base_seed(Some(4242)).await;
assert_eq!(
seeded.first(),
seeded_again.first(),
"the same base_seed must derive the same reset seed (reproducibility)"
);
let unseeded = run_with_base_seed(None).await;
assert_eq!(
unseeded.first(),
Some(&None),
"no base_seed must leave the reset seed unset, got {unseeded:?}"
);
}
#[test]
fn run_local_and_serve_options_cover_all_axes() {
let run = RunLocalOptions::parse("tcp://env:50051")
.unwrap()
.for_episodes(5)
.base_seed(123)
.execution_horizon(8)
.prefetch_lead(2);
assert_eq!(run.max_episodes, Some(5));
assert_eq!(run.base_seed, Some(123));
assert_eq!(run.execution_horizon, 8);
assert_eq!(run.prefetch_lead, 2);
assert_eq!(
run.env_address,
ConnectAddress::parse("tcp://env:50051").unwrap()
);
let default_run = RunLocalOptions::new(ConnectAddress::parse("tcp://env:1").unwrap());
assert_eq!(default_run.max_episodes, None);
assert_eq!(default_run.base_seed, None);
assert_eq!(default_run.execution_horizon, 1);
assert_eq!(default_run.prefetch_lead, 0);
assert_eq!(default_run.execution_horizon(0).execution_horizon, 1);
let serve = ServeModelOptions::parse("tcp://0.0.0.0:50061")
.unwrap()
.token("secret")
.serve_options(ServeOptions {
allow_remote_shutdown: true,
..ServeOptions::default()
});
assert_eq!(serve.token, "secret");
assert!(serve.serve.allow_remote_shutdown);
let default_serve = ServeModelOptions::new(BindAddress::parse("tcp://0.0.0.0:1").unwrap());
assert!(default_serve.token.is_empty());
assert_eq!(default_serve.serve, ServeOptions::default());
}
#[tokio::test]
async fn served_model_resolve_adapter_requires_env_spec() {
let response = handle_model_request(
JoinRequest {
kind: Some(join_request::Kind::ResolveAdapter(ResolveAdapterRequest {
context: Some(AdapterContext {
session_id: "session-1".to_string(),
env_id: "env-1".to_string(),
request_id: "resolve-1".to_string(),
}),
env_spec: None,
..Default::default()
})),
request_id: "resolve-1".to_string(),
},
Arc::new(Mutex::new(SmokeModel {
predicts: Arc::new(AtomicUsize::new(0)),
closes: Arc::new(AtomicUsize::new(0)),
})),
None,
Arc::new(Mutex::new(HashMap::new())),
Admission::now(),
)
.await;
assert!(matches!(response.kind, Some(join_response::Kind::Error(_))));
}
#[tokio::test]
async fn served_model_predict_mirrors_route_context() {
let context = AdapterContext {
session_id: "session-1".to_string(),
env_id: "env-1".to_string(),
request_id: "request-1".to_string(),
};
let response = handle_model_request(
JoinRequest {
kind: Some(join_request::Kind::Predict(PredictRequest {
history: Vec::new(),
step: None,
context: Some(context.clone()),
observation: None,
episode_info: vec![EpisodeInfo {
episode_id: "episode-1".to_string(),
seed: None,
}],
})),
request_id: "request-1".to_string(),
},
Arc::new(Mutex::new(SmokeModel {
predicts: Arc::new(AtomicUsize::new(0)),
closes: Arc::new(AtomicUsize::new(0)),
})),
None,
Arc::new(Mutex::new(HashMap::from([(
"env-1".to_string(),
ModelRouteConfig {
env_contract: Some(std::sync::Arc::new(SmokeEnv::new().env_contract)),
floor: None,
},
)]))),
Admission::now(),
)
.await;
match response.kind {
Some(join_response::Kind::Predict(response)) => {
assert_eq!(response.context, Some(context));
}
other => panic!("expected predict response, got {other:?}"),
}
}
#[tokio::test]
async fn served_model_predict_uses_episode_id_count_as_lane_count() {
let response = handle_model_request(
JoinRequest {
kind: Some(join_request::Kind::Predict(PredictRequest {
history: Vec::new(),
step: None,
context: Some(AdapterContext {
session_id: "session-1".to_string(),
env_id: "env-1".to_string(),
request_id: "request-1".to_string(),
}),
observation: None,
episode_info: vec![
EpisodeInfo {
episode_id: "episode-0".to_string(),
seed: None,
},
EpisodeInfo {
episode_id: "episode-1".to_string(),
seed: None,
},
],
})),
request_id: "request-1".to_string(),
},
Arc::new(Mutex::new(SmokeModel {
predicts: Arc::new(AtomicUsize::new(0)),
closes: Arc::new(AtomicUsize::new(0)),
})),
None,
Arc::new(Mutex::new(HashMap::from([(
"env-1".to_string(),
ModelRouteConfig {
env_contract: Some(std::sync::Arc::new(SmokeEnv::new().env_contract)),
floor: None,
},
)]))),
Admission::now(),
)
.await;
assert!(matches!(
response.kind,
Some(join_response::Kind::Predict(_))
));
}
#[derive(Clone, Default)]
struct PerRouteActionHandler {
seen: Arc<Mutex<Vec<(String, String)>>>,
}
#[async_trait]
impl ModelHandler for PerRouteActionHandler {
async fn predict(&mut self, observation: ModelObservation) -> Result<Vec<spaces::SpaceValue>> {
self.seen.lock().await.push((
observation.route.env_id.clone(),
observation.episode_id().to_string(),
));
let action = if observation.route.env_id == "env-disc" {
spaces::SpaceValue::Discrete(0)
} else {
spaces::SpaceValue::Box(
spaces::Tensor::from_vec(vec![0u8], vec![1], spaces::DType::Uint8).unwrap(),
)
};
Ok((0..observation.num_envs).map(|_| action.clone()).collect())
}
}
fn discrete_action_contract() -> spaces::EnvContract {
rlmesh_grpc::wire::env_spec_from_proto(rlmesh_proto::core::v1::EnvSpec {
id: "Disc-v0".to_string(),
observation_space: None,
action_space: Some(rlmesh_proto::spaces::v1::SpaceSpec {
shape: vec![],
dtype: rlmesh_proto::spaces::v1::DataType::Int64 as i32,
spec: Some(rlmesh_proto::spaces::v1::space_spec::Spec::Discrete(
rlmesh_proto::spaces::v1::DiscreteSpec { n: 1000, start: 0 },
)),
}),
metadata: None,
})
.expect("discrete env spec builds a contract")
}
fn grouped_member(env_id: &str, request_id: &str, episode_id: &str) -> PredictRequest {
PredictRequest {
history: Vec::new(),
step: None,
context: Some(AdapterContext {
session_id: "session-1".to_string(),
env_id: env_id.to_string(),
request_id: request_id.to_string(),
}),
observation: None,
episode_info: vec![EpisodeInfo {
episode_id: episode_id.to_string(),
seed: None,
}],
}
}
fn expect_group_response(
result: &GroupedPredictResult,
) -> rlmesh_proto::model::v1::PredictResponse {
match result.outcome.as_ref().expect("group has an outcome") {
grouped_predict_result::Outcome::Response(response) => response.clone(),
grouped_predict_result::Outcome::Error(error) => {
panic!(
"expected a per-group response, got error: {}",
error.message
)
}
}
}
#[tokio::test]
async fn grouped_predict_processes_each_group_against_its_own_route() {
let seen = Arc::new(Mutex::new(Vec::new()));
let handler = Arc::new(Mutex::new(PerRouteActionHandler {
seen: Arc::clone(&seen),
}));
let configs = Arc::new(Mutex::new(HashMap::from([
(
"env-box".to_string(),
ModelRouteConfig {
env_contract: Some(Arc::new(SmokeEnv::new().env_contract)),
floor: None,
},
),
(
"env-disc".to_string(),
ModelRouteConfig {
env_contract: Some(Arc::new(discrete_action_contract())),
floor: None,
},
),
])));
let response = handle_model_request(
JoinRequest {
kind: Some(join_request::Kind::GroupedPredict(GroupedPredictRequest {
groups: vec![
grouped_member("env-box", "p-box", "ep-box"),
grouped_member("env-disc", "p-disc", "ep-disc"),
],
})),
request_id: "grouped-1".to_string(),
},
Arc::clone(&handler),
None,
Arc::clone(&configs),
Admission::now(),
)
.await;
let results = match response.kind {
Some(join_response::Kind::GroupedPredict(GroupedPredictResponse { results })) => results,
other => panic!("expected grouped predict response, got {other:?}"),
};
assert_eq!(results.len(), 2, "one result per group, in order");
let box_response = expect_group_response(&results[0]);
assert_eq!(box_response.context.as_ref().unwrap().env_id, "env-box");
assert_eq!(
box_response.actions.len(),
1,
"the box group's action was encoded against its own action space (no chunking)"
);
let disc_response = expect_group_response(&results[1]);
assert_eq!(disc_response.context.as_ref().unwrap().env_id, "env-disc");
assert_eq!(disc_response.actions.len(), 1);
assert_eq!(
*seen.lock().await,
vec![
("env-box".to_string(), "ep-box".to_string()),
("env-disc".to_string(), "ep-disc".to_string()),
]
);
}
struct PerRouteChunkHandler;
impl PerRouteChunkHandler {
fn box_action() -> spaces::SpaceValue {
spaces::SpaceValue::Box(
spaces::Tensor::from_vec(vec![0u8], vec![1], spaces::DType::Uint8).unwrap(),
)
}
}
#[async_trait]
impl ModelHandler for PerRouteChunkHandler {
async fn predict(&mut self, observation: ModelObservation) -> Result<Vec<spaces::SpaceValue>> {
Ok((0..observation.num_envs)
.map(|_| Self::box_action())
.collect())
}
async fn predict_chunked(&mut self, observation: ModelObservation) -> Result<PredictFrames> {
let lanes = observation.num_envs;
let frame = || (0..lanes).map(|_| Self::box_action()).collect();
let replay = if observation.route.env_id == "env-chunk" {
vec![frame(), frame()]
} else {
Vec::new()
};
Ok(PredictFrames {
actions: frame(),
replay,
})
}
}
#[tokio::test]
async fn grouped_predict_carries_each_groups_own_chunk_frames() {
let handler = Arc::new(Mutex::new(PerRouteChunkHandler));
let configs = Arc::new(Mutex::new(HashMap::from([
(
"env-chunk".to_string(),
ModelRouteConfig {
env_contract: Some(Arc::new(SmokeEnv::new().env_contract)),
floor: None,
},
),
(
"env-single".to_string(),
ModelRouteConfig {
env_contract: Some(Arc::new(SmokeEnv::new().env_contract)),
floor: None,
},
),
])));
let response = handle_model_request(
JoinRequest {
kind: Some(join_request::Kind::GroupedPredict(GroupedPredictRequest {
groups: vec![
grouped_member("env-chunk", "p-chunk", "ep-chunk"),
grouped_member("env-single", "p-single", "ep-single"),
],
})),
request_id: "grouped-chunk-1".to_string(),
},
Arc::clone(&handler),
None,
Arc::clone(&configs),
Admission::now(),
)
.await;
let results = match response.kind {
Some(join_response::Kind::GroupedPredict(GroupedPredictResponse { results })) => results,
other => panic!("expected grouped predict response, got {other:?}"),
};
assert_eq!(results.len(), 2, "one result per group, in order");
let chunked = expect_group_response(&results[0]);
assert_eq!(chunked.context.as_ref().unwrap().env_id, "env-chunk");
assert_eq!(
chunked.actions.len(),
3,
"the chunking group returns frame 0 plus its two replay frames"
);
let single = expect_group_response(&results[1]);
assert_eq!(single.context.as_ref().unwrap().env_id, "env-single");
assert_eq!(
single.actions.len(),
1,
"the un-chunked group stays a single frame"
);
}
#[tokio::test]
async fn grouped_predict_isolates_a_single_group_failure() {
let seen = Arc::new(Mutex::new(Vec::new()));
let handler = Arc::new(Mutex::new(PerRouteActionHandler {
seen: Arc::clone(&seen),
}));
let configs = Arc::new(Mutex::new(HashMap::from([(
"env-box".to_string(),
ModelRouteConfig {
env_contract: Some(Arc::new(SmokeEnv::new().env_contract)),
floor: None,
},
)])));
let response = handle_model_request(
JoinRequest {
kind: Some(join_request::Kind::GroupedPredict(GroupedPredictRequest {
groups: vec![
grouped_member("env-box", "p-box", "ep-box"),
grouped_member("env-missing", "p-missing", "ep-missing"),
],
})),
request_id: "grouped-2".to_string(),
},
Arc::clone(&handler),
None,
Arc::clone(&configs),
Admission::now(),
)
.await;
let results = match response.kind {
Some(join_response::Kind::GroupedPredict(GroupedPredictResponse { results })) => results,
other => panic!("expected grouped predict response, got {other:?}"),
};
assert_eq!(results.len(), 2);
assert!(
matches!(
results[0].outcome.as_ref().unwrap(),
grouped_predict_result::Outcome::Response(_)
),
"the configured group still succeeds"
);
match results[1].outcome.as_ref().unwrap() {
grouped_predict_result::Outcome::Error(error) => {
assert!(
error.message.contains("was not resolved"),
"got: {}",
error.message
);
}
other => panic!("expected an error for the unresolved group, got {other:?}"),
}
assert_eq!(
*seen.lock().await,
vec![("env-box".to_string(), "ep-box".to_string())]
);
}
type ResetLog = Arc<Mutex<Vec<(String, Vec<String>)>>>;
#[derive(Clone, Default)]
struct ReleaseRecordingSetup {
released: Arc<Mutex<Vec<String>>>,
evicted: ResetLog,
}
#[async_trait]
impl ModelRouteSetup for ReleaseRecordingSetup {
async fn resolve_adapter(
&self,
_env_id: &str,
_env_contract: &spaces::EnvContract,
_options: crate::model::ResolveOptions,
) -> Result<crate::model::RouteNeeds> {
Ok(Default::default())
}
async fn reset_adapter(&self, env_id: &str, episode_ids: &[String]) -> Result<()> {
self.evicted
.lock()
.await
.push((env_id.to_string(), episode_ids.to_vec()));
Ok(())
}
async fn release_adapter(&self, env_id: &str) -> Result<()> {
self.released.lock().await.push(env_id.to_string());
Ok(())
}
}
#[derive(Clone, Default)]
struct EndRecordingHandler {
ended: ResetLog,
}
#[async_trait]
impl ModelHandler for EndRecordingHandler {
async fn predict(&mut self, _observation: ModelObservation) -> Result<Vec<spaces::SpaceValue>> {
Ok(Vec::new())
}
async fn reset_adapter(&mut self, env_id: &str, episode_ids: Vec<String>) -> Result<()> {
self.ended
.lock()
.await
.push((env_id.to_string(), episode_ids));
Ok(())
}
}
#[tokio::test]
async fn reset_adapter_evicts_route_state_without_the_handler_lock() {
let evicted = Arc::new(Mutex::new(Vec::new()));
let route_setup: Arc<dyn ModelRouteSetup> = Arc::new(ReleaseRecordingSetup {
evicted: Arc::clone(&evicted),
..Default::default()
});
let ended = Arc::new(Mutex::new(Vec::new()));
let handler = Arc::new(Mutex::new(EndRecordingHandler {
ended: Arc::clone(&ended),
}));
let route_configs = Arc::new(Mutex::new(HashMap::new()));
let request = |request: JoinRequest| {
handle_model_request(
request,
Arc::clone(&handler),
Some(Arc::clone(&route_setup)),
Arc::clone(&route_configs),
Admission::now(),
)
};
let forward = handler.lock().await;
let mut reset = Box::pin(request(JoinRequest {
kind: Some(join_request::Kind::ResetAdapter(ResetAdapterRequest {
context: Some(AdapterContext {
session_id: "session".to_string(),
env_id: "env-b".to_string(),
request_id: "reset-b".to_string(),
}),
episode_ids: vec!["ep-1".to_string()],
})),
request_id: "reset-b".to_string(),
}));
assert!(
tokio::time::timeout(Duration::from_millis(50), &mut reset)
.await
.is_err(),
"the handler's episode-end hook stays serialized behind the forward"
);
assert_eq!(
*evicted.lock().await,
vec![("env-b".to_string(), vec!["ep-1".to_string()])],
"route-local eviction must not wait for the forward"
);
assert!(ended.lock().await.is_empty());
let predict = request(predict_join_request("env-c", "predict-c", false));
let response = tokio::time::timeout(Duration::from_millis(50), predict)
.await
.expect("decode must not wait for the handler lock");
assert!(
response.queue_ns.is_some(),
"a request that never reached the handler still reports what it waited"
);
match response.kind {
Some(join_response::Kind::Error(error)) => {
assert!(
error.message.contains("was not resolved"),
"{}",
error.message
);
}
other => panic!("expected the unresolved-route error, got {other:?}"),
}
drop(forward);
let response = reset.await;
assert!(matches!(
response.kind,
Some(join_response::Kind::ResetAdapter(_))
));
assert_eq!(
*ended.lock().await,
vec![("env-b".to_string(), vec!["ep-1".to_string()])]
);
}
#[tokio::test]
async fn served_model_release_adapter_tears_down_only_its_env() {
let released = Arc::new(Mutex::new(Vec::new()));
let route_setup: Arc<dyn ModelRouteSetup> = Arc::new(ReleaseRecordingSetup {
released: Arc::clone(&released),
..Default::default()
});
let route_configs = Arc::new(Mutex::new(HashMap::from([
(
"env-1".to_string(),
ModelRouteConfig {
env_contract: None,
floor: None,
},
),
(
"env-2".to_string(),
ModelRouteConfig {
env_contract: None,
floor: None,
},
),
])));
let response = handle_model_request(
JoinRequest {
kind: Some(join_request::Kind::ReleaseAdapter(ReleaseAdapterRequest {
context: Some(AdapterContext {
session_id: "session-1".to_string(),
env_id: "env-1".to_string(),
request_id: "release-1".to_string(),
}),
reason: "env complete".to_string(),
})),
request_id: "release-1".to_string(),
},
Arc::new(Mutex::new(SmokeModel {
predicts: Arc::new(AtomicUsize::new(0)),
closes: Arc::new(AtomicUsize::new(0)),
})),
Some(Arc::clone(&route_setup)),
Arc::clone(&route_configs),
Admission::now(),
)
.await;
assert!(matches!(
response.kind,
Some(join_response::Kind::ReleaseAdapter(_))
));
assert_eq!(*released.lock().await, vec!["env-1".to_string()]);
let route_configs = route_configs.lock().await;
assert!(!route_configs.contains_key("env-1"));
assert!(route_configs.contains_key("env-2"));
}
#[tokio::test]
async fn served_model_close_releases_every_adapter() {
let released = Arc::new(Mutex::new(Vec::new()));
let route_setup: Arc<dyn ModelRouteSetup> = Arc::new(ReleaseRecordingSetup {
released: Arc::clone(&released),
..Default::default()
});
let route_configs = Arc::new(Mutex::new(HashMap::from([
(
"env-1".to_string(),
ModelRouteConfig {
env_contract: None,
floor: None,
},
),
(
"env-2".to_string(),
ModelRouteConfig {
env_contract: None,
floor: None,
},
),
])));
let response = handle_model_request(
JoinRequest {
kind: Some(join_request::Kind::Close(CloseParticipantRequest {
reason: "session complete".to_string(),
})),
request_id: "close-1".to_string(),
},
Arc::new(Mutex::new(SmokeModel {
predicts: Arc::new(AtomicUsize::new(0)),
closes: Arc::new(AtomicUsize::new(0)),
})),
Some(Arc::clone(&route_setup)),
Arc::clone(&route_configs),
Admission::now(),
)
.await;
assert!(matches!(response.kind, Some(join_response::Kind::Close(_))));
let mut released = released.lock().await.clone();
released.sort();
assert_eq!(released, vec!["env-1".to_string(), "env-2".to_string()]);
assert!(route_configs.lock().await.is_empty());
}
#[tokio::test]
async fn served_model_close_detaches_and_shutdown_runs_close_hook_once() {
let listener = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
let port = listener.local_addr().unwrap().port();
drop(listener);
let predicts = Arc::new(AtomicUsize::new(0));
let closes = Arc::new(AtomicUsize::new(0));
let server_predicts = Arc::clone(&predicts);
let server_closes = Arc::clone(&closes);
let server = tokio::spawn(async move {
ModelWorker::new(SmokeModel {
predicts: server_predicts,
closes: server_closes,
})
.serve_async(
ServeModelOptions::new(BindAddress::Tcp {
host: "127.0.0.1".to_string(),
port,
})
.token("route-token")
.serve_options(ServeOptions {
allow_remote_shutdown: true,
..ServeOptions::default()
}),
)
.await
});
let address = format!("tcp://127.0.0.1:{port}");
let connect_options = rlmesh_grpc::ConnectOptions::with_deadline(Duration::from_secs(5))
.backoff(Duration::from_millis(10));
let mut client =
rlmesh_grpc::ModelClient::connect_with_retry(&address, "route-token", &connect_options)
.await
.expect("model server did not start");
client.handshake().await.unwrap();
client.close("client session complete").await.unwrap();
tokio::time::sleep(Duration::from_millis(50)).await;
assert!(!server.is_finished());
assert_eq!(closes.load(Ordering::SeqCst), 0);
let mut second = rlmesh_grpc::ModelClient::connect(&address, "route-token")
.await
.unwrap();
second.handshake().await.unwrap();
let shutdown = second.shutdown("test complete").await.unwrap();
assert!(shutdown.accepted);
shutdown_and_join(server).await;
assert_eq!(predicts.load(Ordering::SeqCst), 0);
assert_eq!(closes.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn public_env_runtime_adapter_drives_a_remote_env_with_telemetry() {
use rlmesh_runtime::RuntimeEnv;
let listener = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
let env_address = format!("tcp://{}", listener.local_addr().unwrap());
let env_server = tokio::spawn(async move {
tonic::transport::Server::builder()
.add_service(rlmesh_grpc::env::env_service(
crate::env::WireLaneAdapter::new(vec![SmokeEnv::new()]).unwrap(),
))
.serve_with_incoming(TcpListenerStream::new(listener))
.await
.unwrap()
});
let mut client = rlmesh_grpc::EnvClient::connect(&env_address).await.unwrap();
client.handshake().await.unwrap();
let mut adapter = crate::EnvClientRuntimeEnv::new(client);
let reset = adapter
.reset(rlmesh_proto::env::v1::ResetRequest {
seeds: vec![7],
options: None,
timeout_ms: 0,
env_indices: vec![],
episode_ids: vec![],
})
.await
.expect("adapter reset must succeed");
assert!(
reset.endpoint_total_ns.is_some(),
"adapter must encapsulate and surface the per-call endpoint duration"
);
env_server.abort();
}
struct AutoresetVectorEnv {
obs_space: spaces::SpaceSpec,
action_space: spaces::SpaceSpec,
env_contract: spaces::EnvContract,
terminal_after: Vec<usize>,
lane_step: Vec<usize>,
pending: Vec<bool>,
}
impl AutoresetVectorEnv {
fn new(terminal_after: Vec<usize>) -> Self {
let n = terminal_after.len();
let obs_space = spaces::spaces::BoxSpaceBuilder::scalar(-1.0, 1.0, vec![1])
.dtype(spaces::DType::Float32)
.build()
.unwrap();
let action_space = spaces::spaces::DiscreteBuilder::new(2).build().unwrap();
let env_contract = spaces::EnvContract {
id: "AutoresetVectorEnv-v0".to_string(),
autoreset_mode: spaces::types::AutoresetMode::NextStep,
observation_space: Some(obs_space.clone()),
action_space: Some(action_space.clone()),
metadata: None,
render_mode: String::new(),
num_envs: n as u32,
};
Self {
obs_space,
action_space,
env_contract,
lane_step: vec![0; n],
pending: vec![false; n],
terminal_after,
}
}
fn lane_obs() -> spaces::SpaceValue {
spaces::SpaceValue::Box(
spaces::Tensor::from_vec(vec![0u8; 4], vec![1], spaces::DType::Float32).unwrap(),
)
}
}
#[async_trait]
impl crate::VectorEnv for AutoresetVectorEnv {
fn observation_space(&self) -> &spaces::SpaceSpec {
&self.obs_space
}
fn action_space(&self) -> &spaces::SpaceSpec {
&self.action_space
}
fn num_envs(&self) -> usize {
self.terminal_after.len()
}
fn env_contract(&self) -> &spaces::EnvContract {
&self.env_contract
}
async fn reset(
&mut self,
_req: crate::VectorResetRequest,
) -> std::result::Result<crate::VectorResetResult, spaces::EnvRuntimeError> {
let n = self.terminal_after.len();
self.lane_step = vec![0; n];
self.pending = vec![false; n];
Ok(crate::VectorResetResult {
observations: (0..n).map(|_| Self::lane_obs()).collect(),
info: None,
episode_ids: Vec::new(),
})
}
async fn step(
&mut self,
_req: crate::VectorStepRequest,
) -> std::result::Result<crate::VectorStepResult, spaces::EnvRuntimeError> {
let n = self.terminal_after.len();
let mut rewards = vec![1.0; n];
let mut terminated = vec![false; n];
for lane in 0..n {
if self.pending[lane] {
self.pending[lane] = false;
self.lane_step[lane] = 0;
rewards[lane] = 0.0;
} else {
self.lane_step[lane] += 1;
if self.lane_step[lane] >= self.terminal_after[lane] {
terminated[lane] = true;
self.pending[lane] = true;
}
}
}
Ok(crate::VectorStepResult {
observations: (0..n).map(|_| Self::lane_obs()).collect(),
rewards,
terminated,
truncated: vec![false; n],
info: None,
completed_episodes: Vec::new(),
episode_ids: Vec::new(),
})
}
async fn render(
&mut self,
_req: spaces::RenderRequest,
) -> std::result::Result<spaces::RenderResult, spaces::EnvRuntimeError> {
Ok(spaces::RenderResult::default())
}
async fn close(
&mut self,
_req: spaces::CloseRequest,
) -> std::result::Result<crate::VectorCloseResult, spaces::EnvRuntimeError> {
Ok(crate::VectorCloseResult::default())
}
}
#[derive(Clone, Default)]
struct IdRecordingHandler {
predict_ids: Arc<Mutex<Vec<Vec<String>>>>,
evicted_ids: Arc<Mutex<Vec<String>>>,
}
#[async_trait]
impl ModelHandler for IdRecordingHandler {
async fn predict(&mut self, observation: ModelObservation) -> Result<Vec<spaces::SpaceValue>> {
self.predict_ids
.lock()
.await
.push(observation.episode_ids());
Ok((0..observation.num_envs)
.map(|_| spaces::SpaceValue::Discrete(0))
.collect())
}
async fn reset_adapter(&mut self, _env_id: &str, episode_ids: Vec<String>) -> Result<()> {
self.evicted_ids.lock().await.extend(episode_ids);
Ok(())
}
}
#[tokio::test]
async fn run_local_refuses_chunking_on_a_vector_env() {
let bound = crate::VectorEnvServer::new(AutoresetVectorEnv::new(vec![2, 3]))
.bind(BindAddress::Tcp {
host: "127.0.0.1".to_string(),
port: 0,
})
.await
.unwrap();
let address = bound.local_addr().to_string();
let server = tokio::spawn(async move { bound.serve().await });
let error = ModelWorker::new(IdRecordingHandler::default())
.run_local_async(
RunLocalOptions::parse(&address)
.unwrap()
.for_episodes(2)
.execution_horizon(2),
)
.await
.expect_err("a chunked vector route is refused");
server.abort();
let message = error.to_string();
assert!(
message.contains("num_envs=2") && message.contains("execution_horizon=2"),
"the refusal should name both numbers: {message}"
);
}
#[tokio::test]
async fn next_step_episode_ids_round_trip_through_the_real_env_server() {
let bound = crate::VectorEnvServer::new(AutoresetVectorEnv::new(vec![2, 3]))
.bind(BindAddress::Tcp {
host: "127.0.0.1".to_string(),
port: 0,
})
.await
.unwrap();
let address = bound.local_addr().to_string();
let server = tokio::spawn(async move { bound.serve().await });
let handler = IdRecordingHandler::default();
let predict_ids = Arc::clone(&handler.predict_ids);
let evicted_ids = Arc::clone(&handler.evicted_ids);
ModelWorker::new(handler)
.run_local_async(RunLocalOptions::parse(&address).unwrap().for_episodes(6))
.await
.unwrap();
server.abort();
let predicts = predict_ids.lock().await.clone();
assert!(!predicts.is_empty(), "the model ran at least one predict");
for row in &predicts {
assert_eq!(row.len(), 2, "two lanes per predict");
for id in row {
assert_eq!(id.len(), 36, "episode id must be a UUID, got {id:?}");
assert_eq!(id.as_bytes()[14], b'7', "UUIDv7 version nibble: {id:?}");
}
}
let lane0: std::collections::BTreeSet<&str> =
predicts.iter().map(|row| row[0].as_str()).collect();
assert!(
lane0.len() >= 2,
"lane 0's episode id must roll across autoreset, saw {lane0:?}"
);
let evicted = evicted_ids.lock().await.clone();
assert!(
!evicted.is_empty(),
"ResetAdapter must evict completed episodes"
);
for id in &evicted {
assert_eq!(id.len(), 36, "evicted a UUID id, got {id:?}");
}
}
#[tokio::test]
async fn model_bind_resolves_port_zero_before_serving() {
let predicts = Arc::new(AtomicUsize::new(0));
let closes = Arc::new(AtomicUsize::new(0));
let bound = ModelWorker::new(SmokeModel {
predicts: Arc::clone(&predicts),
closes: Arc::clone(&closes),
})
.bind_async(
ServeModelOptions::new(BindAddress::Tcp {
host: "127.0.0.1".to_string(),
port: 0,
})
.token("route-token")
.serve_options(ServeOptions {
allow_remote_shutdown: true,
..ServeOptions::default()
}),
)
.await
.unwrap();
let (port, server) = spawn_bound_server(bound);
let address = format!("tcp://127.0.0.1:{port}");
let mut client = rlmesh_grpc::ModelClient::connect(&address, "route-token")
.await
.unwrap();
client.handshake().await.unwrap();
let shutdown = client.shutdown("test complete").await.unwrap();
assert!(shutdown.accepted);
shutdown_and_join(server).await;
assert_eq!(closes.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn a_declared_edition_no_peer_can_run_is_refused_at_the_env_handshake() {
let bound = crate::EnvServer::new(SmokeEnv::new())
.bind_with_options(
BindAddress::Tcp {
host: "127.0.0.1".to_string(),
port: 0,
},
ServeOptions {
workflow_edition: Some("2020.01".to_string()),
..ServeOptions::default()
},
)
.await
.unwrap();
let address = bound.local_addr().to_string();
let server = tokio::spawn(async move { bound.serve().await });
let connect =
crate::RemoteVectorEnv::connect_to(ConnectAddress::parse(&address).unwrap()).await;
server.abort();
let Err(error) = connect else {
panic!("the env declares an edition this build cannot run")
};
let error = error.to_string();
assert!(
error.contains("no mutual workflow edition with the env"),
"{error}"
);
for half in [
"env wants \"2020.01\"",
"runtime wants",
rlmesh_proto::CURRENT_WORKFLOW_EDITION,
] {
assert!(error.contains(half), "expected {half:?} in: {error}");
}
}
#[tokio::test]
async fn a_declared_edition_no_peer_can_run_is_refused_at_the_session_floor() {
let predicts = Arc::new(AtomicUsize::new(0));
let closes = Arc::new(AtomicUsize::new(0));
let bound = ModelWorker::new(SmokeModel {
predicts: Arc::clone(&predicts),
closes: Arc::clone(&closes),
})
.bind_async(
ServeModelOptions::new(BindAddress::Tcp {
host: "127.0.0.1".to_string(),
port: 0,
})
.serve_options(ServeOptions {
workflow_edition: Some("2020.01".to_string()),
..ServeOptions::default()
}),
)
.await
.unwrap();
let (port, server) = spawn_bound_server(bound);
let connect = crate::RemoteModel::connect(
&format!("tcp://127.0.0.1:{port}"),
SmokeEnv::new().env_contract,
)
.await;
server.abort();
let Err(error) = connect else {
panic!("the model declares an edition no tier can run")
};
let error = error.to_string();
assert!(
error.contains("no mutual workflow edition across env, model, and runtime"),
"{error}"
);
for half in ["env ", "model wants \"2020.01\"", "runtime "] {
assert!(error.contains(half), "expected {half:?} in: {error}");
}
assert_eq!(
predicts.load(Ordering::SeqCst),
0,
"the refusal must precede any predict"
);
}
#[tokio::test]
async fn remote_model_connects_resets_and_predicts() {
let predicts = Arc::new(AtomicUsize::new(0));
let closes = Arc::new(AtomicUsize::new(0));
let bound = ModelWorker::new(SmokeModel {
predicts: Arc::clone(&predicts),
closes: Arc::clone(&closes),
})
.bind_async(
ServeModelOptions::new(BindAddress::Tcp {
host: "127.0.0.1".to_string(),
port: 0,
})
.serve_options(ServeOptions {
allow_remote_shutdown: true,
..ServeOptions::default()
}),
)
.await
.unwrap();
let (port, server) = spawn_bound_server(bound);
let address = format!("tcp://127.0.0.1:{port}");
let env_contract = SmokeEnv::new().env_contract;
let mut model = crate::RemoteModel::connect(&address, env_contract)
.await
.expect("model server did not start");
let observe = || {
spaces::SpaceValue::Box(
spaces::Tensor::from_vec(vec![5], vec![1], spaces::DType::Uint8).unwrap(),
)
};
model.reset(None);
let action = model.predict(observe()).await.unwrap();
assert!(matches!(action, spaces::SpaceValue::Box(_)));
model.predict(observe()).await.unwrap();
assert_eq!(predicts.load(Ordering::SeqCst), 2);
model.close().await.unwrap();
drop(model);
let mut shutdown_client = rlmesh_grpc::ModelClient::connect(&address, "")
.await
.unwrap();
shutdown_client.handshake().await.unwrap();
assert!(shutdown_client.shutdown("done").await.unwrap().accepted);
shutdown_and_join(server).await;
}
#[tokio::test]
async fn remote_reset_emits_reset_adapter() {
#[derive(Clone, Default)]
struct ResetRecorder {
predict_ids: Arc<Mutex<Vec<String>>>,
evicted_ids: Arc<Mutex<Vec<String>>>,
}
#[async_trait]
impl ModelHandler for ResetRecorder {
async fn predict(
&mut self,
observation: ModelObservation,
) -> Result<Vec<spaces::SpaceValue>> {
self.predict_ids
.lock()
.await
.extend(observation.episode_ids());
Ok((0..observation.num_envs)
.map(|_| {
spaces::SpaceValue::Box(
spaces::Tensor::from_vec(vec![0u8], vec![1], spaces::DType::Uint8).unwrap(),
)
})
.collect())
}
async fn reset_adapter(&mut self, _env_id: &str, episode_ids: Vec<String>) -> Result<()> {
self.evicted_ids.lock().await.extend(episode_ids);
Ok(())
}
}
let handler = ResetRecorder::default();
let bound = ModelWorker::new(handler.clone())
.bind_async(
ServeModelOptions::new(BindAddress::Tcp {
host: "127.0.0.1".to_string(),
port: 0,
})
.serve_options(ServeOptions {
allow_remote_shutdown: true,
..ServeOptions::default()
}),
)
.await
.unwrap();
let (port, server) = spawn_bound_server(bound);
let address = format!("tcp://127.0.0.1:{port}");
let mut model = crate::RemoteModel::connect(&address, SmokeEnv::new().env_contract)
.await
.expect("model server did not start");
let observe = || {
spaces::SpaceValue::Box(
spaces::Tensor::from_vec(vec![5u8], vec![1], spaces::DType::Uint8).unwrap(),
)
};
model.reset(None);
model.predict(observe()).await.unwrap();
model.reset(None);
model.predict(observe()).await.unwrap();
model.close().await.unwrap();
let predicted = handler.predict_ids.lock().await.clone();
let evicted = handler.evicted_ids.lock().await.clone();
let episode_one = predicted[0].clone();
let episode_two = predicted[1].clone();
assert_ne!(episode_one, episode_two, "reset mints a fresh episode id");
assert_eq!(evicted, vec![episode_one, episode_two]);
drop(model);
let mut shutdown_client = rlmesh_grpc::ModelClient::connect(&address, "")
.await
.unwrap();
shutdown_client.handshake().await.unwrap();
assert!(shutdown_client.shutdown("done").await.unwrap().accepted);
shutdown_and_join(server).await;
}
#[tokio::test]
async fn remote_model_reconciles_three_way_floor_and_pins_route() {
let predicts = Arc::new(AtomicUsize::new(0));
let closes = Arc::new(AtomicUsize::new(0));
let bound = ModelWorker::new(SmokeModel {
predicts: Arc::clone(&predicts),
closes: Arc::clone(&closes),
})
.bind_async(
ServeModelOptions::new(BindAddress::Tcp {
host: "127.0.0.1".to_string(),
port: 0,
})
.serve_options(ServeOptions {
allow_remote_shutdown: true,
..ServeOptions::default()
}),
)
.await
.unwrap();
let (port, server) = spawn_bound_server(bound);
let address = format!("tcp://127.0.0.1:{port}");
let env_offer = rlmesh_proto::SessionOffer::new(&[rlmesh_proto::CURRENT_WORKFLOW_EDITION]);
let mut model = crate::RemoteModel::connect_with_env_offer(
&address,
"",
SmokeEnv::new().env_contract,
env_offer,
)
.await
.expect("model server did not start");
assert_eq!(
model.selected_workflow_edition(),
rlmesh_proto::CURRENT_WORKFLOW_EDITION
);
model.reset(None);
let action = model
.predict(spaces::SpaceValue::Box(
spaces::Tensor::from_vec(vec![5], vec![1], spaces::DType::Uint8).unwrap(),
))
.await
.unwrap();
assert!(matches!(action, spaces::SpaceValue::Box(_)));
drop(model);
let mut shutdown_client = rlmesh_grpc::ModelClient::connect(&address, "")
.await
.unwrap();
shutdown_client.handshake().await.unwrap();
assert!(shutdown_client.shutdown("done").await.unwrap().accepted);
shutdown_and_join(server).await;
}
#[tokio::test]
async fn remote_model_fails_fast_when_no_mutual_edition() {
let predicts = Arc::new(AtomicUsize::new(0));
let closes = Arc::new(AtomicUsize::new(0));
let bound = ModelWorker::new(SmokeModel {
predicts: Arc::clone(&predicts),
closes: Arc::clone(&closes),
})
.bind_async(
ServeModelOptions::new(BindAddress::Tcp {
host: "127.0.0.1".to_string(),
port: 0,
})
.serve_options(ServeOptions {
allow_remote_shutdown: true,
..ServeOptions::default()
}),
)
.await
.unwrap();
let (port, server) = spawn_bound_server(bound);
let address = format!("tcp://127.0.0.1:{port}");
let env_offer = rlmesh_proto::SessionOffer::new(&["2099.01"]);
let message = match crate::RemoteModel::connect_with_env_offer(
&address,
"",
SmokeEnv::new().env_contract,
env_offer,
)
.await
{
Ok(_) => panic!("a session with no mutual edition must fail to connect"),
Err(err) => err.to_string(),
};
assert!(
message.contains("no mutual workflow edition") && message.contains("2099.01"),
"expected an all-tiers edition diagnostic naming the offers, got: {message}"
);
server.abort();
}
#[tokio::test]
async fn remote_model_predict_requires_reset() {
let predicts = Arc::new(AtomicUsize::new(0));
let closes = Arc::new(AtomicUsize::new(0));
let bound = ModelWorker::new(SmokeModel {
predicts: Arc::clone(&predicts),
closes: Arc::clone(&closes),
})
.bind_async(
ServeModelOptions::new(BindAddress::Tcp {
host: "127.0.0.1".to_string(),
port: 0,
})
.serve_options(ServeOptions {
allow_remote_shutdown: true,
..ServeOptions::default()
}),
)
.await
.unwrap();
let (port, server) = spawn_bound_server(bound);
let address = format!("tcp://127.0.0.1:{port}");
let mut model = crate::RemoteModel::connect(&address, SmokeEnv::new().env_contract)
.await
.unwrap();
let err = model
.predict(spaces::SpaceValue::Box(
spaces::Tensor::from_vec(vec![0], vec![1], spaces::DType::Uint8).unwrap(),
))
.await
.unwrap_err();
assert!(err.to_string().contains("reset()"));
assert_eq!(predicts.load(Ordering::SeqCst), 0);
server.abort();
}
#[tokio::test]
async fn two_remote_models_in_one_process_use_distinct_env_keys() {
#[derive(Clone)]
struct KeyRecordingModel {
keys: Arc<Mutex<Vec<String>>>,
}
#[async_trait]
impl ModelHandler for KeyRecordingModel {
async fn predict(
&mut self,
observation: ModelObservation,
) -> Result<Vec<spaces::SpaceValue>> {
self.keys
.lock()
.await
.push(observation.route.env_id.clone());
Ok(vec![spaces::SpaceValue::Box(
spaces::Tensor::from_vec(vec![0u8], vec![1], spaces::DType::Uint8).unwrap(),
)])
}
}
let keys = Arc::new(Mutex::new(Vec::new()));
let bound = ModelWorker::new(KeyRecordingModel {
keys: Arc::clone(&keys),
})
.bind_async(
ServeModelOptions::new(BindAddress::Tcp {
host: "127.0.0.1".to_string(),
port: 0,
})
.serve_options(ServeOptions {
allow_remote_shutdown: true,
..ServeOptions::default()
}),
)
.await
.unwrap();
let (port, server) = spawn_bound_server(bound);
let address = format!("tcp://127.0.0.1:{port}");
let observe = || {
spaces::SpaceValue::Box(
spaces::Tensor::from_vec(vec![5], vec![1], spaces::DType::Uint8).unwrap(),
)
};
let mut first = crate::RemoteModel::connect(&address, SmokeEnv::new().env_contract)
.await
.unwrap();
let mut second = crate::RemoteModel::connect(&address, SmokeEnv::new().env_contract)
.await
.unwrap();
first.reset(None);
first.predict(observe()).await.unwrap();
second.reset(None);
second.predict(observe()).await.unwrap();
let recorded = keys.lock().await.clone();
assert_eq!(recorded.len(), 2);
assert_ne!(
recorded[0], recorded[1],
"two sessions in one process collided on a single route key: {recorded:?}"
);
first.close().await.unwrap();
second.close().await.unwrap();
server.abort();
}
#[tokio::test]
async fn served_model_reports_grpc_health_serving() {
use tonic_health::ServingStatus;
use tonic_health::pb::HealthCheckRequest;
use tonic_health::pb::health_client::HealthClient;
let predicts = Arc::new(AtomicUsize::new(0));
let closes = Arc::new(AtomicUsize::new(0));
let bound = ModelWorker::new(SmokeModel {
predicts: Arc::clone(&predicts),
closes: Arc::clone(&closes),
})
.bind_async(
ServeModelOptions::new(BindAddress::Tcp {
host: "127.0.0.1".to_string(),
port: 0,
})
.token("route-token")
.serve_options(ServeOptions {
allow_remote_shutdown: true,
..ServeOptions::default()
}),
)
.await
.unwrap();
let (port, server) = spawn_bound_server(bound);
let channel = tonic::transport::Endpoint::from_shared(format!("http://127.0.0.1:{port}"))
.unwrap()
.connect()
.await
.unwrap();
let mut health = HealthClient::new(channel);
let response = health
.check(HealthCheckRequest {
service: String::new(),
})
.await
.unwrap()
.into_inner();
assert_eq!(response.status, ServingStatus::Serving as i32);
let mut client =
rlmesh_grpc::ModelClient::connect(&format!("tcp://127.0.0.1:{port}"), "route-token")
.await
.unwrap();
assert!(client.shutdown("done").await.unwrap().accepted);
shutdown_and_join(server).await;
}
#[derive(Clone)]
struct OrderingHandler {
slow_delay: Duration,
predict_order: Arc<Mutex<Vec<(String, String)>>>,
route_events: Arc<Mutex<Vec<(String, String)>>>,
}
#[async_trait]
impl ModelHandler for OrderingHandler {
async fn predict(&mut self, observation: ModelObservation) -> Result<Vec<spaces::SpaceValue>> {
let env_id = observation.route.env_id.clone();
let request_id = observation.route.request_id.clone();
let slow = observation.episode_id().contains("slow");
self.predict_order
.lock()
.await
.push((env_id.clone(), request_id));
self.route_events
.lock()
.await
.push((env_id, "predict".to_string()));
if slow {
tokio::time::sleep(self.slow_delay).await;
}
Ok(vec![spaces::SpaceValue::Discrete(0)])
}
}
async fn bound_ordering_server(
handler: OrderingHandler,
predict_concurrency: Option<usize>,
) -> (
tokio::task::JoinHandle<crate::Result<()>>,
rlmesh_proto::model::v1::model_service_client::ModelServiceClient<tonic::transport::Channel>,
u16,
) {
let bound = ModelWorker::new(handler)
.bind_async(
ServeModelOptions::new(BindAddress::Tcp {
host: "127.0.0.1".to_string(),
port: 0,
})
.serve_options(ServeOptions {
allow_remote_shutdown: true,
predict_concurrency,
..ServeOptions::default()
}),
)
.await
.unwrap();
let (port, server) = spawn_bound_server(bound);
let channel = tonic::transport::Endpoint::from_shared(format!("http://127.0.0.1:{port}"))
.unwrap()
.connect()
.await
.unwrap();
let client = rlmesh_proto::model::v1::model_service_client::ModelServiceClient::new(channel);
(server, client, port)
}
fn resolve_adapter_request(env_id: &str, request_id: &str) -> JoinRequest {
JoinRequest {
kind: Some(join_request::Kind::ResolveAdapter(ResolveAdapterRequest {
context: Some(AdapterContext {
session_id: "session".to_string(),
env_id: env_id.to_string(),
request_id: request_id.to_string(),
}),
env_spec: Some(rlmesh_proto::core::v1::EnvSpec {
id: "Ordering-v0".to_string(),
observation_space: None,
action_space: Some(rlmesh_proto::spaces::v1::SpaceSpec {
shape: vec![],
dtype: rlmesh_proto::spaces::v1::DataType::Int64 as i32,
spec: Some(rlmesh_proto::spaces::v1::space_spec::Spec::Discrete(
rlmesh_proto::spaces::v1::DiscreteSpec { n: 1000, start: 0 },
)),
}),
metadata: None,
}),
..Default::default()
})),
request_id: request_id.to_string(),
}
}
fn predict_join_request(env_id: &str, request_id: &str, slow: bool) -> JoinRequest {
let suffix = if slow { "-slow" } else { "" };
JoinRequest {
kind: Some(join_request::Kind::Predict(PredictRequest {
history: Vec::new(),
step: None,
context: Some(AdapterContext {
session_id: "session".to_string(),
env_id: env_id.to_string(),
request_id: request_id.to_string(),
}),
observation: None,
episode_info: vec![EpisodeInfo {
episode_id: format!("ep-{env_id}-{request_id}{suffix}"),
seed: None,
}],
})),
request_id: request_id.to_string(),
}
}
#[tokio::test]
async fn pipelined_requests_complete_out_of_order() {
let handler = OrderingHandler {
slow_delay: Duration::from_millis(300),
predict_order: Arc::new(Mutex::new(Vec::new())),
route_events: Arc::new(Mutex::new(Vec::new())),
};
let (server, mut client, _port) = bound_ordering_server(handler, None).await;
let (req_tx, req_rx) = mpsc::channel::<JoinRequest>(8);
let request_stream = tokio_stream::wrappers::ReceiverStream::new(req_rx);
let mut responses = client
.join(tonic::Request::new(request_stream))
.await
.unwrap()
.into_inner();
req_tx
.send(resolve_adapter_request("slow", "cfg-slow"))
.await
.unwrap();
let ack = responses.message().await.unwrap().unwrap();
assert!(matches!(
ack.kind,
Some(join_response::Kind::ResolveAdapter(_))
));
req_tx
.send(predict_join_request("slow", "predict-slow", true))
.await
.unwrap();
req_tx
.send(resolve_adapter_request("other", "cfg-other"))
.await
.unwrap();
let first = responses.message().await.unwrap().unwrap();
assert_eq!(
first.request_id, "cfg-other",
"a resolve must not be head-of-line-blocked by an in-flight slow predict"
);
assert!(matches!(
first.kind,
Some(join_response::Kind::ResolveAdapter(_))
));
let second = responses.message().await.unwrap().unwrap();
assert_eq!(second.request_id, "predict-slow");
drop(req_tx);
let _ = responses.message().await;
server.abort();
}
#[tokio::test]
async fn pipelined_predicts_preserve_per_route_order() {
let route_events = Arc::new(Mutex::new(Vec::new()));
let handler = OrderingHandler {
slow_delay: Duration::from_millis(150),
predict_order: Arc::new(Mutex::new(Vec::new())),
route_events: Arc::clone(&route_events),
};
let (server, mut client, _port) = bound_ordering_server(handler, None).await;
let (req_tx, req_rx) = mpsc::channel::<JoinRequest>(8);
let request_stream = tokio_stream::wrappers::ReceiverStream::new(req_rx);
let mut responses = client
.join(tonic::Request::new(request_stream))
.await
.unwrap()
.into_inner();
req_tx
.send(resolve_adapter_request("r", "cfg"))
.await
.unwrap();
let ack = responses.message().await.unwrap().unwrap();
assert!(matches!(
ack.kind,
Some(join_response::Kind::ResolveAdapter(_))
));
req_tx
.send(predict_join_request("r", "p0", true))
.await
.unwrap();
req_tx
.send(predict_join_request("r", "p1", false))
.await
.unwrap();
let first = responses.message().await.unwrap().unwrap();
assert_eq!(first.request_id, "p0");
let second = responses.message().await.unwrap().unwrap();
assert_eq!(second.request_id, "p1");
drop(req_tx);
let _ = responses.message().await;
let events = route_events.lock().await.clone();
let r_events: Vec<&str> = events
.iter()
.filter(|(env_id, _)| env_id == "r")
.map(|(_, event)| event.as_str())
.collect();
assert_eq!(
r_events,
vec!["predict", "predict"],
"per-env predict order must match send order: {events:?}"
);
server.abort();
}
#[tokio::test]
async fn close_drains_after_in_flight_same_route_predict() {
let route_events = Arc::new(Mutex::new(Vec::new()));
let predict_order = Arc::new(Mutex::new(Vec::new()));
let handler = OrderingHandler {
slow_delay: Duration::from_millis(250),
predict_order: Arc::clone(&predict_order),
route_events: Arc::clone(&route_events),
};
let (server, mut client, _port) = bound_ordering_server(handler, None).await;
let (req_tx, req_rx) = mpsc::channel::<JoinRequest>(8);
let request_stream = tokio_stream::wrappers::ReceiverStream::new(req_rx);
let mut responses = client
.join(tonic::Request::new(request_stream))
.await
.unwrap()
.into_inner();
req_tx
.send(resolve_adapter_request("r", "cfg"))
.await
.unwrap();
let _ = responses.message().await.unwrap().unwrap();
req_tx
.send(predict_join_request("r", "p0", true))
.await
.unwrap();
req_tx
.send(JoinRequest {
kind: Some(join_request::Kind::Close(CloseParticipantRequest {
reason: "done".to_string(),
})),
request_id: "close".to_string(),
})
.await
.unwrap();
let first = responses.message().await.unwrap().unwrap();
assert_eq!(
first.request_id, "p0",
"the in-flight predict must complete before Close drains"
);
let second = responses.message().await.unwrap().unwrap();
assert_eq!(second.request_id, "close");
assert!(matches!(second.kind, Some(join_response::Kind::Close(_))));
let order = predict_order.lock().await.clone();
assert_eq!(order, vec![("r".to_string(), "p0".to_string())]);
drop(req_tx);
server.abort();
}
#[tokio::test]
async fn public_client_predict_concurrent_demuxes_overlapping_predicts() {
let handler = OrderingHandler {
slow_delay: Duration::from_millis(100),
predict_order: Arc::new(Mutex::new(Vec::new())),
route_events: Arc::new(Mutex::new(Vec::new())),
};
let bound = ModelWorker::new(handler)
.bind_async(
ServeModelOptions::new(BindAddress::Tcp {
host: "127.0.0.1".to_string(),
port: 0,
})
.token("tok")
.serve_options(ServeOptions {
allow_remote_shutdown: true,
..ServeOptions::default()
}),
)
.await
.unwrap();
let (port, server) = spawn_bound_server(bound);
let address = format!("tcp://127.0.0.1:{port}");
let mut client = rlmesh_grpc::ModelClient::connect(&address, "tok")
.await
.unwrap();
client.handshake().await.unwrap();
client
.resolve_adapter(ResolveAdapterRequest {
context: Some(AdapterContext {
session_id: "s".to_string(),
env_id: "r".to_string(),
request_id: "cfg".to_string(),
}),
env_spec: Some(rlmesh_proto::core::v1::EnvSpec {
id: "Ordering-v0".to_string(),
observation_space: None,
action_space: Some(rlmesh_proto::spaces::v1::SpaceSpec {
shape: vec![],
dtype: rlmesh_proto::spaces::v1::DataType::Int64 as i32,
spec: Some(rlmesh_proto::spaces::v1::space_spec::Spec::Discrete(
rlmesh_proto::spaces::v1::DiscreteSpec { n: 1000, start: 0 },
)),
}),
metadata: None,
}),
..Default::default()
})
.await
.unwrap();
let client = Arc::new(client);
let make_predict = |request_id: &str, slow: bool| PredictRequest {
history: Vec::new(),
step: None,
context: Some(AdapterContext {
session_id: "s".to_string(),
env_id: "r".to_string(),
request_id: request_id.to_string(),
}),
observation: None,
episode_info: vec![EpisodeInfo {
episode_id: format!("ep-{request_id}{}", if slow { "-slow" } else { "" }),
seed: None,
}],
};
let c1 = Arc::clone(&client);
let p1 = make_predict("predict-1", true);
let first = tokio::spawn(async move { c1.predict_concurrent(p1).await });
let c2 = Arc::clone(&client);
let p2 = make_predict("predict-2", false);
let second = tokio::spawn(async move { c2.predict_concurrent(p2).await });
let r1 = first.await.unwrap().unwrap();
let r2 = second.await.unwrap().unwrap();
assert_eq!(r1.context.unwrap().request_id, "predict-1");
assert_eq!(r2.context.unwrap().request_id, "predict-2");
server.abort();
}
#[tokio::test]
async fn pipelined_idle_activity_stays_balanced() {
let handler = OrderingHandler {
slow_delay: Duration::from_millis(20),
predict_order: Arc::new(Mutex::new(Vec::new())),
route_events: Arc::new(Mutex::new(Vec::new())),
};
let bound = ModelWorker::new(handler)
.bind_async(
ServeModelOptions::new(BindAddress::Tcp {
host: "127.0.0.1".to_string(),
port: 0,
})
.serve_options(ServeOptions {
idle_timeout: Some(Duration::from_millis(150)),
..ServeOptions::default()
}),
)
.await
.unwrap();
let (port, server) = spawn_bound_server(bound);
{
let channel = tonic::transport::Endpoint::from_shared(format!("http://127.0.0.1:{port}"))
.unwrap()
.connect()
.await
.unwrap();
let mut client =
rlmesh_proto::model::v1::model_service_client::ModelServiceClient::new(channel);
let (req_tx, req_rx) = mpsc::channel::<JoinRequest>(16);
let request_stream = tokio_stream::wrappers::ReceiverStream::new(req_rx);
let mut responses = client
.join(tonic::Request::new(request_stream))
.await
.unwrap()
.into_inner();
req_tx
.send(resolve_adapter_request("r", "cfg"))
.await
.unwrap();
let _ = responses.message().await.unwrap().unwrap();
for i in 0..5 {
req_tx
.send(predict_join_request("r", &format!("p{i}"), i == 0))
.await
.unwrap();
}
for _ in 0..5 {
let _ = responses.message().await.unwrap().unwrap();
}
drop(req_tx);
let _ = responses.message().await;
}
tokio::time::timeout(Duration::from_secs(2), server)
.await
.expect("server must idle-shut-down (balanced idle activity)")
.unwrap()
.unwrap();
}
#[tokio::test]
async fn idle_shutdown_arms_immediately_and_activity_extends_window() {
let (tx, mut rx) = mpsc::unbounded_channel();
let shutdown = tokio::spawn(async move {
rlmesh_grpc::lifecycle::wait_for_idle_shutdown(&mut rx, Duration::from_millis(25)).await;
});
tokio::time::sleep(Duration::from_millis(10)).await;
assert!(!shutdown.is_finished());
tx.send(rlmesh_grpc::lifecycle::IdleActivity::Started)
.expect("idle activity receiver should be open");
tx.send(rlmesh_grpc::lifecycle::IdleActivity::Finished)
.expect("idle activity receiver should be open");
tokio::time::sleep(Duration::from_millis(20)).await;
assert!(!shutdown.is_finished());
tokio::time::timeout(Duration::from_millis(100), shutdown)
.await
.unwrap()
.unwrap();
}
#[derive(Clone, Default)]
struct ChunkedFramesHandler;
fn smoke_box_action() -> spaces::SpaceValue {
spaces::SpaceValue::Box(
spaces::Tensor::from_vec(vec![0u8], vec![1], spaces::DType::Uint8).unwrap(),
)
}
#[async_trait]
impl ModelHandler for ChunkedFramesHandler {
async fn predict(&mut self, observation: ModelObservation) -> Result<Vec<spaces::SpaceValue>> {
Ok((0..observation.num_envs)
.map(|_| smoke_box_action())
.collect())
}
async fn predict_chunked(&mut self, observation: ModelObservation) -> Result<PredictFrames> {
let lanes = observation.num_envs;
Ok(PredictFrames {
actions: (0..lanes).map(|_| smoke_box_action()).collect(),
replay: vec![(0..lanes).map(|_| smoke_box_action()).collect()],
})
}
}
#[tokio::test]
async fn grouped_predict_carries_chunk_replay_frames() {
let handler = Arc::new(Mutex::new(ChunkedFramesHandler));
let configs = Arc::new(Mutex::new(HashMap::from([
(
"env-a".to_string(),
ModelRouteConfig {
env_contract: Some(Arc::new(SmokeEnv::new().env_contract)),
floor: None,
},
),
(
"env-b".to_string(),
ModelRouteConfig {
env_contract: Some(Arc::new(SmokeEnv::new().env_contract)),
floor: None,
},
),
])));
let response = handle_model_request(
JoinRequest {
kind: Some(join_request::Kind::GroupedPredict(GroupedPredictRequest {
groups: vec![
grouped_member("env-a", "p-a", "ep-a"),
grouped_member("env-b", "p-b", "ep-b"),
],
})),
request_id: "grouped-chunked".to_string(),
},
Arc::clone(&handler),
None,
Arc::clone(&configs),
Admission::now(),
)
.await;
let results = match response.kind {
Some(join_response::Kind::GroupedPredict(GroupedPredictResponse { results })) => results,
other => panic!("expected grouped predict response, got {other:?}"),
};
assert_eq!(results.len(), 2, "one result per group, in order");
for (result, env_id) in results.iter().zip(["env-a", "env-b"]) {
let response = expect_group_response(result);
assert_eq!(response.context.as_ref().unwrap().env_id, env_id);
assert_eq!(
response.actions.len(),
2,
"{env_id}: frame 0 plus one replay frame survive the grouped path"
);
}
}
#[tokio::test]
async fn model_serve_options_token_is_enforced_by_the_server() {
let predicts = Arc::new(AtomicUsize::new(0));
let bound = ModelWorker::new(SmokeModel {
predicts: Arc::clone(&predicts),
closes: Arc::new(AtomicUsize::new(0)),
})
.bind_async(
ServeModelOptions::new(BindAddress::Tcp {
host: "127.0.0.1".to_string(),
port: 0,
})
.serve_options(ServeOptions {
allow_remote_shutdown: true,
token: Some("serve-options-token".to_string()),
..ServeOptions::default()
}),
)
.await
.unwrap();
let (port, server) = spawn_bound_server(bound);
let address = format!("tcp://127.0.0.1:{port}");
let env_contract = SmokeEnv::new().env_contract;
let error = crate::RemoteModel::connect(&address, env_contract.clone())
.await
.err()
.expect("an untokened client must be rejected when ServeOptions sets a token");
assert!(
error.to_string().contains("invalid route token"),
"expected an unauthenticated rejection, got {error}"
);
let mut model =
crate::RemoteModel::connect_with_token(&address, "serve-options-token", env_contract)
.await
.unwrap();
model.reset(None);
model
.predict(spaces::SpaceValue::Box(
spaces::Tensor::from_vec(vec![5u8], vec![1], spaces::DType::Uint8).unwrap(),
))
.await
.unwrap();
assert_eq!(predicts.load(Ordering::SeqCst), 1);
drop(model);
let mut shutdown_client = rlmesh_grpc::ModelClient::connect(&address, "serve-options-token")
.await
.unwrap();
shutdown_client.handshake().await.unwrap();
assert!(shutdown_client.shutdown("done").await.unwrap().accepted);
shutdown_and_join(server).await;
}
#[tokio::test]
async fn run_local_options_token_reaches_a_token_protected_env() {
let bound = crate::EnvServer::new(SmokeEnv::new())
.bind_with_options(
BindAddress::Tcp {
host: "127.0.0.1".to_string(),
port: 0,
},
ServeOptions {
token: Some("env-token".to_string()),
..ServeOptions::default()
},
)
.await
.unwrap();
let address = bound.local_addr().to_string();
let server = tokio::spawn(async move { bound.serve().await });
let untokened = ModelWorker::new(SmokeModel {
predicts: Arc::new(AtomicUsize::new(0)),
closes: Arc::new(AtomicUsize::new(0)),
})
.run_local_async(RunLocalOptions::parse(&address).unwrap().for_episodes(1))
.await
.expect_err("a token-protected env must reject an untokened run_local");
assert!(
untokened.to_string().contains("invalid env token"),
"expected an unauthenticated rejection, got {untokened}"
);
let predicts = Arc::new(AtomicUsize::new(0));
ModelWorker::new(SmokeModel {
predicts: Arc::clone(&predicts),
closes: Arc::new(AtomicUsize::new(0)),
})
.run_local_async(
RunLocalOptions::parse(&address)
.unwrap()
.token("env-token")
.for_episodes(1),
)
.await
.unwrap();
assert!(predicts.load(Ordering::SeqCst) > 0);
server.abort();
}
#[derive(Clone)]
struct CapabilityProbeModel {
capabilities: HashMap<String, String>,
offered_history: Arc<Mutex<Option<bool>>>,
history_rows: Arc<Mutex<Vec<usize>>>,
}
#[tonic::async_trait]
impl rlmesh_proto::model::v1::model_service_server::ModelService for CapabilityProbeModel {
async fn handshake(
&self,
_request: tonic::Request<rlmesh_proto::model::v1::HandshakeRequest>,
) -> std::result::Result<
tonic::Response<rlmesh_proto::model::v1::HandshakeResponse>,
tonic::Status,
> {
Ok(tonic::Response::new(
rlmesh_proto::model::v1::HandshakeResponse {
base: Some(rlmesh_proto::core::v1::HandshakeResponse {
compatible: true,
peer_info: None,
capabilities: self.capabilities.clone(),
supported_workflow_editions: rlmesh_proto::supported_workflow_editions(),
error_message: None,
preferred_workflow_edition: String::new(),
}),
},
))
}
type JoinStream = tokio_stream::wrappers::ReceiverStream<
std::result::Result<rlmesh_proto::model::v1::JoinResponse, tonic::Status>,
>;
async fn join(
&self,
request: tonic::Request<tonic::Streaming<JoinRequest>>,
) -> std::result::Result<tonic::Response<Self::JoinStream>, tonic::Status> {
use rlmesh_proto::model::v1::{
JoinResponse, ObservationHistoryNeeds, PredictResponse, ReleaseAdapterResponse,
ResetAdapterResponse, ResolveAdapterResponse,
};
use tokio_stream::StreamExt;
let mut requests = request.into_inner();
let probe = self.clone();
let (tx, rx) = mpsc::channel(8);
let frame = rlmesh_grpc::wire::encode_batched_partial_values(
&[smoke_box_action()],
&SmokeEnv::new().action_space,
)
.unwrap();
tokio::spawn(async move {
while let Some(Ok(request)) = requests.next().await {
let kind = match request.kind {
Some(join_request::Kind::ResolveAdapter(resolve)) => {
*probe.offered_history.lock().await = Some(resolve.delivers_history);
join_response::Kind::ResolveAdapter(ResolveAdapterResponse {
native_chunk: None,
history: resolve.delivers_history.then(|| ObservationHistoryNeeds {
keys: vec!["observation".to_string()],
prunable: false,
}),
})
}
Some(join_request::Kind::Predict(predict)) => {
probe.history_rows.lock().await.push(predict.history.len());
join_response::Kind::Predict(PredictResponse {
context: predict.context,
actions: vec![frame.clone(), frame.clone()],
})
}
Some(join_request::Kind::ResetAdapter(_)) => {
join_response::Kind::ResetAdapter(ResetAdapterResponse {})
}
Some(join_request::Kind::ReleaseAdapter(_)) => {
join_response::Kind::ReleaseAdapter(ReleaseAdapterResponse {})
}
_ => break,
};
let response = JoinResponse {
request_id: request.request_id,
kind: Some(kind),
..Default::default()
};
if tx.send(Ok(response)).await.is_err() {
break;
}
}
});
Ok(tonic::Response::new(
tokio_stream::wrappers::ReceiverStream::new(rx),
))
}
async fn shutdown(
&self,
_request: tonic::Request<rlmesh_proto::model::v1::ShutdownRequest>,
) -> std::result::Result<
tonic::Response<rlmesh_proto::model::v1::ShutdownResponse>,
tonic::Status,
> {
Err(tonic::Status::unimplemented("probe model has no shutdown"))
}
}
async fn probe_history_offer(capabilities: &[&str]) -> (Option<bool>, Vec<usize>) {
let probe = CapabilityProbeModel {
capabilities: rlmesh_proto::capability_map(capabilities),
offered_history: Arc::new(Mutex::new(None)),
history_rows: Arc::new(Mutex::new(Vec::new())),
};
let listener = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
let address = format!("tcp://{}", listener.local_addr().unwrap());
let service = probe.clone();
let server = tokio::spawn(async move {
tonic::transport::Server::builder()
.add_service(
rlmesh_proto::model::v1::model_service_server::ModelServiceServer::new(service),
)
.serve_with_incoming(TcpListenerStream::new(listener))
.await
.unwrap()
});
let mut model = crate::RemoteModel::connect(&address, SmokeEnv::new().env_contract)
.await
.unwrap();
model.set_execution_horizon(2);
model.reset(None);
let observe = || {
spaces::SpaceValue::Box(
spaces::Tensor::from_vec(vec![5], vec![1], spaces::DType::Uint8).unwrap(),
)
};
for _ in 0..3 {
model.predict(observe()).await.unwrap();
}
model.close().await.unwrap();
drop(model);
server.abort();
let offered = *probe.offered_history.lock().await;
let rows = probe.history_rows.lock().await.clone();
(offered, rows)
}
#[tokio::test]
async fn remote_model_offers_history_only_to_a_model_that_advertises_it() {
let (offered, rows) = probe_history_offer(&[]).await;
assert_eq!(offered, Some(false));
assert_eq!(rows, vec![0, 0], "no history rows without the capability");
let (offered, rows) =
probe_history_offer(&[rlmesh_proto::capabilities::MODEL_OBSERVATION_HISTORY_V1]).await;
assert_eq!(offered, Some(true));
assert_eq!(
rows,
vec![0, 1],
"the replayed step rides the next predict as a history row"
);
}
#[tokio::test]
async fn served_legs_learn_this_builds_ceilings() {
use rlmesh_proto::capabilities::{
ENV_SUBSET_STEP, MODEL_CONCURRENT_PREDICT_V1, MODEL_OBSERVATION_HISTORY_V1,
};
use rlmesh_runtime::PeerCeiling;
let listener = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
let env_address = format!("tcp://{}", listener.local_addr().unwrap());
let env_server = tokio::spawn(async move {
tonic::transport::Server::builder()
.add_service(rlmesh_grpc::env::env_service(
crate::env::WireLaneAdapter::new(vec![SmokeEnv::new()]).unwrap(),
))
.serve_with_incoming(TcpListenerStream::new(listener))
.await
.unwrap()
});
let mut env = rlmesh_grpc::EnvClient::connect(&env_address).await.unwrap();
let env_ceiling = super::local::env_ceiling(&env.handshake().await.unwrap());
let bound = ModelWorker::new(SmokeModel {
predicts: Arc::new(AtomicUsize::new(0)),
closes: Arc::new(AtomicUsize::new(0)),
})
.bind_async(ServeModelOptions::new(BindAddress::Tcp {
host: "127.0.0.1".to_string(),
port: 0,
}))
.await
.unwrap();
let (port, model_server) = spawn_bound_server(bound);
let model = crate::RemoteModel::connect(
&format!("tcp://127.0.0.1:{port}"),
SmokeEnv::new().env_contract,
)
.await
.unwrap();
let model_ceiling = model.ceiling().clone();
drop(model);
model_server.abort();
env_server.abort();
let this_build = |capabilities: &[&str]| {
PeerCeiling::wire_v1(
rlmesh_proto::Edition::current(),
rlmesh_proto::capability_map(capabilities),
rlmesh_grpc::MAX_MESSAGE_SIZE,
)
};
assert_eq!(env_ceiling, this_build(&[ENV_SUBSET_STEP]));
assert_eq!(
model_ceiling,
this_build(&[MODEL_CONCURRENT_PREDICT_V1, MODEL_OBSERVATION_HISTORY_V1])
);
assert_eq!(env_ceiling.dtypes.len(), spaces::DType::ALL.len());
}