#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
use crate::network::client::agent::AlgorithmInitArgs;
use crate::network::client::agent::{ActorInferenceMode, ActorInfo, ClientModes, ModelMode};
use crate::network::client::runtime::actor::LocalModelHandle;
use crate::network::client::runtime::actor::{
Actor, ActorEntity, ActorError, ActorRuntime, ErasedActorRuntime,
};
use crate::network::client::runtime::control::coordinator::{CHANNEL_THROUGHPUT, ClientNamespace};
use crate::network::client::runtime::control::lifecycle_manager::LifecycleManagerError;
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
use crate::network::client::runtime::control::lifecycle_manager::SharedTransportAddresses;
use crate::network::client::runtime::control::scale_manager::RouterNamespace;
use crate::network::client::runtime::data::environments::EnvironmentInterface;
use crate::network::client::runtime::data::environments::EnvironmentInterfaceError;
use crate::network::client::runtime::data::router::{
ControlPayload, RoutedMessage, RoutingProtocol,
};
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
use crate::network::client::runtime::data::sinks::transport_sink::transport_dispatcher::{
InferenceDispatcher, TrainingDispatcher,
};
use crate::network::client::runtime::data::training::{TrainingError, TrainingInterface};
#[cfg(feature = "metrics")]
use crate::utilities::observability::metrics::MetricsManager;
use crossbeam_utils::CachePadded;
use std::path::PathBuf;
use thiserror::Error;
use active_uuid_registry::UuidPoolError;
use relayrl_algorithms::prelude::nn::NeuralNetwork;
use relayrl_algorithms::prelude::ppo::trainer::PPOTrainerSpec;
#[cfg(feature = "tch-backend")]
use relayrl_env_trait::EnvTchDType;
use relayrl_env_trait::{EnvDType, EnvNdArrayDType, Environment};
#[cfg(feature = "tch-backend")]
use relayrl_types::data::tensor::TchDType;
use relayrl_types::data::tensor::{BackendMatcher, DType, DeviceType, NdArrayDType};
use relayrl_types::model::{HotReloadableModel, ModelModule};
use relayrl_types::prelude::tensor::burn::{BasicOps, Numeric, TensorKind};
use active_uuid_registry::registry_uuid::Uuid;
use arc_swap::ArcSwapOption;
use burn_tensor::backend::Backend;
use dashmap::DashMap;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use tokio::sync::RwLock;
use tokio::sync::mpsc;
use tokio::sync::mpsc::{Receiver, Sender};
use tokio::task::JoinHandle;
#[derive(Debug, Error)]
#[allow(clippy::enum_variant_names)]
pub enum StateManagerError {
#[error(transparent)]
UuidPoolError(#[from] UuidPoolError),
#[error("Failed to create reloadable model: {0}")]
FailedToCreateReloadableModelError(String),
#[error("Actor handle not found: {0}")]
ActorHandleNotFoundError(String),
#[error("Actor inbox not found: {0}")]
ActorInboxNotFoundError(String),
#[error("Actor already taken: {0}")]
ActorAlreadyTakenError(String),
#[error("Subscribe shutdown failed: {0}")]
SubscribeShutdownError(#[from] LifecycleManagerError),
#[error("Failed to receive shutdown signal: {0}")]
ReceiveShutdownSignalError(String),
#[error("Shutdown all actors failed: {0}")]
ShutdownAllActorsError(String),
#[error("Set actor ID failed: {0}")]
SetActorIdError(String),
#[error("Set actor nametag failed: {0}")]
SetActorNameTagError(String),
#[error("Set actor model failed: {0}")]
SetActorModelError(String),
#[error("Get actors failed: {0}")]
GetActorsError(String),
#[error("New actor failed: {0}")]
NewActorError(String),
#[error("Remove actor failed: {0}")]
RemoveActorError(String),
#[error("Get config failed: {0}")]
GetConfigError(String),
#[error("Set config failed: {0}")]
SetConfigError(String),
#[error("Set env failed: {0}")]
SetEnvError(String),
#[error("Step env failed: {0}")]
StepEnvError(String),
#[error("Get env info failed: {0}")]
GetEnvInfoError(String),
#[error("Get env count failed: {0}")]
GetEnvCountError(String),
#[error("Increase env count failed: {0}")]
IncreaseEnvCountError(String),
#[error("Decrease env count failed: {0}")]
DecreaseEnvCountError(String),
#[error("Remove envs failed: {0}")]
RemoveEnvError(String),
#[error("Invalid environment kind: {0}")]
InvalidEnvironmentKindError(String),
#[error(transparent)]
EnvironmentInterfaceError(#[from] EnvironmentInterfaceError),
#[error("Tensor conversion failed: {0}")]
TensorConversionError(String),
#[error("Inference request failed: {0}")]
InferenceRequestError(String),
#[error("Environment training error: {0}")]
TrainerError(String),
#[error(transparent)]
AlgorithmError(#[from] relayrl_algorithms::templates::base_algorithm::AlgorithmError),
#[error("Algorithm config init failed: {0}")]
AlgorithmConfigInitError(String),
#[error(transparent)]
TrainingError(#[from] TrainingError),
#[error(transparent)]
ActorError(#[from] crate::network::client::runtime::actor::ActorError),
}
pub type ActorUuid = Uuid;
#[derive(Clone, Debug, Default, PartialEq, Eq, Hash)]
pub struct NameTag {
pub tag: String,
pub duplicate: usize,
}
#[derive(Clone)]
pub(crate) struct ActorRoute {
pub(crate) router_namespace: Option<RouterNamespace>,
pub(crate) inbox: Sender<RoutedMessage>,
}
pub(crate) struct SharedRouterState {
pub(crate) actor_routes: DashMap<ActorUuid, ActorRoute>,
}
pub(crate) struct StateManager<B: Backend + BackendMatcher<Backend = B>> {
client_namespace: ClientNamespace,
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
shared_inference_dispatcher: Option<Arc<InferenceDispatcher<B>>>,
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
shared_training_dispatcher: Option<Arc<TrainingDispatcher<B>>>,
shared_client_modes: Arc<ClientModes>,
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
shared_transport_addresses: Option<Arc<RwLock<SharedTransportAddresses>>>,
shared_local_model_path: Arc<RwLock<PathBuf>>,
default_model: Option<ModelModule<B>>,
shared_local_models: Vec<(DeviceType, LocalModelHandle<B>)>,
#[cfg(feature = "metrics")]
metrics: MetricsManager,
pub(crate) global_dispatcher_tx: Sender<RoutedMessage>,
pub(crate) shared_router_state: Arc<SharedRouterState>,
actor_envs: Arc<DashMap<ActorUuid, EnvironmentInterface>>,
actor_handles: DashMap<ActorUuid, Arc<JoinHandle<()>>>,
actor_devices: DashMap<ActorUuid, DeviceType>,
pub(crate) actor_model_handles: DashMap<ActorUuid, LocalModelHandle<B>>,
pub(crate) actor_runtime_handles: Arc<DashMap<ActorUuid, Arc<dyn ErasedActorRuntime<B>>>>,
pub(crate) shared_actor_count: Arc<CachePadded<AtomicUsize>>,
}
impl<B: Backend + BackendMatcher<Backend = B>> StateManager<B> {
#[allow(clippy::too_many_arguments)]
pub(crate) fn new(
client_namespace: ClientNamespace,
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
shared_inference_dispatcher: Option<Arc<InferenceDispatcher<B>>>,
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
shared_training_dispatcher: Option<Arc<TrainingDispatcher<B>>>,
shared_client_modes: Arc<ClientModes>,
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
shared_transport_addresses: Option<Arc<RwLock<SharedTransportAddresses>>>,
shared_local_model_path: Arc<RwLock<PathBuf>>,
default_model: Option<ModelModule<B>>,
#[cfg(feature = "metrics")] metrics: MetricsManager,
) -> (Self, Receiver<RoutedMessage>) {
let (global_dispatcher_tx, global_dispatcher_rx) =
mpsc::channel::<RoutedMessage>(CHANNEL_THROUGHPUT * 2);
(
Self {
client_namespace,
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
shared_inference_dispatcher,
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
shared_training_dispatcher,
shared_client_modes,
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
shared_transport_addresses,
shared_local_model_path,
default_model,
shared_local_models: Vec::new(),
#[cfg(feature = "metrics")]
metrics,
global_dispatcher_tx,
shared_router_state: Arc::new(SharedRouterState {
actor_routes: DashMap::new(),
}),
actor_envs: Arc::new(DashMap::new()),
actor_handles: DashMap::new(),
actor_devices: DashMap::new(),
actor_model_handles: DashMap::new(),
actor_runtime_handles: Arc::new(DashMap::new()),
shared_actor_count: Arc::new(CachePadded::new(AtomicUsize::new(0))),
},
global_dispatcher_rx,
)
}
async fn load_reloadable_model(
&self,
model_module: Option<ModelModule<B>>,
device: DeviceType,
) -> Result<Option<HotReloadableModel<B>>, StateManagerError> {
if let Some(model) = model_module {
return Ok(Some(
HotReloadableModel::<B>::new_from_module(model, device)
.await
.map_err(|_| {
StateManagerError::FailedToCreateReloadableModelError(
"[StateManager] Failed to create reloadable model from parameter"
.to_string(),
)
})?,
));
}
if let Some(model) = self.default_model.clone() {
return Ok(Some(
HotReloadableModel::<B>::new_from_module(model, device)
.await
.map_err(|_| {
StateManagerError::FailedToCreateReloadableModelError(
"[StateManager] Failed to create reloadable model from cache"
.to_string(),
)
})?,
));
}
let local_model_path = self.shared_local_model_path.read().await;
if !local_model_path.to_str().unwrap_or_default().is_empty() {
return Ok(Some(
HotReloadableModel::<B>::new_from_path(local_model_path.as_path(), device)
.await
.map_err(|_| {
StateManagerError::FailedToCreateReloadableModelError(
"[StateManager] Failed to load model from local_model_path".to_string(),
)
})?,
));
}
Ok(None)
}
async fn get_or_init_model_handle(
&mut self,
default_model: Option<ModelModule<B>>,
device: DeviceType,
) -> Result<(LocalModelHandle<B>, bool), StateManagerError> {
match &self.shared_client_modes.actor_inference_mode {
ActorInferenceMode::Client(ModelMode::Shared) => {
if let Some(idx) = self
.shared_local_models
.iter()
.position(|(d, _)| d == &device)
{
let handle = self.shared_local_models[idx].1.clone();
return Ok((handle, false));
}
let reloadable = self
.load_reloadable_model(default_model, device.clone())
.await?;
let needs_handshake = reloadable.is_none();
let handle: LocalModelHandle<B> =
Arc::new(ArcSwapOption::new(reloadable.map(Arc::new)));
self.shared_local_models.push((device, handle.clone()));
Ok((handle, needs_handshake))
}
_ => {
let reloadable = self
.load_reloadable_model(default_model, device.clone())
.await?;
let needs_handshake = reloadable.is_none();
let handle: LocalModelHandle<B> =
Arc::new(ArcSwapOption::new(reloadable.map(Arc::new)));
Ok((handle, needs_handshake))
}
}
}
#[allow(clippy::too_many_arguments)]
pub(crate) async fn new_actor<const D_IN: usize, const D_OUT: usize>(
&mut self,
actor_id: ActorUuid,
router_namespace: RouterNamespace,
device: DeviceType,
max_traj_length: usize,
nametag: Option<NameTag>,
default_model: Option<ModelModule<B>>,
tx_to_buffer: Sender<RoutedMessage>,
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
algorithm_args: AlgorithmInitArgs,
) -> Result<(), StateManagerError> {
if self.actor_handles.contains_key(&actor_id) {
log::warn!(
"[StateManager] Actor ID {} already exists, replacing existing actor...",
actor_id
);
self.remove_actor(actor_id)?
}
let (tx_to_actor, actor_inbox_rx) = mpsc::channel::<RoutedMessage>(CHANNEL_THROUGHPUT);
self.shared_router_state.actor_routes.insert(
actor_id,
ActorRoute {
router_namespace: Some(router_namespace.clone()),
inbox: tx_to_actor.clone(),
},
);
let shared_local_model_path = self.shared_local_model_path.clone();
let shared_client_modes = self.shared_client_modes.clone();
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
let shared_inference_dispatcher = self.shared_inference_dispatcher.clone();
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
let shared_training_dispatcher = self.shared_training_dispatcher.clone();
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
let shared_transport_addresses = self.shared_transport_addresses.clone();
#[cfg(feature = "metrics")]
let actor_metrics = self.metrics.clone();
#[cfg(feature = "metrics")]
let runtime_metrics = self.metrics.clone();
let client_namespace = self.client_namespace.clone();
let (model_handle, _model_handshake_flag) = self
.get_or_init_model_handle(default_model, device.clone())
.await?;
self.actor_devices.insert(actor_id, device.clone());
self.actor_envs.insert(
actor_id,
EnvironmentInterface::new(self.client_namespace.clone(), device.clone()),
);
self.actor_model_handles
.insert(actor_id, model_handle.clone());
let runtime = Arc::new(
ActorRuntime::<B, D_IN, D_OUT>::new(
actor_id,
nametag,
model_handle.clone(),
max_traj_length,
tx_to_buffer.clone(),
#[cfg(feature = "metrics")]
runtime_metrics,
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
algorithm_args,
)
.await,
);
let erased_runtime: Arc<dyn ErasedActorRuntime<B>> = runtime.clone();
self.actor_runtime_handles.insert(actor_id, erased_runtime);
let handle: Arc<JoinHandle<()>> = Arc::new(tokio::spawn(async move {
let mut actor: Actor<B, D_IN, D_OUT> = Actor::<B, D_IN, D_OUT>::new(
client_namespace,
actor_id,
device.clone(),
runtime,
shared_local_model_path,
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
shared_inference_dispatcher,
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
shared_training_dispatcher,
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
shared_transport_addresses,
actor_inbox_rx,
shared_client_modes,
#[cfg(feature = "metrics")]
actor_metrics,
)
.await;
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
if _model_handshake_flag {
let model_handshake_ms = RoutedMessage {
actor_id,
protocol: RoutingProtocol::Control(ControlPayload::ModelHandshake),
};
let _ = tx_to_actor.send(model_handshake_ms).await;
}
if let Err(e) = actor.spawn_loop().await {
log::error!("[StateManager] Actor {:?} loop error: {}", actor_id, e);
}
}));
self.actor_handles.insert(actor_id, handle);
self.shared_actor_count.fetch_add(1, Ordering::Release);
Ok(())
}
pub(crate) async fn shutdown_all_actors(&self) -> Result<Vec<ActorInfo>, StateManagerError> {
let mut actor_ids = Vec::<ActorUuid>::new();
for entry in self.shared_router_state.actor_routes.iter() {
let actor_id: ActorUuid = *entry.key();
let tx: Sender<RoutedMessage> = entry.value().inbox.clone();
let shutdown_msg = RoutedMessage {
actor_id,
protocol: RoutingProtocol::Control(ControlPayload::Shutdown),
};
let _ = tx.send(shutdown_msg).await;
let handle: Result<
dashmap::mapref::one::Ref<'_, Uuid, Arc<JoinHandle<()>>>,
StateManagerError,
> = self.actor_handles.get(&actor_id).ok_or(
StateManagerError::ActorHandleNotFoundError(
"[StateManager] Actor handle not found".to_string(),
),
);
while active_uuid_registry::interface::list_ids(
self.client_namespace.as_ref(),
crate::network::ACTOR_CONTEXT,
)
.contains(&actor_id)
{
tokio::time::timeout(
std::time::Duration::from_secs(10),
tokio::time::sleep(std::time::Duration::from_secs(1)),
)
.await
.map_err(|_| {
StateManagerError::ShutdownAllActorsError(format!(
"[StateManager] Shutdown all actors timeout: {}",
actor_id
))
})?;
}
actor_ids.push(actor_id);
if let Ok(handle) = handle {
handle.abort();
} else {
continue;
}
}
Ok(actor_ids
.iter()
.filter_map(|id| {
let runtime = self.actor_runtime_handles.get(id)?;
match runtime.get_actor_info() {
Ok(actor_info) => Some(actor_info),
Err(e) => {
log::error!("{}", e);
None
}
}
})
.collect())
}
pub(crate) async fn clear_runtime_components(&mut self) -> Result<(), StateManagerError> {
self.actor_handles.clear();
self.shared_router_state.actor_routes.clear();
self.actor_devices.clear();
self.actor_model_handles.clear();
self.actor_runtime_handles.clear();
self.shared_local_models.clear();
Ok(())
}
pub(crate) fn remove_actor(&mut self, id: Uuid) -> Result<(), StateManagerError> {
if let Some(mut route) = self.shared_router_state.actor_routes.get_mut(&id) {
route.router_namespace = None;
}
if let Some((_, handle)) = self.actor_handles.remove(&id) {
handle.abort();
}
self.actor_envs.remove(&id);
self.actor_devices.remove(&id);
self.actor_model_handles.remove(&id);
self.actor_runtime_handles.remove(&id);
self.shared_router_state.actor_routes.remove(&id);
self.shared_actor_count.fetch_sub(1, Ordering::Release);
self.client_namespace
.remove_id(crate::network::ACTOR_CONTEXT, id)
.map_err(StateManagerError::from)?;
Ok(())
}
pub(crate) fn set_actor_id(
&self,
current_id: ActorUuid,
new_id: ActorUuid,
) -> Result<(), StateManagerError> {
{
let current_id_handle = match StateManager::<B>::get_actor_handle(self, current_id) {
Some(handle) => handle.clone(),
None => {
return Err(StateManagerError::ActorHandleNotFoundError(format!(
"[StateManager] Actor ID {} not found",
current_id
)));
}
};
let current_route = match StateManager::<B>::get_actor_route(self, current_id) {
Some(route) => route,
None => {
return Err(StateManagerError::ActorInboxNotFoundError(format!(
"[StateManager] Actor ID {} not found",
current_id
)));
}
};
if StateManager::<B>::get_actor_handle(self, new_id).is_some()
|| StateManager::<B>::get_actor_route(self, new_id).is_some()
{
return Err(StateManagerError::ActorAlreadyTakenError(format!(
"[StateManager] Actor ID {} already taken",
new_id
)));
}
self.actor_handles.insert(new_id, current_id_handle);
self.actor_handles.remove(¤t_id);
self.shared_router_state
.actor_routes
.insert(new_id, current_route);
self.shared_router_state.actor_routes.remove(¤t_id);
}
if let Some((_, current_device)) = self.actor_devices.remove(¤t_id) {
self.actor_devices.insert(new_id, current_device);
}
if let Some((_, current_env)) = self.actor_envs.remove(¤t_id) {
self.actor_envs.insert(new_id, current_env);
}
if let Some((_, runtime)) = self.actor_runtime_handles.remove(¤t_id) {
runtime
.set_actor_id(new_id)
.map_err(StateManagerError::from)?;
self.actor_runtime_handles.insert(new_id, runtime);
}
self.client_namespace
.replace_id(crate::network::ACTOR_CONTEXT, current_id, new_id)
.map_err(StateManagerError::from)?;
Ok(())
}
pub(crate) fn set_actor_nametag(
&self,
actor_id: ActorUuid,
new_nametag: Option<NameTag>,
) -> Result<(), StateManagerError> {
if let Some(entry) = self.actor_runtime_handles.get(&actor_id) {
entry
.value()
.set_actor_nametag(new_nametag)
.map_err(StateManagerError::from)
} else {
Err(StateManagerError::SetActorNameTagError(format!(
"[StateManager] Failed to change actor id {} to {:?}; actor runtime could not be found",
actor_id, new_nametag
)))
}
}
pub(crate) fn distribute_actors(&self, router_namespaces: Vec<RouterNamespace>) {
if router_namespaces.is_empty() {
return;
}
let mut actor_ids: Vec<ActorUuid> = StateManager::<B>::get_actor_id_list(self);
actor_ids.sort_by_key(|actor_id| actor_id.to_string());
for (i, actor_id) in actor_ids.iter().enumerate() {
let router_namespace = router_namespaces[i % router_namespaces.len()].clone();
if let Some(mut route) = self.shared_router_state.actor_routes.get_mut(actor_id) {
route.router_namespace = Some(router_namespace);
}
}
}
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
pub(crate) fn restore_actor_router_mappings(
&self,
mappings: Vec<(ActorUuid, RouterNamespace)>,
) {
let mappings_by_actor: std::collections::HashMap<ActorUuid, RouterNamespace> =
mappings.into_iter().collect();
for mut route in self.shared_router_state.actor_routes.iter_mut() {
route.router_namespace = mappings_by_actor.get(route.key()).cloned();
}
}
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
pub(crate) fn get_actor_router_mappings(&self) -> Vec<(ActorUuid, RouterNamespace)> {
self.shared_router_state
.actor_routes
.iter()
.filter_map(|entry| {
entry
.value()
.router_namespace
.clone()
.map(|router_namespace| (*entry.key(), router_namespace))
})
.collect()
}
pub(crate) fn get_actor_id_list(&self) -> Vec<ActorUuid> {
self.actor_handles
.iter()
.map(|entry| *entry.key())
.collect()
}
fn sorted_actors_for_model_updates(&self, actors: Option<&[ActorInfo]>) -> Vec<ActorInfo> {
let mut actors: Vec<ActorInfo> = match actors {
Some(infos) => infos
.iter()
.filter(|actor| self.actor_handles.contains_key(&actor.id()))
.cloned()
.collect(),
None => self
.get_actor_id_list()
.iter()
.map(|id| match self.actor_runtime_handles.get(id) {
Some(runtime) => match runtime.get_actor_info() {
Ok(actor_info) => actor_info,
Err(e) => {
log::error!("{}", e);
ActorInfo::new(*id, None)
}
},
None => ActorInfo::new(*id, None),
})
.collect(),
};
actors.sort_by_key(|actor| actor.id().to_string());
actors.dedup_by(|a, b| a.id() == b.id());
actors
}
fn canonical_model_update_target_from_sorted_actors(
&self,
actor: &ActorInfo,
sorted_actors: &[ActorInfo],
) -> ActorInfo {
match &self.shared_client_modes.actor_inference_mode {
ActorInferenceMode::Client(ModelMode::Shared) => {
let Some(actor_device) = self
.actor_devices
.get(&actor.id())
.map(|device_entry| device_entry.value().clone())
else {
return actor.clone();
};
sorted_actors
.iter()
.find(|candidate_actor| {
self.actor_devices
.get(&candidate_actor.id())
.map(|device_entry| device_entry.value() == &actor_device)
.unwrap_or(false)
})
.cloned()
.unwrap_or_else(|| actor.clone())
}
_ => actor.clone(),
}
}
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
pub(crate) fn canonical_model_update_target(&self, actor: &ActorInfo) -> ActorInfo {
let sorted_actors = self.sorted_actors_for_model_updates(None);
self.canonical_model_update_target_from_sorted_actors(actor, &sorted_actors)
}
#[cfg(test)]
pub(crate) fn model_update_dispatch_targets(&self) -> Vec<ActorInfo> {
self.model_update_dispatch_targets_for_subset(None)
}
pub(crate) fn model_update_dispatch_targets_for_subset(
&self,
actors: Option<&[ActorInfo]>,
) -> Vec<ActorInfo> {
let sorted_actors = self.sorted_actors_for_model_updates(actors);
let mut dispatch_targets = Vec::new();
for actor in sorted_actors.iter() {
let canonical_target =
self.canonical_model_update_target_from_sorted_actors(actor, &sorted_actors);
if dispatch_targets.contains(&canonical_target) {
continue;
}
dispatch_targets.push(canonical_target);
}
dispatch_targets
}
fn get_actor_handle(&self, id: Uuid) -> Option<Arc<JoinHandle<()>>> {
self.actor_handles
.get(&id)
.map(|handle| Arc::clone(handle.value()))
}
fn get_actor_route(&self, id: Uuid) -> Option<ActorRoute> {
self.shared_router_state
.actor_routes
.get(&id)
.map(|route| route.value().clone())
}
fn get_actor_runtime(&self, id: Uuid) -> Option<Arc<dyn ErasedActorRuntime<B>>> {
self.actor_runtime_handles
.get(&id)
.map(|runtime| Arc::clone(runtime.value()))
}
pub(crate) fn set_env(
&self,
id: Uuid,
env: Box<dyn Environment>,
count: u32,
) -> Result<(), StateManagerError> {
let device = self
.actor_devices
.get(&id)
.map(|device| device.clone())
.ok_or_else(|| {
StateManagerError::SetEnvError(format!(
"[StateManager] Actor device not found for {}",
id
))
})?;
if let Some(mut env_interface) = self.actor_envs.get_mut(&id) {
env_interface.set_env(Some(env), count as usize)?;
} else {
let mut env_interface =
EnvironmentInterface::new(self.client_namespace.clone(), device);
env_interface.set_env(Some(env), count as usize)?;
self.actor_envs.insert(id, env_interface);
}
Ok(())
}
pub(crate) fn get_env_count(&self, actor_id: ActorUuid) -> Result<u32, StateManagerError> {
self.actor_envs
.get(&actor_id)
.ok_or_else(|| {
StateManagerError::GetEnvCountError(format!(
"[StateManager] Environment interface not found for {}",
actor_id
))
})?
.get_env_count()
.map_err(StateManagerError::from)
}
pub(crate) fn increase_env_count(
&self,
actor_id: ActorUuid,
count: u32,
) -> Result<(), StateManagerError> {
self.actor_envs
.get_mut(&actor_id)
.ok_or_else(|| {
StateManagerError::IncreaseEnvCountError(format!(
"[StateManager] Environment interface not found for {}",
actor_id
))
})?
.increase_env_count(count)
.map_err(StateManagerError::from)
}
pub(crate) fn decrease_env_count(
&self,
actor_id: ActorUuid,
count: u32,
) -> Result<(), StateManagerError> {
self.actor_envs
.get_mut(&actor_id)
.ok_or_else(|| {
StateManagerError::DecreaseEnvCountError(format!(
"[StateManager] Environment interface not found for {}",
actor_id
))
})?
.decrease_env_count(count)
.map_err(StateManagerError::from)
}
pub(crate) fn remove_env(&self, actor_id: ActorUuid) -> Result<(), StateManagerError> {
self.actor_envs
.get_mut(&actor_id)
.ok_or_else(|| {
StateManagerError::RemoveEnvError(format!(
"[StateManager] Environment interface not found for {}",
actor_id
))
})?
.remove_env()
.map_err(StateManagerError::from)
}
#[allow(clippy::type_complexity)]
pub(crate) fn get_run_env_handles(
&self,
actor_id: ActorUuid,
) -> Result<
(
Arc<dyn ErasedActorRuntime<B>>,
Arc<DashMap<ActorUuid, EnvironmentInterface>>,
),
StateManagerError,
> {
if !self.actor_envs.contains_key(&actor_id) {
return Err(StateManagerError::GetEnvInfoError(format!(
"[StateManager] Environment interface not found for {}",
actor_id
)));
}
let runtime = self.get_actor_runtime(actor_id).ok_or_else(|| {
StateManagerError::ActorHandleNotFoundError(format!(
"[StateManager] Actor runtime not found for {}",
actor_id
))
})?;
Ok((runtime, self.actor_envs.clone()))
}
pub(crate) fn run_env_eval_step_loop(
actor_id: ActorUuid,
runtime: Arc<dyn ErasedActorRuntime<B>>,
env_map: Arc<DashMap<ActorUuid, EnvironmentInterface>>,
loop_iters: usize,
) -> Result<(), StateManagerError> {
tokio::task::block_in_place(|| {
tokio::runtime::Handle::current().block_on(async {
let mut env_interface = env_map.get_mut(&actor_id).ok_or_else(|| {
StateManagerError::GetEnvInfoError(format!(
"[StateManager] Environment interface not found for {}",
actor_id
))
})?;
env_interface.ensure_ready()?;
let (n_envs, obs_dim, act_dim) = env_interface.n_envs_dims().ok_or_else(|| {
StateManagerError::StepEnvError(format!(
"[StateManager] Failed to get environment dimensions for {}",
actor_id
))
})?;
let obs_dtype = {
let env_dtype = env_interface
.obs_dtype()
.unwrap_or(EnvDType::NdArray(EnvNdArrayDType::F32));
env_dtype_to_dtype(&env_dtype)?
};
let act_dtype = {
let env_dtype = env_interface
.act_dtype()
.unwrap_or(EnvDType::NdArray(EnvNdArrayDType::F32));
env_dtype_to_dtype(&env_dtype)?
};
let discrete = env_interface.action_is_discrete().unwrap_or(true);
for _ in 0..loop_iters {
let obs_bytes = env_interface.flat_observation_bytes().ok_or_else(|| {
StateManagerError::StepEnvError(
"[StateManager] flat_observation_bytes returned None".to_string(),
)
})?;
let action_bytes = runtime
.perform_local_byte_inference_erased(
&obs_bytes, n_envs, obs_dim, act_dim, &obs_dtype, &act_dtype, discrete,
)
.await
.map_err(|e| StateManagerError::InferenceRequestError(e.to_string()))?;
let _ = env_interface.step_bytes(&action_bytes).ok_or_else(|| {
StateManagerError::GetEnvInfoError(
"[StateManager] step_bytes returned None".to_string(),
)
})?;
}
Ok(())
})
})
}
pub(crate) fn run_env_step_loop_with_ppo<KindIn, KindOut, Pi>(
actor_id: ActorUuid,
shutdown_rx: Option<tokio::sync::broadcast::Receiver<()>>,
runtime: Arc<dyn ErasedActorRuntime<B>>,
env_map: Arc<DashMap<ActorUuid, EnvironmentInterface>>,
loop_iters: usize,
max_traj_length: usize,
trainer_spec: PPOTrainerSpec<B, KindIn, KindOut, Pi>,
) -> Result<ModelModule<B>, StateManagerError>
where
KindIn: TensorKind<B> + BasicOps<B> + Send + 'static,
KindOut: TensorKind<B> + BasicOps<B> + Numeric<B> + Send + 'static,
Pi: NeuralNetwork<B, KindIn, KindOut> + Clone + Send + 'static,
B: Default + Send + Sync + 'static,
{
TrainingInterface::<B>::train_ppo(
actor_id,
shutdown_rx,
runtime,
env_map,
loop_iters,
max_traj_length,
trainer_spec,
)
.map_err(StateManagerError::from)
}
pub(crate) fn run_env_step_loop_with_ippo<KindIn, KindOut, Pi>(
actor_id: ActorUuid,
shutdown_rx: tokio::sync::broadcast::Receiver<()>,
runtime: Arc<dyn ErasedActorRuntime<B>>,
env_map: Arc<DashMap<ActorUuid, EnvironmentInterface>>,
loop_iters: usize,
max_traj_length: usize,
trainer_spec: PPOTrainerSpec<B, KindIn, KindOut, Pi>,
) -> Result<ModelModule<B>, StateManagerError>
where
KindIn: TensorKind<B> + BasicOps<B> + Send + 'static,
KindOut: TensorKind<B> + BasicOps<B> + Numeric<B> + Send + 'static,
Pi: NeuralNetwork<B, KindIn, KindOut> + Send + 'static,
B: Default + Send + Sync + 'static,
{
TrainingInterface::<B>::train_ippo(
actor_id,
shutdown_rx,
runtime,
env_map,
loop_iters,
max_traj_length,
trainer_spec,
)
.map_err(StateManagerError::from)
}
pub(crate) fn run_env_step_loop_with_mappo<KindIn, KindOut, Pi>(
actor_id: ActorUuid,
shutdown_rx: tokio::sync::broadcast::Receiver<()>,
runtime: Arc<dyn ErasedActorRuntime<B>>,
env_map: Arc<DashMap<ActorUuid, EnvironmentInterface>>,
loop_iters: usize,
max_traj_length: usize,
trainer_spec: PPOTrainerSpec<B, KindIn, KindOut, Pi>,
) -> Result<ModelModule<B>, StateManagerError>
where
KindIn: TensorKind<B> + BasicOps<B> + Send + 'static,
KindOut: TensorKind<B> + BasicOps<B> + Numeric<B> + Send + 'static,
Pi: NeuralNetwork<B, KindIn, KindOut> + Send + 'static,
B: Default + Send + Sync + 'static,
{
TrainingInterface::<B>::train_mappo(
actor_id,
shutdown_rx,
runtime,
env_map,
loop_iters,
max_traj_length,
trainer_spec,
)
.map_err(StateManagerError::from)
}
}
pub(crate) fn env_dtype_to_dtype(dtype: &EnvDType) -> Result<DType, ActorError> {
match dtype {
EnvDType::NdArray(nd_type) => Ok(DType::NdArray(match nd_type {
EnvNdArrayDType::F16 => NdArrayDType::F16,
EnvNdArrayDType::F32 => NdArrayDType::F32,
EnvNdArrayDType::F64 => NdArrayDType::F64,
EnvNdArrayDType::I8 => NdArrayDType::I8,
EnvNdArrayDType::I16 => NdArrayDType::I16,
EnvNdArrayDType::I32 => NdArrayDType::I32,
EnvNdArrayDType::I64 => NdArrayDType::I64,
EnvNdArrayDType::Bool => NdArrayDType::Bool,
})),
#[cfg(feature = "tch-backend")]
EnvDType::Tch(tch_type) => Ok(DType::Tch(match tch_type {
EnvTchDType::F16 => TchDType::F16,
EnvTchDType::Bf16 => TchDType::Bf16,
EnvTchDType::F32 => TchDType::F32,
EnvTchDType::F64 => TchDType::F64,
EnvTchDType::I8 => TchDType::I8,
EnvTchDType::I16 => TchDType::I16,
EnvTchDType::I32 => TchDType::I32,
EnvTchDType::I64 => TchDType::I64,
EnvTchDType::U8 => TchDType::U8,
EnvTchDType::Bool => TchDType::Bool,
})),
#[cfg(not(feature = "tch-backend"))]
EnvDType::Tch(_) => Err(ActorError::TypeConversionError(format!(
"Unsupported environment dtype: {dtype:?}"
))),
}
}
pub(crate) fn decode_argmax(
data: &[u8],
dtype: &relayrl_types::data::tensor::DType,
n_envs: usize,
act_dim: usize,
) -> Vec<u8> {
macro_rules! argmax_float {
($T:ty) => {{
let vals: &[$T] = bytemuck::cast_slice(data);
(0..n_envs)
.map(|i| {
vals[i * act_dim..(i + 1) * act_dim]
.iter()
.copied()
.enumerate()
.max_by(|(_, a), (_, b)| {
a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal)
})
.map(|(j, _)| j as u8)
.unwrap_or(0)
})
.collect()
}};
}
macro_rules! argmax_int {
($T:ty) => {{
let vals: &[$T] = bytemuck::cast_slice(data);
(0..n_envs)
.map(|i| {
vals[i * act_dim..(i + 1) * act_dim]
.iter()
.copied()
.enumerate()
.max_by(|(_, a), (_, b)| a.cmp(b))
.map(|(j, _)| j as u8)
.unwrap_or(0)
})
.collect()
}};
}
match dtype {
DType::NdArray(NdArrayDType::F16) => argmax_float!(half::f16),
DType::NdArray(NdArrayDType::F32) => argmax_float!(f32),
DType::NdArray(NdArrayDType::F64) => argmax_float!(f64),
DType::NdArray(NdArrayDType::I8) => argmax_int!(i8),
DType::NdArray(NdArrayDType::I16) => argmax_int!(i16),
DType::NdArray(NdArrayDType::I32) => argmax_int!(i32),
DType::NdArray(NdArrayDType::I64) => argmax_int!(i64),
DType::NdArray(NdArrayDType::Bool) => argmax_int!(u8),
#[cfg(feature = "tch-backend")]
DType::Tch(tch_type) => match tch_type {
TchDType::F16 => argmax_float!(half::f16),
TchDType::Bf16 => argmax_float!(half::bf16),
TchDType::F32 => argmax_float!(f32),
TchDType::F64 => argmax_float!(f64),
TchDType::I8 => argmax_int!(i8),
TchDType::I16 => argmax_int!(i16),
TchDType::I32 => argmax_int!(i32),
TchDType::I64 => argmax_int!(i64),
TchDType::Bool => argmax_int!(u8),
TchDType::U8 => argmax_int!(u8),
},
}
}
pub(crate) fn decode_continuous_bytes(
data: &[u8],
src_dtype: &DType,
count: usize,
tgt_dtype: &DType,
) -> Vec<u8> {
let as_f64: Vec<f64> = match src_dtype {
DType::NdArray(NdArrayDType::F16) => bytemuck::cast_slice::<u8, half::f16>(data)[..count]
.iter()
.map(|&x| f64::from(x))
.collect(),
DType::NdArray(NdArrayDType::F32) => bytemuck::cast_slice::<u8, f32>(data)[..count]
.iter()
.map(|&x| x as f64)
.collect(),
DType::NdArray(NdArrayDType::F64) => {
bytemuck::cast_slice::<u8, f64>(data)[..count].to_vec()
}
DType::NdArray(NdArrayDType::I8) => bytemuck::cast_slice::<u8, i8>(data)[..count]
.iter()
.map(|&x| x as f64)
.collect(),
DType::NdArray(NdArrayDType::I16) => bytemuck::cast_slice::<u8, i16>(data)[..count]
.iter()
.map(|&x| x as f64)
.collect(),
DType::NdArray(NdArrayDType::I32) => bytemuck::cast_slice::<u8, i32>(data)[..count]
.iter()
.map(|&x| x as f64)
.collect(),
DType::NdArray(NdArrayDType::I64) => bytemuck::cast_slice::<u8, i64>(data)[..count]
.iter()
.map(|&x| x as f64)
.collect(),
DType::NdArray(NdArrayDType::Bool) => data[..count]
.iter()
.map(|&x| if x != 0 { 1.0f64 } else { 0.0 })
.collect(),
#[cfg(feature = "tch-backend")]
DType::Tch(tc) => {
use relayrl_types::data::tensor::TchDType;
match tc {
TchDType::F16 => bytemuck::cast_slice::<u8, half::f16>(data)[..count]
.iter()
.map(|&x| f64::from(x))
.collect(),
TchDType::Bf16 => bytemuck::cast_slice::<u8, half::bf16>(data)[..count]
.iter()
.map(|&x| f64::from(x))
.collect(),
TchDType::F32 => bytemuck::cast_slice::<u8, f32>(data)[..count]
.iter()
.map(|&x| x as f64)
.collect(),
TchDType::F64 => bytemuck::cast_slice::<u8, f64>(data)[..count].to_vec(),
TchDType::I8 => bytemuck::cast_slice::<u8, i8>(data)[..count]
.iter()
.map(|&x| x as f64)
.collect(),
TchDType::I16 => bytemuck::cast_slice::<u8, i16>(data)[..count]
.iter()
.map(|&x| x as f64)
.collect(),
TchDType::I32 => bytemuck::cast_slice::<u8, i32>(data)[..count]
.iter()
.map(|&x| x as f64)
.collect(),
TchDType::I64 => bytemuck::cast_slice::<u8, i64>(data)[..count]
.iter()
.map(|&x| x as f64)
.collect(),
TchDType::U8 => data[..count].iter().map(|&x| x as f64).collect(),
TchDType::Bool => data[..count]
.iter()
.map(|&x| if x != 0 { 1.0f64 } else { 0.0 })
.collect(),
}
}
};
match tgt_dtype {
DType::NdArray(NdArrayDType::F16) => {
let v: Vec<half::f16> = as_f64.iter().map(|&x| half::f16::from_f64(x)).collect();
bytemuck::cast_slice::<half::f16, u8>(&v).to_vec()
}
DType::NdArray(NdArrayDType::F32) => {
let v: Vec<f32> = as_f64.iter().map(|&x| x as f32).collect();
bytemuck::cast_slice::<f32, u8>(&v).to_vec()
}
DType::NdArray(NdArrayDType::F64) => bytemuck::cast_slice::<f64, u8>(&as_f64).to_vec(),
DType::NdArray(NdArrayDType::I8) => {
let v: Vec<i8> = as_f64.iter().map(|&x| x as i8).collect();
bytemuck::cast_slice::<i8, u8>(&v).to_vec()
}
DType::NdArray(NdArrayDType::I16) => {
let v: Vec<i16> = as_f64.iter().map(|&x| x as i16).collect();
bytemuck::cast_slice::<i16, u8>(&v).to_vec()
}
DType::NdArray(NdArrayDType::I32) => {
let v: Vec<i32> = as_f64.iter().map(|&x| x as i32).collect();
bytemuck::cast_slice::<i32, u8>(&v).to_vec()
}
DType::NdArray(NdArrayDType::I64) => {
let v: Vec<i64> = as_f64.iter().map(|&x| x as i64).collect();
bytemuck::cast_slice::<i64, u8>(&v).to_vec()
}
DType::NdArray(NdArrayDType::Bool) => as_f64
.iter()
.map(|&x| if x != 0.0 { 1u8 } else { 0u8 })
.collect(),
#[cfg(feature = "tch-backend")]
DType::Tch(tch_type) => match tch_type {
TchDType::F16 => {
let v: Vec<half::f16> = as_f64.iter().map(|&x| half::f16::from_f64(x)).collect();
bytemuck::cast_slice::<half::f16, u8>(&v).to_vec()
}
TchDType::Bf16 => {
let v: Vec<half::bf16> = as_f64.iter().map(|&x| half::bf16::from_f64(x)).collect();
bytemuck::cast_slice::<half::bf16, u8>(&v).to_vec()
}
TchDType::F32 => {
let v: Vec<f32> = as_f64.iter().map(|&x| x as f32).collect();
bytemuck::cast_slice::<f32, u8>(&v).to_vec()
}
TchDType::F64 => {
let v: Vec<f64> = as_f64.iter().map(|&x| x as f64).collect();
bytemuck::cast_slice::<f64, u8>(&v).to_vec()
}
TchDType::I8 => {
let v: Vec<i8> = as_f64.iter().map(|&x| x as i8).collect();
bytemuck::cast_slice::<i8, u8>(&v).to_vec()
}
TchDType::I16 => {
let v: Vec<i16> = as_f64.iter().map(|&x| x as i16).collect();
bytemuck::cast_slice::<i16, u8>(&v).to_vec()
}
TchDType::I32 => {
let v: Vec<i32> = as_f64.iter().map(|&x| x as i32).collect();
bytemuck::cast_slice::<i32, u8>(&v).to_vec()
}
TchDType::I64 => {
let v: Vec<i64> = as_f64.iter().map(|&x| x as i64).collect();
bytemuck::cast_slice::<i64, u8>(&v).to_vec()
}
TchDType::U8 => {
let v: Vec<u8> = as_f64.iter().map(|&x| x as u8).collect();
bytemuck::cast_slice::<u8, u8>(&v).to_vec()
}
TchDType::Bool => as_f64
.iter()
.map(|&x| if x != 0.0 { 1u8 } else { 0u8 })
.collect(),
},
}
}
#[cfg(test)]
mod unit_tests {
use super::*;
use crate::network::client::agent::{
ActorDataMode, ActorInferenceMode, ClientModes, ModelMode,
};
use active_uuid_registry::registry_uuid::Uuid;
use burn_ndarray::NdArray;
use std::path::PathBuf;
use std::sync::Arc;
use tokio::sync::{RwLock, mpsc};
type TestBackend = NdArray<f32>;
fn disabled_modes() -> Arc<ClientModes> {
Arc::new(ClientModes {
actor_inference_mode: ActorInferenceMode::Client(ModelMode::Independent),
actor_data_mode: ActorDataMode::Disabled,
})
}
fn shared_modes() -> Arc<ClientModes> {
Arc::new(ClientModes {
actor_inference_mode: ActorInferenceMode::Client(ModelMode::Shared),
actor_data_mode: ActorDataMode::Disabled,
})
}
fn make_state_manager(
modes: Arc<ClientModes>,
) -> (
StateManager<TestBackend>,
tokio::sync::mpsc::Receiver<RoutedMessage>,
) {
let namespace_str = format!("test-sm-{}", Uuid::new_v4());
let namespace_handle =
active_uuid_registry::interface::reserve_owned_namespace(&namespace_str)
.expect("reserve owned test namespace");
let namespace = ClientNamespace::new(namespace_handle, Arc::from(namespace_str));
StateManager::<TestBackend>::new(
namespace,
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
None,
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
None,
modes,
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
None,
Arc::new(RwLock::new(PathBuf::new())),
None,
#[cfg(feature = "metrics")]
test_metrics(),
)
}
#[cfg(feature = "metrics")]
fn test_metrics() -> MetricsManager {
MetricsManager::new(
Arc::new(RwLock::new((
"test-state-manager".to_string(),
String::new(),
))),
("test-state-manager".to_string(), String::new()),
None,
)
}
fn deterministic_actor_id(last_byte: u8) -> Uuid {
let mut bytes = [0_u8; 16];
bytes[15] = last_byte;
Uuid::from_bytes(bytes)
}
fn actor_info(id: Uuid) -> ActorInfo {
ActorInfo::new(id, None)
}
#[tokio::test]
async fn distribute_actors_round_robin_2_routers() {
let (sm, _rx) = make_state_manager(disabled_modes());
let actor_ids: Vec<Uuid> = (0..4).map(|_| Uuid::new_v4()).collect();
let ns1: RouterNamespace = Arc::from("r1");
let ns2: RouterNamespace = Arc::from("r2");
for id in &actor_ids {
let handle = Arc::new(tokio::spawn(async {}));
sm.actor_handles.insert(*id, handle);
let (tx_to_actor, _actor_inbox_rx) = mpsc::channel::<RoutedMessage>(CHANNEL_THROUGHPUT);
sm.shared_router_state.actor_routes.insert(
*id,
ActorRoute {
router_namespace: None,
inbox: tx_to_actor,
},
);
}
sm.distribute_actors(vec![ns1.clone(), ns2.clone()]);
let mut sorted_actor_ids = actor_ids.clone();
sorted_actor_ids.sort_by_key(|actor_id| actor_id.to_string());
assert_eq!(sm.shared_router_state.actor_routes.len(), 4);
for (index, id) in sorted_actor_ids.iter().enumerate() {
let assigned = sm.shared_router_state.actor_routes.get(id).unwrap();
let expected_namespace = if index % 2 == 0 {
ns1.clone()
} else {
ns2.clone()
};
assert_eq!(
assigned.router_namespace,
Some(expected_namespace),
"Actor {} assigned to unexpected namespace",
id
);
}
}
#[tokio::test]
async fn distribute_actors_empty_namespaces_is_noop() {
let (sm, _rx) = make_state_manager(disabled_modes());
let actor_id = Uuid::new_v4();
let original_ns: RouterNamespace = Arc::from("original");
sm.actor_handles
.insert(actor_id, Arc::new(tokio::spawn(async {})));
let (tx_to_actor, _actor_inbox_rx) = mpsc::channel::<RoutedMessage>(CHANNEL_THROUGHPUT);
sm.shared_router_state.actor_routes.insert(
actor_id,
ActorRoute {
router_namespace: Some(original_ns.clone()),
inbox: tx_to_actor,
},
);
sm.distribute_actors(vec![]);
let assigned = sm.shared_router_state.actor_routes.get(&actor_id).unwrap();
assert_eq!(
assigned.router_namespace,
Some(original_ns),
"Namespace should not change"
);
}
#[tokio::test]
async fn distribute_actors_single_namespace() {
let (sm, _rx) = make_state_manager(disabled_modes());
let actor_ids: Vec<Uuid> = (0..3).map(|_| Uuid::new_v4()).collect();
let ns: RouterNamespace = Arc::from("only");
for id in &actor_ids {
sm.actor_handles
.insert(*id, Arc::new(tokio::spawn(async {})));
let (tx_to_actor, _actor_inbox_rx) = mpsc::channel::<RoutedMessage>(CHANNEL_THROUGHPUT);
sm.shared_router_state.actor_routes.insert(
*id,
ActorRoute {
router_namespace: None,
inbox: tx_to_actor,
},
);
}
sm.distribute_actors(vec![ns.clone()]);
for id in &actor_ids {
let assigned = sm.shared_router_state.actor_routes.get(id).unwrap();
assert_eq!(assigned.router_namespace, Some(ns.clone()));
}
}
#[test]
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
fn restore_replaces_all_mappings() {
let (sm, _rx) = make_state_manager(disabled_modes());
let old_id = Uuid::new_v4();
let new_id = Uuid::new_v4();
let old_ns: RouterNamespace = Arc::from("old");
let new_ns: RouterNamespace = Arc::from("new");
let (old_tx, _old_rx) = mpsc::channel::<RoutedMessage>(CHANNEL_THROUGHPUT);
let (new_tx, _new_rx) = mpsc::channel::<RoutedMessage>(CHANNEL_THROUGHPUT);
sm.shared_router_state.actor_routes.insert(
old_id,
ActorRoute {
router_namespace: Some(old_ns),
inbox: old_tx,
},
);
sm.shared_router_state.actor_routes.insert(
new_id,
ActorRoute {
router_namespace: None,
inbox: new_tx,
},
);
sm.restore_actor_router_mappings(vec![(new_id, new_ns.clone())]);
assert!(
matches!(
sm.shared_router_state.actor_routes.get(&old_id),
Some(route) if route.router_namespace.is_none()
),
"Old mapping should be cleared while preserving the inbox"
);
let assigned = sm.shared_router_state.actor_routes.get(&new_id).unwrap();
assert_eq!(assigned.router_namespace, Some(new_ns));
}
#[test]
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
fn restore_with_empty_clears_all() {
let (sm, _rx) = make_state_manager(disabled_modes());
let actor_id = Uuid::new_v4();
let (tx_to_actor, _actor_inbox_rx) = mpsc::channel::<RoutedMessage>(CHANNEL_THROUGHPUT);
sm.shared_router_state.actor_routes.insert(
actor_id,
ActorRoute {
router_namespace: Some(Arc::from("ns")),
inbox: tx_to_actor,
},
);
sm.restore_actor_router_mappings(vec![]);
assert!(matches!(
sm.shared_router_state.actor_routes.get(&actor_id),
Some(route) if route.router_namespace.is_none()
));
}
#[test]
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
fn get_returns_all_inserted_mappings() {
let (sm, _rx) = make_state_manager(disabled_modes());
let id1 = Uuid::new_v4();
let id2 = Uuid::new_v4();
let id3 = Uuid::new_v4();
let ns1: RouterNamespace = Arc::from("r1");
let ns2: RouterNamespace = Arc::from("r2");
let ns3: RouterNamespace = Arc::from("r3");
let (tx1, _rx1) = mpsc::channel::<RoutedMessage>(CHANNEL_THROUGHPUT);
let (tx2, _rx2) = mpsc::channel::<RoutedMessage>(CHANNEL_THROUGHPUT);
let (tx3, _rx3) = mpsc::channel::<RoutedMessage>(CHANNEL_THROUGHPUT);
sm.shared_router_state.actor_routes.insert(
id1,
ActorRoute {
router_namespace: Some(ns1.clone()),
inbox: tx1,
},
);
sm.shared_router_state.actor_routes.insert(
id2,
ActorRoute {
router_namespace: Some(ns2.clone()),
inbox: tx2,
},
);
sm.shared_router_state.actor_routes.insert(
id3,
ActorRoute {
router_namespace: Some(ns3.clone()),
inbox: tx3,
},
);
let result = sm.get_actor_router_mappings();
assert_eq!(result.len(), 3);
assert!(result.iter().any(|(id, ns)| *id == id1 && *ns == ns1));
assert!(result.iter().any(|(id, ns)| *id == id2 && *ns == ns2));
assert!(result.iter().any(|(id, ns)| *id == id3 && *ns == ns3));
}
#[tokio::test]
async fn get_actor_id_list_reflects_inserted_handles() {
let (sm, _rx) = make_state_manager(disabled_modes());
let ids: Vec<Uuid> = (0..3).map(|_| Uuid::new_v4()).collect();
for id in &ids {
sm.actor_handles
.insert(*id, Arc::new(tokio::spawn(async {})));
}
let list = sm.get_actor_id_list();
assert_eq!(list.len(), 3);
for id in &ids {
assert!(list.contains(id), "Actor {} not in list", id);
}
}
#[tokio::test]
async fn remove_actor_clears_device_and_router_metadata() {
let (mut sm, _rx) = make_state_manager(disabled_modes());
let actor_id = sm
.client_namespace
.reserve_id_with(crate::network::ACTOR_CONTEXT, 117, 100)
.unwrap();
let (tx_to_actor, _actor_inbox_rx) = mpsc::channel::<RoutedMessage>(CHANNEL_THROUGHPUT);
sm.actor_handles
.insert(actor_id, Arc::new(tokio::spawn(async {})));
sm.shared_router_state.actor_routes.insert(
actor_id,
ActorRoute {
router_namespace: Some(Arc::from("router-a")),
inbox: tx_to_actor,
},
);
sm.actor_devices.insert(actor_id, DeviceType::Cpu);
sm.remove_actor(actor_id).unwrap();
assert!(sm.actor_handles.get(&actor_id).is_none());
assert!(sm.shared_router_state.actor_routes.get(&actor_id).is_none());
assert!(sm.actor_devices.get(&actor_id).is_none());
}
#[tokio::test]
async fn set_actor_id_moves_device_and_router_metadata() {
let (sm, _rx) = make_state_manager(disabled_modes());
let current_id = sm
.client_namespace
.reserve_id_with(crate::network::ACTOR_CONTEXT, 117, 100)
.unwrap();
let new_id = Uuid::new_v4();
let (tx_to_actor, _actor_inbox_rx) = mpsc::channel::<RoutedMessage>(CHANNEL_THROUGHPUT);
sm.actor_handles
.insert(current_id, Arc::new(tokio::spawn(async {})));
sm.shared_router_state.actor_routes.insert(
current_id,
ActorRoute {
router_namespace: Some(Arc::from("router-a")),
inbox: tx_to_actor,
},
);
sm.actor_devices.insert(current_id, DeviceType::Cpu);
sm.set_actor_id(current_id, new_id).unwrap();
assert!(sm.actor_handles.get(¤t_id).is_none());
assert!(
sm.shared_router_state
.actor_routes
.get(¤t_id)
.is_none()
);
assert!(sm.actor_devices.get(¤t_id).is_none());
assert!(sm.actor_handles.get(&new_id).is_some());
assert!(sm.shared_router_state.actor_routes.get(&new_id).is_some());
assert!(matches!(
sm.actor_devices.get(&new_id),
Some(device) if *device == DeviceType::Cpu
));
assert!(matches!(
sm.shared_router_state.actor_routes.get(&new_id),
Some(route) if route.router_namespace == Some(Arc::<str>::from("router-a"))
));
}
#[tokio::test]
async fn model_update_dispatch_targets_returns_all_actor_ids_in_independent_mode() {
let (sm, _rx) = make_state_manager(disabled_modes());
let ids: Vec<Uuid> = (0..3).map(|_| Uuid::new_v4()).collect();
for id in &ids {
sm.actor_handles
.insert(*id, Arc::new(tokio::spawn(async {})));
}
let mut expected: Vec<ActorInfo> = ids.iter().copied().map(actor_info).collect();
expected.sort_by_key(|actor| actor.id().to_string());
assert_eq!(sm.model_update_dispatch_targets(), expected);
}
#[tokio::test]
async fn model_update_dispatch_targets_deduplicates_shared_mode_by_device() {
let (sm, _rx) = make_state_manager(shared_modes());
let ids: Vec<Uuid> = (0..3).map(|_| Uuid::new_v4()).collect();
for id in &ids {
sm.actor_handles
.insert(*id, Arc::new(tokio::spawn(async {})));
sm.actor_devices.insert(*id, DeviceType::Cpu);
}
let expected_target = ids
.iter()
.min_by_key(|actor_id| actor_id.to_string())
.copied()
.unwrap();
assert_eq!(
sm.model_update_dispatch_targets(),
vec![actor_info(expected_target)]
);
}
#[tokio::test]
async fn model_update_dispatch_targets_for_subset_returns_known_actor_ids_in_independent_mode()
{
let (sm, _rx) = make_state_manager(disabled_modes());
let id1 = deterministic_actor_id(1);
let id2 = deterministic_actor_id(2);
let id3 = deterministic_actor_id(3);
let unknown_id = deterministic_actor_id(9);
for actor_id in [id1, id2, id3] {
sm.actor_handles
.insert(actor_id, Arc::new(tokio::spawn(async {})));
}
let subset = vec![
actor_info(id3),
actor_info(unknown_id),
actor_info(id1),
actor_info(id3),
];
assert_eq!(
sm.model_update_dispatch_targets_for_subset(Some(&subset)),
vec![actor_info(id1), actor_info(id3)]
);
}
#[tokio::test]
async fn model_update_dispatch_targets_for_subset_ignores_unknown_actor_ids() {
let (sm, _rx) = make_state_manager(disabled_modes());
let known_id = deterministic_actor_id(1);
let unknown_id = deterministic_actor_id(2);
sm.actor_handles
.insert(known_id, Arc::new(tokio::spawn(async {})));
let subset = vec![actor_info(unknown_id)];
assert!(
sm.model_update_dispatch_targets_for_subset(Some(&subset))
.is_empty()
);
let subset = vec![actor_info(unknown_id), actor_info(known_id)];
assert_eq!(
sm.model_update_dispatch_targets_for_subset(Some(&subset)),
vec![actor_info(known_id)]
);
}
#[tokio::test]
#[cfg(feature = "tch-backend")]
async fn model_update_dispatch_targets_for_subset_deduplicates_selected_shared_devices() {
let (sm, _rx) = make_state_manager(shared_modes());
let cpu_small = deterministic_actor_id(1);
let cpu_large = deterministic_actor_id(2);
let cuda_small = deterministic_actor_id(3);
let cuda_large = deterministic_actor_id(4);
for (actor_id, device) in [
(cpu_small, DeviceType::Cpu),
(cpu_large, DeviceType::Cpu),
(cuda_small, DeviceType::Cuda(0)),
(cuda_large, DeviceType::Cuda(0)),
] {
sm.actor_handles
.insert(actor_id, Arc::new(tokio::spawn(async {})));
sm.actor_devices.insert(actor_id, device);
}
let subset = vec![
actor_info(cuda_large),
actor_info(cpu_large),
actor_info(cpu_small),
];
assert_eq!(
sm.model_update_dispatch_targets_for_subset(Some(&subset)),
vec![actor_info(cpu_small), actor_info(cuda_large)]
);
}
#[tokio::test]
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
async fn canonical_model_update_target_uses_shared_device_representative() {
let (sm, _rx) = make_state_manager(shared_modes());
let ids: Vec<Uuid> = (0..3).map(|_| Uuid::new_v4()).collect();
for id in &ids {
sm.actor_handles
.insert(*id, Arc::new(tokio::spawn(async {})));
sm.actor_devices.insert(*id, DeviceType::Cpu);
}
let expected_target = ids
.iter()
.min_by_key(|actor_id| actor_id.to_string())
.copied()
.unwrap();
for id in &ids {
assert_eq!(
sm.canonical_model_update_target(&actor_info(*id)),
actor_info(expected_target)
);
}
}
#[tokio::test]
#[cfg(any(feature = "nats-transport", feature = "zmq-transport"))]
async fn canonical_model_update_target_preserves_independent_actor_ids() {
let (sm, _rx) = make_state_manager(disabled_modes());
let ids: Vec<Uuid> = (0..3).map(|_| Uuid::new_v4()).collect();
for id in &ids {
sm.actor_handles
.insert(*id, Arc::new(tokio::spawn(async {})));
}
for id in &ids {
assert_eq!(
sm.canonical_model_update_target(&actor_info(*id)),
actor_info(*id)
);
}
}
#[tokio::test]
async fn shared_mode_second_actor_reuses_same_arc() {
let (mut sm, _rx) = make_state_manager(shared_modes());
let (h1, needs1) = sm
.get_or_init_model_handle(None, DeviceType::Cpu)
.await
.unwrap();
let (h2, needs2) = sm
.get_or_init_model_handle(None, DeviceType::Cpu)
.await
.unwrap();
assert!(
Arc::ptr_eq(&h1, &h2),
"Shared mode should reuse the same Arc"
);
assert!(needs1, "First call should need handshake (no model)");
assert!(!needs2, "Second call should NOT need handshake");
}
#[tokio::test]
async fn independent_mode_each_actor_gets_fresh_arc() {
let (mut sm, _rx) = make_state_manager(disabled_modes());
let (h1, _) = sm
.get_or_init_model_handle(None, DeviceType::Cpu)
.await
.unwrap();
let (h2, _) = sm
.get_or_init_model_handle(None, DeviceType::Cpu)
.await
.unwrap();
assert!(
!Arc::ptr_eq(&h1, &h2),
"Independent mode should create fresh Arc each time"
);
}
#[tokio::test]
async fn no_model_and_empty_path_sets_needs_handshake() {
let (mut sm, _rx) = make_state_manager(disabled_modes());
let (_, needs_handshake) = sm
.get_or_init_model_handle(None, DeviceType::Cpu)
.await
.unwrap();
assert!(
needs_handshake,
"No model available → needs_handshake must be true"
);
}
}