use std::collections::HashMap;
use std::sync::Arc;
use std::time::{Duration, Instant};
use rlmesh_grpc::lifecycle::{
ActivityFinishedGuard, IdleActivity, await_close_with_timeout, start_idle_shutdown,
};
use rlmesh_grpc::wire::env_spec_from_proto;
use rlmesh_proto::model::v1::{
CloseParticipantResponse, GroupedPredictRequest, GroupedPredictResponse, GroupedPredictResult,
HandshakeRequest, HandshakeResponse, JoinRequest, JoinResponse, ModelError, ModelErrorCode,
ObservationHistoryNeeds, PredictRequest, PredictResponse, ReleaseAdapterResponse,
ResetAdapterResponse, ResolveAdapterRequest, ResolveAdapterResponse, ShutdownRequest,
ShutdownResponse, grouped_predict_result, join_request, join_response,
model_service_server::{ModelService as ModelServiceTrait, ModelServiceServer},
};
use rlmesh_proto::{
Edition, EndpointPhases, capabilities, capability_map, declared_workflow_edition, elapsed_ns,
evaluate_handshake, generation_mismatch_message, peer_info, supported_workflow_editions,
};
use tokio::sync::{Mutex, mpsc};
use tokio_stream::StreamExt;
use tonic::{Request, Response, Status, Streaming};
use super::handler::{ModelHandler, ModelRouteSetup, PredictFrames, ResolveOptions, RouteNeeds};
use super::types::{ModelObservation, ModelRouteContext};
use super::wire::{
ModelAction, check_actions_conform, encode_replay_frames, model_action_to_endpoint_response,
model_endpoint_total_ns, model_error, model_error_from_error, model_error_value,
model_observation_from_endpoint_request,
};
use crate::bound::BoundListener;
use crate::{BindAddress, Error, Result, ServeOptions, spaces};
pub struct BoundModelServer {
listener: BoundListener,
router: tonic::transport::server::Router,
shutdown: rlmesh_grpc::lifecycle::ShutdownTrigger,
handler: Arc<Mutex<dyn ModelHandler>>,
local_addr: BindAddress,
drain_timeout: Option<Duration>,
close_timeout: Option<Duration>,
}
impl BoundModelServer {
pub fn local_addr(&self) -> &BindAddress {
&self.local_addr
}
pub fn shutdown_trigger(&self) -> rlmesh_grpc::lifecycle::ShutdownTrigger {
self.shutdown.clone()
}
pub async fn serve(self) -> Result<()> {
let serve_result = self
.listener
.serve(self.router, self.shutdown, self.drain_timeout)
.await;
let close_result = close_model(self.handler, self.close_timeout).await;
crate::error::join_results(serve_result, close_result, "model server failed")
}
}
pub(super) async fn bind_model_with_options<H>(
handler: H,
address: BindAddress,
token: &str,
options: ServeOptions,
) -> Result<BoundModelServer>
where
H: ModelHandler + 'static,
{
let handler = Arc::new(Mutex::new(handler));
let route_setup = handler.lock().await.route_setup();
let shutdown = rlmesh_grpc::lifecycle::ShutdownTrigger::new();
let activity_tx = start_idle_shutdown(options.idle_timeout, shutdown.clone());
let drain_timeout = options.drain_timeout;
let close_timeout = options.close_timeout;
let service = model_service(
Arc::clone(&handler),
route_setup,
token.to_string(),
activity_tx,
shutdown.clone(),
options,
);
let listener = BoundListener::bind(address).await?;
let local_addr = listener.local_addr()?;
let (_health_reporter, health_service) = rlmesh_grpc::health::serving_health_service().await;
let router = tonic::transport::Server::builder()
.add_service(health_service)
.add_service(service);
let handler: Arc<Mutex<dyn ModelHandler>> = handler;
Ok(BoundModelServer {
listener,
router,
shutdown,
handler,
local_addr,
drain_timeout,
close_timeout,
})
}
async fn close_model(
handler: Arc<Mutex<dyn ModelHandler>>,
close_timeout: Option<Duration>,
) -> Result<()> {
let close = async { handler.lock().await.on_close().await };
await_close_with_timeout(close, close_timeout)
.await
.map_err(Error::Timeout)?
}
struct ServedModelServer<H> {
handler: Arc<Mutex<H>>,
route_setup: Option<Arc<dyn ModelRouteSetup>>,
route_configs: Arc<Mutex<HashMap<String, ModelRouteConfig>>>,
token: String,
activity_tx: Option<mpsc::UnboundedSender<IdleActivity>>,
shutdown: rlmesh_grpc::lifecycle::ShutdownTrigger,
serve_options: ServeOptions,
declared_workflow_edition: Option<String>,
}
#[derive(Debug, Clone)]
pub(super) struct ModelRouteConfig {
pub(super) env_contract: Option<Arc<spaces::EnvContract>>,
pub(super) floor: Option<RouteFloor>,
}
#[derive(Debug, Clone)]
pub(super) struct RouteFloor {
pub(super) edition: Edition,
pub(super) selected_workflow_edition: String,
}
fn declared_workflow_edition_of(serve_options: &ServeOptions) -> Option<String> {
serve_options
.workflow_edition
.as_deref()
.map(str::trim)
.filter(|edition| !edition.is_empty())
.map(str::to_string)
}
fn model_service<H>(
handler: Arc<Mutex<H>>,
route_setup: Option<Arc<dyn ModelRouteSetup>>,
token: String,
activity_tx: Option<mpsc::UnboundedSender<IdleActivity>>,
shutdown: rlmesh_grpc::lifecycle::ShutdownTrigger,
serve_options: ServeOptions,
) -> ModelServiceServer<ServedModelServer<H>>
where
H: ModelHandler + 'static,
{
let declared_workflow_edition = declared_workflow_edition_of(&serve_options);
ModelServiceServer::new(ServedModelServer {
handler,
route_setup,
route_configs: Arc::new(Mutex::new(HashMap::new())),
token,
activity_tx,
shutdown,
serve_options,
declared_workflow_edition,
})
.max_decoding_message_size(rlmesh_grpc::MAX_MESSAGE_SIZE)
.max_encoding_message_size(rlmesh_grpc::MAX_MESSAGE_SIZE)
}
#[tonic::async_trait]
impl<H> ModelServiceTrait for ServedModelServer<H>
where
H: ModelHandler + 'static,
{
async fn handshake(
&self,
request: Request<HandshakeRequest>,
) -> std::result::Result<Response<HandshakeResponse>, Status> {
self.authenticate(&request)?;
let request = request
.into_inner()
.base
.ok_or_else(|| Status::invalid_argument("handshake request missing base"))?;
let compatible = evaluate_handshake(&request.protocol_generation);
Ok(Response::new(HandshakeResponse {
base: Some(rlmesh_proto::core::v1::HandshakeResponse {
compatible,
peer_info: Some(peer_info("rlmesh-model")),
error_message: (!compatible)
.then(|| generation_mismatch_message(&request.protocol_generation)),
capabilities: capability_map(&[
capabilities::MODEL_CONCURRENT_PREDICT_V1,
capabilities::MODEL_OBSERVATION_HISTORY_V1,
]),
supported_workflow_editions: supported_workflow_editions(),
preferred_workflow_edition: declared_workflow_edition(
self.declared_workflow_edition.as_deref(),
)
.to_string(),
}),
}))
}
type JoinStream =
tokio_stream::wrappers::ReceiverStream<std::result::Result<JoinResponse, Status>>;
async fn join(
&self,
request: Request<Streaming<JoinRequest>>,
) -> std::result::Result<Response<Self::JoinStream>, Status> {
self.authenticate(&request)?;
let mut request_stream = request.into_inner();
let handler = Arc::clone(&self.handler);
let route_setup = self.route_setup.clone();
let route_configs = Arc::clone(&self.route_configs);
let activity_tx = self.activity_tx.clone();
let concurrency = self
.serve_options
.predict_concurrency
.unwrap_or(rlmesh_grpc::DEFAULT_PREDICT_CONCURRENCY)
.max(1);
let semaphore = Arc::new(tokio::sync::Semaphore::new(concurrency));
let (tx, rx) = tokio::sync::mpsc::channel::<std::result::Result<JoinResponse, Status>>(64);
let (read_tx, mut read_rx) = tokio::sync::mpsc::channel::<(JoinRequest, Instant)>(1);
tokio::spawn(async move {
while let Some(request_result) = request_stream.next().await {
let request = match request_result {
Ok(request) => request,
Err(error) => {
log_join_stream_error(&error);
break;
}
};
let close_after = matches!(request.kind, Some(join_request::Kind::Close(_)));
if read_tx.send((request, Instant::now())).await.is_err() || close_after {
break;
}
}
});
tokio::spawn(async move {
let mut route_tails = RouteTails::new();
while let Some((request, arrived_at)) = read_rx.recv().await {
let close_after = matches!(request.kind, Some(join_request::Kind::Close(_)));
let route_key = join_request_route_key(&request);
let (gate, dones): (RequestGate, Vec<tokio::sync::oneshot::Sender<()>>) =
if close_after {
route_tails.close_all_gate()
} else if let Some(keys) = grouped_predict_route_keys(&request) {
route_tails.next_multi_keyed_gate(&keys)
} else {
route_tails.next_keyed_gate(route_key.as_deref())
};
let permit = match Arc::clone(&semaphore).acquire_owned().await {
Ok(permit) => permit,
Err(_) => break,
};
if let Some(activity_tx) = &activity_tx {
let _ = activity_tx.send(IdleActivity::Started);
}
let activity_guard = ActivityFinishedGuard::new(activity_tx.clone());
let handler = Arc::clone(&handler);
let route_setup = route_setup.clone();
let route_configs = Arc::clone(&route_configs);
let tx = tx.clone();
let semaphore = Arc::clone(&semaphore);
tokio::spawn(async move {
let _permit = permit;
let _activity_guard = activity_guard;
gate.wait().await;
let admission = Admission {
arrived_at,
in_flight: (concurrency - semaphore.available_permits()) as u32,
};
let response = handle_model_request(
request,
handler,
route_setup,
route_configs,
admission,
)
.await;
for done in dones {
let _ = done.send(());
}
if tx.send(Ok(response)).await.is_err() {
tracing::warn!(
"model join response receiver closed before response could be delivered"
);
}
});
if close_after {
break;
}
}
});
Ok(Response::new(tokio_stream::wrappers::ReceiverStream::new(
rx,
)))
}
async fn shutdown(
&self,
request: Request<ShutdownRequest>,
) -> std::result::Result<Response<ShutdownResponse>, Status> {
self.authenticate(&request)?;
let request = request
.into_inner()
.base
.ok_or_else(|| Status::invalid_argument("shutdown request missing base"))?;
if !self.serve_options.allow_remote_shutdown {
return Ok(Response::new(ShutdownResponse {
base: Some(rlmesh_proto::core::v1::ShutdownResponse {
accepted: false,
message: "remote shutdown is disabled for this model endpoint".to_string(),
}),
}));
}
self.shutdown.trigger(if request.reason.is_empty() {
"remote shutdown".to_string()
} else {
request.reason.clone()
});
Ok(Response::new(ShutdownResponse {
base: Some(rlmesh_proto::core::v1::ShutdownResponse {
accepted: true,
message: if request.reason.is_empty() {
"shutdown accepted".to_string()
} else {
format!("shutdown accepted: {}", request.reason)
},
}),
}))
}
}
#[derive(Clone, Copy)]
pub(super) struct Admission {
arrived_at: Instant,
in_flight: u32,
}
impl Admission {
fn depth_only(self) -> EndpointPhases {
EndpointPhases {
in_flight: self.in_flight,
..EndpointPhases::default()
}
}
#[cfg(test)]
pub(super) fn now() -> Self {
Self {
arrived_at: Instant::now(),
in_flight: 0,
}
}
}
pub(super) async fn handle_model_request<H: ModelHandler + 'static>(
request: JoinRequest,
handler: Arc<Mutex<H>>,
route_setup: Option<Arc<dyn ModelRouteSetup>>,
route_configs: Arc<Mutex<HashMap<String, ModelRouteConfig>>>,
admission: Admission,
) -> JoinResponse {
let request_id = request.request_id.clone();
let started_at = Instant::now();
let mut phases = EndpointPhases::default();
let kind = match request.kind {
Some(join_request::Kind::ResolveAdapter(request)) => {
handle_resolve_adapter(request, route_setup.as_deref(), route_configs).await
}
Some(join_request::Kind::Predict(request)) => {
let (kind, measured) =
handle_predict(request, handler, route_configs, admission.arrived_at).await;
phases = measured;
kind
}
Some(join_request::Kind::GroupedPredict(request)) => {
let (kind, measured) =
handle_grouped_predict(request, handler, route_configs, admission.arrived_at).await;
phases = measured;
kind
}
Some(join_request::Kind::ResetAdapter(request)) => {
let env_id = request
.context
.as_ref()
.map(|context| context.env_id.clone());
match env_id.filter(|env_id| !env_id.is_empty()) {
Some(env_id) => {
let result = async {
if let Some(route_setup) = route_setup.as_deref() {
route_setup
.reset_adapter(&env_id, &request.episode_ids)
.await?;
}
handler
.lock()
.await
.reset_adapter(&env_id, request.episode_ids)
.await
}
.await;
match result {
Ok(()) => Some(join_response::Kind::ResetAdapter(ResetAdapterResponse {})),
Err(error) => Some(model_error_from_error(&error)),
}
}
None => Some(model_error("reset_adapter missing env_id")),
}
}
Some(join_request::Kind::ReleaseAdapter(request)) => {
let env_id = request.context.as_ref().and_then(route_config_key);
match env_id {
Some(env_id) => {
route_configs.lock().await.remove(&env_id);
if let Some(route_setup) = route_setup.as_deref()
&& let Err(error) = route_setup.release_adapter(&env_id).await
{
return model_join_response(
Some(model_error(error.to_string())),
started_at,
admission.depth_only(),
request_id,
);
}
Some(join_response::Kind::ReleaseAdapter(
ReleaseAdapterResponse {},
))
}
_ => Some(model_error("release_adapter missing env_id")),
}
}
Some(join_request::Kind::Close(_request)) => {
let env_ids: Vec<String> = route_configs.lock().await.keys().cloned().collect();
if let Some(route_setup) = route_setup.as_deref() {
for env_id in &env_ids {
if let Err(error) = route_setup.release_adapter(env_id).await {
return model_join_response(
Some(model_error_from_error(&error)),
started_at,
admission.depth_only(),
request_id,
);
}
}
}
route_configs.lock().await.clear();
Some(join_response::Kind::Close(CloseParticipantResponse {}))
}
None => Some(join_response::Kind::Error(ModelError {
code: ModelErrorCode::Unsupported as i32,
message: "empty or unrecognized Join request (a newer-edition arm this build does \
not implement?)"
.to_string(),
is_recoverable: true,
debug_info: String::new(),
})),
};
phases.in_flight = admission.in_flight;
model_join_response(kind, started_at, phases, request_id)
}
fn model_join_response(
kind: Option<join_response::Kind>,
started_at: Instant,
phases: EndpointPhases,
request_id: String,
) -> JoinResponse {
JoinResponse {
kind,
endpoint_total_ns: Some(model_endpoint_total_ns(started_at)),
decode_ns: EndpointPhases::reported(phases.decode_ns),
user_ns: EndpointPhases::reported(phases.user_ns),
encode_ns: EndpointPhases::reported(phases.encode_ns),
queue_ns: EndpointPhases::reported(phases.queue_ns),
in_flight: (phases.in_flight != 0).then_some(phases.in_flight),
adapter_ns: EndpointPhases::reported(phases.adapter_ns),
held_episodes: phases.held_episodes,
held_state_bytes: phases.held_state_bytes,
request_id,
}
}
async fn handle_resolve_adapter(
request: ResolveAdapterRequest,
route_setup: Option<&dyn ModelRouteSetup>,
route_configs: Arc<Mutex<HashMap<String, ModelRouteConfig>>>,
) -> Option<join_response::Kind> {
let route = match request.context {
Some(route) => route,
None => return Some(model_error("resolve_adapter missing adapter context")),
};
let env_id = route.env_id.clone();
if env_id.is_empty() {
return Some(model_error("resolve_adapter env_id is empty"));
}
let route_key = match route_config_key(&route) {
Some(route_key) => route_key,
None => return Some(model_error("resolve_adapter missing env_id")),
};
let env_spec = match request.env_spec {
Some(env_spec) => env_spec,
None => return Some(model_error("resolve_adapter missing env_spec")),
};
let floor = if request.selected_workflow_edition.is_empty() {
let retained_bases: std::collections::HashSet<Edition> =
rlmesh_proto::SUPPORTED_WORKFLOW_EDITIONS
.iter()
.filter_map(|spelling| Edition::parse(spelling).ok())
.collect();
if retained_bases.len() > 1 {
return Some(model_error(
"resolve_adapter arrived without a workflow edition pin; the runtime must \
select the session floor once more than one edition is supported",
));
}
None
} else {
let floor = match route_floor(request.selected_workflow_edition) {
Ok(floor) => floor,
Err(error) => return Some(model_error(error)),
};
tracing::debug!(
env_id = %env_id,
selected_workflow_edition = %floor.selected_workflow_edition,
workflow_edition_base = %floor.edition,
"model adapter pinned to runtime-selected edition"
);
Some(floor)
};
let env_contract = match env_spec_from_proto(env_spec) {
Ok(env_contract) => env_contract,
Err(error) => return Some(model_error(error.to_string())),
};
if env_contract.action_space.is_none() {
return Some(model_error(
"env EnvSpec has no action_space; a model worker cannot encode actions without it",
));
}
let execution_horizon = request.execution_horizon;
let mut needs = RouteNeeds::default();
if let Some(route_setup) = route_setup {
match route_setup
.resolve_adapter(
&route_key,
&env_contract,
ResolveOptions {
execution_horizon,
delivers_history: request.delivers_history,
},
)
.await
{
Ok(resolved) => needs = resolved,
Err(error) => return Some(model_error_from_error(&error)),
}
}
route_configs.lock().await.insert(
route_key,
ModelRouteConfig {
env_contract: Some(Arc::new(env_contract)),
floor,
},
);
Some(join_response::Kind::ResolveAdapter(
ResolveAdapterResponse {
native_chunk: needs.native_chunk,
history: needs.history.map(|history| ObservationHistoryNeeds {
keys: history.keys,
prunable: history.prunable,
}),
},
))
}
fn route_floor(pinned: String) -> std::result::Result<RouteFloor, String> {
let edition = Edition::parse(&pinned)
.ok()
.filter(|edition| rlmesh_proto::is_retained_edition(*edition))
.ok_or_else(|| {
format!(
"runtime pinned this route to workflow edition {pinned:?}, which this model build \
does not implement (implements {:?})",
rlmesh_proto::SUPPORTED_WORKFLOW_EDITIONS,
)
})?;
Ok(RouteFloor {
edition,
selected_workflow_edition: pinned,
})
}
struct PreparedPredict {
observation: ModelObservation,
action_space: spaces::SpaceSpec,
num_envs: usize,
route: ModelRouteContext,
}
async fn prepare_predict(
request: PredictRequest,
route_configs: &Arc<Mutex<HashMap<String, ModelRouteConfig>>>,
) -> Result<PreparedPredict> {
let mut observation = model_observation_from_endpoint_request(request)?;
let route = ModelRouteContext {
session_id: observation.route.session_id.clone(),
env_id: observation.route.env_id.clone(),
request_id: observation.route.request_id.clone(),
..Default::default()
};
let config = route_configs
.lock()
.await
.get(&route.env_id)
.cloned()
.ok_or_else(|| Error::model("model env adapter was not resolved"))?;
observation.env_contract = config.env_contract;
if let Some(floor) = config.floor.as_ref() {
tracing::trace!(
selected_workflow_edition = %floor.selected_workflow_edition,
"predict on adapter pinned to session floor"
);
}
let num_envs = observation.num_envs;
let action_space = observation
.env_contract
.as_ref()
.and_then(|contract| contract.action_space.clone())
.ok_or_else(|| Error::model("model route contract missing action space"))?;
Ok(PreparedPredict {
observation,
action_space,
num_envs,
route,
})
}
fn finish_predict(
frames: PredictFrames,
num_envs: usize,
action_space: &spaces::SpaceSpec,
route: ModelRouteContext,
) -> Result<PredictResponse> {
let PredictFrames { actions, replay } = frames;
if actions.len() != num_envs {
return Err(Error::model(format!(
"predict returned {} actions for {num_envs} lanes",
actions.len()
)));
}
check_actions_conform(action_space, &actions)?;
let frame0 = rlmesh_grpc::wire::encode_batched_partial_values(&actions, action_space)
.map_err(|err| Error::model(err.to_string()))?;
let mut wire_actions = Vec::with_capacity(1 + replay.len());
wire_actions.push(frame0);
wire_actions.extend(encode_replay_frames(&replay, num_envs, action_space)?);
Ok(model_action_to_endpoint_response(ModelAction {
actions: wire_actions,
route,
}))
}
async fn handle_predict<H: ModelHandler + 'static>(
request: PredictRequest,
handler: Arc<Mutex<H>>,
route_configs: Arc<Mutex<HashMap<String, ModelRouteConfig>>>,
arrived_at: Instant,
) -> (Option<join_response::Kind>, EndpointPhases) {
let mut phases = EndpointPhases::default();
let result = async {
phases.queue_ns = elapsed_ns(arrived_at);
let decode_started = Instant::now();
let prepared = prepare_predict(request, &route_configs).await?;
let PreparedPredict {
observation,
action_space,
num_envs,
route,
} = prepared;
phases.decode_ns = elapsed_ns(decode_started);
let frames = {
let mut handler = handler.lock().await;
phases.queue_ns = elapsed_ns(arrived_at).saturating_sub(phases.decode_ns);
let call_started = Instant::now();
let frames = handler.predict_chunked(observation).await;
phases.user_ns = elapsed_ns(call_started);
phases.adapter_ns = handler.take_adapter_ns();
let held = handler.held_state();
phases.held_episodes = held.map(|held| held.episodes.min(u64::from(u32::MAX)) as u32);
phases.held_state_bytes = held.map(|held| held.bytes);
frames?
};
let encode_started = Instant::now();
let response = finish_predict(frames, num_envs, &action_space, route)?;
phases.encode_ns = elapsed_ns(encode_started);
Ok(response)
}
.await;
match result {
Ok(response) => (Some(join_response::Kind::Predict(response)), phases),
Err(error) => (Some(model_error_from_error(&error)), phases),
}
}
async fn handle_grouped_predict<H: ModelHandler + 'static>(
request: GroupedPredictRequest,
handler: Arc<Mutex<H>>,
route_configs: Arc<Mutex<HashMap<String, ModelRouteConfig>>>,
arrived_at: Instant,
) -> (Option<join_response::Kind>, EndpointPhases) {
enum Finisher {
Failed(Error),
Pending {
num_envs: usize,
action_space: spaces::SpaceSpec,
route: ModelRouteContext,
},
}
let decode_started = Instant::now();
let mut batch: Vec<ModelObservation> = Vec::with_capacity(request.groups.len());
let mut finishers: Vec<Finisher> = Vec::with_capacity(request.groups.len());
for group in request.groups {
match prepare_predict(group, &route_configs).await {
Ok(prepared) => {
let PreparedPredict {
observation,
action_space,
num_envs,
route,
} = prepared;
batch.push(observation);
finishers.push(Finisher::Pending {
num_envs,
action_space,
route,
});
}
Err(error) => finishers.push(Finisher::Failed(error)),
}
}
let decode_ns = elapsed_ns(decode_started);
let mut handler = handler.lock().await;
let queue_ns = elapsed_ns(arrived_at).saturating_sub(decode_ns);
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();
let results = finishers
.into_iter()
.map(|finisher| {
let outcome = match finisher {
Finisher::Failed(error) => Err(error),
Finisher::Pending {
num_envs,
action_space,
route,
} => match frames.next() {
Some(Ok(frames)) => finish_predict(frames, num_envs, &action_space, route),
Some(Err(error)) => Err(error),
None => Err(Error::model(
"predict_grouped returned fewer results than prepared groups",
)),
},
};
grouped_predict_result(outcome)
})
.collect();
(
Some(join_response::Kind::GroupedPredict(
GroupedPredictResponse { results },
)),
EndpointPhases {
decode_ns,
user_ns,
encode_ns: elapsed_ns(encode_started),
queue_ns,
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()
},
)
}
fn grouped_predict_result(outcome: Result<PredictResponse>) -> GroupedPredictResult {
GroupedPredictResult {
outcome: Some(match outcome {
Ok(response) => grouped_predict_result::Outcome::Response(response),
Err(error) => grouped_predict_result::Outcome::Error(model_error_value(&error)),
}),
}
}
#[derive(Default)]
struct RouteTails {
tails: HashMap<String, tokio::sync::oneshot::Receiver<()>>,
}
impl RouteTails {
fn new() -> Self {
Self::default()
}
fn next_keyed_gate(
&mut self,
route_key: Option<&str>,
) -> (RequestGate, Vec<tokio::sync::oneshot::Sender<()>>) {
let prev = route_key.and_then(|key| self.tails.remove(key));
let (done_tx, done_rx) = tokio::sync::oneshot::channel();
if let Some(key) = route_key {
self.tails.insert(key.to_string(), done_rx);
}
self.reap_completed();
(RequestGate::Prev(prev), vec![done_tx])
}
fn next_multi_keyed_gate(
&mut self,
keys: &[String],
) -> (RequestGate, Vec<tokio::sync::oneshot::Sender<()>>) {
let mut prev = Vec::with_capacity(keys.len());
let mut dones = Vec::with_capacity(keys.len());
for key in keys {
if let Some(rx) = self.tails.remove(key) {
prev.push(rx);
}
let (done_tx, done_rx) = tokio::sync::oneshot::channel();
self.tails.insert(key.clone(), done_rx);
dones.push(done_tx);
}
self.reap_completed();
(RequestGate::All(prev), dones)
}
fn close_all_gate(&mut self) -> (RequestGate, Vec<tokio::sync::oneshot::Sender<()>>) {
let prev = self.tails.drain().map(|(_, rx)| rx).collect::<Vec<_>>();
(RequestGate::All(prev), Vec::new())
}
fn reap_completed(&mut self) {
self.tails.retain(|_, rx| {
matches!(
rx.try_recv(),
Err(tokio::sync::oneshot::error::TryRecvError::Empty)
)
});
}
#[cfg(test)]
fn len(&self) -> usize {
self.tails.len()
}
}
enum RequestGate {
Prev(Option<tokio::sync::oneshot::Receiver<()>>),
All(Vec<tokio::sync::oneshot::Receiver<()>>),
}
impl RequestGate {
async fn wait(self) {
match self {
RequestGate::Prev(prev) => {
if let Some(rx) = prev {
let _ = rx.await;
}
}
RequestGate::All(prev) => {
for rx in prev {
let _ = rx.await;
}
}
}
}
}
fn join_request_route_key(request: &JoinRequest) -> Option<String> {
let context = match request.kind.as_ref()? {
join_request::Kind::ResolveAdapter(request) => request.context.as_ref()?,
join_request::Kind::Predict(request) => request.context.as_ref()?,
join_request::Kind::ResetAdapter(request) => request.context.as_ref()?,
join_request::Kind::ReleaseAdapter(request) => request.context.as_ref()?,
join_request::Kind::Close(_) => return None,
join_request::Kind::GroupedPredict(_) => return None,
};
route_config_key(context)
}
fn grouped_predict_route_keys(request: &JoinRequest) -> Option<Vec<String>> {
let join_request::Kind::GroupedPredict(grouped) = request.kind.as_ref()? else {
return None;
};
let mut keys = Vec::new();
for group in &grouped.groups {
if let Some(key) = group.context.as_ref().and_then(route_config_key)
&& !keys.contains(&key)
{
keys.push(key);
}
}
Some(keys)
}
fn route_config_key(context: &rlmesh_proto::model::v1::AdapterContext) -> Option<String> {
if context.env_id.is_empty() {
return None;
}
Some(context.env_id.clone())
}
fn log_join_stream_error(error: &Status) {
tracing::debug!("model join stream closed: {}", error);
}
impl<H> ServedModelServer<H> {
fn authenticate<T>(&self, request: &Request<T>) -> std::result::Result<(), Status> {
let provided = request
.metadata()
.get("authorization")
.and_then(|value| value.to_str().ok())
.unwrap_or("");
if rlmesh_grpc::helpers::bearer_token_matches(&self.token, provided) {
Ok(())
} else {
Err(Status::unauthenticated("invalid route token"))
}
}
}
#[cfg(test)]
mod tests {
use std::sync::Mutex as StdMutex;
use async_trait::async_trait;
use rlmesh_proto::CURRENT_WORKFLOW_EDITION;
use tracing::field::{Field, Visit};
use tracing::subscriber::with_default;
use tracing_subscriber::layer::{Context, SubscriberExt};
use tracing_subscriber::{Layer, Registry};
use super::*;
#[derive(Clone, Default)]
struct CaptureLayer {
messages: Arc<StdMutex<Vec<String>>>,
}
struct MessageVisitor<'a>(&'a mut Vec<String>);
impl Visit for MessageVisitor<'_> {
fn record_debug(&mut self, field: &Field, value: &dyn std::fmt::Debug) {
if field.name() == "message" {
self.0.push(format!("{value:?}"));
}
}
}
impl<S: tracing::Subscriber> Layer<S> for CaptureLayer {
fn on_event(&self, event: &tracing::Event<'_>, _ctx: Context<'_, S>) {
let mut messages = self.messages.lock().unwrap();
let mut visitor = MessageVisitor(&mut messages);
event.record(&mut visitor);
}
}
#[test]
fn inbound_join_stream_error_is_logged_not_swallowed() {
let layer = CaptureLayer::default();
let messages = Arc::clone(&layer.messages);
let subscriber = Registry::default().with(layer);
with_default(subscriber, || {
log_join_stream_error(&Status::aborted("client went away"));
});
let messages = messages.lock().unwrap();
assert!(
messages
.iter()
.any(|message| message.contains("model join stream closed")
&& message.contains("client went away")),
"expected a diagnostic log for the inbound stream error, got {messages:?}"
);
}
#[derive(Default)]
struct NoopModelHandler;
#[async_trait]
impl ModelHandler for NoopModelHandler {
async fn predict(
&mut self,
_observation: super::super::types::ModelObservation,
) -> Result<Vec<spaces::SpaceValue>> {
Ok(Vec::new())
}
}
fn test_server() -> ServedModelServer<NoopModelHandler> {
ServedModelServer {
handler: Arc::new(Mutex::new(NoopModelHandler)),
route_setup: None,
route_configs: Arc::new(Mutex::new(HashMap::new())),
token: String::new(),
activity_tx: None,
shutdown: rlmesh_grpc::lifecycle::ShutdownTrigger::new(),
serve_options: ServeOptions::default(),
declared_workflow_edition: None,
}
}
fn handshake_request(offered_editions: &[&str]) -> HandshakeRequest {
HandshakeRequest {
base: Some(rlmesh_proto::core::v1::HandshakeRequest {
protocol_generation: rlmesh_proto::PROTOCOL_GENERATION.to_string(),
peer_info: Some(peer_info("rlmesh-model-test-client")),
capabilities: Default::default(),
supported_workflow_editions: offered_editions
.iter()
.map(|edition| edition.to_string())
.collect(),
preferred_workflow_edition: String::new(),
}),
}
}
#[test]
fn route_floor_refuses_an_unimplemented_edition() {
let supported =
route_floor(CURRENT_WORKFLOW_EDITION.to_string()).expect("the build's own edition");
assert_eq!(
supported.edition,
Edition::parse(CURRENT_WORKFLOW_EDITION).expect("the current edition parses")
);
assert_eq!(
supported.selected_workflow_edition,
CURRENT_WORKFLOW_EDITION
);
let error = route_floor("2099.01".to_string()).expect_err("must refuse");
assert!(
error.contains("2099.01") && error.contains("does not implement"),
"expected the membership-refusal error, got: {error}"
);
}
#[tokio::test]
async fn handshake_selects_highest_mutual_edition() {
let server = test_server();
for offer in [
&[CURRENT_WORKFLOW_EDITION][..],
&["2025.01", CURRENT_WORKFLOW_EDITION, "2031.12"][..],
] {
let response =
ModelServiceTrait::handshake(&server, Request::new(handshake_request(offer)))
.await
.unwrap()
.into_inner();
let base = response.base.expect("handshake response includes base");
assert!(base.compatible, "offer {offer:?} must be accepted");
assert_eq!(
base.supported_workflow_editions,
supported_workflow_editions()
);
assert_eq!(base.preferred_workflow_edition, CURRENT_WORKFLOW_EDITION);
}
}
#[tokio::test]
async fn serve_options_workflow_edition_is_the_declared_want() {
for (declared, expected) in [
(Some(" 2026.06 "), "2026.06"),
(Some(" "), CURRENT_WORKFLOW_EDITION),
(None, CURRENT_WORKFLOW_EDITION),
] {
let serve_options = ServeOptions {
workflow_edition: declared.map(str::to_string),
..ServeOptions::default()
};
let declared_workflow_edition = declared_workflow_edition_of(&serve_options);
let server = ServedModelServer {
serve_options,
declared_workflow_edition,
..test_server()
};
let response = ModelServiceTrait::handshake(
&server,
Request::new(handshake_request(&[CURRENT_WORKFLOW_EDITION])),
)
.await
.unwrap()
.into_inner();
let base = response.base.expect("handshake response includes base");
assert_eq!(base.preferred_workflow_edition, expected);
assert_eq!(
base.supported_workflow_editions,
supported_workflow_editions()
);
}
}
fn fire(dones: Vec<tokio::sync::oneshot::Sender<()>>) {
for done in dones {
let _ = done.send(());
}
}
#[tokio::test]
async fn route_tails_reaps_closed_routes_so_the_map_stays_bounded() {
let mut tails = RouteTails::new();
for episode in 0..1_000 {
let key = format!("session:{episode}");
let (gate, configure_done) = tails.next_keyed_gate(Some(&key));
assert!(matches!(gate, RequestGate::Prev(None)));
let (gate, predict_done) = tails.next_keyed_gate(Some(&key));
assert!(matches!(gate, RequestGate::Prev(Some(_))));
let (gate, close_done) = tails.next_keyed_gate(Some(&key));
assert!(matches!(gate, RequestGate::Prev(Some(_))));
fire(configure_done);
fire(predict_done);
fire(close_done);
assert!(
tails.len() <= 1,
"episode {episode}: route tails grew to {} entries",
tails.len()
);
}
let (_gate, _done) = tails.next_keyed_gate(Some("session:final"));
assert_eq!(
tails.len(),
1,
"only the in-flight route should remain after reaping"
);
}
#[tokio::test]
async fn route_tails_reopen_after_close_still_sequences() {
let mut tails = RouteTails::new();
let key = "session:route";
let (_g, configure_done) = tails.next_keyed_gate(Some(key));
let (_g, close_done) = tails.next_keyed_gate(Some(key));
fire(configure_done);
let (reopen_gate, _reopen_done) = tails.next_keyed_gate(Some(key));
let mut reopen_prev = match reopen_gate {
RequestGate::Prev(Some(rx)) => rx,
RequestGate::Prev(None) => {
panic!("reopen must gate on the in-flight ReleaseAdapter, got an ungated request")
}
RequestGate::All(_) => panic!("a keyed request must never produce an All gate"),
};
assert!(matches!(
reopen_prev.try_recv(),
Err(tokio::sync::oneshot::error::TryRecvError::Empty)
));
fire(close_done);
assert!(reopen_prev.try_recv().is_ok());
}
#[tokio::test]
async fn route_tails_reopen_after_reaped_close_is_ungated() {
let mut tails = RouteTails::new();
let key = "session:route";
let (_g, configure_done) = tails.next_keyed_gate(Some(key));
let (_g, close_done) = tails.next_keyed_gate(Some(key));
fire(configure_done);
fire(close_done);
let (_g, _d) = tails.next_keyed_gate(Some("other:route"));
let (reopen_gate, _d) = tails.next_keyed_gate(Some(key));
assert!(
matches!(reopen_gate, RequestGate::Prev(None)),
"reaped-then-reopened route should be ungated"
);
}
#[tokio::test]
async fn route_tails_close_drains_every_route() {
let mut tails = RouteTails::new();
let (_g, _d0) = tails.next_keyed_gate(Some("a:1"));
let (_g, _d1) = tails.next_keyed_gate(Some("b:1"));
assert_eq!(tails.len(), 2);
let (gate, _close_done) = tails.close_all_gate();
match gate {
RequestGate::All(prev) => assert_eq!(prev.len(), 2),
RequestGate::Prev(_) => panic!("Close must produce an All gate over every route"),
}
assert_eq!(tails.len(), 0, "Close must clear every route tail");
}
#[tokio::test]
async fn route_tails_grouped_predict_gates_each_route_and_chains_successors() {
let mut tails = RouteTails::new();
let (_g, _a_prev) = tails.next_keyed_gate(Some("s:a"));
let (_g, _b_prev) = tails.next_keyed_gate(Some("s:b"));
let (gate, grouped_done) =
tails.next_multi_keyed_gate(&["s:a".to_owned(), "s:b".to_owned()]);
match gate {
RequestGate::All(prev) => assert_eq!(prev.len(), 2, "must gate on both routes"),
RequestGate::Prev(_) => panic!("a grouped predict must produce an All gate"),
}
let (close_gate, _close_done) = tails.next_keyed_gate(Some("s:a"));
let mut close_prev = match close_gate {
RequestGate::Prev(Some(rx)) => rx,
_ => panic!("ReleaseAdapter after a grouped predict must gate on it, not run ungated"),
};
assert!(matches!(
close_prev.try_recv(),
Err(tokio::sync::oneshot::error::TryRecvError::Empty)
));
fire(grouped_done);
assert!(close_prev.try_recv().is_ok());
}
#[test]
fn grouped_predict_route_keys_dedups_referenced_routes() {
let group = |env_id: &str| PredictRequest {
context: Some(rlmesh_proto::model::v1::AdapterContext {
session_id: "s".to_owned(),
env_id: env_id.to_owned(),
..Default::default()
}),
..Default::default()
};
let grouped = JoinRequest {
kind: Some(join_request::Kind::GroupedPredict(GroupedPredictRequest {
groups: vec![group("a"), group("b"), group("a")],
})),
..Default::default()
};
assert_eq!(
grouped_predict_route_keys(&grouped),
Some(vec!["a".to_owned(), "b".to_owned()])
);
let predict = JoinRequest {
kind: Some(join_request::Kind::Predict(group("a"))),
..Default::default()
};
assert_eq!(grouped_predict_route_keys(&predict), None);
}
#[tokio::test]
async fn close_tears_down_every_route_config() {
let server = test_server();
{
let mut configs = server.route_configs.lock().await;
for key in ["session:a", "session:b"] {
configs.insert(
key.to_string(),
ModelRouteConfig {
env_contract: None,
floor: None,
},
);
}
}
let response = handle_model_request(
JoinRequest {
request_id: "close-1".to_string(),
kind: Some(join_request::Kind::Close(
rlmesh_proto::model::v1::CloseParticipantRequest::default(),
)),
},
Arc::clone(&server.handler),
server.route_setup.clone(),
Arc::clone(&server.route_configs),
Admission::now(),
)
.await;
assert!(matches!(response.kind, Some(join_response::Kind::Close(_))));
assert!(
server.route_configs.lock().await.is_empty(),
"Close must clear every route config"
);
}
#[tokio::test]
async fn handshake_accepts_any_generation_compatible_offer() {
let server = test_server();
for offer in [&[][..], &["2026"][..], &["2026.11", "next"][..]] {
let response =
ModelServiceTrait::handshake(&server, Request::new(handshake_request(offer)))
.await
.unwrap()
.into_inner();
let base = response.base.expect("handshake response includes base");
assert!(
base.compatible,
"generation-ok offer {offer:?} is compatible"
);
assert!(base.error_message.is_none());
}
}
}