use std::collections::BTreeMap;
use std::net::SocketAddr;
use std::time::Duration;
use aion_core::{ActivityId, ContentType, Payload, WorkflowId};
use aion_proto::generated::worker_protocol_client::WorkerProtocolClient;
use aion_proto::generated::{self, server_to_worker, worker_to_server};
use aion_server::ServerState;
use aion_server::api::worker_grpc::worker_service;
use aion_server::config::{
AuthConfig, AuthoringConfig, DeployConfig, ListenConfig, MetricsConfig, NamespaceConfig,
NamespaceMode, OpsConsoleAssetSource, OpsConsoleConfig, RuntimeConfig, WebSocketConfig,
WorkerConfig,
};
use aion_server::worker::{ActivityDispatcher, ConnectedWorkerRegistry, ScheduledActivity};
use aion_server::{NamespaceResolver, StaticScheduleNamespaces, StaticWorkflowNamespaces};
use tokio::net::TcpListener;
use tokio_stream::wrappers::{ReceiverStream, TcpListenerStream};
type TestError = Box<dyn std::error::Error>;
const ACTIVITY_TYPE: &str = "charge";
fn runtime_config() -> RuntimeConfig {
RuntimeConfig {
listen: ListenConfig {
grpc: SocketAddr::from(([127, 0, 0, 1], 0)),
http: SocketAddr::from(([127, 0, 0, 1], 0)),
},
tls: None,
auth: AuthConfig {
enabled: false,
jwks_url: None,
jwks_refresh_seconds: 300,
},
ops_console: OpsConsoleConfig {
source: OpsConsoleAssetSource::Embedded,
},
namespace: NamespaceConfig {
mode: NamespaceMode::SharedEngine,
},
worker: WorkerConfig {
heartbeat_window: Duration::from_secs(30),
..Default::default()
},
websocket: WebSocketConfig {
outbound_buffer_bound: 32,
event_broadcast_capacity: Some(64),
cluster_broadcast_capacity: Some(64),
},
workflow_packages: Vec::new(),
deploy: DeployConfig::default(),
authoring: AuthoringConfig::default(),
dev: aion_server::config::DevConfig::default(),
outbox: aion_server::config::OutboxConfig::default(),
observability: aion_server::config::ObservabilityConfig::with_flush_policy(64, 0),
mcp: aion_server::config::ResolvedMcpConfig::default(),
assistant: aion_server::config::ResolvedAssistantConfig::default(),
scheduler_threads: 1,
stop_drain_timeout: Some(std::time::Duration::from_secs(5)),
jit_threshold: None,
query_timeout: Some(Duration::from_secs(10)),
workloop_sweep_interval: Some(std::time::Duration::from_millis(50)),
default_namespace: "default".to_owned(),
auto_create: aion_server::config::AutoCreate::Open,
max_in_flight_activities: aion_server::config::DEFAULT_MAX_IN_FLIGHT_ACTIVITIES,
drain_timeout: Duration::from_secs(30),
metrics: MetricsConfig { enabled: false },
owned_shards: Vec::new(),
cors_allowed_origins: Vec::new(),
}
}
struct Cluster {
address: SocketAddr,
dispatcher: ActivityDispatcher,
server: tokio::task::JoinHandle<Result<(), tonic::transport::Error>>,
}
impl Cluster {
async fn start() -> Result<Self, TestError> {
let listener = TcpListener::bind("127.0.0.1:0").await?;
let address = listener.local_addr()?;
let registry = ConnectedWorkerRegistry::default();
let resolver = NamespaceResolver::authorization_only(
NamespaceMode::SharedEngine,
StaticWorkflowNamespaces::default(),
StaticScheduleNamespaces::default(),
);
let state =
ServerState::from_parts_with_registry(resolver, runtime_config(), registry.clone());
let server = tokio::spawn(
tonic::transport::Server::builder()
.add_service(worker_service(state.clone()))
.serve_with_incoming(TcpListenerStream::new(listener)),
);
let dispatcher =
ActivityDispatcher::new(registry.clone()).with_drain_state(state.drain_state().clone());
Ok(Self {
address,
dispatcher,
server,
})
}
async fn connect_worker(
&self,
label: &'static str,
namespaces: &[&str],
task_queue: &str,
node: &str,
) -> Result<TestWorker, TestError> {
let mut client = WorkerProtocolClient::connect(format!("http://{}", self.address)).await?;
let (worker_tx, worker_rx) = tokio::sync::mpsc::channel::<generated::WorkerToServer>(8);
worker_tx
.send(generated::WorkerToServer {
message: Some(worker_to_server::Message::Register(
generated::RegisterWorker {
namespaces: namespaces.iter().map(|ns| (*ns).to_owned()).collect(),
activity_types: vec![ACTIVITY_TYPE.to_owned()],
task_queue: task_queue.to_owned(),
node: node.to_owned(),
activities: Vec::new(),
identity: label.to_owned(),
instance: None,
max_concurrency: Some(4),
},
)),
})
.await?;
let mut request = tonic::Request::new(ReceiverStream::new(worker_rx));
request
.metadata_mut()
.insert("x-aion-subject", "tester".parse()?);
request
.metadata_mut()
.insert("x-aion-namespaces", namespaces.join(",").parse()?);
let mut inbound = client.stream_worker(request).await?.into_inner();
let first = inbound
.message()
.await?
.and_then(|frame| frame.message)
.ok_or("response stream ended before the RegisterAck")?;
let server_to_worker::Message::RegisterAck(ack) = first else {
return Err(format!("first response frame must be RegisterAck, got {first:?}").into());
};
Ok(TestWorker {
label,
worker_id: ack.worker_id,
_worker_tx: worker_tx,
inbound,
})
}
fn scheduled(namespace: &str, task_queue: &str, node: Option<&str>) -> ScheduledActivity {
ScheduledActivity {
namespace: namespace.to_owned(),
task_queue: task_queue.to_owned(),
activity_type: ACTIVITY_TYPE.to_owned(),
node: node.map(str::to_owned),
workflow_id: WorkflowId::new(uuid::Uuid::new_v4()),
activity_id: ActivityId::from_sequence_position(0),
run_id: Some(aion_core::RunId::new_v4()),
input: Payload::new(ContentType::Json, b"{}".to_vec()),
attempt: 1,
labels: BTreeMap::new(),
origin: aion_server::worker::dispatch::DispatchOrigin::Engine,
}
}
fn server(&self) {
self.server.abort();
}
}
struct TestWorker {
label: &'static str,
worker_id: u64,
_worker_tx: tokio::sync::mpsc::Sender<generated::WorkerToServer>,
inbound: tonic::Streaming<generated::ServerToWorker>,
}
impl TestWorker {
async fn expect_task(&mut self) -> Result<generated::ActivityTask, TestError> {
let frame = tokio::time::timeout(Duration::from_secs(5), async {
while let Some(message) = self.inbound.message().await? {
if let Some(server_to_worker::Message::Task(task)) = message.message {
return Ok::<_, TestError>(Some(task));
}
}
Ok(None)
})
.await
.map_err(|_| format!("worker {} received no task within the deadline", self.label))??;
frame.ok_or_else(|| {
format!("worker {} stream closed before a task arrived", self.label).into()
})
}
async fn expect_no_task(&mut self) -> Result<(), TestError> {
let outcome = tokio::time::timeout(Duration::from_millis(300), async {
while let Some(message) = self.inbound.message().await? {
if let Some(server_to_worker::Message::Task(task)) = message.message {
return Ok::<_, TestError>(Some(task));
}
}
Ok(None)
})
.await;
match outcome {
Err(_) | Ok(Ok(None)) => Ok(()),
Ok(Ok(Some(task))) => Err(format!(
"worker {} wrongly received a task ({}); routing leaked across a \
dimension it should have isolated",
self.label, task.activity_type
)
.into()),
Ok(Err(error)) => Err(error),
}
}
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn namespace_routing_isolates_and_a_set_worker_serves_both() -> Result<(), TestError> {
let cluster = Cluster::start().await?;
let mut worker_a = cluster
.connect_worker("ns-a", &["a"], "default", "")
.await?;
let mut worker_b = cluster
.connect_worker("ns-b", &["b"], "default", "")
.await?;
let mut worker_both = cluster
.connect_worker("ns-ab", &["a", "b"], "queue-ab", "")
.await?;
cluster
.dispatcher
.dispatch(&Cluster::scheduled("a", "default", None))
.await?;
worker_a.expect_task().await?;
worker_b.expect_no_task().await?;
worker_both.expect_no_task().await?;
cluster
.dispatcher
.dispatch(&Cluster::scheduled("a", "queue-ab", None))
.await?;
worker_both.expect_task().await?;
worker_a.expect_no_task().await?;
worker_b.expect_no_task().await?;
cluster
.dispatcher
.dispatch(&Cluster::scheduled("b", "queue-ab", None))
.await?;
worker_both.expect_task().await?;
worker_a.expect_no_task().await?;
worker_b.expect_no_task().await?;
cluster.server();
Ok(())
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn task_queue_routing_selects_the_named_pool() -> Result<(), TestError> {
let cluster = Cluster::start().await?;
let mut gpu = cluster
.connect_worker("gpu", &["tenant"], "gpu", "")
.await?;
let mut cpu = cluster
.connect_worker("cpu", &["tenant"], "cpu", "")
.await?;
cluster
.dispatcher
.dispatch(&Cluster::scheduled("tenant", "gpu", None))
.await?;
gpu.expect_task().await?;
cpu.expect_no_task().await?;
cluster
.dispatcher
.dispatch(&Cluster::scheduled("tenant", "cpu", None))
.await?;
cpu.expect_task().await?;
gpu.expect_no_task().await?;
cluster.server();
Ok(())
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn node_pinned_dispatch_reaches_only_the_pinned_node() -> Result<(), TestError> {
let cluster = Cluster::start().await?;
let mut n1 = cluster
.connect_worker("n1", &["tenant"], "pool", "n1")
.await?;
let mut n2 = cluster
.connect_worker("n2", &["tenant"], "pool", "n2")
.await?;
cluster
.dispatcher
.dispatch(&Cluster::scheduled("tenant", "pool", Some("n1")))
.await?;
n1.expect_task().await?;
n2.expect_no_task().await?;
cluster
.dispatcher
.dispatch(&Cluster::scheduled("tenant", "pool", Some("n2")))
.await?;
n2.expect_task().await?;
n1.expect_no_task().await?;
cluster.server();
Ok(())
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn unpinned_dispatch_round_robins_across_the_pool() -> Result<(), TestError> {
let cluster = Cluster::start().await?;
let mut n1 = cluster
.connect_worker("n1", &["tenant"], "pool", "n1")
.await?;
let mut n2 = cluster
.connect_worker("n2", &["tenant"], "pool", "n2")
.await?;
cluster
.dispatcher
.dispatch(&Cluster::scheduled("tenant", "pool", None))
.await?;
let first_served = served_exactly_one(&mut n1, &mut n2).await?;
cluster
.dispatcher
.dispatch(&Cluster::scheduled("tenant", "pool", None))
.await?;
let second_served = served_exactly_one(&mut n1, &mut n2).await?;
assert_ne!(
first_served, second_served,
"two unpinned dispatches must round-robin across both pool members, \
proving the pool is shared and not pinned to one worker"
);
cluster.server();
Ok(())
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn two_workers_sharing_a_node_are_both_eligible_when_pinned() -> Result<(), TestError> {
let cluster = Cluster::start().await?;
let mut a = cluster
.connect_worker("shared-a", &["tenant"], "pool", "shared")
.await?;
let mut b = cluster
.connect_worker("shared-b", &["tenant"], "pool", "shared")
.await?;
cluster
.dispatcher
.dispatch(&Cluster::scheduled("tenant", "pool", Some("shared")))
.await?;
let first = served_exactly_one(&mut a, &mut b).await?;
cluster
.dispatcher
.dispatch(&Cluster::scheduled("tenant", "pool", Some("shared")))
.await?;
let second = served_exactly_one(&mut a, &mut b).await?;
assert_ne!(
first, second,
"both workers sharing the pinned node must be eligible; the rotation \
must serve each once across two pinned dispatches"
);
cluster.server();
Ok(())
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn combined_address_lands_on_exact_match_and_is_refused_by_every_mismatch()
-> Result<(), TestError> {
let cluster = Cluster::start().await?;
let mut exact = cluster
.connect_worker("exact", &["ns"], "tq", "node")
.await?;
let mut wrong_ns = cluster
.connect_worker("wrong-ns", &["other"], "tq", "node")
.await?;
let mut wrong_tq = cluster
.connect_worker("wrong-tq", &["ns"], "other-tq", "node")
.await?;
let mut wrong_node = cluster
.connect_worker("wrong-node", &["ns"], "tq", "other-node")
.await?;
let mut no_node = cluster.connect_worker("no-node", &["ns"], "tq", "").await?;
cluster
.dispatcher
.dispatch(&Cluster::scheduled("ns", "tq", Some("node")))
.await?;
let task = exact.expect_task().await?;
assert_eq!(task.activity_type, ACTIVITY_TYPE);
wrong_ns.expect_no_task().await?;
wrong_tq.expect_no_task().await?;
wrong_node.expect_no_task().await?;
no_node.expect_no_task().await?;
cluster.server();
Ok(())
}
async fn served_exactly_one(
first: &mut TestWorker,
second: &mut TestWorker,
) -> Result<&'static str, TestError> {
tokio::select! {
result = first.expect_task() => {
result?;
second.expect_no_task().await?;
Ok(first.label)
}
result = second.expect_task() => {
result?;
first.expect_no_task().await?;
Ok(second.label)
}
}
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn distinct_workers_get_distinct_ids() -> Result<(), TestError> {
let cluster = Cluster::start().await?;
let one = cluster
.connect_worker("one", &["tenant"], "pool", "n1")
.await?;
let two = cluster
.connect_worker("two", &["tenant"], "pool", "n2")
.await?;
assert_ne!(
one.worker_id, two.worker_id,
"each registered stream must receive a distinct worker id"
);
cluster.server();
Ok(())
}