use std::sync::Arc;
use std::time::Duration;
use aion_core::{ActivityId, ContentType, Payload, RunId, WorkflowId};
use liminal::protocol::WorkerRegistration;
use liminal_sdk::{PushClient, PushedFrame};
use serde::{Deserialize, Serialize};
use crate::activity::ActivityRegistry;
use crate::config::WorkerConfig;
use crate::context::ActivityContext;
use crate::error::WorkerError;
use crate::protocol::ActivityTask;
use crate::runtime::liminal_redial::ServeResult;
use crate::runtime::loop_::{ActivityDispatcher, DispatchOutcome};
#[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>,
}
#[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>,
}
const RECV_POLL: Duration = Duration::from_millis(100);
pub struct LiminalActivityWorker {
client: PushClient,
registry: Arc<ActivityRegistry>,
}
impl std::fmt::Debug for LiminalActivityWorker {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("LiminalActivityWorker")
.field("client", &self.client)
.finish_non_exhaustive()
}
}
impl LiminalActivityWorker {
pub fn connect(
address: &str,
config: &WorkerConfig,
registry: Arc<ActivityRegistry>,
) -> Result<Self, WorkerError> {
let registration = registration_from(config, ®istry);
let client = PushClient::connect_with_registration(address, registration)
.map_err(|error| transport_error(&error))?;
Ok(Self { client, registry })
}
pub async fn serve_one(&self) -> Result<bool, WorkerError> {
match self.client.recv_timeout(RECV_POLL) {
Ok(frame) => {
self.handle_pushed_frame(frame).await?;
Ok(true)
}
Err(error) if is_recv_timeout(&error) => Ok(false),
Err(error) => Err(transport_error(&error)),
}
}
pub async fn serve_until<Stop>(&self, mut stop: Stop) -> Result<(), WorkerError>
where
Stop: FnMut() -> bool + Send,
{
while !stop() {
self.serve_one().await?;
}
Ok(())
}
pub(crate) async fn serve_until_drop<Stop>(&self, mut stop: Stop) -> ServeResult
where
Stop: FnMut() -> bool + Send,
{
let mut served_work = false;
while !stop() {
match self.serve_one().await {
Ok(true) => served_work = true,
Ok(false) => {}
Err(_) => return ServeResult::Dropped { served_work },
}
}
ServeResult::Stopped
}
async fn handle_pushed_frame(&self, frame: PushedFrame) -> Result<(), WorkerError> {
let correlation_id = frame.correlation_id();
let request: DispatchRequest =
serde_json::from_slice(frame.payload()).map_err(WorkerError::decode)?;
let response = self.execute(&request).await?;
let payload = serde_json::to_vec(&response).map_err(WorkerError::encode)?;
self.client
.reply(correlation_id, payload)
.map_err(|error| transport_error(&error))
}
async fn execute(&self, request: &DispatchRequest) -> Result<DispatchResponse, WorkerError> {
let activity_id = ActivityId::from_sequence_position(request.ordinal);
let input = Payload::new(ContentType::Json, request.input.clone());
let task = ActivityTask {
workflow_id: request.workflow_id.clone(),
activity_id: activity_id.clone(),
run_id: request.run_id.clone(),
activity_type: request.activity_type.clone(),
attempt: 1,
input,
labels: std::collections::BTreeMap::new(),
};
let (context, cancellation) = ActivityContext::new(activity_id, task.attempt);
drop(cancellation);
let outcome = match self.registry.dispatch(task, context).await {
Ok(outcome) => outcome,
Err(error) => {
return Ok(DispatchResponse {
workflow_id: request.workflow_id.clone(),
ordinal: request.ordinal,
run_id: request.run_id.clone(),
outcome: Err(error.to_string()),
});
}
};
let outcome = match outcome {
DispatchOutcome::Completed { output } => Ok(result_string(&output)),
DispatchOutcome::Failed { failure } => Err(failure.message),
};
Ok(DispatchResponse {
workflow_id: request.workflow_id.clone(),
ordinal: request.ordinal,
run_id: request.run_id.clone(),
outcome,
})
}
}
fn result_string(output: &Payload) -> String {
String::from_utf8_lossy(output.bytes()).into_owned()
}
fn is_recv_timeout(error: &liminal_sdk::SdkError) -> bool {
error
.to_string()
.contains("no server push arrived within the timeout")
}
fn registration_from(config: &WorkerConfig, registry: &ActivityRegistry) -> WorkerRegistration {
let node = if config.node.is_empty() {
None
} else {
Some(config.node.clone())
};
WorkerRegistration {
namespaces: config.namespaces.clone(),
task_queue: config.task_queue.clone(),
node,
activity_types: registry.activity_types().into_iter().collect(),
identity: config.identity.clone(),
}
}
fn transport_error(error: &liminal_sdk::SdkError) -> WorkerError {
WorkerError::Transport {
source: tonic::Status::unavailable(format!("liminal worker transport error: {error}")),
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use super::{DispatchRequest, DispatchResponse, registration_from};
use crate::activity::ActivityRegistry;
use crate::config::WorkerConfig;
use aion_core::{RunId, WorkflowId};
use uuid::Uuid;
fn worker_config(node: &str) -> Result<WorkerConfig, Box<dyn std::error::Error>> {
Ok(WorkerConfig::builder()
.endpoint("127.0.0.1:0")
.task_queue("gpu")
.identity("worker-a")
.max_concurrency(1)
.reconnect_initial_backoff(Duration::from_millis(5))
.reconnect_max_backoff(Duration::from_millis(20))
.reconnect_max_attempts(3)
.namespaces([String::from("remote"), String::from("payments")])
.node(node)
.build()?)
}
fn two_activity_registry() -> Result<ActivityRegistry, Box<dyn std::error::Error>> {
let registry = ActivityRegistry::new()
.register_activity("charge-card", |_input: serde_json::Value, _ctx| {
Box::pin(async move { Ok(serde_json::json!({})) })
})?
.register_activity("refund", |_input: serde_json::Value, _ctx| {
Box::pin(async move { Ok(serde_json::json!({})) })
})?;
Ok(registry)
}
#[test]
fn registration_carries_config_and_registry_activity_types()
-> Result<(), Box<dyn std::error::Error>> {
let config = worker_config("box-7")?;
let registry = two_activity_registry()?;
let registration = registration_from(&config, ®istry);
assert_eq!(
registration.namespaces,
vec![String::from("remote"), String::from("payments")]
);
assert_eq!(registration.task_queue, "gpu");
assert_eq!(registration.node, Some(String::from("box-7")));
assert_eq!(registration.identity, "worker-a");
assert_eq!(
registration.activity_types,
vec![String::from("charge-card"), String::from("refund")],
"activity types come from the bound registry, sorted"
);
Ok(())
}
#[test]
fn registration_empty_node_is_unpinned() -> Result<(), Box<dyn std::error::Error>> {
let config = worker_config("")?;
let registry = two_activity_registry()?;
let registration = registration_from(&config, ®istry);
assert_eq!(registration.node, None);
Ok(())
}
#[test]
fn dispatch_request_round_trips_through_json() -> Result<(), Box<dyn std::error::Error>> {
let request = DispatchRequest {
activity_type: "charge-card".to_owned(),
workflow_id: WorkflowId::new(Uuid::new_v4()),
ordinal: 7,
run_id: Some(RunId::new(Uuid::new_v4())),
input: br#"{"amount":42}"#.to_vec(),
};
let bytes = serde_json::to_vec(&request)?;
let decoded: DispatchRequest = serde_json::from_slice(&bytes)?;
assert_eq!(decoded, request);
let json = String::from_utf8(bytes)?;
for field in ["activity_type", "workflow_id", "ordinal", "run_id", "input"] {
assert!(json.contains(field), "wire JSON must carry `{field}`");
}
Ok(())
}
#[test]
fn dispatch_response_round_trips_both_outcomes() -> Result<(), Box<dyn std::error::Error>> {
let workflow_id = WorkflowId::new(Uuid::new_v4());
let ok = DispatchResponse {
workflow_id: workflow_id.clone(),
ordinal: 0,
run_id: None,
outcome: Ok(r#"{"charged":true}"#.to_owned()),
};
let err = DispatchResponse {
workflow_id,
ordinal: 1,
run_id: None,
outcome: Err("boom".to_owned()),
};
for response in [ok, err] {
let bytes = serde_json::to_vec(&response)?;
let decoded: DispatchResponse = serde_json::from_slice(&bytes)?;
assert_eq!(decoded, response);
}
Ok(())
}
}