use aion_core::Payload;
use aion_proto::{ProtoPayload, ProtoStartWorkflowRequest};
use serde::Serialize;
use crate::client::Client;
use crate::error::ClientError;
use crate::handle::WorkflowHandle;
use crate::ops::{decode_required_run_id, decode_required_workflow_id, operation_namespace};
use crate::payload::to_payload;
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct StartOptions {
pub namespace: Option<String>,
pub idempotency_key: Option<String>,
pub routing_key: Option<String>,
pub task_queue: Option<String>,
pub display_name: Option<String>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct StartOutcome {
pub handle: WorkflowHandle,
pub display_name_not_applied: Option<DisplayNameNotApplied>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct DisplayNameNotApplied {
pub requested: String,
pub standing: Option<String>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub(crate) struct StartFingerprint {
namespace: String,
workflow_type: String,
content_type: aion_core::ContentType,
bytes: Vec<u8>,
routing_key: Option<String>,
task_queue: Option<String>,
idempotency_key: String,
}
impl StartFingerprint {
fn new(
namespace: String,
workflow_type: String,
input: &Payload,
routing_key: Option<String>,
task_queue: Option<String>,
idempotency_key: String,
) -> Self {
Self {
namespace,
workflow_type,
content_type: input.content_type().clone(),
bytes: input.bytes().to_vec(),
routing_key,
task_queue,
idempotency_key,
}
}
pub(crate) fn key(&self) -> &str {
&self.idempotency_key
}
}
#[derive(Clone, Debug)]
pub(crate) struct CachedStart {
fingerprint: StartFingerprint,
handle: WorkflowHandle,
display_name: Option<String>,
}
impl CachedStart {
pub(crate) const fn new(
fingerprint: StartFingerprint,
handle: WorkflowHandle,
display_name: Option<String>,
) -> Self {
Self {
fingerprint,
handle,
display_name,
}
}
pub(crate) const fn fingerprint(&self) -> &StartFingerprint {
&self.fingerprint
}
pub(crate) fn into_handle(self) -> WorkflowHandle {
self.handle
}
pub(crate) fn display_name_not_applied(
&self,
requested: Option<&str>,
) -> Option<DisplayNameNotApplied> {
let requested = requested?;
if self.display_name.as_deref() == Some(requested) {
return None;
}
Some(DisplayNameNotApplied {
requested: requested.to_owned(),
standing: self.display_name.clone(),
})
}
}
fn validate_start_options(opts: &StartOptions) -> Result<(), ClientError> {
if opts
.idempotency_key
.as_ref()
.is_some_and(std::string::String::is_empty)
{
return Err(ClientError::invalid_argument(
"idempotency_key must not be empty",
));
}
if opts
.display_name
.as_ref()
.is_some_and(|name| name.trim().is_empty())
{
return Err(ClientError::invalid_argument(
"display_name must not be blank; omit it to start the run unnamed",
));
}
Ok(())
}
impl Client {
pub async fn start(
&self,
workflow_type: impl Into<String>,
input: Payload,
opts: StartOptions,
) -> Result<StartOutcome, ClientError> {
validate_start_options(&opts)?;
let idempotency_key = opts.idempotency_key.clone();
let routing_key = opts.routing_key.clone();
let task_queue = opts.task_queue.clone();
let display_name = opts
.display_name
.as_deref()
.map(str::trim)
.map(str::to_owned);
let namespace = operation_namespace(self, opts.namespace);
let workflow_type = workflow_type.into();
let fingerprint = idempotency_key.as_ref().map(|key| {
StartFingerprint::new(
namespace.clone(),
workflow_type.clone(),
&input,
routing_key.clone(),
task_queue.clone(),
key.clone(),
)
});
if let Some(fingerprint) = &fingerprint
&& let Some(cached) = self.cached_start(fingerprint).await?
{
let not_applied = cached.display_name_not_applied(display_name.as_deref());
return Ok(StartOutcome {
handle: cached.into_handle(),
display_name_not_applied: not_applied,
});
}
let response = self
.transport
.start_workflow(ProtoStartWorkflowRequest {
namespace,
workflow_type,
input: Some(ProtoPayload::from(input)),
routing_key,
task_queue,
display_name: display_name.clone(),
})
.await?;
let workflow_id = decode_required_workflow_id(response.workflow_id, "start response")?;
let run_id = decode_required_run_id(response.run_id, "start response")?;
let handle = WorkflowHandle::from_ids(self.clone(), workflow_id, run_id);
if let Some(fingerprint) = fingerprint {
self.record_start(CachedStart::new(fingerprint, handle.clone(), display_name))
.await?;
}
Ok(StartOutcome {
handle,
display_name_not_applied: None,
})
}
pub async fn start_typed<T>(
&self,
workflow_type: impl Into<String>,
input: &T,
opts: StartOptions,
) -> Result<StartOutcome, ClientError>
where
T: Serialize + ?Sized,
{
self.start(workflow_type, to_payload(input)?, opts).await
}
}