use std::collections::{BTreeMap, BTreeSet, HashMap};
use std::pin::Pin;
use std::sync::{Arc, Mutex, MutexGuard};
use std::task::{Context, Poll};
use std::time::{Duration, Instant};
use aion_core::{
ClusterEvent, DeploymentAssociation, InterventionCapabilities, WorkerDeathReason,
WorkerTransport,
};
use aion_proto::{ProtoActivityTask, ProtoCancelActivity, ProtoLivenessPing, ProtoRegisterWorker};
use aion_store::{NamespaceOrigin, NamespacePlacement, NamespaceStore, WorkerDeploymentStore};
use tokio::sync::{Notify, mpsc};
use crate::cluster_publisher::ClusterEventPublisher;
use crate::config::AutoCreate;
use crate::error::ServerError;
use crate::namespace::{CallerIdentity, NamespaceGuard, NamespaceMinter, NamespaceOperation};
use crate::observability::Metrics;
use crate::worker::admission_audit::AdmissionAudit;
use crate::worker::heartbeat::DispatchExclusion;
pub use aion_core::DEFAULT_TASK_QUEUE;
pub type WorkerTaskSender = mpsc::Sender<WorkerMessage>;
#[derive(Clone, Debug)]
pub enum WorkerDelivery {
Grpc(WorkerTaskSender),
#[cfg(feature = "liminal-transport")]
Liminal(crate::worker::liminal_transport::LiminalWorkerDelivery),
}
impl WorkerDelivery {
#[must_use]
pub const fn transport(&self) -> WorkerTransport {
match self {
Self::Grpc(_) => WorkerTransport::Grpc,
#[cfg(feature = "liminal-transport")]
Self::Liminal(_) => WorkerTransport::Liminal,
}
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub enum WorkerMessage {
ActivityTask(Box<ProtoActivityTask>),
DrainRequest,
LivenessPing(ProtoLivenessPing),
CancelActivity(ProtoCancelActivity),
}
#[derive(Clone, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)]
pub struct PoolAddress {
namespace: String,
task_queue: String,
}
impl PoolAddress {
#[must_use]
pub fn new(namespace: impl Into<String>, task_queue: impl Into<String>) -> Self {
let task_queue = task_queue.into();
let task_queue = if task_queue.is_empty() {
String::from(DEFAULT_TASK_QUEUE)
} else {
task_queue
};
Self {
namespace: namespace.into(),
task_queue,
}
}
#[must_use]
pub fn namespace(&self) -> &str {
&self.namespace
}
#[must_use]
pub fn task_queue(&self) -> &str {
&self.task_queue
}
}
#[derive(Clone, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)]
struct ActivityKey {
pool: PoolAddress,
activity_type: String,
}
impl ActivityKey {
fn new(pool: PoolAddress, activity_type: impl Into<String>) -> Self {
Self {
pool,
activity_type: activity_type.into(),
}
}
}
type WorkerMap = HashMap<WorkerId, WorkerHandle>;
type RegistryMap = HashMap<ActivityKey, WorkerMap>;
#[derive(Clone, Copy, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)]
pub struct WorkerId(u64);
impl WorkerId {
#[must_use]
pub const fn from_value(value: u64) -> Self {
Self(value)
}
#[must_use]
pub const fn value(self) -> u64 {
self.0
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct WorkerInstanceIdentity {
pub deployment: String,
pub instance_id: String,
pub association: DeploymentAssociation,
}
struct RegistrationOptions {
instance: Option<WorkerInstanceIdentity>,
intervention_capabilities: InterventionCapabilities,
}
#[derive(Clone, Debug)]
pub struct WorkerHandle {
id: WorkerId,
namespaces: BTreeSet<String>,
task_queue: String,
node: Option<String>,
activity_types: BTreeSet<String>,
instance: Option<WorkerInstanceIdentity>,
delivery: WorkerDelivery,
intervention_capabilities: InterventionCapabilities,
}
impl WorkerHandle {
#[must_use]
pub const fn id(&self) -> WorkerId {
self.id
}
#[must_use]
pub const fn namespaces(&self) -> &BTreeSet<String> {
&self.namespaces
}
#[must_use]
pub fn task_queue(&self) -> &str {
&self.task_queue
}
#[must_use]
pub fn node(&self) -> Option<&str> {
self.node.as_deref()
}
#[must_use]
pub fn activity_types(&self) -> &BTreeSet<String> {
&self.activity_types
}
#[must_use]
pub const fn instance(&self) -> Option<&WorkerInstanceIdentity> {
self.instance.as_ref()
}
#[must_use]
pub const fn delivery(&self) -> &WorkerDelivery {
&self.delivery
}
#[must_use]
pub const fn intervention_capabilities(&self) -> &InterventionCapabilities {
&self.intervention_capabilities
}
#[must_use]
pub fn sender(&self) -> Option<&WorkerTaskSender> {
match &self.delivery {
WorkerDelivery::Grpc(sender) => Some(sender),
#[cfg(feature = "liminal-transport")]
WorkerDelivery::Liminal(_) => None,
}
}
}
#[derive(Debug)]
struct RegistryState {
next_worker_id: u64,
workers: BTreeMap<WorkerId, WorkerHandle>,
by_activity: RegistryMap,
rotation: HashMap<ActivityKey, usize>,
last_departure: HashMap<ActivityKey, BTreeMap<Option<String>, Instant>>,
dispatch_ineligible: BTreeMap<WorkerId, DispatchExclusion>,
}
impl Default for RegistryState {
fn default() -> Self {
Self {
next_worker_id: 1,
workers: BTreeMap::new(),
by_activity: HashMap::new(),
rotation: HashMap::new(),
last_departure: HashMap::new(),
dispatch_ineligible: BTreeMap::new(),
}
}
}
#[derive(Clone)]
pub struct ConnectedWorkerRegistry {
inner: Arc<Mutex<RegistryState>>,
metrics: Option<Metrics>,
cluster_publisher: Option<ClusterEventPublisher>,
minter: Option<NamespaceMinter>,
deployment_store: Option<Arc<dyn WorkerDeploymentStore>>,
worker_arrived: Arc<Notify>,
audit: Arc<AdmissionAudit>,
}
impl std::fmt::Debug for ConnectedWorkerRegistry {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("ConnectedWorkerRegistry")
.field(
"deployment_store_attached",
&self.deployment_store.is_some(),
)
.finish_non_exhaustive()
}
}
impl Default for ConnectedWorkerRegistry {
fn default() -> Self {
Self {
inner: Arc::new(Mutex::new(RegistryState::default())),
metrics: None,
cluster_publisher: None,
minter: None,
deployment_store: None,
worker_arrived: Arc::new(Notify::new()),
audit: Arc::new(AdmissionAudit::new()),
}
}
}
#[must_use = "a WorkerArrival that is constructed and dropped is a subscription thrown away"]
pub struct WorkerArrival {
notified: Pin<Box<tokio::sync::futures::OwnedNotified>>,
}
impl std::fmt::Debug for WorkerArrival {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.debug_struct("WorkerArrival").finish()
}
}
impl std::future::Future for WorkerArrival {
type Output = ();
fn poll(self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<()> {
std::future::Future::poll(self.get_mut().notified.as_mut(), context)
}
}
impl ConnectedWorkerRegistry {
#[must_use]
pub fn with_metrics(metrics: Metrics) -> Self {
Self {
inner: Arc::new(Mutex::new(RegistryState::default())),
metrics: Some(metrics),
cluster_publisher: None,
minter: None,
deployment_store: None,
worker_arrived: Arc::new(Notify::new()),
audit: Arc::new(AdmissionAudit::new()),
}
}
#[must_use]
pub fn admission_audit(&self) -> &AdmissionAudit {
&self.audit
}
#[must_use]
pub fn with_cluster_publisher(mut self, publisher: ClusterEventPublisher) -> Self {
self.cluster_publisher = Some(publisher);
self
}
#[must_use]
pub fn with_worker_deployment_store(mut self, store: Arc<dyn WorkerDeploymentStore>) -> Self {
self.deployment_store = Some(store);
self
}
#[must_use]
pub fn with_namespace_minting(
mut self,
store: Arc<dyn NamespaceStore>,
policy: AutoCreate,
) -> Self {
let minter = NamespaceMinter::new(store, policy);
let minter = match &self.cluster_publisher {
Some(publisher) => minter.with_cluster_publisher(publisher.clone()),
None => minter,
};
self.minter = Some(minter);
self
}
#[must_use]
pub fn with_namespace_routing(mut self, routing: crate::namespace::NamespaceRouting) -> Self {
self.minter = self.minter.map(|minter| minter.with_routing(routing));
self
}
pub async fn accept_registration(
&self,
guard: &NamespaceGuard,
caller: &CallerIdentity,
registration: &ProtoRegisterWorker,
sender: WorkerTaskSender,
) -> Result<WorkerRegistration, ServerError> {
guard
.scope(caller, &NamespaceOperation::register_worker(registration))
.await?;
let namespaces = guard.scope_worker_namespaces(caller, ®istration.namespaces)?;
self.mint_or_gate_namespaces(&namespaces).await?;
let node = optional_node(®istration.node);
self.enforce_pinned_placement(&namespaces, node.as_deref())
.await?;
let instance = self
.resolve_instance_identity(registration.instance.as_ref())
.await?;
self.register_delivery(
namespaces,
registration.task_queue.clone(),
node,
instance,
registration.activity_types.iter(),
WorkerDelivery::Grpc(sender),
)
}
pub async fn resolve_instance_identity(
&self,
instance: Option<&aion_proto::ProtoWorkerInstanceIdentity>,
) -> Result<Option<WorkerInstanceIdentity>, ServerError> {
let Some(instance) = instance else {
return Ok(None);
};
let association = match &self.deployment_store {
Some(store) => {
if store
.get_worker_deployment(&instance.deployment)
.await
.map_err(ServerError::from)?
.is_some()
{
DeploymentAssociation::Known
} else {
DeploymentAssociation::Absent
}
}
None => DeploymentAssociation::Unchecked,
};
Ok(Some(WorkerInstanceIdentity {
deployment: instance.deployment.clone(),
instance_id: instance.instance_id.clone(),
association,
}))
}
async fn mint_or_gate_namespaces(&self, namespaces: &[String]) -> Result<(), ServerError> {
let Some(minter) = &self.minter else {
return Ok(());
};
minter
.mint_or_gate(namespaces, NamespaceOrigin::WorkerMint)
.await
}
async fn enforce_pinned_placement(
&self,
namespaces: &[String],
node: Option<&str>,
) -> Result<(), ServerError> {
let Some(minter) = &self.minter else {
return Ok(());
};
for namespace in namespaces {
let NamespacePlacement::Pinned { nodes } = minter.placement_of(namespace).await? else {
continue;
};
let admitted = node.is_some_and(|n| nodes.contains(n));
if !admitted {
return Err(ServerError::placement_admission_denied(
namespace, node, &nodes,
));
}
}
Ok(())
}
pub fn register<'a>(
&self,
namespace: impl Into<String>,
activity_types: impl IntoIterator<Item = &'a String>,
sender: WorkerTaskSender,
) -> Result<WorkerRegistration, ServerError> {
self.register_namespaces(
[namespace.into()],
String::from(DEFAULT_TASK_QUEUE),
None,
activity_types,
sender,
)
}
pub fn register_pool<'a>(
&self,
pool: PoolAddress,
activity_types: impl IntoIterator<Item = &'a String>,
sender: WorkerTaskSender,
) -> Result<WorkerRegistration, ServerError> {
let PoolAddress {
namespace,
task_queue,
} = pool;
self.register_namespaces([namespace], task_queue, None, activity_types, sender)
}
pub fn register_namespaces<'a>(
&self,
namespaces: impl IntoIterator<Item = String>,
task_queue: impl Into<String>,
node: Option<String>,
activity_types: impl IntoIterator<Item = &'a String>,
sender: WorkerTaskSender,
) -> Result<WorkerRegistration, ServerError> {
self.register_delivery(
namespaces,
task_queue,
node,
None,
activity_types,
WorkerDelivery::Grpc(sender),
)
}
pub fn register_delivery<'a>(
&self,
namespaces: impl IntoIterator<Item = String>,
task_queue: impl Into<String>,
node: Option<String>,
instance: Option<WorkerInstanceIdentity>,
activity_types: impl IntoIterator<Item = &'a String>,
delivery: WorkerDelivery,
) -> Result<WorkerRegistration, ServerError> {
self.register_delivery_inner(
namespaces,
task_queue,
node,
activity_types,
delivery,
RegistrationOptions {
instance,
intervention_capabilities: InterventionCapabilities::none(),
},
)
}
pub fn register_delivery_with_capabilities<'a>(
&self,
namespaces: impl IntoIterator<Item = String>,
task_queue: impl Into<String>,
node: Option<String>,
activity_types: impl IntoIterator<Item = &'a String>,
delivery: WorkerDelivery,
intervention_capabilities: InterventionCapabilities,
) -> Result<WorkerRegistration, ServerError> {
self.register_delivery_inner(
namespaces,
task_queue,
node,
activity_types,
delivery,
RegistrationOptions {
instance: None,
intervention_capabilities,
},
)
}
fn register_delivery_inner<'a>(
&self,
namespaces: impl IntoIterator<Item = String>,
task_queue: impl Into<String>,
node: Option<String>,
activity_types: impl IntoIterator<Item = &'a String>,
delivery: WorkerDelivery,
options: RegistrationOptions,
) -> Result<WorkerRegistration, ServerError> {
let namespaces = namespaces.into_iter().collect::<BTreeSet<_>>();
let task_queue = task_queue.into();
let activity_types = activity_types.into_iter().cloned().collect::<BTreeSet<_>>();
let mut state = self.state()?;
let worker_id = WorkerId(state.next_worker_id);
state.next_worker_id = state.next_worker_id.saturating_add(1);
let node_for_event = node.clone();
let instance_for_event = options.instance.clone();
let handle = WorkerHandle {
id: worker_id,
namespaces: namespaces.clone(),
task_queue: task_queue.clone(),
node,
activity_types: activity_types.clone(),
instance: options.instance,
delivery,
intervention_capabilities: options.intervention_capabilities,
};
for namespace in &namespaces {
let pool = PoolAddress::new(namespace.clone(), task_queue.clone());
for activity_type in &activity_types {
let key = ActivityKey::new(pool.clone(), activity_type.clone());
if let Some(by_node) = state.last_departure.get_mut(&key) {
by_node.remove(&handle.node);
if by_node.is_empty() {
state.last_departure.remove(&key);
}
}
state
.by_activity
.entry(key)
.or_default()
.insert(worker_id, handle.clone());
}
}
let transport = handle.delivery.transport();
state.workers.insert(worker_id, handle);
drop(state);
if let Some(metrics) = &self.metrics {
for namespace in &namespaces {
metrics.worker_connected(namespace);
}
}
if let Some(publisher) = &self.cluster_publisher {
let namespaces_vec: Vec<String> = namespaces.iter().cloned().collect();
let task_queue_owned = task_queue.clone();
drop(publisher.emit(|meta| {
ClusterEvent::WorkerConnected {
meta,
worker_id: worker_id.value().to_string(),
namespaces: namespaces_vec,
task_queue: task_queue_owned,
transport,
node: node_for_event,
deployment: instance_for_event
.as_ref()
.map(|identity| identity.deployment.clone()),
deployment_association: instance_for_event.map(|identity| identity.association),
}
}));
}
self.worker_arrived.notify_waiters();
Ok(WorkerRegistration {
registry: self.clone(),
parts: Some(WorkerRegistrationParts {
worker_id,
namespaces,
task_queue,
activity_types,
}),
})
}
pub fn worker_arrival(&self) -> WorkerArrival {
let mut notified = Box::pin(Arc::clone(&self.worker_arrived).notified_owned());
notified.as_mut().enable();
WorkerArrival { notified }
}
pub fn workers_for(
&self,
namespace: &str,
task_queue: &str,
activity_type: &str,
node: Option<&str>,
) -> Result<Vec<WorkerHandle>, ServerError> {
let mut state = self.state()?;
let key = ActivityKey::new(PoolAddress::new(namespace, task_queue), activity_type);
Ok(eligible_candidates_in_rotation(&mut state, key, node))
}
pub fn all_workers(&self) -> Result<Vec<WorkerHandle>, ServerError> {
let state = self.state()?;
Ok(state.workers.values().cloned().collect())
}
pub fn worker_by_id(&self, worker_id: WorkerId) -> Result<Option<WorkerHandle>, ServerError> {
Ok(self.state()?.workers.get(&worker_id).cloned())
}
pub fn set_intervention_capabilities(
&self,
worker_id: WorkerId,
capabilities: &InterventionCapabilities,
) -> Result<bool, ServerError> {
let mut state = self.state()?;
if !state.workers.contains_key(&worker_id) {
return Ok(false);
}
if let Some(handle) = state.workers.get_mut(&worker_id) {
handle.intervention_capabilities = capabilities.clone();
}
for workers in state.by_activity.values_mut() {
if let Some(handle) = workers.get_mut(&worker_id) {
handle.intervention_capabilities = capabilities.clone();
}
}
Ok(true)
}
pub fn broadcast_drain(&self) -> Result<usize, ServerError> {
let workers = self.all_workers()?;
let mut delivered = 0usize;
for worker in workers {
if self.drain_worker(worker.id())? {
delivered = delivered.saturating_add(1);
}
}
Ok(delivered)
}
pub fn drain_worker(&self, worker_id: WorkerId) -> Result<bool, ServerError> {
let worker = {
let mut state = self.state()?;
let Some(worker) = state.workers.get(&worker_id).cloned() else {
return Ok(false);
};
Self::remove_worker_from_service(&mut state, &worker);
worker
};
match worker.delivery() {
WorkerDelivery::Grpc(sender) => {
if sender.try_send(WorkerMessage::DrainRequest).is_ok() {
tracing::info!(worker_id = worker_id.value(), "worker drain requested");
Ok(true)
} else {
tracing::error!(
worker_id = worker_id.value(),
"worker drain signal failed; force-deregistering closed transport"
);
self.deregister(worker_id)?;
Ok(false)
}
}
#[cfg(feature = "liminal-transport")]
WorkerDelivery::Liminal(delivery) => {
tracing::error!(
worker_id = worker_id.value(),
connection_pid = delivery.pid(),
"liminal transport has no drain control channel; worker fenced and \
deregistered at drain start"
);
self.deregister(worker_id)?;
Ok(false)
}
}
}
pub fn select_worker(
&self,
namespace: &str,
task_queue: &str,
activity_type: &str,
node: Option<&str>,
) -> Result<Option<WorkerHandle>, ServerError> {
let mut state = self.state()?;
let key = ActivityKey::new(PoolAddress::new(namespace, task_queue), activity_type);
Ok(eligible_candidates_in_rotation(&mut state, key, node)
.into_iter()
.next())
}
pub fn set_dispatch_ineligible(
&self,
unreachable: BTreeMap<WorkerId, DispatchExclusion>,
) -> Result<(), ServerError> {
self.state()?.dispatch_ineligible = unreachable;
self.worker_arrived.notify_waiters();
Ok(())
}
pub fn transports_of(
&self,
workers: impl IntoIterator<Item = WorkerId>,
) -> Result<BTreeMap<WorkerId, WorkerTransport>, ServerError> {
let state = self.state()?;
Ok(workers
.into_iter()
.filter_map(|worker_id| {
state
.workers
.get(&worker_id)
.map(|worker| (worker_id, worker.delivery.transport()))
})
.collect())
}
pub fn is_dispatch_ineligible(&self, worker_id: WorkerId) -> Result<bool, ServerError> {
Ok(self.state()?.dispatch_ineligible.contains_key(&worker_id))
}
pub fn grpc_liveness_targets(
&self,
) -> Result<Vec<super::grpc_liveness::GrpcLivenessTarget>, ServerError> {
Ok(self
.state()?
.workers
.values()
.filter_map(|worker| match &worker.delivery {
WorkerDelivery::Grpc(sender) => Some(super::grpc_liveness::GrpcLivenessTarget {
worker_id: worker.id,
sender: sender.clone(),
}),
#[cfg(feature = "liminal-transport")]
WorkerDelivery::Liminal(_) => None,
})
.collect())
}
pub fn ineligible_workers_over_tiers(
&self,
namespace: &str,
task_queue: &str,
activity_type: &str,
tiers: &[Option<String>],
) -> Result<usize, ServerError> {
let state = self.state()?;
let key = ActivityKey::new(PoolAddress::new(namespace, task_queue), activity_type);
Ok(state.by_activity.get(&key).map_or(0, |workers| {
workers
.values()
.filter(|worker| state.dispatch_ineligible.contains_key(&worker.id))
.filter(|worker| {
tiers
.iter()
.any(|tier| worker_matches_node(worker, tier.as_deref()))
})
.count()
}))
}
pub fn dispatch_ineligible(
&self,
) -> Result<BTreeMap<WorkerId, DispatchExclusion>, ServerError> {
Ok(self.state()?.dispatch_ineligible.clone())
}
pub fn pool_census(
&self,
namespace: &str,
task_queue: &str,
activity_type: &str,
node: Option<&str>,
) -> Result<super::queue_service::PoolCensus, ServerError> {
let state = self.state()?;
let key = ActivityKey::new(PoolAddress::new(namespace, task_queue), activity_type);
let workers_in_pool = state
.workers
.values()
.filter(|worker| {
worker.task_queue == task_queue && worker.namespaces.contains(namespace)
})
.count();
let serving = state.by_activity.get(&key);
let workers_serving_activity = serving.map_or(0, HashMap::len);
let compatible: Vec<_> = serving.map_or_else(Vec::new, |workers| {
workers
.values()
.filter(|worker| worker_matches_node(worker, node))
.collect()
});
let compatible_workers = compatible.len();
let eligible_compatible_workers = compatible
.iter()
.filter(|worker| !state.dispatch_ineligible.contains_key(&worker.id))
.count();
let compatible_workers_reachability_lost = compatible
.iter()
.filter(|worker| {
matches!(
state.dispatch_ineligible.get(&worker.id),
Some(&DispatchExclusion::ReachabilityLost)
)
})
.count();
let last_compatible_poller_age = if compatible_workers > 0 {
Some(Duration::ZERO)
} else {
state
.last_departure
.get(&key)
.and_then(|by_node| {
by_node
.iter()
.filter(|(departed_node, _)| match node {
None => true,
Some(node) => departed_node.as_deref() == Some(node),
})
.map(|(_, departed_at)| *departed_at)
.max()
})
.map(|departed_at| departed_at.elapsed())
};
Ok(super::queue_service::PoolCensus {
workers_in_pool,
workers_serving_activity,
compatible_workers,
eligible_compatible_workers,
compatible_workers_reachability_lost,
last_compatible_poller_age,
})
}
pub fn is_registered(&self, worker_id: WorkerId) -> Result<bool, ServerError> {
Ok(self.state()?.workers.contains_key(&worker_id))
}
pub fn deregister(&self, worker_id: WorkerId) -> Result<(), ServerError> {
self.deregister_with_reason(worker_id, WorkerDeathReason::Disconnect)
}
pub fn deregister_with_reason(
&self,
worker_id: WorkerId,
reason: WorkerDeathReason,
) -> Result<(), ServerError> {
let mut state = self.state()?;
let removed_namespaces = Self::remove_worker(&mut state, worker_id);
drop(state);
let Some(namespaces) = removed_namespaces else {
return Ok(());
};
if let Some(metrics) = &self.metrics {
for namespace in &namespaces {
metrics.worker_disconnected(namespace);
}
}
self.emit_worker_disconnected(worker_id, &namespaces, reason);
Ok(())
}
fn emit_worker_disconnected(
&self,
worker_id: WorkerId,
namespaces: &BTreeSet<String>,
reason: WorkerDeathReason,
) {
if let Some(publisher) = &self.cluster_publisher {
let namespaces_vec: Vec<String> = namespaces.iter().cloned().collect();
drop(publisher.emit(|meta| ClusterEvent::WorkerDisconnected {
meta,
worker_id: worker_id.value().to_string(),
namespaces: namespaces_vec,
reason,
}));
}
}
fn remove_worker(state: &mut RegistryState, worker_id: WorkerId) -> Option<BTreeSet<String>> {
let handle = state.workers.remove(&worker_id)?;
Self::remove_worker_from_service(state, &handle);
Some(handle.namespaces)
}
fn remove_worker_from_service(state: &mut RegistryState, handle: &WorkerHandle) {
let departed_at = Instant::now();
for namespace in &handle.namespaces {
let pool = PoolAddress::new(namespace.clone(), handle.task_queue.clone());
for activity_type in &handle.activity_types {
let key = ActivityKey::new(pool.clone(), activity_type.clone());
state
.last_departure
.entry(key.clone())
.or_default()
.insert(handle.node.clone(), departed_at);
if let Some(workers) = state.by_activity.get_mut(&key) {
workers.remove(&handle.id);
if workers.is_empty() {
state.by_activity.remove(&key);
state.rotation.remove(&key);
}
}
}
}
}
fn state(&self) -> Result<MutexGuard<'_, RegistryState>, ServerError> {
self.inner
.lock()
.map_err(|_| ServerError::lock_poisoned("connected worker registry"))
}
}
pub(crate) fn optional_node(node: &str) -> Option<String> {
if node.is_empty() {
None
} else {
Some(node.to_owned())
}
}
fn eligible_candidates_in_rotation(
state: &mut RegistryState,
key: ActivityKey,
node: Option<&str>,
) -> Vec<WorkerHandle> {
let mut workers: Vec<WorkerHandle> = state
.by_activity
.get(&key)
.map(|workers| {
workers
.values()
.filter(|worker| worker_matches_node(worker, node))
.filter(|worker| !state.dispatch_ineligible.contains_key(&worker.id))
.cloned()
.collect()
})
.unwrap_or_default();
if workers.is_empty() {
return workers;
}
workers.sort_by_key(WorkerHandle::id);
let cursor = state.rotation.entry(key).or_insert(0);
let start = *cursor % workers.len();
*cursor = cursor.wrapping_add(1);
let mut rotated = Vec::with_capacity(workers.len());
rotated.extend_from_slice(&workers[start..]);
rotated.extend_from_slice(&workers[..start]);
rotated
}
fn worker_matches_node(worker: &WorkerHandle, node: Option<&str>) -> bool {
match node {
None => true,
Some(node) => worker.node() == Some(node),
}
}
#[derive(Clone, Debug)]
struct WorkerRegistrationParts {
worker_id: WorkerId,
namespaces: BTreeSet<String>,
task_queue: String,
activity_types: BTreeSet<String>,
}
#[derive(Debug)]
pub struct WorkerRegistration {
registry: ConnectedWorkerRegistry,
parts: Option<WorkerRegistrationParts>,
}
impl WorkerRegistration {
#[must_use]
pub fn worker_id(&self) -> Option<WorkerId> {
self.parts.as_ref().map(|parts| parts.worker_id)
}
#[must_use]
pub fn namespaces(&self) -> Option<&BTreeSet<String>> {
self.parts.as_ref().map(|parts| &parts.namespaces)
}
#[must_use]
pub fn task_queue(&self) -> Option<&str> {
self.parts.as_ref().map(|parts| parts.task_queue.as_str())
}
#[must_use]
pub fn activity_types(&self) -> Option<&BTreeSet<String>> {
self.parts.as_ref().map(|parts| &parts.activity_types)
}
pub fn deregister(mut self) -> Result<(), ServerError> {
let Some(parts) = self.parts.take() else {
return Ok(());
};
self.registry.deregister(parts.worker_id)
}
}
impl Drop for WorkerRegistration {
fn drop(&mut self) {
let Some(parts) = self.parts.take() else {
return;
};
let removed_namespaces = self.registry.inner.lock().ok().and_then(|mut state| {
ConnectedWorkerRegistry::remove_worker(&mut state, parts.worker_id)
});
if let Some(namespaces) = removed_namespaces {
if let Some(metrics) = &self.registry.metrics {
for namespace in &namespaces {
metrics.worker_disconnected(namespace);
}
}
self.registry.emit_worker_disconnected(
parts.worker_id,
&namespaces,
WorkerDeathReason::Disconnect,
);
}
}
}
#[cfg(test)]
mod tests {
use crate::config::NamespaceMode;
use crate::namespace::{NamespaceResolver, StaticScheduleNamespaces, StaticWorkflowNamespaces};
use crate::worker::heartbeat::{DISPATCH_PROBATION_PINGS, HeartbeatTracker};
use super::*;
fn guard() -> NamespaceGuard {
NamespaceGuard::new(NamespaceResolver::authorization_only(
NamespaceMode::SharedEngine,
StaticWorkflowNamespaces::default(),
StaticScheduleNamespaces::default(),
))
}
fn caller(namespace: &str) -> CallerIdentity {
CallerIdentity::new("worker", [namespace.to_owned()])
}
fn test_failure(message: &str) -> ServerError {
ServerError::worker_dispatch("default".to_owned(), "test".to_owned(), message.to_owned())
}
fn registration(namespace: &str, activity_types: &[&str]) -> ProtoRegisterWorker {
registration_with_queue(namespace, "", activity_types)
}
fn registration_with_queue(
namespace: &str,
task_queue: &str,
activity_types: &[&str],
) -> ProtoRegisterWorker {
registration_full(&[namespace], task_queue, "", activity_types)
}
fn registration_full(
namespaces: &[&str],
task_queue: &str,
node: &str,
activity_types: &[&str],
) -> ProtoRegisterWorker {
ProtoRegisterWorker {
namespaces: namespaces.iter().map(|value| (*value).to_owned()).collect(),
activity_types: activity_types
.iter()
.map(|value| (*value).to_owned())
.collect(),
task_queue: task_queue.to_owned(),
node: node.to_owned(),
activities: Vec::new(),
identity: String::new(),
instance: None,
}
}
fn multi_caller(namespaces: &[&str]) -> CallerIdentity {
CallerIdentity::new("worker", namespaces.iter().map(|value| (*value).to_owned()))
}
#[tokio::test]
async fn set_intervention_capabilities_updates_live_worker() -> Result<(), ServerError> {
let registry = ConnectedWorkerRegistry::default();
let (sender, _receiver) = mpsc::channel(1);
let types = ["scout".to_owned()];
let guard = registry.register_delivery_with_capabilities(
["default".to_owned()],
"default",
None,
types.iter(),
WorkerDelivery::Grpc(sender),
InterventionCapabilities::none(),
)?;
let Some(worker_id) = guard.worker_id() else {
return Err(test_failure("registration carries an id"));
};
let announced = InterventionCapabilities {
supported: vec![aion_core::InterventionPrimitive::InjectMessage],
};
assert!(
registry.set_intervention_capabilities(worker_id, &announced)?,
"a live worker's capabilities must be updatable"
);
let Some(handle) = registry.worker_by_id(worker_id)? else {
return Err(test_failure("worker stays registered"));
};
assert_eq!(handle.intervention_capabilities(), &announced);
assert!(
!registry.set_intervention_capabilities(WorkerId(u64::MAX), &announced)?,
"an unknown worker reports false, never an error"
);
Ok(())
}
#[tokio::test]
async fn register_and_deregister_are_namespace_isolated() -> Result<(), ServerError> {
let registry = ConnectedWorkerRegistry::default();
let (tenant_a_tx, _tenant_a_rx) = mpsc::channel(1);
let (tenant_b_tx, _tenant_b_rx) = mpsc::channel(1);
let tenant_a = registry
.accept_registration(
&guard(),
&caller("tenant-a"),
®istration("tenant-a", &["charge", "charge"]),
tenant_a_tx,
)
.await?;
let tenant_b = registry
.accept_registration(
&guard(),
&caller("tenant-b"),
®istration("tenant-b", &["charge"]),
tenant_b_tx,
)
.await?;
let tq = DEFAULT_TASK_QUEUE;
assert_eq!(
registry.workers_for("tenant-a", tq, "charge", None)?.len(),
1
);
assert_eq!(
registry.workers_for("tenant-b", tq, "charge", None)?.len(),
1
);
assert!(
registry
.workers_for("tenant-a", tq, "missing", None)?
.is_empty()
);
let tenant_a_id = tenant_a.worker_id();
tenant_a.deregister()?;
assert!(
registry
.workers_for("tenant-a", tq, "charge", None)?
.is_empty()
);
assert_eq!(
registry.workers_for("tenant-b", tq, "charge", None)?.len(),
1
);
assert_ne!(tenant_a_id, tenant_b.worker_id());
tenant_b.deregister()?;
assert!(
registry
.workers_for("tenant-b", tq, "charge", None)?
.is_empty()
);
Ok(())
}
#[tokio::test]
async fn denied_namespace_is_not_registered() -> Result<(), ServerError> {
let registry = ConnectedWorkerRegistry::default();
let (tx, _rx) = mpsc::channel(1);
let denied = registry
.accept_registration(
&guard(),
&caller("tenant-a"),
®istration("tenant-b", &["charge"]),
tx,
)
.await;
assert!(denied.is_err());
assert!(
registry
.workers_for("tenant-b", DEFAULT_TASK_QUEUE, "charge", None)?
.is_empty()
);
Ok(())
}
#[tokio::test]
async fn task_queues_partition_disjoint_pools_within_one_namespace() -> Result<(), ServerError>
{
let registry = ConnectedWorkerRegistry::default();
let (norn_tx, _norn_rx) = mpsc::channel(1);
let (claude_a_tx, _claude_a_rx) = mpsc::channel(1);
let (claude_b_tx, _claude_b_rx) = mpsc::channel(1);
let norn = registry
.accept_registration(
&guard(),
&caller("local"),
®istration_with_queue("local", "norn", &["dev"]),
norn_tx,
)
.await?;
let claude_a = registry
.accept_registration(
&guard(),
&caller("local"),
®istration_with_queue("local", "claude", &["dev"]),
claude_a_tx,
)
.await?;
let claude_b = registry
.accept_registration(
&guard(),
&caller("local"),
®istration_with_queue("local", "claude", &["dev"]),
claude_b_tx,
)
.await?;
let norn_pool = registry.workers_for("local", "norn", "dev", None)?;
assert_eq!(norn_pool.len(), 1, "norn pool has exactly its one worker");
let norn_id = norn.worker_id().ok_or_else(missing_id)?;
assert_eq!(norn_pool[0].id(), norn_id);
let claude_pool = registry.workers_for("local", "claude", "dev", None)?;
assert_eq!(
claude_pool.len(),
2,
"claude pool sees only its two workers"
);
let claude_ids: BTreeSet<WorkerId> = claude_pool.iter().map(WorkerHandle::id).collect();
assert!(
!claude_ids.contains(&norn_id),
"the norn worker must never appear in the claude pool"
);
assert!(
!registry
.workers_for("local", "norn", "dev", None)?
.iter()
.any(|worker| claude_ids.contains(&worker.id()))
);
let first = registry.workers_for("local", "claude", "dev", None)?[0].id();
let second = registry.workers_for("local", "claude", "dev", None)?[0].id();
assert_ne!(
first, second,
"claude pool round-robins across both workers"
);
assert_eq!(
registry.workers_for("local", "norn", "dev", None)?[0].id(),
norn_id,
"the norn pool rotation is unaffected by claude traffic"
);
norn.deregister()?;
claude_a.deregister()?;
claude_b.deregister()?;
Ok(())
}
#[tokio::test]
async fn same_task_queue_in_different_namespaces_is_isolated() -> Result<(), ServerError> {
let registry = ConnectedWorkerRegistry::default();
let (local_tx, _local_rx) = mpsc::channel(1);
let (remote_tx, _remote_rx) = mpsc::channel(1);
let local = registry
.accept_registration(
&guard(),
&caller("local"),
®istration_with_queue("local", "gpu", &["render"]),
local_tx,
)
.await?;
let remote = registry
.accept_registration(
&guard(),
&caller("remote"),
®istration_with_queue("remote", "gpu", &["render"]),
remote_tx,
)
.await?;
let local_pool = registry.workers_for("local", "gpu", "render", None)?;
let remote_pool = registry.workers_for("remote", "gpu", "render", None)?;
assert_eq!(local_pool.len(), 1);
assert_eq!(remote_pool.len(), 1);
assert_ne!(
local_pool[0].id(),
remote_pool[0].id(),
"a shared task_queue string does not merge two namespaces"
);
local.deregister()?;
assert!(
registry
.workers_for("local", "gpu", "render", None)?
.is_empty(),
"deregistering the local worker leaves the remote namespace untouched"
);
assert_eq!(
registry.workers_for("remote", "gpu", "render", None)?.len(),
1
);
remote.deregister()?;
Ok(())
}
#[tokio::test]
async fn worker_serving_a_namespace_set_is_reachable_in_each() -> Result<(), ServerError> {
let registry = ConnectedWorkerRegistry::default();
let (ab_tx, _ab_rx) = mpsc::channel(1);
let (a_tx, _a_rx) = mpsc::channel(1);
let worker_ab = registry
.accept_registration(
&guard(),
&multi_caller(&["a", "b"]),
®istration_full(&["a", "b"], "default", "", &["dev"]),
ab_tx,
)
.await?;
let worker_a = registry
.accept_registration(
&guard(),
&caller("a"),
®istration_full(&["a"], "default", "", &["dev"]),
a_tx,
)
.await?;
let in_a = registry.workers_for("a", "default", "dev", None)?;
let in_b = registry.workers_for("b", "default", "dev", None)?;
let both_id = worker_ab.worker_id().ok_or_else(missing_id)?;
let only_a_id = worker_a.worker_id().ok_or_else(missing_id)?;
let a_ids: BTreeSet<WorkerId> = in_a.iter().map(WorkerHandle::id).collect();
assert_eq!(a_ids, BTreeSet::from([both_id, only_a_id]));
assert_eq!(in_b.len(), 1, "only the {{a, b}} worker is reachable in b");
assert_eq!(in_b[0].id(), both_id);
assert!(
!in_b.iter().any(|worker| worker.id() == only_a_id),
"the {{a}}-only worker must not be reachable in b"
);
worker_ab.deregister()?;
assert!(
registry
.workers_for("b", "default", "dev", None)?
.is_empty()
);
assert_eq!(registry.workers_for("a", "default", "dev", None)?.len(), 1);
worker_a.deregister()?;
Ok(())
}
#[tokio::test]
async fn node_pin_filters_within_pool() -> Result<(), ServerError> {
let registry = ConnectedWorkerRegistry::default();
let (n1_tx, _n1_rx) = mpsc::channel(1);
let (n2_tx, _n2_rx) = mpsc::channel(1);
let on_n1 = registry
.accept_registration(
&guard(),
&caller("ns"),
®istration_full(&["ns"], "tq", "n1", &["dev"]),
n1_tx,
)
.await?;
let on_n2 = registry
.accept_registration(
&guard(),
&caller("ns"),
®istration_full(&["ns"], "tq", "n2", &["dev"]),
n2_tx,
)
.await?;
let n1_id = on_n1.worker_id().ok_or_else(missing_id)?;
let n2_id = on_n2.worker_id().ok_or_else(missing_id)?;
let unpinned = registry.workers_for("ns", "tq", "dev", None)?;
assert_eq!(unpinned.len(), 2, "unpinned reaches the whole pool");
let first = registry.workers_for("ns", "tq", "dev", None)?[0].id();
let second = registry.workers_for("ns", "tq", "dev", None)?[0].id();
assert_ne!(first, second, "unpinned round-robins across both nodes");
let pinned_n1 = registry.workers_for("ns", "tq", "dev", Some("n1"))?;
assert_eq!(pinned_n1.len(), 1);
assert_eq!(pinned_n1[0].id(), n1_id);
let pinned_n2 = registry.workers_for("ns", "tq", "dev", Some("n2"))?;
assert_eq!(pinned_n2.len(), 1);
assert_eq!(pinned_n2[0].id(), n2_id);
assert_eq!(
registry
.select_worker("ns", "tq", "dev", Some("n1"))?
.map(|worker| worker.id()),
Some(n1_id)
);
assert!(
registry
.workers_for("ns", "tq", "dev", Some("absent"))?
.is_empty(),
"a pin to a node with no worker yields no candidate"
);
assert!(
registry
.select_worker("ns", "tq", "dev", Some("absent"))?
.is_none()
);
on_n1.deregister()?;
on_n2.deregister()?;
Ok(())
}
#[tokio::test]
async fn shared_node_id_round_robins_across_workers() -> Result<(), ServerError> {
let registry = ConnectedWorkerRegistry::default();
let (a_tx, _a_rx) = mpsc::channel(1);
let (b_tx, _b_rx) = mpsc::channel(1);
let worker_a = registry
.accept_registration(
&guard(),
&caller("ns"),
®istration_full(&["ns"], "tq", "shared", &["dev"]),
a_tx,
)
.await?;
let worker_b = registry
.accept_registration(
&guard(),
&caller("ns"),
®istration_full(&["ns"], "tq", "shared", &["dev"]),
b_tx,
)
.await?;
let a_id = worker_a.worker_id().ok_or_else(missing_id)?;
let b_id = worker_b.worker_id().ok_or_else(missing_id)?;
let pinned = registry.workers_for("ns", "tq", "dev", Some("shared"))?;
assert_eq!(
pinned.len(),
2,
"both workers on the shared node are candidates"
);
let pinned_ids: BTreeSet<WorkerId> = pinned.iter().map(WorkerHandle::id).collect();
assert_eq!(pinned_ids, BTreeSet::from([a_id, b_id]));
let first = registry.workers_for("ns", "tq", "dev", Some("shared"))?[0].id();
let second = registry.workers_for("ns", "tq", "dev", Some("shared"))?[0].id();
assert_ne!(
first, second,
"a pin to a shared node round-robins across both workers on it"
);
worker_a.deregister()?;
worker_b.deregister()?;
Ok(())
}
#[tokio::test]
async fn unpinned_selection_rotates_across_nodes() -> Result<(), ServerError> {
let registry = ConnectedWorkerRegistry::default();
let (n1_tx, _n1_rx) = mpsc::channel(1);
let (n2_tx, _n2_rx) = mpsc::channel(1);
let on_n1 = registry
.accept_registration(
&guard(),
&caller("ns"),
®istration_full(&["ns"], "tq", "n1", &["dev"]),
n1_tx,
)
.await?;
let on_n2 = registry
.accept_registration(
&guard(),
&caller("ns"),
®istration_full(&["ns"], "tq", "n2", &["dev"]),
n2_tx,
)
.await?;
let mut visited = Vec::new();
for _ in 0..4 {
let selected = registry
.select_worker("ns", "tq", "dev", None)?
.ok_or_else(|| test_failure("an unpinned selection must find a worker"))?;
visited.push(selected.node().map(str::to_owned));
}
assert_eq!(
visited,
vec![
Some(String::from("n1")),
Some(String::from("n2")),
Some(String::from("n1")),
Some(String::from("n2")),
],
"unpinned selection must rotate across the nodes rather than pin itself to the \
lowest worker id"
);
on_n1.deregister()?;
on_n2.deregister()?;
Ok(())
}
#[test]
fn an_ineligible_worker_is_never_selected_by_either_selector() -> Result<(), ServerError> {
let registry = ConnectedWorkerRegistry::default();
let types = [String::from("dev")];
let mut receivers = Vec::new();
let mut registrations = Vec::new();
for _ in 0..3 {
let (tx, rx) = mpsc::channel(1);
receivers.push(rx);
registrations.push(registry.register_namespaces(
[String::from("ns")],
"tq",
None,
types.iter(),
tx,
)?);
}
let mut ids = Vec::new();
for registration in ®istrations {
ids.push(registration.worker_id().ok_or_else(missing_id)?);
}
ids.sort_unstable();
let excluded = *ids.first().ok_or_else(missing_id)?;
registry.set_dispatch_ineligible(
[(excluded, DispatchExclusion::ReachabilityLost)]
.into_iter()
.collect(),
)?;
for _ in 0..4 {
let selected = registry
.select_worker("ns", "tq", "dev", None)?
.ok_or_else(|| test_failure("two eligible workers remain in the pool"))?;
assert_ne!(
selected.id(),
excluded,
"select_worker must never return the worker the liveness verdict excluded"
);
let candidates = registry.workers_for("ns", "tq", "dev", None)?;
assert_eq!(
candidates.len(),
2,
"workers_for must offer the two ELIGIBLE workers, not all three registered ones"
);
assert!(
candidates.iter().all(|worker| worker.id() != excluded),
"workers_for must not list the excluded worker in ANY position: the push \
dispatcher walks the whole list"
);
}
for registration in registrations {
registration.deregister()?;
}
Ok(())
}
#[test]
fn an_all_ineligible_pool_selects_nobody_while_the_census_still_reads_served()
-> Result<(), ServerError> {
let registry = ConnectedWorkerRegistry::default();
let types = [String::from("dev")];
let (tx, _rx) = mpsc::channel(1);
let worker =
registry.register_namespaces([String::from("ns")], "tq", None, types.iter(), tx)?;
let worker_id = worker.worker_id().ok_or_else(missing_id)?;
registry.set_dispatch_ineligible(
[(
worker_id,
DispatchExclusion::OpeningProbation { answers: 1 },
)]
.into_iter()
.collect(),
)?;
assert!(
registry.workers_for("ns", "tq", "dev", None)?.is_empty(),
"the push dispatcher's candidate list must exclude an unreachable worker"
);
assert!(
registry.select_worker("ns", "tq", "dev", None)?.is_none(),
"and so must the single selector"
);
let census = registry.pool_census("ns", "tq", "dev", None)?;
assert_eq!(
census.compatible_workers, 1,
"the census counts the REGISTERED worker: #197 R3 needs the count to separate an \
empty pool from an excluded one"
);
assert!(
census.is_served(),
"so `classify` reads this address as served and returns None while selection has \
nobody — the disagreement dispatch_to_node must park on rather than spin through"
);
assert!(
census.will_be_served(),
"and the park is the RIGHT outcome here: a probation clears itself within seconds, \
so nobody should be warned and nothing should be published. The pool that has LOST \
reachability is the one that must not reach this state — see \
`the_exclusion_cause_is_what_decides`"
);
assert_eq!(
census.compatible_workers_reachability_lost, 0,
"precondition for the assertion above: this fixture's exclusion is a probation"
);
worker.deregister()?;
Ok(())
}
#[test]
fn an_arrival_subscription_retains_a_wake_taken_before_it() -> Result<(), ServerError> {
use std::future::Future;
use std::task::Waker;
let registry = ConnectedWorkerRegistry::default();
let mut context = Context::from_waker(Waker::noop());
let (first_tx, _first_rx) = mpsc::channel(1);
let first = registry.register("ns", [String::from("dev")].iter(), first_tx)?;
let mut too_late = std::pin::pin!(registry.worker_arrived.notified());
assert!(
matches!(too_late.as_mut().poll(&mut context), Poll::Pending),
"a wait constructed AFTER the registration cannot have retained it: \
`notify_waiters` stored no permit and there was no waiter in the list to \
broadcast to. This is the finding, reproduced."
);
assert!(
matches!(too_late.as_mut().poll(&mut context), Poll::Pending),
"and it stays parked — without this control the arm above would pass on a wait \
that merely reports Pending on its first poll for registration reasons"
);
let arrival = registry.worker_arrival();
let (second_tx, _second_rx) = mpsc::channel(1);
let second = registry.register("ns", [String::from("dev")].iter(), second_tx)?;
let mut arrival = std::pin::pin!(arrival);
assert!(
matches!(arrival.as_mut().poll(&mut context), Poll::Ready(())),
"a subscription taken BEFORE the registry read must retain the registration that \
landed during it: the dispatch that missed by a microsecond must not park"
);
assert!(
matches!(too_late.as_mut().poll(&mut context), Poll::Ready(())),
"the control's own wait was in the list by now, so the SECOND registration wakes \
it — proving the control arm above was parked on the lost wake and not on a \
registry that never notified at all"
);
first.deregister()?;
second.deregister()?;
Ok(())
}
#[test]
fn publishing_a_reachability_verdict_wakes_the_selection_wait() -> Result<(), ServerError> {
use std::future::Future;
use std::task::Waker;
let registry = ConnectedWorkerRegistry::default();
let mut waiting = std::pin::pin!(registry.worker_arrival());
let mut context = Context::from_waker(Waker::noop());
assert!(
matches!(waiting.as_mut().poll(&mut context), Poll::Pending),
"the wait parks until something changes what selection can see"
);
assert!(
matches!(waiting.as_mut().poll(&mut context), Poll::Pending),
"and stays parked while nothing has been published: without this control the \
assertion below would pass on a wait that simply completes on a second poll"
);
registry.set_dispatch_ineligible(BTreeMap::new())?;
assert!(
matches!(waiting.as_mut().poll(&mut context), Poll::Ready(())),
"a published verdict must wake a dispatch parked on eligibility; no registration is \
coming for a worker that never left the registry"
);
Ok(())
}
#[tokio::test]
async fn rotation_cursor_is_pruned_when_last_worker_leaves() -> Result<(), ServerError> {
let registry = ConnectedWorkerRegistry::default();
let (tx, _rx) = mpsc::channel(1);
let worker = registry
.accept_registration(
&guard(),
&caller("ns"),
®istration_full(&["ns"], "tq", "", &["dev"]),
tx,
)
.await?;
let _ = registry.workers_for("ns", "tq", "dev", None)?;
let key = ActivityKey::new(PoolAddress::new("ns", "tq"), "dev");
assert!(
registry.state()?.rotation.contains_key(&key),
"a lookup must have created the rotation cursor"
);
worker.deregister()?;
let state = registry.state()?;
assert!(
!state.rotation.contains_key(&key),
"the rotation cursor must be pruned once the last worker leaves"
);
assert!(
!state.by_activity.contains_key(&key),
"the activity bucket must also be gone"
);
Ok(())
}
fn missing_id() -> ServerError {
ServerError::lock_poisoned("registration unexpectedly missing a worker id")
}
fn namespace_store() -> Arc<dyn NamespaceStore> {
Arc::new(aion_store::InMemoryStore::default())
}
fn minting_registry(
store: &Arc<dyn NamespaceStore>,
policy: AutoCreate,
) -> ConnectedWorkerRegistry {
ConnectedWorkerRegistry::default().with_namespace_minting(Arc::clone(store), policy)
}
#[tokio::test]
async fn open_register_mints_durable_record_and_is_idempotent() -> Result<(), ServerError> {
let store = namespace_store();
let registry = minting_registry(&store, AutoCreate::Open);
let (tx_one, _rx_one) = mpsc::channel(1);
let first = registry
.accept_registration(
&guard(),
&caller("orders"),
®istration("orders", &["charge"]),
tx_one,
)
.await?;
let record = store
.get_namespace("orders")
.await?
.ok_or_else(|| ServerError::namespace_denied("expected a minted record"))?;
assert_eq!(record.name, "orders");
assert_eq!(record.origin, NamespaceOrigin::WorkerMint);
let (tx_two, _rx_two) = mpsc::channel(1);
registry
.accept_registration(
&guard(),
&caller("orders"),
®istration("orders", &["refund"]),
tx_two,
)
.await?;
let all = store.list_namespaces().await?;
assert_eq!(
all.iter().filter(|r| r.name == "orders").count(),
1,
"re-register must not create a duplicate namespace row"
);
drop(first);
Ok(())
}
#[tokio::test]
async fn open_register_mints_each_namespace_in_a_multi_namespace_worker()
-> Result<(), ServerError> {
let store = namespace_store();
let registry = minting_registry(&store, AutoCreate::Open);
let (tx, _rx) = mpsc::channel(1);
registry
.accept_registration(
&guard(),
&multi_caller(&["alpha", "beta"]),
®istration_full(&["alpha", "beta"], "", "", &["charge"]),
tx,
)
.await?;
assert!(store.get_namespace("alpha").await?.is_some());
assert!(store.get_namespace("beta").await?.is_some());
Ok(())
}
async fn pinned_registry(
store: &Arc<dyn NamespaceStore>,
namespace: &str,
nodes: &[&str],
) -> Result<ConnectedWorkerRegistry, ServerError> {
store
.register_namespace(namespace, NamespaceOrigin::Explicit)
.await?;
store
.set_namespace_placement(
namespace,
NamespacePlacement::Pinned {
nodes: nodes.iter().map(|n| (*n).to_owned()).collect(),
},
)
.await?;
Ok(minting_registry(store, AutoCreate::Open))
}
#[tokio::test]
async fn pinned_admits_a_worker_on_a_required_node() -> Result<(), ServerError> {
let store = namespace_store();
let registry = pinned_registry(&store, "iso", &["n1"]).await?;
let (tx, _rx) = mpsc::channel(1);
let _registration = registry
.accept_registration(
&guard(),
&caller("iso"),
®istration_full(&["iso"], "", "n1", &["charge"]),
tx,
)
.await?;
assert_eq!(
registry
.workers_for("iso", DEFAULT_TASK_QUEUE, "charge", Some("n1"))?
.len(),
1,
"an n1 worker must be admitted into the Pinned{{n1}} namespace's pool"
);
Ok(())
}
#[tokio::test]
async fn pinned_rejects_a_wrong_node_worker() -> Result<(), ServerError> {
let store = namespace_store();
let registry = pinned_registry(&store, "iso", &["n1"]).await?;
let (tx, _rx) = mpsc::channel(1);
let denied = registry
.accept_registration(
&guard(),
&caller("iso"),
®istration_full(&["iso"], "", "n2", &["charge"]),
tx,
)
.await;
assert!(
matches!(denied, Err(ServerError::Namespace { .. })),
"a wrong-node (n2) worker must be rejected from a Pinned{{n1}} namespace"
);
assert!(
registry
.workers_for("iso", DEFAULT_TASK_QUEUE, "charge", None)?
.is_empty(),
"a rejected registration must not insert a worker on any node"
);
Ok(())
}
#[tokio::test]
async fn pinned_rejects_a_node_less_worker() -> Result<(), ServerError> {
let store = namespace_store();
let registry = pinned_registry(&store, "iso", &["n1"]).await?;
let (tx, _rx) = mpsc::channel(1);
let denied = registry
.accept_registration(
&guard(),
&caller("iso"),
®istration_full(&["iso"], "", "", &["charge"]),
tx,
)
.await;
assert!(
matches!(denied, Err(ServerError::Namespace { .. })),
"a node-less worker must be rejected from a Pinned{{n1}} namespace"
);
assert!(
registry
.workers_for("iso", DEFAULT_TASK_QUEUE, "charge", None)?
.is_empty(),
"a rejected node-less registration must not insert a worker"
);
Ok(())
}
#[tokio::test]
async fn pinned_violation_rejects_the_whole_multi_namespace_registration()
-> Result<(), ServerError> {
let store = namespace_store();
let registry = pinned_registry(&store, "iso", &["n1"]).await?;
let (tx, _rx) = mpsc::channel(1);
let denied = registry
.accept_registration(
&guard(),
&multi_caller(&["free", "iso"]),
®istration_full(&["free", "iso"], "", "n2", &["charge"]),
tx,
)
.await;
assert!(
matches!(denied, Err(ServerError::Namespace { .. })),
"a wrong-node worker serving a Pinned namespace fails the WHOLE registration"
);
assert!(
registry
.workers_for("free", DEFAULT_TASK_QUEUE, "charge", None)?
.is_empty(),
"the compliant namespace must NOT be partially admitted"
);
Ok(())
}
#[tokio::test]
async fn unplaced_and_prefer_admission_is_unaffected_by_the_pinned_gate()
-> Result<(), ServerError> {
let store = namespace_store();
store
.register_namespace("pref", NamespaceOrigin::Explicit)
.await?;
store
.set_namespace_placement(
"pref",
NamespacePlacement::Prefer {
nodes: ["n1".to_owned()].into_iter().collect(),
},
)
.await?;
let registry = minting_registry(&store, AutoCreate::Open);
let (tx_a, _rx_a) = mpsc::channel(1);
let _reg_a = registry
.accept_registration(
&guard(),
&caller("unpl"),
®istration_full(&["unpl"], "", "", &["charge"]),
tx_a,
)
.await?;
let (tx_b, _rx_b) = mpsc::channel(1);
let _reg_b = registry
.accept_registration(
&guard(),
&caller("pref"),
®istration_full(&["pref"], "", "", &["charge"]),
tx_b,
)
.await?;
assert_eq!(
registry
.workers_for("unpl", DEFAULT_TASK_QUEUE, "charge", None)?
.len(),
1,
"an Unplaced namespace admits a node-less worker unchanged"
);
assert_eq!(
registry
.workers_for("pref", DEFAULT_TASK_QUEUE, "charge", None)?
.len(),
1,
"a Prefer namespace admits a node-less worker unchanged (only Pinned gates)"
);
Ok(())
}
#[tokio::test]
async fn no_minter_registry_skips_the_placement_gate() -> Result<(), ServerError> {
let registry = ConnectedWorkerRegistry::default();
let (tx, _rx) = mpsc::channel(1);
let _registration = registry
.accept_registration(
&guard(),
&caller("plain"),
®istration_full(&["plain"], "", "", &["charge"]),
tx,
)
.await?;
assert_eq!(
registry
.workers_for("plain", DEFAULT_TASK_QUEUE, "charge", None)?
.len(),
1,
"with no minter the placement gate is a no-op — registration is unchanged"
);
Ok(())
}
#[tokio::test]
async fn concurrent_registrations_for_a_new_namespace_create_exactly_one_record()
-> Result<(), ServerError> {
let store = namespace_store();
let registry = minting_registry(&store, AutoCreate::Open);
let mut handles = Vec::new();
for _ in 0..8 {
let registry = registry.clone();
handles.push(tokio::spawn(async move {
let (tx, rx) = mpsc::channel(1);
let outcome = registry
.accept_registration(
&guard(),
&caller("rush"),
®istration("rush", &["charge"]),
tx,
)
.await;
drop(rx);
outcome.map(|registration| registration.worker_id())
}));
}
for handle in handles {
handle
.await
.map_err(|_| ServerError::lock_poisoned("registration task panicked"))??;
}
let all = store.list_namespaces().await?;
assert_eq!(
all.iter().filter(|r| r.name == "rush").count(),
1,
"racing registrations must converge on exactly one durable record"
);
Ok(())
}
#[tokio::test]
async fn closed_rejects_unknown_namespace_and_does_not_create_it() -> Result<(), ServerError> {
let store = namespace_store();
let registry = minting_registry(&store, AutoCreate::Closed);
let (tx, _rx) = mpsc::channel(1);
let denied = registry
.accept_registration(
&guard(),
&caller("ghost"),
®istration("ghost", &["charge"]),
tx,
)
.await;
assert!(
matches!(denied, Err(ServerError::Namespace { .. })),
"closed policy must reject an unknown namespace"
);
assert!(
store.get_namespace("ghost").await?.is_none(),
"closed policy must NOT create the namespace it rejected"
);
let tq = DEFAULT_TASK_QUEUE;
assert!(
registry
.workers_for("ghost", tq, "charge", None)?
.is_empty(),
"a rejected registration must not insert a worker"
);
Ok(())
}
#[tokio::test]
async fn closed_admits_a_known_namespace() -> Result<(), ServerError> {
let store = namespace_store();
store
.register_namespace("known", NamespaceOrigin::Explicit)
.await?;
let registry = minting_registry(&store, AutoCreate::Closed);
let (tx, _rx) = mpsc::channel(1);
let _registration = registry
.accept_registration(
&guard(),
&caller("known"),
®istration("known", &["charge"]),
tx,
)
.await?;
let tq = DEFAULT_TASK_QUEUE;
assert_eq!(
registry.workers_for("known", tq, "charge", None)?.len(),
1,
"a known namespace must register under closed policy"
);
Ok(())
}
#[tokio::test]
async fn no_minter_leaves_registration_untouched() -> Result<(), ServerError> {
let registry = ConnectedWorkerRegistry::default();
let (tx, _rx) = mpsc::channel(1);
let _registration = registry
.accept_registration(
&guard(),
&caller("orders"),
®istration("orders", &["charge"]),
tx,
)
.await?;
let tq = DEFAULT_TASK_QUEUE;
assert_eq!(registry.workers_for("orders", tq, "charge", None)?.len(), 1);
Ok(())
}
#[test]
fn a_never_served_address_censuses_empty_with_no_poller_age() -> Result<(), ServerError> {
let registry = ConnectedWorkerRegistry::default();
let census = registry.pool_census("default", "general", "greet", None)?;
assert_eq!(census.workers_in_pool, 0);
assert_eq!(census.workers_serving_activity, 0);
assert_eq!(census.compatible_workers, 0);
assert_eq!(census.last_compatible_poller_age, None);
assert!(!census.is_served());
Ok(())
}
#[test]
fn the_census_separates_pool_activity_and_node_coverage() -> Result<(), ServerError> {
let registry = ConnectedWorkerRegistry::default();
let (tx, _rx) = mpsc::channel(1);
let _worker = registry.register_namespaces(
[String::from("default")],
"general",
Some(String::from("n1")),
[String::from("greet")].iter(),
tx,
)?;
let unpinned = registry.pool_census("default", "general", "greet", None)?;
assert_eq!(unpinned.workers_in_pool, 1);
assert_eq!(unpinned.workers_serving_activity, 1);
assert_eq!(unpinned.compatible_workers, 1);
assert_eq!(unpinned.last_compatible_poller_age, Some(Duration::ZERO));
let other_activity = registry.pool_census("default", "general", "settle", None)?;
assert_eq!(other_activity.workers_in_pool, 1);
assert_eq!(other_activity.workers_serving_activity, 0);
assert_eq!(other_activity.compatible_workers, 0);
let wrong_node = registry.pool_census("default", "general", "greet", Some("n2"))?;
assert_eq!(wrong_node.workers_serving_activity, 1);
assert_eq!(wrong_node.compatible_workers, 0);
assert_eq!(wrong_node.last_compatible_poller_age, None);
Ok(())
}
#[test]
fn a_pump_alive_worker_the_server_cannot_reach_is_not_selected() -> Result<(), ServerError> {
const WINDOW: Duration = Duration::from_secs(30);
let registry = ConnectedWorkerRegistry::default();
let tracker = HeartbeatTracker::new(WINDOW);
let (tx, _rx) = mpsc::channel(1);
let worker = registry.register_namespaces(
[String::from("default")],
"general",
None,
[String::from("greet")].iter(),
tx,
)?;
let Some(worker_id) = worker.worker_id() else {
return Err(test_failure("registration carries an id"));
};
let start = Instant::now();
tracker.register_connection(worker_id, start)?;
let publish = |now: Instant| -> Result<(), ServerError> {
registry.set_dispatch_ineligible(
tracker
.unreachable_workers(now)?
.into_iter()
.map(|excluded| (excluded.worker_id, excluded.exclusion))
.collect(),
)
};
for _ in 0..DISPATCH_PROBATION_PINGS {
assert!(
tracker.record_dispatch_reachability(worker_id, start)?,
"the worker is tracked while it serves its probation"
);
}
publish(start)?;
assert!(
registry
.select_worker("default", "general", "greet", None)?
.is_some(),
"a worker that has answered a full run of pings must be selectable: this is the \
control, and without it a registry selecting NOBODY would satisfy every assertion \
below"
);
let much_later = start + WINDOW * 4;
assert!(
tracker.record_connection_activity(worker_id, much_later)?,
"the worker is still tracked; its process is plainly alive"
);
assert!(
tracker.record_dispatch_unreachable(worker_id)?,
"the probe fired and went unanswered"
);
publish(much_later)?;
assert!(
registry.is_dispatch_ineligible(worker_id)?,
"a worker the server cannot reach must be marked ineligible however alive its \
process looks"
);
assert!(
registry
.select_worker("default", "general", "greet", None)?
.is_none(),
"an ineligible worker must not be SELECTED: selecting one produces a dispatch that \
can only fail, and on the liminal transport it fails by consuming connection \
capacity — making the unreachability worse"
);
assert!(
tracker.record_dispatch_reachability(worker_id, much_later)?,
"the worker is still tracked"
);
publish(much_later)?;
assert!(
registry.is_dispatch_ineligible(worker_id)?,
"one answer part-way through a fresh probation must NOT restore eligibility: a link \
answering one probe in three would otherwise flap in and out of selection"
);
for _ in 1..DISPATCH_PROBATION_PINGS {
assert!(
tracker.record_dispatch_reachability(worker_id, much_later)?,
"the worker is still tracked"
);
}
publish(much_later)?;
assert!(
!registry.is_dispatch_ineligible(worker_id)?,
"an answered ping must clear the exclusion"
);
assert!(
registry
.select_worker("default", "general", "greet", None)?
.is_some(),
"and the worker must be selectable again"
);
Ok(())
}
#[test]
fn a_departed_worker_leaves_a_last_compatible_poller_age() -> Result<(), ServerError> {
let registry = ConnectedWorkerRegistry::default();
let (tx, _rx) = mpsc::channel(1);
let worker = registry.register_namespaces(
[String::from("default")],
"general",
None,
[String::from("greet")].iter(),
tx,
)?;
worker.deregister()?;
let census = registry.pool_census("default", "general", "greet", None)?;
assert_eq!(census.workers_in_pool, 0);
assert_eq!(census.compatible_workers, 0);
let age = census
.last_compatible_poller_age
.ok_or_else(|| test_failure("a departed worker must leave an age behind"))?;
assert!(
age < Duration::from_secs(60),
"the recorded departure is implausibly old: {age:?}"
);
let (tx, _rx) = mpsc::channel(1);
let _back = registry.register_namespaces(
[String::from("default")],
"general",
None,
[String::from("greet")].iter(),
tx,
)?;
let served = registry.pool_census("default", "general", "greet", None)?;
assert_eq!(served.last_compatible_poller_age, Some(Duration::ZERO));
assert!(served.is_served());
Ok(())
}
}