use std::borrow::Borrow;
use std::fmt;
use std::ops::Deref;
use std::str::FromStr;
use serde::{Deserialize, Deserializer, Serialize};
use crate::error::{Error, Result};
pub const HEADER_RUN_ID: &str = "workflow.run_id";
pub const HEADER_STEP: &str = "workflow.step";
pub const RESERVED_HEADER_PREFIX: &str = "workflow.";
pub const RESERVED_KV_PREFIX: &str = "workflow/";
pub const HEADER_TERMINAL: &str = "workflow.terminal";
pub const HEADER_SIGNAL_WAIT: &str = "workflow.signal_wait";
pub const HEADER_SIGNAL_DELIVERED: &str = "workflow.signal_delivered";
pub const HEADER_GROUP: &str = "workflow.group";
pub const HEADER_GROUP_KEY: &str = "workflow.group_key";
pub(crate) const DEDUP_PREFIX: &str = "run:";
pub const MAX_RUN_ID_LEN: usize = 128;
pub(crate) const RUN_KV_PREFIX: &[u8] = b"workflow/runs/";
pub(crate) const STEP_KV_PREFIX: &[u8] = b"workflow/steps/";
pub(crate) const SIGNAL_WAIT_KV_PREFIX: &[u8] = b"workflow/signal-wait/";
pub(crate) const SIGNAL_BUF_KV_PREFIX: &[u8] = b"workflow/signal-buf/";
pub(crate) const SIGNAL_DELIVERED_KV_PREFIX: &[u8] = b"workflow/signal-delivered/";
pub(crate) const TERMINAL_KV_PREFIX: &[u8] = b"workflow/terminals/";
pub(crate) const OUTCOME_KV_PREFIX: &[u8] = b"workflow/outcomes/";
pub(crate) const GROUP_KV_PREFIX: &[u8] = b"workflow/groups/";
pub(crate) fn group_members_kv_prefix(group_id: &RunId) -> Vec<u8> {
prefixed(GROUP_KV_PREFIX, &format!("{group_id}/"))
}
pub(crate) fn group_member_kv_key(group_id: &RunId, key: &str) -> Vec<u8> {
prefixed(&group_members_kv_prefix(group_id), key)
}
pub(crate) const GROUP_TERMINAL_KV_PREFIX: &[u8] = b"workflow/group-terminals/";
pub(crate) fn hash_input(input: &[u8]) -> [u8; 32] {
use sha2::{Digest, Sha256};
Sha256::digest(input).into()
}
pub(crate) fn hex_sha256(parts: &[&[u8]]) -> String {
use sha2::{Digest, Sha256};
use std::fmt::Write;
let mut hasher = Sha256::new();
for part in parts {
hasher.update(part);
}
let mut hex = String::with_capacity(64);
for byte in hasher.finalize() {
let _ = write!(&mut hex, "{byte:02x}");
}
hex
}
#[derive(Debug, Clone, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize)]
#[serde(transparent)]
pub struct RunId(String);
impl RunId {
pub fn new(id: impl Into<String>) -> Result<Self> {
let id = id.into();
let reason = if id.is_empty() {
"run id must not be empty"
} else if id.len() > MAX_RUN_ID_LEN {
"run id exceeds maximum length of 128 bytes"
} else if !id
.bytes()
.all(|b| b.is_ascii_alphanumeric() || b == b'_' || b == b'-')
{
"run id must contain only `[A-Za-z0-9_-]`"
} else {
return Ok(Self(id));
};
Err(Error::InvalidRunId { run_id: id, reason })
}
pub(crate) fn generate() -> Self {
Self(ulid::Ulid::new().to_string())
}
pub(crate) fn digest(parts: &[&[u8]]) -> Self {
Self(hex_sha256(parts))
}
pub fn as_str(&self) -> &str {
&self.0
}
pub fn into_string(self) -> String {
self.0
}
}
impl Deref for RunId {
type Target = str;
fn deref(&self) -> &str {
&self.0
}
}
impl AsRef<str> for RunId {
fn as_ref(&self) -> &str {
&self.0
}
}
impl Borrow<str> for RunId {
fn borrow(&self) -> &str {
&self.0
}
}
impl fmt::Display for RunId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(&self.0)
}
}
impl FromStr for RunId {
type Err = Error;
fn from_str(id: &str) -> Result<Self> {
Self::new(id)
}
}
impl PartialEq<str> for RunId {
fn eq(&self, other: &str) -> bool {
self.0 == other
}
}
impl PartialEq<&str> for RunId {
fn eq(&self, other: &&str) -> bool {
self.0 == *other
}
}
impl PartialEq<String> for RunId {
fn eq(&self, other: &String) -> bool {
&self.0 == other
}
}
impl From<RunId> for String {
fn from(id: RunId) -> Self {
id.0
}
}
impl<'de> Deserialize<'de> for RunId {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> std::result::Result<Self, D::Error> {
let id = String::deserialize(deserializer)?;
Self::new(id).map_err(serde::de::Error::custom)
}
}
fn prefixed(prefix: &[u8], suffix: &str) -> Vec<u8> {
let mut k = Vec::with_capacity(prefix.len() + suffix.len());
k.extend_from_slice(prefix);
k.extend_from_slice(suffix.as_bytes());
k
}
pub(crate) fn run_kv_key(run_id: &RunId) -> Vec<u8> {
prefixed(RUN_KV_PREFIX, run_id)
}
pub(crate) fn step_kv_key(run_id: &RunId) -> Vec<u8> {
prefixed(STEP_KV_PREFIX, run_id)
}
pub(crate) fn outcome_kv_key(run_id: &RunId) -> Vec<u8> {
prefixed(OUTCOME_KV_PREFIX, run_id)
}
pub(crate) fn signal_wait_kv_key(correlation_key: &str) -> Vec<u8> {
prefixed(SIGNAL_WAIT_KV_PREFIX, correlation_key)
}
pub(crate) fn signal_buf_kv_key(correlation_key: &str) -> Vec<u8> {
prefixed(SIGNAL_BUF_KV_PREFIX, correlation_key)
}
pub(crate) fn signal_delivered_kv_key(run_id: &RunId, step_number: u32) -> Vec<u8> {
prefixed(
SIGNAL_DELIVERED_KV_PREFIX,
&format!("{run_id}/{step_number}"),
)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn internal_kv_prefixes_are_under_the_reserved_prefix() {
for prefix in [
RUN_KV_PREFIX,
STEP_KV_PREFIX,
SIGNAL_WAIT_KV_PREFIX,
SIGNAL_BUF_KV_PREFIX,
SIGNAL_DELIVERED_KV_PREFIX,
TERMINAL_KV_PREFIX,
OUTCOME_KV_PREFIX,
GROUP_KV_PREFIX,
GROUP_TERMINAL_KV_PREFIX,
] {
assert!(
prefix.starts_with(RESERVED_KV_PREFIX.as_bytes()),
"internal kv prefix `{}` is outside the reserved prefix",
String::from_utf8_lossy(prefix),
);
}
}
#[test]
fn run_id_rejects_empty_long_and_unsafe_ids() {
for bad in [
"",
"run/1",
"run 1",
"run:1",
&"a".repeat(MAX_RUN_ID_LEN + 1),
] {
assert!(
matches!(RunId::new(bad), Err(Error::InvalidRunId { .. })),
"`{bad}` must be rejected",
);
}
assert!(RunId::new("a".repeat(MAX_RUN_ID_LEN)).is_ok());
assert!(rmp_serde::from_slice::<RunId>(&rmp_serde::to_vec("").unwrap()).is_err());
assert_eq!(
rmp_serde::from_slice::<RunId>(&rmp_serde::to_vec("run-1").unwrap()).unwrap(),
"run-1"
);
}
}