use std::collections::HashSet;
use std::fmt::Write as _;
use super::activity_model::ResolvedActivity;
use super::model::{EnumDef, Field, GleamType, RecordDef, TypeDef};
const RUST_HEADER: &str =
"//! Generated by aion generate — do not edit; regenerate from the activity declarations.";
pub(crate) fn emit(package_name: &str, activities: &[&ResolvedActivity]) -> String {
let mut out = String::new();
out.push_str(RUST_HEADER);
let _ = write!(
out,
"\n//!\n//! Remote Rust worker for the `{package_name}` package. Serves its `RemoteRust`\n\
//! activities; the side-effecting bodies live in `handlers.rs`, which this module\n\
//! dispatches to. Configure the connection through the `AION_WORKER_*` and\n\
//! `AION_RECONNECT_*` environment variables.\n\n"
);
out.push_str("use std::time::Duration;\n\n");
out.push_str("use aion_worker::{Worker, WorkerConfig};\n");
out.push_str("use serde::{Deserialize, Serialize};\n\n");
out.push_str("mod handlers;\n");
for def in ordered_type_defs(activities) {
match def {
TypeDef::Record(record) => emit_struct(&mut out, record),
TypeDef::Enum(definition) => emit_enum(&mut out, definition),
}
}
emit_main(&mut out, activities);
out
}
fn ordered_type_defs<'a>(activities: &[&'a ResolvedActivity<'a>]) -> Vec<&'a TypeDef> {
let mut seen: HashSet<&str> = HashSet::new();
let mut defs: Vec<&TypeDef> = Vec::new();
for activity in activities {
for boundary in [activity.input.boundary, activity.output.boundary] {
for def in &boundary.defs {
if seen.insert(def.type_name()) {
defs.push(def);
}
}
}
}
defs
}
fn emit_struct(out: &mut String, record: &RecordDef) {
let _ = write!(
out,
"\n#[derive(Clone, Debug, Serialize, Deserialize)]\npub struct {} {{",
record.type_name
);
if record.fields.is_empty() {
out.push_str("}\n");
return;
}
out.push('\n');
for field in &record.fields {
emit_struct_field(out, field);
}
out.push_str("}\n");
}
fn emit_struct_field(out: &mut String, field: &Field) {
let (ident, renamed) = rust_field_ident(&field.wire);
let mut attrs: Vec<String> = Vec::new();
if renamed {
attrs.push(format!("rename = \"{}\"", field.wire));
}
if !field.required {
attrs.push("default".to_owned());
attrs.push("skip_serializing_if = \"Option::is_none\"".to_owned());
}
if !attrs.is_empty() {
let _ = writeln!(out, " #[serde({})]", attrs.join(", "));
}
let ty = rust_type(&field.ty);
let ty = if field.required {
ty
} else {
format!("Option<{ty}>")
};
let _ = writeln!(out, " pub {ident}: {ty},");
}
fn emit_enum(out: &mut String, definition: &EnumDef) {
let _ = write!(
out,
"\n#[derive(Clone, Debug, Serialize, Deserialize)]\npub enum {} {{\n",
definition.type_name
);
for variant in &definition.variants {
let _ = writeln!(out, " #[serde(rename = \"{}\")]", variant.wire);
let _ = writeln!(out, " {},", variant.constructor);
}
out.push_str("}\n");
}
fn emit_main(out: &mut String, activities: &[&ResolvedActivity]) {
out.push_str("\n#[tokio::main]\nasync fn main() -> Result<(), Box<dyn std::error::Error>> {\n");
out.push_str(" let config = WorkerConfig::builder()\n");
out.push_str(" .endpoint(require_var(\"AION_WORKER_ENDPOINT\")?)\n");
out.push_str(" .task_queue(require_var(\"AION_TASK_QUEUE\")?)\n");
out.push_str(" .identity(require_var(\"AION_WORKER_IDENTITY\")?)\n");
out.push_str(" .max_concurrency(require_parse(\"AION_WORKER_CONCURRENCY\")?)\n");
out.push_str(
" .reconnect_initial_backoff(Duration::from_secs_f64(require_parse(\n \
\"AION_RECONNECT_INITIAL_BACKOFF_SECONDS\",\n )?))\n",
);
out.push_str(
" .reconnect_max_backoff(Duration::from_secs_f64(require_parse(\n \
\"AION_RECONNECT_MAX_BACKOFF_SECONDS\",\n )?))\n",
);
out.push_str(
" .reconnect_max_attempts(require_parse(\"AION_RECONNECT_MAX_ATTEMPTS\")?)\n",
);
out.push_str(" .build()?;\n\n");
out.push_str(" Worker::builder(config)\n");
for activity in activities {
let name = &activity.declaration.name;
let _ = writeln!(
out,
" .register_activity(\"{name}\", handlers::{name})?"
);
}
out.push_str(" .build()?\n .run()\n .await?;\n\n Ok(())\n}\n");
emit_env_helpers(out);
}
fn emit_env_helpers(out: &mut String) {
out.push_str(
"\n/// Reads a required environment variable, erroring if it is unset.\n\
fn require_var(name: &str) -> Result<String, Box<dyn std::error::Error>> {\n \
std::env::var(name)\n \
.map_err(|_| format!(\"required environment variable `{name}` is not set\").into())\n\
}\n",
);
out.push_str(
"\n/// Reads and parses a required environment variable.\n\
fn require_parse<T>(name: &str) -> Result<T, Box<dyn std::error::Error>>\n\
where\n \
T: std::str::FromStr,\n \
T::Err: std::fmt::Display,\n\
{\n \
require_var(name)?\n \
.parse::<T>()\n \
.map_err(|error| format!(\"environment variable `{name}` is invalid: {error}\").into())\n\
}\n",
);
}
fn rust_type(ty: &GleamType) -> String {
match ty {
GleamType::String => "String".to_owned(),
GleamType::Int => "i64".to_owned(),
GleamType::Float => "f64".to_owned(),
GleamType::Bool => "bool".to_owned(),
GleamType::List(inner) => format!("Vec<{}>", rust_type(inner)),
GleamType::Named { type_name, .. } => type_name.clone(),
}
}
fn rust_field_ident(wire: &str) -> (String, bool) {
if matches!(wire, "self" | "crate" | "super") {
(format!("{wire}_"), true)
} else if is_rust_keyword(wire) {
(format!("r#{wire}"), false)
} else {
(wire.to_owned(), false)
}
}
fn is_rust_keyword(word: &str) -> bool {
matches!(
word,
"as" | "break"
| "const"
| "continue"
| "crate"
| "else"
| "enum"
| "extern"
| "false"
| "fn"
| "for"
| "if"
| "impl"
| "in"
| "let"
| "loop"
| "match"
| "mod"
| "move"
| "mut"
| "pub"
| "ref"
| "return"
| "self"
| "static"
| "struct"
| "super"
| "trait"
| "true"
| "type"
| "unsafe"
| "use"
| "where"
| "while"
| "async"
| "await"
| "dyn"
| "abstract"
| "become"
| "box"
| "do"
| "final"
| "macro"
| "override"
| "priv"
| "typeof"
| "unsized"
| "virtual"
| "yield"
| "try"
)
}
#[cfg(test)]
mod tests {
use std::path::PathBuf;
use super::{emit, is_rust_keyword, rust_field_ident, rust_type};
use crate::codegen::activity_model::{ResolvedActivity, ResolvedType};
use crate::codegen::declaration::{ActivityDeclaration, Tier};
use crate::codegen::model::{
BoundaryType, EnumDef, EnumVariant, Field, GleamType, RecordDef, TypeDef,
};
fn record(type_name: &str, fields: Vec<Field>) -> BoundaryType {
BoundaryType {
file: PathBuf::from(format!("schemas/{}.json", type_name.to_lowercase())),
stem: type_name.to_lowercase(),
root: GleamType::Named {
type_name: type_name.to_owned(),
fn_prefix: type_name.to_lowercase(),
},
defs: vec![TypeDef::Record(RecordDef {
type_name: type_name.to_owned(),
fn_prefix: type_name.to_lowercase(),
fields,
})],
}
}
fn field(wire: &str, ty: GleamType, required: bool) -> Field {
Field {
wire: wire.to_owned(),
ty,
required,
}
}
fn declaration(name: &str, input: &str, output: &str) -> ActivityDeclaration {
ActivityDeclaration {
name: name.to_owned(),
tier: Tier::RemoteRust,
input_type: input.to_owned(),
output_type: output.to_owned(),
}
}
fn resolved<'a>(
declaration: &'a ActivityDeclaration,
input: &'a BoundaryType,
output: &'a BoundaryType,
) -> ResolvedActivity<'a> {
ResolvedActivity {
declaration,
input: ResolvedType {
gleam_type: declaration.input_type.clone(),
fn_prefix: declaration.input_type.to_lowercase(),
boundary: input,
},
output: ResolvedType {
gleam_type: declaration.output_type.clone(),
fn_prefix: declaration.output_type.to_lowercase(),
boundary: output,
},
}
}
#[test]
fn type_mapping_covers_every_scalar_list_and_named() {
assert_eq!(rust_type(&GleamType::String), "String");
assert_eq!(rust_type(&GleamType::Int), "i64");
assert_eq!(rust_type(&GleamType::Float), "f64");
assert_eq!(rust_type(&GleamType::Bool), "bool");
assert_eq!(
rust_type(&GleamType::List(Box::new(GleamType::List(Box::new(
GleamType::Int
))))),
"Vec<Vec<i64>>"
);
assert_eq!(
rust_type(&GleamType::Named {
type_name: "OrderInput".to_owned(),
fn_prefix: "order_input".to_owned(),
}),
"OrderInput"
);
}
#[test]
fn keyword_field_names_are_escaped_or_mangled() {
assert_eq!(rust_field_ident("order_id"), ("order_id".to_owned(), false));
assert_eq!(rust_field_ident("type"), ("r#type".to_owned(), false));
assert_eq!(rust_field_ident("match"), ("r#match".to_owned(), false));
assert_eq!(rust_field_ident("self"), ("self_".to_owned(), true));
assert_eq!(rust_field_ident("crate"), ("crate_".to_owned(), true));
assert!(is_rust_keyword("move") && !is_rust_keyword("item"));
}
#[test]
fn worker_emits_structs_registration_and_config_in_order() {
let order = record(
"OrderInput",
vec![
field("order_id", GleamType::String, true),
field("quantity", GleamType::Int, true),
field("note", GleamType::String, false),
],
);
let receipt = record(
"Receipt",
vec![field("payment_id", GleamType::String, true)],
);
let shipment = record(
"Shipment",
vec![field("shipment_id", GleamType::String, true)],
);
let charge = declaration("charge_payment", "OrderInput", "Receipt");
let ship = declaration("ship_order", "OrderInput", "Shipment");
let activities = [
resolved(&charge, &order, &receipt),
resolved(&ship, &order, &shipment),
];
let refs: Vec<&ResolvedActivity> = activities.iter().collect();
let module = emit("order_saga", &refs);
assert!(module.starts_with(super::RUST_HEADER));
assert!(module.contains("use aion_worker::{Worker, WorkerConfig};"));
assert!(
!module.contains("ActivityContext") && !module.contains("HandlerFuture"),
"generated worker must not import handler-signature types it never names"
);
assert!(module.contains("mod handlers;\n"));
assert_eq!(module.matches("pub struct OrderInput").count(), 1);
let order_at = module.find("pub struct OrderInput");
let receipt_at = module.find("pub struct Receipt");
let shipment_at = module.find("pub struct Shipment");
assert!(order_at.is_some() && receipt_at.is_some() && shipment_at.is_some());
assert!(order_at < receipt_at && receipt_at < shipment_at);
assert!(module.contains(" pub order_id: String,\n"));
assert!(module.contains(" pub quantity: i64,\n"));
assert!(module.contains(
" #[serde(default, skip_serializing_if = \"Option::is_none\")]\n pub note: Option<String>,\n"
));
assert!(
module.contains(".register_activity(\"charge_payment\", handlers::charge_payment)?")
);
assert!(module.contains(".register_activity(\"ship_order\", handlers::ship_order)?"));
let charge_at = module.find("charge_payment\", handlers");
let ship_at = module.find("ship_order\", handlers");
assert!(charge_at < ship_at);
assert!(module.contains(".endpoint(require_var(\"AION_WORKER_ENDPOINT\")?)"));
assert!(module.contains(".task_queue(require_var(\"AION_TASK_QUEUE\")?)"));
assert!(module.contains(".identity(require_var(\"AION_WORKER_IDENTITY\")?)"));
assert!(module.contains(".max_concurrency(require_parse(\"AION_WORKER_CONCURRENCY\")?)"));
assert!(module.contains("require_parse(\"AION_RECONNECT_MAX_ATTEMPTS\")?"));
assert!(module.contains("\"AION_RECONNECT_INITIAL_BACKOFF_SECONDS\""));
assert!(module.contains("\"AION_RECONNECT_MAX_BACKOFF_SECONDS\""));
assert!(module.contains("fn require_var(name: &str)"));
assert!(module.contains("fn require_parse<T>(name: &str)"));
assert!(
!module.contains("unwrap_or"),
"generated worker must invent no connection default (ADR-001)"
);
assert!(
!module.contains("Duration::from_secs(5)")
&& !module.contains("Duration::from_millis(500)"),
"reconnect values must not be hardcoded"
);
}
#[test]
fn enums_and_keyword_fields_render_with_serde_attrs() {
let mut tier = record(
"Job",
vec![
field("order_id", GleamType::String, true),
field("type", GleamType::String, true),
],
);
tier.defs.push(TypeDef::Enum(EnumDef {
type_name: "JobKind".to_owned(),
fn_prefix: "job_kind".to_owned(),
variants: vec![
EnumVariant {
constructor: "JobKindFast".to_owned(),
wire: "fast".to_owned(),
},
EnumVariant {
constructor: "JobKindSlowRun".to_owned(),
wire: "slow_run".to_owned(),
},
],
}));
let out = record("Done", vec![field("ok", GleamType::Bool, true)]);
let run = declaration("run_job", "Job", "Done");
let activities = [resolved(&run, &tier, &out)];
let refs: Vec<&ResolvedActivity> = activities.iter().collect();
let module = emit("demo", &refs);
assert!(module.contains(" pub r#type: String,\n"));
assert!(!module.contains("rename = \"type\""));
assert!(module.contains("pub enum JobKind {"));
assert!(module.contains(" #[serde(rename = \"fast\")]\n JobKindFast,\n"));
assert!(module.contains(" #[serde(rename = \"slow_run\")]\n JobKindSlowRun,\n"));
}
}