use std::sync::Arc;
use std::time::{Duration, Instant};
use async_trait::async_trait;
use rlmesh_grpc::wire::{
encode_batched_partial_values, env_contract_from_proto, env_contract_to_proto,
};
use rlmesh_proto::model::v1::{PredictRequest, ResetAdapterRequest};
use rlmesh_proto::{EndpointPhases, elapsed_ns};
use rlmesh_runtime::{
PeerCeiling, RuntimeDriver, RuntimeEnv, RuntimeEnvReset, RuntimeEnvStep, RuntimeError,
RuntimeHooks, RuntimeModel, RuntimeModelPrediction, RuntimeReport, RuntimeSessionSpec,
};
use super::handler::{ModelHandler, PredictFrames};
use super::wire::{
ModelAction, check_actions_conform, encode_replay_frames, model_action_to_endpoint_response,
model_observation_from_endpoint_request,
};
use crate::{Error, Result, spaces};
pub(super) async fn run_local<H>(
handler: &mut H,
options: crate::RunLocalOptions,
cancellation: tokio_util::sync::CancellationToken,
hooks: Arc<dyn RuntimeHooks>,
) -> Result<RuntimeReport>
where
H: ModelHandler + 'static,
{
let mut env = rlmesh_grpc::EnvClient::connect_with_token(
&options.env_address.to_string(),
&options.token,
)
.await
.map_err(Error::from)?;
env.declare_workflow_edition(options.workflow_edition.clone());
let handshake = env.handshake().await.map_err(Error::from)?;
let env_ceiling = env_ceiling(&handshake);
let subset_step = rlmesh_proto::has_capability(
&handshake.capabilities,
rlmesh_proto::capabilities::ENV_SUBSET_STEP,
);
let env_contract = env_contract_from_proto(handshake.env_contract)
.map_err(|err| Error::Internal(format!("invalid spaces spec from env: {err}")))?;
if env_contract.action_space.is_none() {
return Err(Error::Internal(
"env contract has no action_space; a model cannot encode actions without it"
.to_string(),
));
}
let env_id = crate::mint_id();
let num_envs = handshake.num_envs;
let session_id = format!("local-{}", std::process::id());
if num_envs > 1 && !subset_step && options.execution_horizon > 1 {
return Err(Error::Internal(format!(
"execution_horizon={} cannot be combined with a lockstep vector env (num_envs={num_envs}): \
chunk replay is whole-batch, so one lane's episode end discards every lane's \
buffered frames. Use num_envs=1, a lane endpoint, or execution_horizon=1.",
options.execution_horizon,
)));
}
let mut wants_history = false;
if let Some(route_setup) = handler.route_setup() {
let needs = route_setup
.resolve_adapter(
&env_id,
&env_contract,
crate::model::ResolveOptions {
execution_horizon: options.execution_horizon,
delivers_history: true,
},
)
.await?;
wants_history = needs.history.is_some();
}
let spec = RuntimeSessionSpec {
session_id,
env_id,
env_component_id: "local-env".to_string(),
model_component_id: "local-model".to_string(),
workflow_edition: handshake.workflow_edition,
env_contract: env_contract_to_proto(&env_contract),
num_envs,
base_seed: options.base_seed,
episode_seeds: options.episode_seeds,
max_episodes: options.max_episodes,
trial_index_base: options.trial_index_base,
max_episode_steps: options.max_episode_steps,
max_episode_seconds: options.max_episode_seconds,
close_env_on_end: options.close_env,
subset_step,
limits: Default::default(),
env_ceiling: Some(env_ceiling),
model_ceiling: None,
};
let env = EnvClientRuntimeEnv::new(env);
let model = ModelHandlerRuntimeModel::new(handler, env_contract).with_history(wants_history);
RuntimeDriver::new(spec, env, model, hooks)
.with_prefetch(options.prefetch_lead)
.run_with_cancellation_reason(cancellation, "interrupted by the host (signal)")
.await
.map_err(run_error)
}
pub(super) fn env_ceiling(handshake: &rlmesh_grpc::EnvHandshake) -> PeerCeiling {
PeerCeiling::wire_v1(
PeerCeiling::highest_shared_edition(&handshake.supported_workflow_editions)
.unwrap_or(handshake.workflow_edition),
handshake.capabilities.clone(),
rlmesh_grpc::MAX_MESSAGE_SIZE,
)
}
fn model_rpc(error: Error) -> RuntimeError {
RuntimeError::model_rpc_with_recoverability("local-model", error.is_recoverable(), error)
}
fn run_error(error: RuntimeError) -> Error {
let recoverable = error.is_recoverable();
let message = error.to_string();
match error {
RuntimeError::ModelRpc { source, .. } => source
.and_then(|source| source.downcast::<Error>().ok())
.map_or_else(
|| {
if recoverable {
Error::model_recoverable(message)
} else {
Error::model(message)
}
},
|error| *error,
),
RuntimeError::EnvRpc { source, .. } => match source
.and_then(|source| source.downcast::<rlmesh_grpc::error::Error>().ok())
.map(|error| Error::from(*error))
{
Some(Error::Connection(_)) => Error::Connection(message),
Some(Error::Timeout(timeout)) => Error::Timeout(timeout),
Some(Error::Environment(env)) => {
Error::Environment(crate::EnvironmentError { message, ..env })
}
_ => Error::Environment(crate::EnvironmentError {
code: crate::ErrorCode::Internal,
message,
is_recoverable: recoverable,
}),
},
RuntimeError::OperationTimeout { timeout, .. } => Error::Timeout(timeout),
_ => Error::Internal(message),
}
}
#[derive(Clone)]
pub struct EnvClientRuntimeEnv {
inner: rlmesh_grpc::EnvClient,
}
impl EnvClientRuntimeEnv {
pub fn new(client: rlmesh_grpc::EnvClient) -> Self {
Self { inner: client }
}
pub fn into_inner(self) -> rlmesh_grpc::EnvClient {
self.inner
}
}
#[async_trait]
impl RuntimeEnv for EnvClientRuntimeEnv {
async fn reset(
&mut self,
request: rlmesh_proto::env::v1::ResetRequest,
) -> std::result::Result<RuntimeEnvReset, rlmesh_runtime::RuntimeError> {
let response = self.inner.reset(request).await.map_err(|err| {
let recoverable = err.is_recoverable();
rlmesh_runtime::RuntimeError::env_rpc_with_recoverability(
"env.reset",
0,
recoverable,
err,
)
})?;
Ok(RuntimeEnvReset {
response,
endpoint_total_ns: self.inner.take_last_endpoint_total_ns(),
phases: self.inner.take_last_phases(),
})
}
async fn step(
&mut self,
request: rlmesh_proto::env::v1::StepRequest,
) -> std::result::Result<RuntimeEnvStep, rlmesh_runtime::RuntimeError> {
let response = self.inner.step(request).await.map_err(|err| {
let recoverable = err.is_recoverable();
rlmesh_runtime::RuntimeError::env_rpc_with_recoverability(
"env.step",
0,
recoverable,
err,
)
})?;
Ok(RuntimeEnvStep {
response,
endpoint_total_ns: self.inner.take_last_endpoint_total_ns(),
phases: self.inner.take_last_phases(),
})
}
async fn close(&mut self, timeout: Duration) -> std::result::Result<(), String> {
let close = self.inner.close();
tokio::time::timeout(timeout, close)
.await
.map_err(|err| err.to_string())?
.map(|_| ())
.map_err(|err| err.to_string())
}
}
pub struct ModelHandlerRuntimeModel<'a, H> {
handler: Arc<tokio::sync::Mutex<&'a mut H>>,
env_contract: Arc<spaces::EnvContract>,
wants_history: bool,
}
impl<'a, H> ModelHandlerRuntimeModel<'a, H> {
pub fn new(handler: &'a mut H, env_contract: spaces::EnvContract) -> Self {
Self {
handler: Arc::new(tokio::sync::Mutex::new(handler)),
env_contract: Arc::new(env_contract),
wants_history: false,
}
}
pub fn with_history(mut self, wants_history: bool) -> Self {
self.wants_history = wants_history;
self
}
}
struct PreparedGroup {
route: crate::model::types::ModelRouteContext,
num_envs: usize,
}
#[async_trait]
impl<H> RuntimeModel for ModelHandlerRuntimeModel<'_, H>
where
H: ModelHandler + 'static,
{
fn wants_history(&self) -> bool {
self.wants_history
}
fn fuses_predicts(&self) -> bool {
true
}
async fn predict(
&self,
request: PredictRequest,
) -> std::result::Result<RuntimeModelPrediction, rlmesh_runtime::RuntimeError> {
self.predict_group(vec![request])
.await
.pop()
.unwrap_or_else(|| {
Err(rlmesh_runtime::RuntimeError::model_rpc(
"local-model",
Error::model("predict_grouped returned no result"),
))
})
}
async fn predict_group(
&self,
requests: Vec<PredictRequest>,
) -> Vec<std::result::Result<RuntimeModelPrediction, rlmesh_runtime::RuntimeError>> {
let started = Instant::now();
let model_err = model_rpc;
let Some(action_space) = self.env_contract.action_space.clone() else {
return requests
.iter()
.map(|_| {
Err(model_err(Error::model(
"model route contract missing action space",
)))
})
.collect();
};
let mut batch = Vec::with_capacity(requests.len());
let mut prepared: Vec<std::result::Result<PreparedGroup, rlmesh_runtime::RuntimeError>> =
Vec::with_capacity(requests.len());
for request in requests {
match model_observation_from_endpoint_request(request) {
Ok(mut observation) => {
let num_envs = observation.route.episodes.len().max(1);
observation.env_contract = Some(Arc::clone(&self.env_contract));
observation.num_envs = num_envs;
prepared.push(Ok(PreparedGroup {
route: observation.route.clone(),
num_envs,
}));
batch.push(observation);
}
Err(err) => prepared.push(Err(model_err(err))),
}
}
let decode_ns = elapsed_ns(started);
let mut handler = self.handler.lock().await;
let call_started = Instant::now();
let mut frames = handler.predict_grouped(batch).await.into_iter();
let user_ns = elapsed_ns(call_started);
let adapter_ns = handler.take_adapter_ns();
let held = handler.held_state();
drop(handler);
let encode_started = Instant::now();
prepared
.into_iter()
.map(|group| {
let PreparedGroup { route, num_envs } = group?;
let PredictFrames { actions, replay } = frames
.next()
.ok_or_else(|| {
model_err(Error::model(
"predict_grouped returned fewer results than prepared groups",
))
})?
.map_err(model_err)?;
if actions.len() != num_envs {
return Err(model_err(Error::model(format!(
"predict returned {} actions for {num_envs} lanes",
actions.len()
))));
}
check_actions_conform(&action_space, &actions).map_err(model_err)?;
let frame0 = encode_batched_partial_values(&actions, &action_space)
.map_err(|err| model_err(Error::model(err.to_string())))?;
let replay_frames =
encode_replay_frames(&replay, num_envs, &action_space).map_err(model_err)?;
let mut wire_actions = Vec::with_capacity(1 + replay_frames.len());
wire_actions.push(frame0);
wire_actions.extend(replay_frames);
Ok(RuntimeModelPrediction {
response: model_action_to_endpoint_response(ModelAction {
actions: wire_actions,
route,
}),
endpoint_total_ns: Some(elapsed_ns(started)),
phases: EndpointPhases {
decode_ns,
user_ns,
encode_ns: elapsed_ns(encode_started),
adapter_ns,
held_episodes: held
.map(|held| held.episodes.min(u64::from(u32::MAX)) as u32),
held_state_bytes: held.map(|held| held.bytes),
..EndpointPhases::default()
},
group_size: None,
})
})
.collect()
}
async fn reset_adapter(
&self,
request: ResetAdapterRequest,
) -> std::result::Result<(), RuntimeError> {
let env_id = request
.context
.map(|context| context.env_id)
.unwrap_or_default();
let mut handler = self.handler.lock().await;
if let Some(route_setup) = handler.route_setup() {
route_setup
.reset_adapter(&env_id, &request.episode_ids)
.await
.map_err(model_rpc)?;
}
handler
.reset_adapter(&env_id, request.episode_ids)
.await
.map_err(model_rpc)
}
}