use std::collections::{BTreeMap, BTreeSet, HashMap};
use std::sync::{Arc, Mutex, MutexGuard};
use aion_proto::{ProtoActivityTask, ProtoRegisterWorker};
use tokio::sync::{Notify, mpsc};
use crate::error::ServerError;
use crate::namespace::{CallerIdentity, NamespaceGuard, 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,
}
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 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>,
worker_arrived: Arc<Notify>,
}
impl Default for ConnectedWorkerRegistry {
fn default() -> Self {
Self {
inner: Arc::new(Mutex::new(RegistryState::default())),
metrics: 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),
worker_arrived: Arc::new(Notify::new()),
}
}
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)?;
let node = optional_node(®istration.node);
self.register_namespaces(
namespaces,
registration.task_queue.clone(),
node,
registration.activity_types.iter(),
sender,
)
}
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> {
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 handle = WorkerHandle {
id: worker_id,
namespaces: namespaces.clone(),
task_queue: task_queue.clone(),
node,
activity_types: activity_types.clone(),
delivery,
};
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());
}
}
state.workers.insert(worker_id, handle);
drop(state);
if let Some(metrics) = &self.metrics {
for namespace in &namespaces {
metrics.worker_connected(namespace);
}
}
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 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> {
let mut state = self.state()?;
let removed_namespaces = Self::remove_worker(&mut state, worker_id);
drop(state);
if let (Some(namespaces), Some(metrics)) = (removed_namespaces, &self.metrics) {
for namespace in &namespaces {
metrics.worker_disconnected(namespace);
}
}
Ok(())
}
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"))
}
}
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;
};
if let Ok(mut state) = self.registry.inner.lock() {
let removed_namespaces =
ConnectedWorkerRegistry::remove_worker(&mut state, parts.worker_id);
if let (Some(namespaces), Some(metrics)) = (removed_namespaces, &self.registry.metrics)
{
for namespace in &namespaces {
metrics.worker_disconnected(namespace);
}
}
}
}
}
#[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 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 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")
}
}