use std::collections::{BTreeMap, BTreeSet, HashMap};
use std::sync::{Arc, Mutex, MutexGuard};
use aion_core::{ClusterEvent, InterventionCapabilities, WorkerDeathReason, WorkerTransport};
use aion_proto::{ProtoActivityTask, ProtoRegisterWorker};
use aion_store::{NamespaceOrigin, NamespacePlacement, NamespaceStore};
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;
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),
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub enum WorkerMessage {
ActivityTask(ProtoActivityTask),
DrainRequest,
}
#[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 value(self) -> u64 {
self.0
}
}
#[derive(Clone, Debug)]
pub struct WorkerHandle {
id: WorkerId,
namespaces: BTreeSet<String>,
task_queue: String,
node: Option<String>,
activity_types: BTreeSet<String>,
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 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>,
}
impl Default for RegistryState {
fn default() -> Self {
Self {
next_worker_id: 1,
workers: BTreeMap::new(),
by_activity: HashMap::new(),
rotation: HashMap::new(),
}
}
}
#[derive(Clone, Debug)]
pub struct ConnectedWorkerRegistry {
inner: Arc<Mutex<RegistryState>>,
metrics: Option<Metrics>,
cluster_publisher: Option<ClusterEventPublisher>,
minter: Option<NamespaceMinter>,
worker_arrived: Arc<Notify>,
}
impl Default for ConnectedWorkerRegistry {
fn default() -> Self {
Self {
inner: Arc::new(Mutex::new(RegistryState::default())),
metrics: None,
cluster_publisher: None,
minter: None,
worker_arrived: Arc::new(Notify::new()),
}
}
}
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,
worker_arrived: Arc::new(Notify::new()),
}
}
#[must_use]
pub fn with_cluster_publisher(mut self, publisher: ClusterEventPublisher) -> Self {
self.cluster_publisher = Some(publisher);
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
}
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?;
self.register_namespaces(
namespaces,
registration.task_queue.clone(),
node,
registration.activity_types.iter(),
sender,
)
}
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,
activity_types,
WorkerDelivery::Grpc(sender),
)
}
pub fn register_delivery<'a>(
&self,
namespaces: impl IntoIterator<Item = String>,
task_queue: impl Into<String>,
node: Option<String>,
activity_types: impl IntoIterator<Item = &'a String>,
delivery: WorkerDelivery,
) -> Result<WorkerRegistration, ServerError> {
self.register_delivery_with_capabilities(
namespaces,
task_queue,
node,
activity_types,
delivery,
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> {
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 handle = WorkerHandle {
id: worker_id,
namespaces: namespaces.clone(),
task_queue: task_queue.clone(),
node,
activity_types: activity_types.clone(),
delivery,
intervention_capabilities,
};
for namespace in &namespaces {
let pool = PoolAddress::new(namespace.clone(), task_queue.clone());
for activity_type in &activity_types {
state
.by_activity
.entry(ActivityKey::new(pool.clone(), activity_type.clone()))
.or_default()
.insert(worker_id, handle.clone());
}
}
let transport = transport_of(&handle.delivery);
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,
}));
}
self.worker_arrived.notify_waiters();
Ok(WorkerRegistration {
registry: self.clone(),
parts: Some(WorkerRegistrationParts {
worker_id,
namespaces,
task_queue,
activity_types,
}),
})
}
pub async fn wait_for_worker(&self) {
self.worker_arrived.notified().await;
}
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);
let mut workers: Vec<WorkerHandle> = state
.by_activity
.get(&key)
.map(|workers| {
workers
.values()
.filter(|worker| worker_matches_node(worker, node))
.cloned()
.collect()
})
.unwrap_or_default();
if workers.is_empty() {
return Ok(workers);
}
workers.sort_by_key(WorkerHandle::id);
let idx = state.rotation.entry(key).or_insert(0);
let start = *idx % workers.len();
*idx = idx.wrapping_add(1);
let mut rotated = Vec::with_capacity(workers.len());
rotated.extend_from_slice(&workers[start..]);
rotated.extend_from_slice(&workers[..start]);
Ok(rotated)
}
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 {
let Some(sender) = worker.sender() else {
continue;
};
if sender.try_send(WorkerMessage::DrainRequest).is_ok() {
delivered = delivered.saturating_add(1);
} else {
self.deregister(worker.id())?;
}
}
Ok(delivered)
}
pub fn select_worker(
&self,
namespace: &str,
task_queue: &str,
activity_type: &str,
node: Option<&str>,
) -> Result<Option<WorkerHandle>, ServerError> {
let state = self.state()?;
let key = ActivityKey::new(PoolAddress::new(namespace, task_queue), activity_type);
Ok(state.by_activity.get(&key).and_then(|workers| {
workers
.values()
.filter(|worker| worker_matches_node(worker, node))
.min_by_key(|worker| worker.id)
.cloned()
}))
}
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)?;
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());
if let Some(workers) = state.by_activity.get_mut(&key) {
workers.remove(&worker_id);
if workers.is_empty() {
state.by_activity.remove(&key);
state.rotation.remove(&key);
}
}
}
}
Some(handle.namespaces)
}
fn state(&self) -> Result<MutexGuard<'_, RegistryState>, ServerError> {
self.inner
.lock()
.map_err(|_| ServerError::lock_poisoned("connected worker registry"))
}
}
const fn transport_of(delivery: &WorkerDelivery) -> WorkerTransport {
match delivery {
WorkerDelivery::Grpc(_) => WorkerTransport::Grpc,
#[cfg(feature = "liminal-transport")]
WorkerDelivery::Liminal(_) => WorkerTransport::Liminal,
}
}
fn optional_node(node: &str) -> Option<String> {
if node.is_empty() {
None
} else {
Some(node.to_owned())
}
}
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 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(),
}
}
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 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(())
}
}