use std::collections::HashMap;
use std::sync::{Arc, Mutex, OnceLock};
use std::time::Duration;
use aion_core::{ActivityId, ContentType, Payload, RunId, WorkflowId};
use aion_store::OutboxRow;
use async_trait::async_trait;
use liminal::protocol::WorkerRegistration as WireWorkerRegistration;
use liminal_sdk::{SchemaMetadata, SchemaValidate};
use liminal_server::ServerError as LiminalServerError;
use liminal_server::server::connection::{
ConnectionNotifier, ConnectionSupervisor, PushReplyAwaiter,
};
use serde::{Deserialize, Serialize};
use super::bridge::OutboxDeliveryCallback;
use super::outbox_dispatcher::OutboxRowDispatch;
use super::registry::{ConnectedWorkerRegistry, WorkerDelivery, WorkerHandle, WorkerRegistration};
use crate::error::ServerError;
const PUSH_REPLY_TIMEOUT: Duration = Duration::from_secs(30);
const BRIDGE_REPLY_POLL: Duration = Duration::from_secs(1);
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
pub struct DispatchRequest {
pub activity_type: String,
pub workflow_id: WorkflowId,
pub ordinal: u64,
pub run_id: Option<RunId>,
pub input: Vec<u8>,
#[serde(default = "first_attempt")]
pub attempt: u32,
#[serde(default)]
pub labels: std::collections::BTreeMap<String, String>,
#[serde(default)]
pub heartbeat_window_ms: u64,
}
const fn first_attempt() -> u32 {
1
}
impl SchemaValidate for DispatchRequest {
fn schema_metadata() -> SchemaMetadata {
SchemaMetadata::new(
"aion.outbox.dispatch.request",
"1",
br#"{"type":"object"}"#.as_slice(),
)
}
}
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
pub struct DispatchResponse {
pub workflow_id: WorkflowId,
pub ordinal: u64,
pub run_id: Option<RunId>,
pub outcome: Result<String, String>,
}
impl SchemaValidate for DispatchResponse {
fn schema_metadata() -> SchemaMetadata {
SchemaMetadata::new(
"aion.outbox.dispatch.response",
"1",
br#"{"type":"object"}"#.as_slice(),
)
}
}
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
pub struct InterventionRequest {
pub intervention: aion_core::InterventionCommand,
}
impl SchemaValidate for InterventionRequest {
fn schema_metadata() -> SchemaMetadata {
SchemaMetadata::new(
"aion.intervention.request",
"1",
br#"{"type":"object"}"#.as_slice(),
)
}
}
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
pub struct InterventionReply {
pub outcome: aion_core::InterventionOutcome,
}
impl SchemaValidate for InterventionReply {
fn schema_metadata() -> SchemaMetadata {
SchemaMetadata::new(
"aion.intervention.reply",
"1",
br#"{"type":"object"}"#.as_slice(),
)
}
}
pub const WORKER_LIVENESS_CHANNEL: &str = "aion.worker.liveness";
pub const WORKER_CAPABILITIES_CHANNEL: &str = "aion.worker.capabilities";
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
pub struct WorkerCapabilitiesAnnouncement {
pub capabilities: aion_core::InterventionCapabilities,
}
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
pub struct WorkerLivenessBeat {
pub workflow_id: WorkflowId,
pub ordinal: u64,
}
#[must_use]
pub fn request_for_row(row: &OutboxRow) -> DispatchRequest {
DispatchRequest {
activity_type: row.activity_type.clone(),
workflow_id: row.workflow_id.clone(),
ordinal: row.ordinal,
run_id: row.run_id.clone(),
input: row.input.bytes().to_vec(),
attempt: row.attempt.saturating_add(1),
labels: std::collections::BTreeMap::new(),
heartbeat_window_ms: 0,
}
}
const SEGMENT_SEPARATOR: char = '.';
const SEGMENT_ESCAPE: char = '%';
fn encode_segment(segment: &str) -> String {
if !segment.contains([SEGMENT_SEPARATOR, SEGMENT_ESCAPE]) {
return segment.to_owned();
}
let mut encoded = String::with_capacity(segment.len());
for ch in segment.chars() {
match ch {
SEGMENT_ESCAPE => encoded.push_str("%25"),
SEGMENT_SEPARATOR => encoded.push_str("%2E"),
other => encoded.push(other),
}
}
encoded
}
#[must_use]
pub fn dispatch_channel_name(namespace: &str, task_queue: &str, node: Option<&str>) -> String {
let namespace = encode_segment(namespace);
let task_queue = encode_segment(task_queue);
match node {
Some(node) => {
let node = encode_segment(node);
format!("aion.dispatch.{namespace}.{task_queue}.{node}")
}
None => format!("aion.dispatch.{namespace}.{task_queue}"),
}
}
#[must_use]
pub fn channel_for_row(row: &OutboxRow) -> String {
dispatch_channel_name(&row.namespace, &row.task_queue, row.node.as_deref())
}
fn dispatch_error(channel: &str, reason: String) -> ServerError {
ServerError::WorkerDispatch {
namespace: "liminal".to_owned(),
activity_type: channel.to_owned(),
reason,
}
}
pub struct LiminalCompletionSource {
callback: Arc<dyn OutboxDeliveryCallback>,
}
impl std::fmt::Debug for LiminalCompletionSource {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("LiminalCompletionSource")
.finish_non_exhaustive()
}
}
impl LiminalCompletionSource {
#[must_use]
pub fn new(callback: Arc<dyn OutboxDeliveryCallback>) -> Self {
Self { callback }
}
pub fn deliver(&self, response: &DispatchResponse) -> Result<bool, ServerError> {
let activity_id = ActivityId::from_sequence_position(response.ordinal);
match &response.outcome {
Ok(result) => self.callback.deliver_completion(
&response.workflow_id,
&activity_id,
response.run_id.as_ref(),
result.clone(),
),
Err(reason) => self.callback.deliver_failure(
&response.workflow_id,
&activity_id,
response.run_id.as_ref(),
reason.clone(),
),
}
}
}
#[must_use]
pub fn payload_from_request(request: &DispatchRequest) -> Payload {
Payload::new(ContentType::Json, request.input.clone())
}
#[derive(Clone)]
pub struct LiminalWorkerDelivery {
supervisor: ConnectionSupervisor,
pid: u64,
}
impl std::fmt::Debug for LiminalWorkerDelivery {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("LiminalWorkerDelivery")
.field("pid", &self.pid)
.finish_non_exhaustive()
}
}
impl LiminalWorkerDelivery {
#[must_use]
pub const fn new(supervisor: ConnectionSupervisor, pid: u64) -> Self {
Self { supervisor, pid }
}
#[must_use]
pub const fn pid(&self) -> u64 {
self.pid
}
pub fn dispatch(&self, request: &DispatchRequest) -> Result<DispatchResponse, ServerError> {
let awaiter = self.push_dispatch(request)?;
let reply = awaiter.receive(PUSH_REPLY_TIMEOUT).map_err(|error| {
if is_connection_closed_reply_error(&error) {
ServerError::worker_connection_lost(
"liminal-push",
format!("worker connection closed before reply: {error}"),
)
} else {
dispatch_error("liminal-push", format!("worker reply failed: {error}"))
}
})?;
decode_dispatch_response(&reply)
}
pub(crate) fn push_dispatch(
&self,
request: &DispatchRequest,
) -> Result<PushReplyAwaiter, ServerError> {
let payload = serde_json::to_vec(request).map_err(|error| {
dispatch_error("liminal-push", format!("request serialize failed: {error}"))
})?;
self.supervisor
.push_to_connection(self.pid, payload)
.map_err(|error| {
ServerError::worker_connection_lost(
"liminal-push",
format!("push to worker failed: {error}"),
)
})
}
pub fn push_intervention(
&self,
request: &InterventionRequest,
) -> Result<InterventionReply, ServerError> {
let payload = serde_json::to_vec(request).map_err(|error| {
dispatch_error(
"liminal-push",
format!("intervention serialize failed: {error}"),
)
})?;
let awaiter = self
.supervisor
.push_to_connection(self.pid, payload)
.map_err(|error| {
ServerError::worker_connection_lost(
"liminal-push",
format!("push intervention to worker failed: {error}"),
)
})?;
let reply = awaiter.receive(PUSH_REPLY_TIMEOUT).map_err(|error| {
if is_connection_closed_reply_error(&error) {
ServerError::worker_connection_lost(
"liminal-push",
format!("worker connection closed before intervention ack: {error}"),
)
} else {
dispatch_error("liminal-push", format!("intervention ack failed: {error}"))
}
})?;
serde_json::from_slice(&reply).map_err(|error| {
dispatch_error(
"liminal-push",
format!("intervention ack decode failed: {error}"),
)
})
}
}
fn decode_dispatch_response(reply: &[u8]) -> Result<DispatchResponse, ServerError> {
serde_json::from_slice(reply).map_err(|error| {
dispatch_error(
"liminal-push",
format!("worker reply decode failed: {error}"),
)
})
}
pub(crate) fn receive_bridge_reply(
awaiter: &PushReplyAwaiter,
keep_waiting: impl Fn() -> bool,
) -> Result<Option<DispatchResponse>, ServerError> {
loop {
match awaiter.receive(BRIDGE_REPLY_POLL) {
Ok(reply) => return decode_dispatch_response(&reply).map(Some),
Err(LiminalServerError::PushReplyTimeout { .. }) => {
if !keep_waiting() {
return Ok(None);
}
}
Err(error) if is_connection_closed_reply_error(&error) => {
return Err(ServerError::worker_connection_lost(
"liminal-push",
format!("worker connection closed before reply: {error}"),
));
}
Err(error) => {
return Err(dispatch_error(
"liminal-push",
format!("worker reply failed: {error}"),
));
}
}
}
}
fn is_connection_closed_reply_error(error: &LiminalServerError) -> bool {
matches!(error, LiminalServerError::PushReplyDisconnected { .. })
}
pub struct RegistryLiminalDispatch {
registry: ConnectedWorkerRegistry,
completion: LiminalCompletionSource,
placement_cache: Option<crate::worker::PlacementCache>,
attempt_owners: Option<super::intervention::AttemptOwnerIndex>,
}
impl std::fmt::Debug for RegistryLiminalDispatch {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("RegistryLiminalDispatch")
.field("placement_cache", &self.placement_cache.is_some())
.finish_non_exhaustive()
}
}
impl RegistryLiminalDispatch {
#[must_use]
pub fn new(
registry: ConnectedWorkerRegistry,
callback: Arc<dyn OutboxDeliveryCallback>,
) -> Self {
Self {
registry,
completion: LiminalCompletionSource::new(callback),
placement_cache: None,
attempt_owners: None,
}
}
#[must_use]
pub fn with_attempt_owners(
mut self,
attempt_owners: super::intervention::AttemptOwnerIndex,
) -> Self {
self.attempt_owners = Some(attempt_owners);
self
}
#[must_use]
pub fn with_placement_cache(mut self, cache: crate::worker::PlacementCache) -> Self {
self.placement_cache = Some(cache);
self
}
async fn select_liminal_worker(
&self,
row: &OutboxRow,
) -> Result<Option<WorkerHandle>, ServerError> {
let (Some(cache), None) = (&self.placement_cache, &row.node) else {
return self.registry.select_worker(
&row.namespace,
&row.task_queue,
&row.activity_type,
row.node.as_deref(),
);
};
let placement = cache.placement(&row.namespace).await;
match crate::worker::worker_selection_for(&placement) {
crate::worker::WorkerSelection::PreferTiers(tiers) => {
self.select_over_tiers(row, tiers.iter().map(Option::as_deref))
}
crate::worker::WorkerSelection::Required(required) => self.select_over_tiers(
row,
required.iter().map(|label| Some(String::as_str(label))),
),
}
}
fn select_over_tiers<'a>(
&self,
row: &OutboxRow,
tiers: impl Iterator<Item = Option<&'a str>>,
) -> Result<Option<WorkerHandle>, ServerError> {
for tier in tiers {
let selected = self.registry.select_worker(
&row.namespace,
&row.task_queue,
&row.activity_type,
tier,
)?;
if selected.is_some() {
return Ok(selected);
}
}
Ok(None)
}
}
#[async_trait]
impl OutboxRowDispatch for RegistryLiminalDispatch {
async fn dispatch(&self, row: &OutboxRow) -> Result<(), ServerError> {
let worker = self.select_liminal_worker(row).await?.ok_or_else(|| {
dispatch_error(
&channel_for_row(row),
"no liminal worker registered for the row's pool".to_owned(),
)
})?;
let delivery = match worker.delivery() {
WorkerDelivery::Liminal(delivery) => delivery.clone(),
WorkerDelivery::Grpc(_) => {
return Err(dispatch_error(
&channel_for_row(row),
"selected worker is not delivered over liminal".to_owned(),
));
}
};
let _owner_guard = self.attempt_owners.as_ref().map(|owners| {
AttemptOwnerGuard::bind(
owners.clone(),
super::intervention::AttemptKey::new(
row.workflow_id.clone(),
ActivityId::from_sequence_position(row.ordinal),
row.attempt.saturating_add(1),
),
worker.id(),
)
});
let request = request_for_row(row);
let response = tokio::task::spawn_blocking(move || delivery.dispatch(&request))
.await
.map_err(|error| {
dispatch_error(
&channel_for_row(row),
format!("dispatch task join failed: {error}"),
)
})??;
self.completion.deliver(&response)?;
Ok(())
}
}
pub(crate) struct AttemptOwnerGuard {
owners: super::intervention::AttemptOwnerIndex,
key: super::intervention::AttemptKey,
}
impl AttemptOwnerGuard {
pub(crate) fn bind(
owners: super::intervention::AttemptOwnerIndex,
key: super::intervention::AttemptKey,
worker: super::registry::WorkerId,
) -> Self {
owners.bind(key.clone(), worker);
Self { owners, key }
}
}
impl Drop for AttemptOwnerGuard {
fn drop(&mut self) {
self.owners.release(&self.key);
}
}
fn normalize_wire_node(node: Option<&str>) -> Option<String> {
node.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
}
pub struct LiminalConnectionNotifier {
registry: ConnectedWorkerRegistry,
supervisor: OnceLock<ConnectionSupervisor>,
guards: Mutex<HashMap<u64, WorkerRegistration>>,
intervention_capabilities: aion_core::InterventionCapabilities,
transcript: Option<TranscriptTap>,
heartbeat_tracker: Option<super::heartbeat::HeartbeatTracker>,
}
#[derive(Clone)]
struct TranscriptTap {
publisher: crate::activity_publisher::ActivityEventPublisher,
runtime: tokio::runtime::Handle,
}
impl std::fmt::Debug for LiminalConnectionNotifier {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("LiminalConnectionNotifier")
.field("supervisor_bound", &self.supervisor.get().is_some())
.finish_non_exhaustive()
}
}
impl LiminalConnectionNotifier {
#[must_use]
pub fn new(registry: ConnectedWorkerRegistry) -> Self {
Self {
registry,
supervisor: OnceLock::new(),
guards: Mutex::new(HashMap::new()),
intervention_capabilities: aion_core::InterventionCapabilities::none(),
transcript: None,
heartbeat_tracker: None,
}
}
#[must_use]
pub fn with_heartbeat_tracker(mut self, tracker: super::heartbeat::HeartbeatTracker) -> Self {
self.heartbeat_tracker = Some(tracker);
self
}
#[must_use]
pub fn with_transcript_publisher(
mut self,
publisher: crate::activity_publisher::ActivityEventPublisher,
) -> Self {
self.transcript = Some(TranscriptTap {
publisher,
runtime: tokio::runtime::Handle::current(),
});
self
}
#[must_use]
pub fn with_intervention_capabilities(
mut self,
capabilities: aion_core::InterventionCapabilities,
) -> Self {
self.intervention_capabilities = capabilities;
self
}
pub fn bind_supervisor(&self, supervisor: ConnectionSupervisor) -> bool {
self.supervisor.set(supervisor).is_ok()
}
fn record_liveness_beat(&self, pid: u64, payload: &[u8]) {
let Some(tracker) = &self.heartbeat_tracker else {
return;
};
let beat: WorkerLivenessBeat = match serde_json::from_slice(payload) {
Ok(beat) => beat,
Err(error) => {
tracing::warn!(%error, "liveness tap: malformed WorkerLivenessBeat payload");
return;
}
};
let worker_id = match self.guards.lock() {
Ok(guards) => guards.get(&pid).and_then(WorkerRegistration::worker_id),
Err(poisoned) => poisoned
.into_inner()
.get(&pid)
.and_then(WorkerRegistration::worker_id),
};
let Some(worker_id) = worker_id else {
tracing::warn!(
connection_pid = pid,
"liveness tap: beat from a connection with no registered worker"
);
return;
};
let activity_id = ActivityId::from_sequence_position(beat.ordinal);
if let Err(error) = tracker.record_liveness(
worker_id,
&beat.workflow_id,
&activity_id,
std::time::Instant::now(),
) {
tracing::error!(
%error,
connection_pid = pid,
"liveness tap: heartbeat tracker refresh failed"
);
}
}
fn record_capabilities_announcement(&self, pid: u64, payload: &[u8]) {
let announcement: WorkerCapabilitiesAnnouncement = match serde_json::from_slice(payload) {
Ok(announcement) => announcement,
Err(error) => {
tracing::warn!(
%error,
"capabilities tap: malformed WorkerCapabilitiesAnnouncement payload"
);
return;
}
};
let worker_id = match self.guards.lock() {
Ok(guards) => guards.get(&pid).and_then(WorkerRegistration::worker_id),
Err(poisoned) => poisoned
.into_inner()
.get(&pid)
.and_then(WorkerRegistration::worker_id),
};
let Some(worker_id) = worker_id else {
tracing::warn!(
connection_pid = pid,
"capabilities tap: announcement from a connection with no registered worker"
);
return;
};
match self
.registry
.set_intervention_capabilities(worker_id, &announcement.capabilities)
{
Ok(true) => {}
Ok(false) => tracing::warn!(
connection_pid = pid,
worker_id = ?worker_id,
"capabilities tap: announcement raced the worker's deregistration"
),
Err(error) => tracing::error!(
%error,
connection_pid = pid,
"capabilities tap: registry capability update failed"
),
}
}
}
impl ConnectionNotifier for LiminalConnectionNotifier {
fn on_worker_registered(
&self,
pid: u64,
registration: &WireWorkerRegistration,
) -> Result<(), LiminalServerError> {
let supervisor =
self.supervisor
.get()
.ok_or_else(|| LiminalServerError::ListenerAccept {
message: format!(
"liminal worker registration for connection {pid} rejected: \
notifier supervisor handle not yet bound"
),
})?;
let delivery = WorkerDelivery::Liminal(LiminalWorkerDelivery::new(supervisor.clone(), pid));
let node = normalize_wire_node(registration.node.as_deref());
let guard = self
.registry
.register_delivery_with_capabilities(
registration.namespaces.iter().cloned(),
registration.task_queue.clone(),
node,
registration.activity_types.iter(),
delivery,
self.intervention_capabilities.clone(),
)
.map_err(|error| LiminalServerError::ListenerAccept {
message: format!(
"liminal worker registration for connection {pid} rejected: {error}"
),
})?;
let mut guards = self.guards.lock().map_err(|_| {
LiminalServerError::ListenerAccept {
message: format!(
"liminal worker registration for connection {pid} rejected: \
notifier guard map poisoned"
),
}
})?;
guards.insert(pid, guard);
tracing::info!(
connection_pid = pid,
identity = %registration.identity,
task_queue = %registration.task_queue,
"registered liminal worker in-band"
);
Ok(())
}
fn on_worker_unregistered(&self, pid: u64) {
let removed = match self.guards.lock() {
Ok(mut guards) => guards.remove(&pid),
Err(poisoned) => poisoned.into_inner().remove(&pid),
};
if removed.is_some() {
tracing::info!(
connection_pid = pid,
"deregistered liminal worker on disconnect"
);
}
}
fn on_channel_publish(&self, pid: u64, channel: &str, payload: &[u8]) -> bool {
if channel == WORKER_LIVENESS_CHANNEL {
self.record_liveness_beat(pid, payload);
return true;
}
if channel == WORKER_CAPABILITIES_CHANNEL {
self.record_capabilities_announcement(pid, payload);
return true;
}
if channel != liminal_sdk::OBSERVABILITY_CHANNEL {
return false;
}
let Some(tap) = &self.transcript else {
return true;
};
let event: aion_core::ActivityEvent = match serde_json::from_slice(payload) {
Ok(event) => event,
Err(error) => {
tracing::warn!(%error, "observability tap: malformed ActivityEvent payload");
return true;
}
};
let publisher = tap.publisher.clone();
tap.runtime.spawn(async move {
if let Err(error) = publisher.publish(&event).await {
tracing::warn!(%error, "observability tap: transcript publish failed");
}
});
true
}
}
#[derive(Clone, Debug, Default)]
pub struct LiminalInterventionTransport;
#[async_trait]
impl super::intervention::InterventionTransport for LiminalInterventionTransport {
async fn push(
&self,
worker: &super::registry::WorkerHandle,
command: aion_core::InterventionCommand,
) -> Result<aion_core::InterventionOutcome, ServerError> {
let delivery = match worker.delivery() {
WorkerDelivery::Liminal(delivery) => delivery.clone(),
WorkerDelivery::Grpc(_) => {
return Err(ServerError::worker_connection_lost(
"liminal-push",
"owning worker is not delivered over liminal".to_owned(),
));
}
};
let request = InterventionRequest {
intervention: command,
};
let reply = tokio::task::spawn_blocking(move || delivery.push_intervention(&request))
.await
.map_err(|error| {
dispatch_error(
"liminal-push",
format!("intervention task join failed: {error}"),
)
})??;
Ok(reply.outcome)
}
}
#[cfg(test)]
mod tests {
use super::{channel_for_row, dispatch_channel_name, normalize_wire_node};
use aion_core::{ActivityId, ContentType, Payload, WorkflowId};
use aion_store::{OutboxRow, OutboxStatus};
use chrono::Utc;
use uuid::Uuid;
#[tokio::test]
async fn attempt_owner_guard_releases_on_drop() -> Result<(), Box<dyn std::error::Error>> {
use super::super::intervention::{AttemptKey, AttemptOwnerIndex};
use super::super::registry::{ConnectedWorkerRegistry, WorkerDelivery};
use super::AttemptOwnerGuard;
let registry = ConnectedWorkerRegistry::default();
let (tx, _rx) = tokio::sync::mpsc::channel(1);
let types = [String::from("agent")];
let registration = registry.register_delivery_with_capabilities(
[String::from("default")],
String::from("default"),
None,
types.iter(),
WorkerDelivery::Grpc(tx),
aion_core::InterventionCapabilities::none(),
)?;
let worker = registration
.worker_id()
.ok_or("registration must assign a worker id")?;
let owners = AttemptOwnerIndex::new();
let key = AttemptKey::new(
WorkflowId::new(Uuid::nil()),
ActivityId::from_sequence_position(3),
1,
);
owners.bind(key.clone(), worker);
assert_eq!(
owners.owner(&key),
Some(worker),
"owner bound before the guard"
);
{
let _guard = AttemptOwnerGuard {
owners: owners.clone(),
key: key.clone(),
};
assert_eq!(
owners.owner(&key),
Some(worker),
"still bound while in flight"
);
}
assert_eq!(
owners.owner(&key),
None,
"owner released when the dispatch returns"
);
Ok(())
}
#[test]
fn channel_format_is_pinned() {
assert_eq!(
dispatch_channel_name("remote", "gpu", None),
"aion.dispatch.remote.gpu"
);
assert_eq!(
dispatch_channel_name("local", "norn", None),
"aion.dispatch.local.norn"
);
}
#[test]
fn node_pinned_channel_appends_node_subsegment() {
assert_eq!(
dispatch_channel_name("remote", "gpu", Some("box-7")),
"aion.dispatch.remote.gpu.box-7"
);
}
#[test]
fn channel_derivation_is_stable() {
assert_eq!(
dispatch_channel_name("default", "default", None),
dispatch_channel_name("default", "default", None)
);
assert_eq!(
dispatch_channel_name("default", "default", Some("box-1")),
dispatch_channel_name("default", "default", Some("box-1"))
);
}
#[test]
fn distinct_pools_get_distinct_channels() {
assert_ne!(
dispatch_channel_name("remote", "gpu", None),
dispatch_channel_name("local", "norn", None)
);
}
#[test]
fn node_pin_separates_channels() {
let unpinned = dispatch_channel_name("remote", "gpu", None);
let box7 = dispatch_channel_name("remote", "gpu", Some("box-7"));
let box8 = dispatch_channel_name("remote", "gpu", Some("box-8"));
assert_ne!(
unpinned, box7,
"pinned dispatch must not reach unpinned pool"
);
assert_ne!(box7, box8, "distinct nodes must not collide");
}
#[test]
fn dotted_fields_do_not_collide_across_segments() {
assert_ne!(
dispatch_channel_name("a.b", "c", None),
dispatch_channel_name("a", "b.c", None),
"a '.' in a field must not bleed across the segment separator"
);
}
#[test]
fn node_subsegment_does_not_collide_with_dotted_fields() {
assert_ne!(
dispatch_channel_name("a", "b", Some("c")),
dispatch_channel_name("a", "b.c", None),
"a node sub-segment must not collide with a dotted task_queue"
);
assert_ne!(
dispatch_channel_name("a.b", "c", None),
dispatch_channel_name("a", "b", Some("c")),
"a dotted namespace must not collide with a node-pinned channel"
);
}
#[test]
fn reserved_char_shifts_stay_distinct() {
assert_ne!(
dispatch_channel_name("ns.", "tq", None),
dispatch_channel_name("ns", ".tq", None)
);
assert_ne!(
dispatch_channel_name("", "a.b", None),
dispatch_channel_name(".a", "b", None)
);
assert_ne!(
dispatch_channel_name("%2E", "x", None),
dispatch_channel_name(".", "x", None)
);
}
#[test]
fn encoding_is_injective_over_reserved_char_triples() {
let fields = ["a", "a.b", "a.", ".a", ".", "", "%", "%2E", "a%b", "%2."];
let nodes = [
None,
Some("a"),
Some("a.b"),
Some("."),
Some(""),
Some("%2E"),
];
let mut channels = std::collections::HashSet::new();
for ns in fields {
for tq in fields {
for node in nodes {
let channel = dispatch_channel_name(ns, tq, node);
assert!(
channels.insert(channel.clone()),
"collision on ({ns:?}, {tq:?}, {node:?}) -> {channel}"
);
}
}
}
}
fn row(namespace: &str, task_queue: &str) -> OutboxRow {
let workflow_id = WorkflowId::new(Uuid::new_v4());
OutboxRow {
dispatch_key: format!("{workflow_id}:0"),
workflow_id,
ordinal: 0,
run_id: None,
namespace: namespace.to_owned(),
task_queue: task_queue.to_owned(),
node: None,
activity_type: "charge-card".to_owned(),
input: Payload::new(ContentType::Json, Vec::new()),
status: OutboxStatus::Pending,
attempt: 0,
visible_after: Utc::now(),
claimed_at: None,
}
}
#[test]
fn channel_for_row_uses_namespace_and_task_queue_only() {
let remote_gpu = row("remote", "gpu");
let local_norn = row("local", "norn");
assert_eq!(channel_for_row(&remote_gpu), "aion.dispatch.remote.gpu");
assert_eq!(channel_for_row(&local_norn), "aion.dispatch.local.norn");
assert_ne!(channel_for_row(&remote_gpu), channel_for_row(&local_norn));
let mut other_activity = row("remote", "gpu");
other_activity.activity_type = "refund".to_owned();
assert_eq!(
channel_for_row(&remote_gpu),
channel_for_row(&other_activity),
"activity_type must not affect the channel"
);
}
#[test]
fn channel_for_row_derives_node_subchannel_when_pinned() {
let mut pinned = row("remote", "gpu");
pinned.node = Some("box-7".to_owned());
assert_eq!(channel_for_row(&pinned), "aion.dispatch.remote.gpu.box-7");
let unpinned = row("remote", "gpu");
assert_eq!(channel_for_row(&unpinned), "aion.dispatch.remote.gpu");
assert_ne!(channel_for_row(&pinned), channel_for_row(&unpinned));
}
#[test]
fn request_for_row_stamps_one_based_attempt_and_no_window() {
let mut retried = row("remote", "gpu");
retried.attempt = 2;
let request = super::request_for_row(&retried);
assert_eq!(
request.attempt, 3,
"zero-based row attempt goes one-based on the wire"
);
assert!(request.labels.is_empty());
assert_eq!(request.heartbeat_window_ms, 0);
let fresh = row("remote", "gpu");
assert_eq!(super::request_for_row(&fresh).attempt, 1);
}
#[test]
fn wire_node_normalizes_empty_to_none() {
assert_eq!(normalize_wire_node(None), None);
assert_eq!(normalize_wire_node(Some("")), None);
assert_eq!(normalize_wire_node(Some("box-7")), Some("box-7".to_owned()));
}
mod placement_selection {
use std::collections::BTreeSet;
use std::sync::Arc;
use std::time::Duration;
use aion_core::{ActivityId, Payload, RunId, WorkflowId};
use aion_store::{
InMemoryStore, NamespaceOrigin, NamespacePlacement, NamespaceStore, OutboxRow,
};
use crate::error::ServerError;
use crate::worker::PlacementCache;
use crate::worker::bridge::OutboxDeliveryCallback;
use crate::worker::registry::{ConnectedWorkerRegistry, WorkerMessage, WorkerRegistration};
use super::super::RegistryLiminalDispatch;
struct NoopCallback;
impl OutboxDeliveryCallback for NoopCallback {
fn deliver_completion(
&self,
_workflow_id: &WorkflowId,
_activity_id: &ActivityId,
_run_id: Option<&RunId>,
_result: String,
) -> Result<bool, ServerError> {
Ok(false)
}
fn deliver_failure(
&self,
_workflow_id: &WorkflowId,
_activity_id: &ActivityId,
_run_id: Option<&RunId>,
_reason: String,
) -> Result<bool, ServerError> {
Ok(false)
}
}
fn labels(values: &[&str]) -> BTreeSet<String> {
values.iter().map(|v| (*v).to_owned()).collect()
}
fn register_node_worker(
registry: &ConnectedWorkerRegistry,
namespace: &str,
node: &str,
) -> Result<WorkerRegistration, ServerError> {
let (tx, _rx) = tokio::sync::mpsc::channel::<WorkerMessage>(1);
let types = [String::from("charge")];
registry.register_namespaces(
[namespace.to_owned()],
String::from("default"),
Some(node.to_owned()),
types.iter(),
tx,
)
}
fn unpinned_row(namespace: &str) -> OutboxRow {
OutboxRow::pending(
WorkflowId::new_v4(),
0,
String::from("charge"),
Payload::from_json(&serde_json::json!({}))
.unwrap_or_else(|_| Payload::new(aion_core::ContentType::Json, Vec::new())),
chrono::Utc::now(),
)
.with_namespace(namespace)
.with_task_queue("default")
}
async fn prefer_store(
namespace: &str,
nodes: &[&str],
) -> Result<Arc<dyn NamespaceStore>, ServerError> {
let store: Arc<dyn NamespaceStore> = Arc::new(InMemoryStore::default());
store
.register_namespace(namespace, NamespaceOrigin::Explicit)
.await?;
store
.set_namespace_placement(
namespace,
NamespacePlacement::Prefer {
nodes: labels(nodes),
},
)
.await?;
Ok(store)
}
async fn pinned_store(
namespace: &str,
nodes: &[&str],
) -> Result<Arc<dyn NamespaceStore>, ServerError> {
let store: Arc<dyn NamespaceStore> = Arc::new(InMemoryStore::default());
store
.register_namespace(namespace, NamespaceOrigin::Explicit)
.await?;
store
.set_namespace_placement(
namespace,
NamespacePlacement::Pinned {
nodes: labels(nodes),
},
)
.await?;
Ok(store)
}
fn liminal_dispatch(
registry: &ConnectedWorkerRegistry,
ns_store: Arc<dyn NamespaceStore>,
) -> RegistryLiminalDispatch {
let cache = PlacementCache::new(ns_store, Duration::ZERO);
RegistryLiminalDispatch::new(registry.clone(), Arc::new(NoopCallback))
.with_placement_cache(cache)
}
#[tokio::test]
async fn prefer_selects_preferred_node_worker_on_liminal_path()
-> Result<(), Box<dyn std::error::Error>> {
let ns_store = prefer_store("t", &["n1"]).await?;
let registry = ConnectedWorkerRegistry::default();
let _n1 = register_node_worker(®istry, "t", "n1")?;
let _n2 = register_node_worker(®istry, "t", "n2")?;
let dispatch = liminal_dispatch(®istry, Arc::clone(&ns_store));
let row = unpinned_row("t");
let selected = dispatch
.select_liminal_worker(&row)
.await?
.ok_or("a worker must be selected")?;
assert_eq!(
selected.node(),
Some("n1"),
"the liminal path prefers the n1 worker while it is live"
);
assert_eq!(row.node, None, "placement must never mutate the row's node");
Ok(())
}
#[tokio::test]
async fn prefer_spills_to_any_live_worker_on_liminal_path()
-> Result<(), Box<dyn std::error::Error>> {
let ns_store = prefer_store("t", &["n1"]).await?;
let registry = ConnectedWorkerRegistry::default();
let _n2 = register_node_worker(®istry, "t", "n2")?;
let dispatch = liminal_dispatch(®istry, Arc::clone(&ns_store));
let row = unpinned_row("t");
let selected = dispatch
.select_liminal_worker(&row)
.await?
.ok_or("the spill must select the live n2 worker")?;
assert_eq!(
selected.node(),
Some("n2"),
"with no n1 worker live, the liminal selection spills to the live n2 worker"
);
assert_eq!(row.node, None, "spill must never mutate the row's node");
Ok(())
}
#[tokio::test]
async fn placement_never_mutates_recorded_row_node_on_liminal_path()
-> Result<(), Box<dyn std::error::Error>> {
let ns_store = prefer_store("t", &["n1"]).await?;
let registry = ConnectedWorkerRegistry::default();
let dispatch = liminal_dispatch(®istry, Arc::clone(&ns_store));
let n1 = register_node_worker(®istry, "t", "n1")?;
let row_a = unpinned_row("t");
let selected_a = dispatch
.select_liminal_worker(&row_a)
.await?
.ok_or("routing A must select a worker")?;
assert_eq!(selected_a.node(), Some("n1"));
n1.deregister()?;
let _n2 = register_node_worker(®istry, "t", "n2")?;
let row_b = unpinned_row("t");
let selected_b = dispatch
.select_liminal_worker(&row_b)
.await?
.ok_or("routing B must spill to a worker")?;
assert_eq!(selected_b.node(), Some("n2"));
assert_eq!(row_a.node, None);
assert_eq!(row_b.node, None);
assert_eq!(
row_a.node, row_b.node,
"the recorded row node is identical regardless of which worker was selected"
);
Ok(())
}
#[tokio::test]
async fn pinned_selects_required_node_worker_on_liminal_path()
-> Result<(), Box<dyn std::error::Error>> {
let ns_store = pinned_store("t", &["n1"]).await?;
let registry = ConnectedWorkerRegistry::default();
let _n1 = register_node_worker(®istry, "t", "n1")?;
let _n2 = register_node_worker(®istry, "t", "n2")?;
let dispatch = liminal_dispatch(®istry, Arc::clone(&ns_store));
let row = unpinned_row("t");
let selected = dispatch
.select_liminal_worker(&row)
.await?
.ok_or("the required n1 worker must be selected")?;
assert_eq!(selected.node(), Some("n1"));
assert_eq!(row.node, None, "placement must never mutate the row's node");
Ok(())
}
#[tokio::test]
async fn pinned_never_spills_to_a_wrong_node_worker_on_liminal_path()
-> Result<(), Box<dyn std::error::Error>> {
let ns_store = pinned_store("t", &["n1"]).await?;
let registry = ConnectedWorkerRegistry::default();
let _n2 = register_node_worker(®istry, "t", "n2")?;
let dispatch = liminal_dispatch(®istry, Arc::clone(&ns_store));
let row = unpinned_row("t");
let selected = dispatch.select_liminal_worker(&row).await?;
assert!(
selected.is_none(),
"Pinned{{n1}} must NOT spill to the live n2 worker — it selects nothing \
so the outbox retries/stalls until an n1 worker returns"
);
assert_eq!(row.node, None, "placement must never mutate the row's node");
Ok(())
}
#[tokio::test]
async fn authored_node_pin_wins_over_namespace_prefer_on_liminal_path()
-> Result<(), Box<dyn std::error::Error>> {
let ns_store = prefer_store("t", &["n1"]).await?;
let registry = ConnectedWorkerRegistry::default();
let _n1 = register_node_worker(®istry, "t", "n1")?;
let _n2 = register_node_worker(®istry, "t", "n2")?;
let dispatch = liminal_dispatch(®istry, Arc::clone(&ns_store));
let row = unpinned_row("t").with_node(Some(String::from("n2")));
let selected = dispatch
.select_liminal_worker(&row)
.await?
.ok_or("the authored pin must select the n2 worker")?;
assert_eq!(
selected.node(),
Some("n2"),
"the authored Some(n2) pin is honoured regardless of the namespace Prefer{{n1}}"
);
assert_eq!(row.node.as_deref(), Some("n2"));
Ok(())
}
#[tokio::test]
async fn unplaced_namespace_selects_any_worker_on_liminal_path()
-> Result<(), Box<dyn std::error::Error>> {
let ns_store: Arc<dyn NamespaceStore> = Arc::new(InMemoryStore::default());
ns_store
.register_namespace("t", NamespaceOrigin::Explicit)
.await?;
let registry = ConnectedWorkerRegistry::default();
let _n2 = register_node_worker(®istry, "t", "n2")?;
let dispatch = liminal_dispatch(®istry, Arc::clone(&ns_store));
let selected = dispatch
.select_liminal_worker(&unpinned_row("t"))
.await?
.ok_or("an Unplaced namespace still selects a live worker")?;
assert_eq!(
selected.node(),
Some("n2"),
"an Unplaced namespace reaches any live worker, exactly as before"
);
Ok(())
}
#[tokio::test]
async fn no_cache_selection_is_byte_identical_to_pre_163()
-> Result<(), Box<dyn std::error::Error>> {
let _ns_store = prefer_store("t", &["n1"]).await?;
let registry = ConnectedWorkerRegistry::default();
let _n2 = register_node_worker(®istry, "t", "n2")?;
let dispatch = RegistryLiminalDispatch::new(registry.clone(), Arc::new(NoopCallback));
let selected = dispatch
.select_liminal_worker(&unpinned_row("t"))
.await?
.ok_or("without a cache the unpinned row still selects any worker")?;
assert_eq!(
selected.node(),
Some("n2"),
"with no placement cache the selection is the unchanged any-worker path"
);
Ok(())
}
}
}