use std::collections::{BTreeSet, HashMap};
use std::sync::{Arc, Mutex};
use async_trait::async_trait;
use rlmesh_adapters::v1::{
FrameBuffers, ObsPlan, Value, apply_actions, assemble_obs, space_value_to_obs_map, split_chunk,
};
use super::handler::{ModelHandler, ModelRouteSetup, PredictFrames};
use super::predict_fn::{PredictFn, RouteConfig, RouteResolver};
use super::types::ModelObservation;
use crate::spaces::{EnvContract, SpaceKind, SpaceValue};
use crate::{Error, Result};
struct RouteEntry {
config: RouteConfig,
buffers: FrameBuffers,
}
type Routes = Arc<Mutex<HashMap<String, Arc<Mutex<RouteEntry>>>>>;
pub struct AdaptedModelHandler {
predict: Arc<dyn PredictFn>,
resolver: Option<Arc<dyn RouteResolver>>,
routes: Routes,
}
impl AdaptedModelHandler {
pub fn new(predict: Arc<dyn PredictFn>, resolver: Option<Arc<dyn RouteResolver>>) -> Self {
Self {
predict,
resolver,
routes: Arc::new(Mutex::new(HashMap::new())),
}
}
fn entry(&self, env_id: &str) -> Option<Arc<Mutex<RouteEntry>>> {
self.routes
.lock()
.expect("routes map poisoned")
.get(env_id)
.cloned()
}
}
fn obs_keys(config: &RouteConfig) -> BTreeSet<String> {
let referenced = config.adapter.referenced_obs_keys();
let has_customs = config
.adapter
.obs_plans
.iter()
.any(|plan| matches!(plan, ObsPlan::Custom(_)));
if !has_customs {
return referenced;
}
match config.observation_space.spec.as_ref() {
Some(SpaceKind::Dict(dict)) => dict.keys.iter().cloned().collect(),
_ => [".".to_owned()].into_iter().collect(),
}
}
fn predict_route(
entry: &Arc<Mutex<RouteEntry>>,
predict: &Arc<dyn PredictFn>,
observation: ModelObservation,
) -> Result<PredictFrames> {
let episode_ids = &observation.route.episode_ids;
let num_envs = observation.num_envs;
observation.ensure_decodable()?;
let mut guard = entry.lock().expect("route entry poisoned");
let RouteEntry { config, buffers } = &mut *guard;
let referenced = obs_keys(config);
let horizon = config.execution_horizon;
let customs: &dyn rlmesh_adapters::v1::CustomTransform = config.customs.as_ref();
let encodings: &dyn rlmesh_adapters::v1::EncodingTransform = config.encodings.as_ref();
let has_chunk = predict.has_chunk();
let decoded = observation.decoded_lanes()?;
let mut inputs: Vec<Value> = Vec::with_capacity(num_envs);
for (index, lane) in decoded.iter().enumerate() {
let episode_id = episode_ids
.get(index)
.map(String::as_str)
.filter(|id| !id.is_empty())
.ok_or_else(|| {
Error::model(format!(
"predict request lane {index} has no episode_id (num_envs={num_envs}, \
episode_ids={}); every lane must carry a non-empty episode_id",
episode_ids.len()
))
})?;
let raw = space_value_to_obs_map(lane, &config.observation_space, &referenced)?;
inputs.push(assemble_obs(
&config.adapter,
&raw,
episode_id,
buffers,
customs,
encodings,
)?);
}
let lane_raw_steps: Vec<Vec<Value>> = if horizon > 1 && predict.has_chunk_batch() {
let chunks = predict.predict_chunk_batch(inputs, horizon)?;
if chunks.len() != num_envs {
return Err(Error::model(format!(
"predict_chunk_batch returned {} chunks for {num_envs} lanes",
chunks.len()
)));
}
chunks
.into_iter()
.map(|chunk| -> Result<Vec<Value>> {
Ok(split_chunk(chunk)?
.into_iter()
.take(horizon as usize)
.collect())
})
.collect::<Result<Vec<_>>>()?
} else if horizon > 1 && has_chunk {
inputs
.into_iter()
.map(|input| -> Result<Vec<Value>> {
let chunk = predict.predict_chunk(input, horizon)?.ok_or_else(|| {
Error::model(
"model reports a chunk corner (has_chunk) but predict_chunk returned None",
)
})?;
Ok(split_chunk(chunk)?
.into_iter()
.take(horizon as usize)
.collect())
})
.collect::<Result<Vec<_>>>()?
} else if predict.has_batch() {
let actions = predict.predict_batch(inputs)?;
if actions.len() != num_envs {
return Err(Error::model(format!(
"predict_batch returned {} actions for {num_envs} lanes",
actions.len()
)));
}
actions.into_iter().map(|action| vec![action]).collect()
} else {
inputs
.into_iter()
.map(|input| -> Result<Vec<Value>> { Ok(vec![predict.predict(input)?]) })
.collect::<Result<Vec<_>>>()?
};
let mut frame0 = Vec::with_capacity(num_envs);
let mut lane_replays: Vec<Vec<SpaceValue>> = Vec::with_capacity(num_envs);
for raw_steps in lane_raw_steps {
let mut applied = raw_steps
.into_iter()
.map(|raw_action| {
apply_actions(&config.adapter, raw_action, &config.action_space, encodings)
})
.collect::<std::result::Result<Vec<SpaceValue>, _>>()?
.into_iter();
let first = applied
.next()
.ok_or_else(|| Error::model("a chunked model returned an empty action chunk"))?;
frame0.push(first);
lane_replays.push(applied.collect());
}
let replay_len = lane_replays.iter().map(Vec::len).min().unwrap_or(0);
let mut replay = Vec::with_capacity(replay_len);
for step in 0..replay_len {
let mut per_lane = Vec::with_capacity(num_envs);
for lane in &lane_replays {
per_lane.push(lane[step].clone());
}
replay.push(per_lane);
}
Ok(PredictFrames {
actions: frame0,
replay,
})
}
const PROBE_SEED_A: u64 = 0x5EED_000A;
const PROBE_SEED_B: u64 = 0x5EED_000B;
const PROBE_TOLERANCE: f64 = 8.0;
const PROBE_ATOL: f64 = 1e-6;
fn probe_model_internal_state(predict: &Arc<dyn PredictFn>, config: &RouteConfig) -> Result<()> {
let space = &config.observation_space;
let referenced = obs_keys(config);
let assemble = |seed: u64, episode: &str| -> Result<rlmesh_adapters::v1::Value> {
let sampled = rlmesh_spaces::sample_seeded(space, seed)
.map_err(|err| Error::Internal(format!("probe sample failed: {err}")))?;
let mut scratch = FrameBuffers::new();
let raw = space_value_to_obs_map(&sampled, space, &referenced)?;
Ok(assemble_obs(
&config.adapter,
&raw,
episode,
&mut scratch,
config.customs.as_ref(),
config.encodings.as_ref(),
)?)
};
let a = assemble(PROBE_SEED_A, "probe-a")?;
let b = assemble(PROBE_SEED_B, "probe-b")?;
let floor = rlmesh_adapters::v1::value_max_abs_diff(
&predict.predict(a.clone())?,
&predict.predict(a.clone())?,
)
.unwrap_or(f64::INFINITY);
let before = predict.predict(a.clone())?;
let _ = predict.predict(b)?;
let after = predict.predict(a)?;
let delta = rlmesh_adapters::v1::value_max_abs_diff(&before, &after).unwrap_or(0.0);
if delta > (floor * PROBE_TOLERANCE).max(PROBE_ATOL) {
return Err(Error::model(
"this model carries internal state across predict() calls, so it cannot be \
served against a vectorized route (num_envs>1): one shared model instance \
across lanes would interleave their state. Serve it against num_envs=1, or \
make predict() pure (move per-step state into the adapter).",
));
}
Ok(())
}
#[async_trait]
impl ModelHandler for AdaptedModelHandler {
async fn predict(&mut self, observation: ModelObservation) -> Result<Vec<SpaceValue>> {
Ok(self.predict_chunked(observation).await?.actions)
}
async fn predict_chunked(&mut self, observation: ModelObservation) -> Result<PredictFrames> {
let entry = self.entry(&observation.route.env_id);
let predict = Arc::clone(&self.predict);
tokio::task::spawn_blocking(move || match entry {
Some(entry) => predict_route(&entry, &predict, observation),
None => Ok(PredictFrames {
actions: predict.predict_spec_less(observation)?,
replay: Vec::new(),
}),
})
.await
.map_err(|err| Error::Internal(format!("predict task panicked: {err}")))?
}
fn route_setup(&self) -> Option<Arc<dyn ModelRouteSetup>> {
let resolver = self.resolver.clone()?;
Some(Arc::new(AdaptedRouteSetup {
resolver,
routes: Arc::clone(&self.routes),
predict: Arc::clone(&self.predict),
}))
}
async fn reset_adapter(&mut self, env_id: &str, episode_ids: Vec<String>) -> Result<()> {
if let Some(entry) = self.entry(env_id) {
let mut guard = entry.lock().expect("route entry poisoned");
if episode_ids.is_empty() {
guard.buffers.clear();
} else {
for episode_id in &episode_ids {
guard.buffers.evict(episode_id);
}
}
}
let count = episode_ids.len();
if count == 0 {
return Ok(());
}
let predict = Arc::clone(&self.predict);
tokio::task::spawn_blocking(move || {
for _ in 0..count {
predict.on_episode_end()?;
}
Ok(())
})
.await
.map_err(|err| Error::Internal(format!("on_episode_end task panicked: {err}")))?
}
async fn on_close(&mut self) -> Result<()> {
for entry in self.routes.lock().expect("routes map poisoned").values() {
let mut guard = entry.lock().expect("route entry poisoned");
guard.buffers.clear();
}
let predict = Arc::clone(&self.predict);
tokio::task::spawn_blocking(move || predict.on_close())
.await
.map_err(|err| Error::Internal(format!("on_close task panicked: {err}")))?
}
}
struct AdaptedRouteSetup {
resolver: Arc<dyn RouteResolver>,
routes: Routes,
predict: Arc<dyn PredictFn>,
}
#[async_trait]
impl ModelRouteSetup for AdaptedRouteSetup {
async fn resolve_adapter(
&self,
env_id: &str,
env_contract: &EnvContract,
execution_horizon: u32,
) -> Result<()> {
let Some(mut config) = self.resolver.resolve(env_id, env_contract).await? else {
return Ok(());
};
for note in config.adapter.advisories() {
tracing::warn!(env_id = %env_id, "adapter advisory: {note}");
}
config.execution_horizon = execution_horizon.max(1);
if config.execution_horizon > 1 && !self.predict.has_chunk() {
tracing::warn!(
env_id = %env_id,
execution_horizon = config.execution_horizon,
"runtime pinned execution_horizon > 1 but the model defines no chunk corner \
(predict_chunk); chunking is inactive — the model re-plans every step",
);
}
if config.execution_horizon > 1
&& let Some((key, depth)) = config.adapter.stacks().into_iter().next()
{
return Err(Error::model(format!(
"frame-stacking (input '{key}' stack={depth}) cannot be combined with action \
chunking (execution_horizon={}): the engine assembles observations once per chunk, \
so the frame window would hold only decision-point frames. Use stack=1 or \
execution_horizon=1.",
config.execution_horizon,
)));
}
let config = if env_contract.num_envs > 1 {
let predict = Arc::clone(&self.predict);
tokio::task::spawn_blocking(move || -> Result<RouteConfig> {
probe_model_internal_state(&predict, &config)?;
Ok(config)
})
.await
.map_err(|err| Error::Internal(format!("probe task panicked: {err}")))??
} else {
config
};
let entry = Arc::new(Mutex::new(RouteEntry {
config,
buffers: FrameBuffers::new(),
}));
self.routes
.lock()
.expect("routes map poisoned")
.insert(env_id.to_string(), entry);
Ok(())
}
async fn release_adapter(&self, env_id: &str) -> Result<()> {
self.routes
.lock()
.expect("routes map poisoned")
.remove(env_id);
Ok(())
}
}