use std::collections::HashSet;
use nasa_macro_support::runtime_root;
use proc_macro::TokenStream;
use proc_macro2::TokenStream as TokenStream2;
use quote::{format_ident, quote};
use syn::{
parse_macro_input, punctuated::Punctuated, Expr, FnArg, GenericArgument, ItemFn, ItemImpl, Lit,
LitStr, Meta, Path, PathArguments, ReturnType, Token, Type,
};
#[proc_macro_attribute]
pub fn application(attr: TokenStream, item: TokenStream) -> TokenStream {
let components =
parse_macro_input!(attr with Punctuated::<LitStr, Token![,]>::parse_terminated);
let function = parse_macro_input!(item as ItemFn);
match expand_application(components.into_iter().collect(), function) {
Ok(expanded) => expanded.into(),
Err(error) => error.to_compile_error().into(),
}
}
#[proc_macro_attribute]
pub fn initializer(attr: TokenStream, item: TokenStream) -> TokenStream {
let metas = parse_macro_input!(attr with Punctuated::<Meta, Token![,]>::parse_terminated);
let item_impl = parse_macro_input!(item as ItemImpl);
match expand_initializer(metas, item_impl) {
Ok(expanded) => expanded.into(),
Err(error) => error.to_compile_error().into(),
}
}
#[proc_macro_attribute]
pub fn redis_job(attr: TokenStream, item: TokenStream) -> TokenStream {
let metas = parse_macro_input!(attr with Punctuated::<Meta, Token![,]>::parse_terminated);
let function = parse_macro_input!(item as ItemFn);
match expand_redis_job(metas, function) {
Ok(expanded) => expanded.into(),
Err(error) => error.to_compile_error().into(),
}
}
#[derive(Default)]
struct RedisJobArgs {
name: Option<LitStr>,
qualifier: Option<LitStr>,
worker: Option<LitStr>,
trigger: Option<LitStr>,
cron: Option<LitStr>,
zone: Option<LitStr>,
fixed_rate_ms: Option<syn::LitInt>,
fixed_delay_ms: Option<syn::LitInt>,
concurrency: Option<LitStr>,
misfire: Option<LitStr>,
timeout_ms: Option<syn::LitInt>,
max_attempts: Option<syn::LitInt>,
retry_delay_ms: Option<syn::LitInt>,
contract_revision: Option<syn::LitInt>,
schema: Option<LitStr>,
codecs: Option<Vec<LitStr>>,
fanout_receipt_timeout_ms: Option<syn::LitInt>,
fanout_receipt_max_retries: Option<syn::LitInt>,
fanout_failure_policy: Option<LitStr>,
definition_revision: Option<syn::LitInt>,
}
enum RedisJobParameter {
Application,
Context,
Payload(Box<Type>),
}
fn expand_redis_job(
metas: Punctuated<Meta, Token![,]>,
function: ItemFn,
) -> syn::Result<TokenStream2> {
if function.sig.asyncness.is_none() {
return Err(syn::Error::new_spanned(
function.sig.fn_token,
"redis_job function must be async",
));
}
if function.sig.receiver().is_some() || !function.sig.generics.params.is_empty() {
return Err(syn::Error::new_spanned(
&function.sig,
"redis_job function cannot have a receiver or generic parameters",
));
}
if function.sig.inputs.len() > 3 {
return Err(syn::Error::new_spanned(
&function.sig.inputs,
"redis_job function accepts at most Application, JobContext, and one payload",
));
}
let args = parse_redis_job_args(&metas)?;
validate_redis_job_args(&args)?;
let runtime = runtime_root("application", "napp")
.map_err(|message| syn::Error::new_spanned(&function.sig.ident, message))?;
let mut parameters = Vec::new();
let mut has_application = false;
let mut has_context = false;
let mut has_payload = false;
for input in &function.sig.inputs {
let FnArg::Typed(argument) = input else {
return Err(syn::Error::new_spanned(
input,
"redis_job does not accept self",
));
};
if matches!(argument.ty.as_ref(), Type::Reference(_)) {
return Err(syn::Error::new_spanned(
&argument.ty,
"redis_job parameters must be owned values",
));
}
let kind = classify_redis_job_parameter(&argument.ty);
match &kind {
RedisJobParameter::Application if has_application => {
return Err(syn::Error::new_spanned(
&argument.ty,
"Application can only be injected once",
));
}
RedisJobParameter::Context if has_context => {
return Err(syn::Error::new_spanned(
&argument.ty,
"JobContext can only be injected once",
));
}
RedisJobParameter::Payload(_) if has_payload => {
return Err(syn::Error::new_spanned(
&argument.ty,
"redis_job accepts only one payload",
));
}
RedisJobParameter::Application => has_application = true,
RedisJobParameter::Context => has_context = true,
RedisJobParameter::Payload(_) => has_payload = true,
}
parameters.push(kind);
}
let function_name = &function.sig.ident;
let derived_name = function_name.to_string();
let name = args.name.clone().unwrap_or_else(|| {
LitStr::new(
derived_name.strip_prefix("r#").unwrap_or(&derived_name),
function_name.span(),
)
});
let builder_steps = redis_job_builder_steps(&args, &runtime)?;
let handler_field = has_application.then(|| quote!(application: #runtime::WeakApplication,));
let handler_value = if has_application {
quote!(__NasaRedisJobHandler {
application: application.downgrade()
})
} else {
quote!({
let _ = application;
__NasaRedisJobHandler {}
})
};
let application_clone = has_application.then(|| {
quote! {
let application = match self.application.upgrade() {
::std::option::Option::Some(application) => application,
::std::option::Option::None => {
return ::std::boxed::Box::pin(async {
#runtime::__private::nadis::job::JobOutcome::retry(
"Application lifecycle ended before RedisJob invocation"
)
});
}
};
}
});
let declared_codecs = args
.codecs
.clone()
.unwrap_or_else(|| vec![LitStr::new("json", proc_macro2::Span::call_site())]);
let codec_wires = declared_codecs
.iter()
.map(|codec| match codec.value().to_ascii_lowercase().as_str() {
"json" => Ok(LitStr::new("JSON", codec.span())),
"protobuf" => Ok(LitStr::new("PROTOBUF", codec.span())),
"raw" => Ok(LitStr::new("RAW", codec.span())),
_ => Err(syn::Error::new_spanned(codec, "redis_job codec is unknown")),
})
.collect::<syn::Result<Vec<_>>>()?;
let codec_guard = quote! {
if !matches!(context.wire_codec.as_str(), #(#codec_wires)|*) {
return #runtime::__private::nadis::job::JobOutcome::fail_permanent(
"payload codec is outside the static Worker contract"
);
}
};
let mut call_arguments = Vec::new();
let mut payload_decode = TokenStream2::new();
for parameter in parameters {
match parameter {
RedisJobParameter::Application => call_arguments.push(quote!(application)),
RedisJobParameter::Context => call_arguments.push(quote!(context.clone())),
RedisJobParameter::Payload(payload_type) => {
payload_decode = if declared_codecs.len() == 1 {
match declared_codecs[0].value().to_ascii_lowercase().as_str() {
"json" => quote! {
let payload: #payload_type = match #runtime::__private::nadis::job::decode_json_payload(context.payload()) {
::std::result::Result::Ok(value) => value,
::std::result::Result::Err(error) => {
return #runtime::__private::nadis::job::JobOutcome::fail_permanent(
::std::format!("payload JSON decode rejected: {error}")
);
}
};
},
"protobuf" => quote! {
let payload: #payload_type = match
<#payload_type as #runtime::__private::prost::Message>::decode(context.payload())
{
::std::result::Result::Ok(value) => value,
::std::result::Result::Err(error) => {
return #runtime::__private::nadis::job::JobOutcome::fail_permanent(
::std::format!("payload Protobuf decode rejected: {error}")
);
}
};
},
"raw" => quote! {
let payload: #payload_type = match
<#payload_type as ::std::convert::TryFrom<::std::vec::Vec<u8>>>::try_from(
context.payload().to_vec()
)
{
::std::result::Result::Ok(value) => value,
::std::result::Result::Err(error) => {
return #runtime::__private::nadis::job::JobOutcome::fail_permanent(
::std::format!("payload RAW decode rejected: {error}")
);
}
};
},
_ => unreachable!("codec 已在宏期校验"),
}
} else {
quote! {
let codec = match #runtime::__private::nadis::job::JobWireCodec::parse(
context.wire_codec.as_str()
) {
::std::option::Option::Some(codec) => codec,
::std::option::Option::None => {
return #runtime::__private::nadis::job::JobOutcome::fail_permanent(
"payload codec is not a known wire value"
);
}
};
if codec == #runtime::__private::nadis::job::JobWireCodec::Json {
if let ::std::result::Result::Err(error) =
#runtime::__private::nadis::job::validate_json_payload(context.payload())
{
return #runtime::__private::nadis::job::JobOutcome::fail_permanent(
::std::format!("payload JSON validation rejected: {error}")
);
}
}
let payload: #payload_type = match
<#payload_type as #runtime::__private::nadis::job::JobParameter>::decode(
codec,
context.payload(),
)
{
::std::result::Result::Ok(value) => value,
::std::result::Result::Err(error) => {
return #runtime::__private::nadis::job::JobOutcome::fail_permanent(
::std::format!("payload decode rejected: {error}")
);
}
};
}
};
call_arguments.push(quote!(payload));
}
}
}
Ok(quote! {
#function
const _: () = {
struct __NasaRedisJobHandler { #handler_field }
impl #runtime::__private::nadis::job::JobHandler for __NasaRedisJobHandler {
fn handle<'a>(
&'a self,
execution: &'a #runtime::__private::nadis::job::JobExecution,
) -> #runtime::__private::nadis::job::JobHandlerFuture<'a> {
use #runtime::__private::nadis::job::IntoJobHandlerResult as _;
let context = execution.clone();
#application_clone
::std::boxed::Box::pin(async move {
#codec_guard
#payload_decode
#function_name(#(#call_arguments),*)
.await
.into_job_handler_result()
})
}
}
fn __nasa_redis_job_factory(
application: #runtime::Application,
) -> #runtime::ApplicationResult<(
#runtime::__private::nadis::job::JobDefinition,
::std::sync::Arc<dyn #runtime::__private::nadis::job::JobHandler>,
)> {
let mut builder = #runtime::__private::nadis::job::JobDefinition::builder(#name);
#(#builder_steps)*
let definition = builder.build().map_err(|error| {
#runtime::ApplicationError::with_source(
#runtime::ComponentId::RedisJob,
#runtime::ApplicationPhase::Prepare,
"invalid #[redis_job] definition",
error,
)
})?;
let handler: ::std::sync::Arc<dyn #runtime::__private::nadis::job::JobHandler> =
::std::sync::Arc::new(#handler_value);
::std::result::Result::Ok((definition, handler))
}
#[#runtime::__private::linkme::distributed_slice(#runtime::COLLECTED_REDIS_JOBS)]
#[linkme(crate = #runtime::__private::linkme)]
static __NASA_REDIS_JOB_DESCRIPTOR: #runtime::RedisJobDescriptor =
#runtime::RedisJobDescriptor::__new(
__nasa_redis_job_factory,
concat!(module_path!(), ":", file!(), ":", line!()),
);
};
})
}
fn classify_redis_job_parameter(ty: &Type) -> RedisJobParameter {
if let Type::Path(path) = ty {
let segments: Vec<String> = path
.path
.segments
.iter()
.map(|segment| segment.ident.to_string())
.collect();
let unqualified = segments.len() == 1;
if let Some(name) = segments.last() {
if name == "Application"
&& (unqualified
|| matches!(segments.as_slice(), [root, value] if matches!(root.as_str(), "nasa" | "napp") && value == "Application"))
{
return RedisJobParameter::Application;
}
let framework_job_path = segments.len() >= 2
&& segments
.get(segments.len() - 2)
.is_some_and(|segment| segment == "job");
if matches!(name.as_str(), "JobContext" | "JobExecution")
&& (unqualified || framework_job_path)
{
return RedisJobParameter::Context;
}
}
}
RedisJobParameter::Payload(Box::new(ty.clone()))
}
fn parse_redis_job_args(metas: &Punctuated<Meta, Token![,]>) -> syn::Result<RedisJobArgs> {
let mut result = RedisJobArgs::default();
let mut seen = HashSet::new();
for meta in metas {
let Meta::NameValue(value) = meta else {
return Err(syn::Error::new_spanned(
meta,
"redis_job attributes must use `key = value`",
));
};
let key = value
.path
.get_ident()
.map(ToString::to_string)
.ok_or_else(|| {
syn::Error::new_spanned(
&value.path,
"redis_job attribute key must be an identifier",
)
})?;
if !seen.insert(key.clone()) {
return Err(syn::Error::new_spanned(
meta,
"redis_job attribute key is repeated",
));
}
macro_rules! string_field {
($field:ident) => {{
result.$field = Some(parse_job_string(&value.value, &key)?);
}};
}
macro_rules! integer_field {
($field:ident) => {{
result.$field = Some(parse_job_integer(&value.value, &key)?);
}};
}
match key.as_str() {
"name" => string_field!(name),
"qualifier" => string_field!(qualifier),
"worker" => string_field!(worker),
"trigger" => string_field!(trigger),
"cron" => string_field!(cron),
"zone" => string_field!(zone),
"fixed_rate_ms" => integer_field!(fixed_rate_ms),
"fixed_delay_ms" => integer_field!(fixed_delay_ms),
"concurrency" => string_field!(concurrency),
"misfire" => string_field!(misfire),
"timeout_ms" => integer_field!(timeout_ms),
"max_attempts" => integer_field!(max_attempts),
"retry_delay_ms" => integer_field!(retry_delay_ms),
"contract_revision" => integer_field!(contract_revision),
"schema" => string_field!(schema),
"fanout_receipt_timeout_ms" => integer_field!(fanout_receipt_timeout_ms),
"fanout_receipt_max_retries" => integer_field!(fanout_receipt_max_retries),
"fanout_failure_policy" => string_field!(fanout_failure_policy),
"definition_revision" => integer_field!(definition_revision),
"codecs" => {
let Expr::Array(array) = &value.value else {
return Err(syn::Error::new_spanned(
&value.value,
"redis_job codecs must be an array of strings",
));
};
result.codecs = Some(
array
.elems
.iter()
.map(|item| parse_job_string(item, "codecs"))
.collect::<syn::Result<Vec<_>>>()?,
);
}
_ => {
return Err(syn::Error::new_spanned(
&value.path,
format!("unknown redis_job attribute `{key}`"),
))
}
}
}
Ok(result)
}
fn validate_redis_job_args(args: &RedisJobArgs) -> syn::Result<()> {
if args.codecs.as_ref().is_some_and(Vec::is_empty) {
return Err(syn::Error::new(
proc_macro2::Span::call_site(),
"redis_job codecs must not be empty",
));
}
let schedule_count = usize::from(args.cron.is_some())
+ usize::from(args.fixed_rate_ms.is_some())
+ usize::from(args.fixed_delay_ms.is_some());
if schedule_count > 1 {
return Err(syn::Error::new(
proc_macro2::Span::call_site(),
"redis_job cron, fixed_rate_ms, and fixed_delay_ms are mutually exclusive",
));
}
if args
.trigger
.as_ref()
.is_some_and(|trigger| trigger.value().eq_ignore_ascii_case("fanout_only"))
&& schedule_count != 0
{
return Err(syn::Error::new_spanned(
args.trigger.as_ref().expect("trigger 已确认存在"),
"redis_job fanout_only must not declare a schedule",
));
}
for (field, value) in [
("fixed_rate_ms", args.fixed_rate_ms.as_ref()),
("fixed_delay_ms", args.fixed_delay_ms.as_ref()),
("timeout_ms", args.timeout_ms.as_ref()),
("max_attempts", args.max_attempts.as_ref()),
("retry_delay_ms", args.retry_delay_ms.as_ref()),
("contract_revision", args.contract_revision.as_ref()),
(
"fanout_receipt_timeout_ms",
args.fanout_receipt_timeout_ms.as_ref(),
),
("definition_revision", args.definition_revision.as_ref()),
] {
if value.is_some_and(|literal| matches!(literal.base10_parse::<u64>(), Ok(0))) {
return Err(syn::Error::new_spanned(
value.expect("value 已确认存在"),
format!("redis_job {field} must be greater than zero"),
));
}
}
Ok(())
}
fn parse_job_string(expression: &Expr, field: &str) -> syn::Result<LitStr> {
match expression {
Expr::Lit(value) => match &value.lit {
Lit::Str(literal) => Ok(literal.clone()),
_ => Err(syn::Error::new_spanned(
expression,
format!("redis_job {field} must be a string literal"),
)),
},
_ => Err(syn::Error::new_spanned(
expression,
format!("redis_job {field} must be a string literal"),
)),
}
}
fn parse_job_integer(expression: &Expr, field: &str) -> syn::Result<syn::LitInt> {
match expression {
Expr::Lit(value) => match &value.lit {
Lit::Int(literal) => Ok(literal.clone()),
_ => Err(syn::Error::new_spanned(
expression,
format!("redis_job {field} must be a non-negative integer literal"),
)),
},
_ => Err(syn::Error::new_spanned(
expression,
format!("redis_job {field} must be a non-negative integer literal"),
)),
}
}
fn redis_job_builder_steps(
args: &RedisJobArgs,
runtime: &TokenStream2,
) -> syn::Result<Vec<TokenStream2>> {
let mut steps = Vec::new();
macro_rules! scalar_step {
($field:ident, $method:ident) => { if let Some(value) = &args.$field { steps.push(quote!(builder = builder.$method(#value);)); } };
}
scalar_step!(qualifier, qualifier);
scalar_step!(worker, worker_name);
if let Some(trigger) = &args.trigger {
match trigger.value().to_ascii_lowercase().as_str() {
"scheduled" => {}
"fanout_only" => steps.push(quote!(builder = builder.fanout_only();)),
_ => {
return Err(syn::Error::new_spanned(
trigger,
"redis_job trigger must be `scheduled` or `fanout_only`",
))
}
}
}
if let Some(cron) = &args.cron {
let zone = args
.zone
.clone()
.unwrap_or_else(|| LitStr::new("UTC", cron.span()));
steps.push(quote! {
builder = builder.cron(#cron, #zone).map_err(|error| {
#runtime::ApplicationError::with_source(
#runtime::ComponentId::RedisJob,
#runtime::ApplicationPhase::Prepare,
"invalid #[redis_job] cron declaration",
error,
)
})?;
});
} else if let Some(zone) = &args.zone {
return Err(syn::Error::new_spanned(
zone,
"redis_job zone requires cron",
));
}
scalar_step!(fixed_rate_ms, fixed_rate_ms);
scalar_step!(fixed_delay_ms, fixed_delay_ms);
if let Some(value) = &args.concurrency {
let variant = match value.value().to_ascii_lowercase().as_str() {
"serial_queue" => quote!(#runtime::__private::nadis::job::JobConcurrency::SerialQueue),
"discard_if_running" => {
quote!(#runtime::__private::nadis::job::JobConcurrency::DiscardIfRunning)
}
"parallel" => quote!(#runtime::__private::nadis::job::JobConcurrency::Parallel),
_ => {
return Err(syn::Error::new_spanned(
value,
"redis_job concurrency is unknown",
))
}
};
steps.push(quote!(builder = builder.concurrency(#variant);));
}
if let Some(value) = &args.misfire {
let variant = match value.value().to_ascii_lowercase().as_str() {
"do_nothing" => quote!(#runtime::__private::nadis::job::JobMisfire::DoNothing),
"fire_once_now" => quote!(#runtime::__private::nadis::job::JobMisfire::FireOnceNow),
"catch_up" => quote!(#runtime::__private::nadis::job::JobMisfire::CatchUp),
_ => {
return Err(syn::Error::new_spanned(
value,
"redis_job misfire is unknown",
))
}
};
steps.push(quote!(builder = builder.misfire(#variant);));
}
scalar_step!(timeout_ms, timeout_ms);
scalar_step!(max_attempts, max_attempts);
scalar_step!(retry_delay_ms, retry_delay_ms);
scalar_step!(contract_revision, contract_revision);
scalar_step!(schema, schema_id);
if let Some(codecs) = &args.codecs {
let mut variants = Vec::new();
for codec in codecs {
variants.push(match codec.value().to_ascii_lowercase().as_str() {
"json" => quote!(#runtime::__private::nadis::job::JobWireCodec::Json),
"protobuf" => quote!(#runtime::__private::nadis::job::JobWireCodec::Protobuf),
"raw" => quote!(#runtime::__private::nadis::job::JobWireCodec::Raw),
_ => return Err(syn::Error::new_spanned(codec, "redis_job codec is unknown")),
});
}
steps.push(quote!(builder = builder.codecs([#(#variants),*]);));
}
if args.fanout_receipt_timeout_ms.is_some() || args.fanout_receipt_max_retries.is_some() {
let timeout = args
.fanout_receipt_timeout_ms
.clone()
.unwrap_or_else(|| syn::LitInt::new("2000", proc_macro2::Span::call_site()));
let retries = args
.fanout_receipt_max_retries
.clone()
.unwrap_or_else(|| syn::LitInt::new("3", proc_macro2::Span::call_site()));
steps.push(quote!(builder = builder.fanout_receipt(#timeout, #retries);));
}
if let Some(value) = &args.fanout_failure_policy {
let variant = match value.value().to_ascii_lowercase().as_str() {
"reassign_on_failure" => {
quote!(#runtime::__private::nadis::job::JobFanoutFailurePolicy::ReassignOnFailure)
}
"strict_snapshot" => {
quote!(#runtime::__private::nadis::job::JobFanoutFailurePolicy::StrictSnapshot)
}
"best_effort" => {
quote!(#runtime::__private::nadis::job::JobFanoutFailurePolicy::BestEffort)
}
_ => {
return Err(syn::Error::new_spanned(
value,
"redis_job fanout_failure_policy is unknown",
))
}
};
steps.push(quote!(builder = builder.fanout_failure_policy(#variant);));
}
scalar_step!(definition_revision, definition_revision);
Ok(steps)
}
struct InitializerArgs {
name: LitStr,
order: Option<i32>,
requires: Vec<LitStr>,
kind: InitializerKindArg,
factory: Option<Path>,
}
enum InitializerKindArg {
OneShot,
Hosted,
}
fn expand_initializer(
metas: Punctuated<Meta, Token![,]>,
item_impl: ItemImpl,
) -> syn::Result<TokenStream2> {
verify_initializer_impl(&item_impl)?;
let args = parse_initializer_args(&metas, &item_impl.self_ty)?;
let runtime = runtime_root("application", "napp")
.map_err(|message| syn::Error::new_spanned(&item_impl.self_ty, message))?;
let initializer_type = (*item_impl.self_ty).clone();
let name = args.name;
let order = args
.order
.map(|value| quote!(#value))
.unwrap_or_else(|| quote!(#runtime::DEFAULT_INITIALIZER_ORDER));
let requires = args.requires;
let kind = match args.kind {
InitializerKindArg::OneShot => quote!(#runtime::InitializerKind::OneShot),
InitializerKindArg::Hosted => quote!(#runtime::InitializerKind::Hosted),
};
let construct = match args.factory {
Some(factory) => quote! {
let result: #runtime::ApplicationResult<::std::option::Option<#initializer_type>> =
#factory(application).await;
let initializer = result?;
::std::result::Result::Ok(initializer.map(|value| {
::std::boxed::Box::new(value)
as ::std::boxed::Box<dyn #runtime::Initialization>
}))
},
None => quote! {
let _ = application;
let value: #initializer_type =
<#initializer_type as ::std::default::Default>::default();
::std::result::Result::Ok(::std::option::Option::Some(
::std::boxed::Box::new(value)
as ::std::boxed::Box<dyn #runtime::Initialization>
))
},
};
Ok(quote! {
#item_impl
const _: () = {
fn __nasa_initializer_factory(
application: #runtime::Application,
) -> #runtime::ApplicationFuture<
'static,
::std::option::Option<
::std::boxed::Box<dyn #runtime::Initialization>
>,
> {
::std::boxed::Box::pin(async move { #construct })
}
#[#runtime::__private::linkme::distributed_slice(#runtime::COLLECTED_INITIALIZERS)]
#[linkme(crate = #runtime::__private::linkme)]
static __NASA_INITIALIZER_DESCRIPTOR: #runtime::InitializerDescriptor =
#runtime::InitializerDescriptor::__new(
#name,
#order,
&[#(#requires),*],
#kind,
__nasa_initializer_factory,
concat!(module_path!(), ":", file!(), ":", line!()),
);
};
})
}
fn parse_initializer_args(
metas: &Punctuated<Meta, Token![,]>,
self_type: &Type,
) -> syn::Result<InitializerArgs> {
let mut name = None;
let mut order = None;
let mut requires = None;
let mut kind = None;
let mut factory = None;
for meta in metas {
let Meta::NameValue(value) = meta else {
return Err(syn::Error::new_spanned(
meta,
"initializer attributes must use `key = value` syntax",
));
};
let key = value
.path
.get_ident()
.map(ToString::to_string)
.ok_or_else(|| {
syn::Error::new_spanned(
&value.path,
"initializer attribute key must be an identifier",
)
})?;
match key.as_str() {
"name" => set_once(&mut name, parse_string_expr(&value.value, "name")?, meta)?,
"order" => {
let parsed = parse_initializer_order(&value.value)?;
set_once(&mut order, parsed, meta)?;
}
"requires" => {
let Expr::Array(array) = &value.value else {
return Err(syn::Error::new_spanned(
&value.value,
"initializer requires must be an array of string literals",
));
};
let mut parsed = Vec::with_capacity(array.elems.len());
for element in &array.elems {
parsed.push(parse_string_expr(element, "requires entry")?);
}
set_once(&mut requires, parsed, meta)?;
}
"kind" => {
let literal = parse_string_expr(&value.value, "kind")?;
let parsed = match literal.value().as_str() {
"one-shot" => InitializerKindArg::OneShot,
"hosted" => InitializerKindArg::Hosted,
_ => {
return Err(syn::Error::new_spanned(
literal,
"initializer kind must be `one-shot` or `hosted`",
));
}
};
set_once(&mut kind, parsed, meta)?;
}
"factory" => {
let Expr::Path(path) = &value.value else {
return Err(syn::Error::new_spanned(
&value.value,
"initializer factory must be a function path",
));
};
set_once(&mut factory, path.path.clone(), meta)?;
}
_ => {
return Err(syn::Error::new_spanned(
&value.path,
"unknown initializer attribute key",
));
}
}
}
let name = match name {
Some(name) => name,
None => default_initializer_name(self_type)?,
};
validate_initializer_name(&name, "initializer name")?;
let requires = requires.unwrap_or_default();
if requires.len() > 32 {
return Err(syn::Error::new_spanned(
&name,
"initializer requires cannot contain more than 32 entries",
));
}
let mut seen = HashSet::new();
for required in &requires {
validate_initializer_name(required, "initializer dependency")?;
if required.value() == name.value() {
return Err(syn::Error::new_spanned(
required,
"initializer cannot require itself",
));
}
if !seen.insert(required.value()) {
return Err(syn::Error::new_spanned(
required,
"initializer dependency is repeated",
));
}
}
Ok(InitializerArgs {
name,
order,
requires,
kind: kind.unwrap_or(InitializerKindArg::OneShot),
factory,
})
}
fn default_initializer_name(self_type: &Type) -> syn::Result<LitStr> {
let Type::Path(path) = self_type else {
return Err(syn::Error::new_spanned(
self_type,
"initializer name cannot be derived from this type; declare `name` explicitly",
));
};
if path.qself.is_some() {
return Err(syn::Error::new_spanned(
self_type,
"initializer name cannot be derived from a qualified self type; declare `name` explicitly",
));
}
let segment = path.path.segments.last().ok_or_else(|| {
syn::Error::new_spanned(
self_type,
"initializer name cannot be derived from this type; declare `name` explicitly",
)
})?;
let identifier = segment.ident.to_string();
let name = canonicalize_type_name(&identifier).ok_or_else(|| {
syn::Error::new_spanned(
&segment.ident,
"initializer type name cannot form a canonical name; declare `name` explicitly",
)
})?;
let name = LitStr::new(&name, segment.ident.span());
validate_initializer_name(&name, "derived initializer name")?;
Ok(name)
}
fn canonicalize_type_name(identifier: &str) -> Option<String> {
let identifier = identifier.strip_prefix("r#").unwrap_or(identifier);
let bytes = identifier.as_bytes();
let mut output = String::with_capacity(bytes.len());
let mut pending_separator = false;
for (index, byte) in bytes.iter().copied().enumerate() {
if byte == b'_' {
pending_separator = !output.is_empty();
continue;
}
if !byte.is_ascii_alphanumeric() {
return None;
}
let previous = index
.checked_sub(1)
.and_then(|value| bytes.get(value))
.copied();
let next = bytes.get(index + 1).copied();
let word_boundary = byte.is_ascii_uppercase()
&& (previous.is_some_and(|value| value.is_ascii_lowercase() || value.is_ascii_digit())
|| (previous.is_some_and(|value| value.is_ascii_uppercase())
&& next.is_some_and(|value| value.is_ascii_lowercase())));
if (pending_separator || word_boundary) && !output.is_empty() && !output.ends_with('-') {
output.push('-');
}
output.push(byte.to_ascii_lowercase() as char);
pending_separator = false;
}
while output.ends_with('-') {
output.pop();
}
(!output.is_empty()).then_some(output)
}
fn parse_initializer_order(expression: &Expr) -> syn::Result<i32> {
let invalid = || {
syn::Error::new_spanned(
expression,
"initializer order must be an i32 integer literal",
)
};
let signed = match expression {
Expr::Lit(expr) => match &expr.lit {
Lit::Int(value) => value.base10_parse::<i64>().map_err(|_| invalid())?,
_ => return Err(invalid()),
},
Expr::Unary(expr) if matches!(expr.op, syn::UnOp::Neg(_)) => match expr.expr.as_ref() {
Expr::Lit(expr) => match &expr.lit {
Lit::Int(value) => value
.base10_parse::<i64>()
.ok()
.and_then(i64::checked_neg)
.ok_or_else(invalid)?,
_ => return Err(invalid()),
},
_ => return Err(invalid()),
},
_ => return Err(invalid()),
};
i32::try_from(signed).map_err(|_| invalid())
}
fn verify_initializer_impl(item_impl: &ItemImpl) -> syn::Result<()> {
if item_impl.unsafety.is_some() {
return Err(syn::Error::new_spanned(
item_impl.unsafety,
"initializer cannot annotate an unsafe impl",
));
}
if !item_impl.generics.params.is_empty() {
return Err(syn::Error::new_spanned(
&item_impl.generics,
"initializer impl cannot declare generic parameters",
));
}
let Some((polarity, trait_path, _)) = &item_impl.trait_ else {
return Err(syn::Error::new_spanned(
&item_impl.self_ty,
"initializer must annotate an Initialization trait impl",
));
};
if polarity.is_some() {
return Err(syn::Error::new_spanned(
polarity,
"initializer cannot annotate a negative impl",
));
}
if trait_path
.segments
.last()
.is_none_or(|segment| segment.ident != "Initialization")
{
return Err(syn::Error::new_spanned(
trait_path,
"initializer trait path must end with Initialization",
));
}
Ok(())
}
fn parse_string_expr(expression: &Expr, field: &str) -> syn::Result<LitStr> {
match expression {
Expr::Lit(expr) => match &expr.lit {
Lit::Str(value) => Ok(value.clone()),
_ => Err(syn::Error::new_spanned(
expression,
format!("initializer {field} must be a string literal"),
)),
},
_ => Err(syn::Error::new_spanned(
expression,
format!("initializer {field} must be a string literal"),
)),
}
}
fn set_once<T>(slot: &mut Option<T>, value: T, meta: &Meta) -> syn::Result<()> {
if slot.is_some() {
return Err(syn::Error::new_spanned(
meta,
"initializer attribute key is repeated",
));
}
*slot = Some(value);
Ok(())
}
fn validate_initializer_name(name: &LitStr, field: &str) -> syn::Result<()> {
let value = name.value();
if value.is_empty() || value.len() > 128 {
return Err(syn::Error::new_spanned(
name,
format!("{field} must contain between 1 and 128 bytes"),
));
}
if !value.bytes().all(|byte| {
byte.is_ascii_lowercase() || byte.is_ascii_digit() || matches!(byte, b'_' | b'-' | b'.')
}) {
return Err(syn::Error::new_spanned(
name,
format!("{field} must contain only lowercase ASCII letters, digits, `_`, `-`, or `.`"),
));
}
Ok(())
}
fn expand_application(
components: Vec<LitStr>,
mut function: ItemFn,
) -> syn::Result<proc_macro2::TokenStream> {
validate_function(&function)?;
let component_names = validate_components(&components)?;
let runtime = runtime_root("application", "napp")
.map_err(|message| syn::Error::new_spanned(&function.sig.ident, message))?;
let has_web = component_names.iter().any(|name| name == "web");
let component_variants = component_names
.iter()
.map(|name| component_variant(name))
.collect::<syn::Result<Vec<_>>>()?;
let feature_modules = component_names
.iter()
.map(|name| component_feature_module(name))
.collect::<syn::Result<Vec<_>>>()?;
let accepts_application = function.sig.inputs.len() == 1;
function.sig.ident = format_ident!("__nasa_user_main");
let hook = if accepts_application {
quote!(|application| __nasa_user_main(application))
} else {
quote!(|_application| __nasa_user_main())
};
let web_items = if has_web {
quote! {
#runtime::__private::naweb::mvc_router!(#runtime::Application);
fn __nasa_route_meta() -> ::std::vec::Vec<#runtime::RouteMeta> {
crate::__mvc::ROUTES
.iter()
.map(|entry| #runtime::RouteMeta {
method: entry.method,
path: entry.path,
handler: entry.handler,
produces: entry.produces,
consumes: entry.consumes,
request_schema: entry.request_schema,
response_schema: entry.response_schema,
query_parameters: entry.query_parameters,
header_parameters: entry.header_parameters,
success_status: entry.success_status,
additional_responses: entry.additional_responses,
streaming: entry.streaming,
auth_required: ::core::matches!(
entry.policy.auth,
#runtime::__private::naweb::AuthRequirement::Required
),
})
.collect()
}
fn __nasa_build_router(
context: #runtime::WebBuildContext,
) -> #runtime::ApplicationResult<
#runtime::__private::axum::Router<#runtime::Application>,
> {
context.build(|router, mapping_runtime, mapping_plan, application| {
crate::__mvc::try_register_all(
router,
mapping_runtime,
mapping_plan,
application,
)
})
}
}
} else {
quote! {}
};
let spec_web = if has_web {
quote!(
.with_web_route_meta(__nasa_route_meta)
.with_web_factory(__nasa_build_router)
)
} else {
quote! {}
};
Ok(quote! {
#[doc(hidden)]
pub mod __nasa_application_must_be_at_crate_root {}
use crate::__nasa_application_must_be_at_crate_root as _;
#(
const _: () = #runtime::components::#feature_modules::FEATURE_CHECK;
)*
#web_items
#function
fn __nasa_require_user_hook<F, Fut, E>(hook: F) -> F
where
F: ::std::ops::FnOnce(#runtime::Application) -> Fut + ::std::marker::Send + 'static,
Fut: ::std::future::Future<Output = ::std::result::Result<(), E>>
+ ::std::marker::Send
+ 'static,
E: ::std::convert::Into<#runtime::__private::anyhow::Error> + 'static,
{
hook
}
fn main() -> ::std::process::ExitCode {
#runtime::run(
#runtime::ApplicationSpec::new(&[
#(#runtime::ComponentId::#component_variants),*
])
.with_default_name(env!("CARGO_PKG_NAME"))
#spec_web,
__nasa_require_user_hook(#hook),
)
}
})
}
fn validate_function(function: &ItemFn) -> syn::Result<()> {
if function.sig.ident != "main" {
return Err(syn::Error::new_spanned(
&function.sig.ident,
"application attribute must be attached to the crate main function",
));
}
if function.sig.asyncness.is_none() {
return Err(syn::Error::new_spanned(
function.sig.fn_token,
"application main must be async and must not use another runtime entry attribute",
));
}
if !function.sig.generics.params.is_empty() {
return Err(syn::Error::new_spanned(
&function.sig.generics,
"application main cannot declare generics",
));
}
if function.sig.inputs.len() > 1 {
return Err(syn::Error::new_spanned(
&function.sig.inputs,
"application main accepts at most one Application parameter",
));
}
if let Some(argument) = function.sig.inputs.first() {
validate_application_parameter(argument)?;
}
validate_return_type(&function.sig.output)?;
for attribute in &function.attrs {
let segments = attribute
.path()
.segments
.iter()
.map(|segment| segment.ident.to_string())
.collect::<Vec<_>>();
if segments.first().is_some_and(|name| name == "tokio")
&& segments.last().is_some_and(|name| name == "main")
{
return Err(syn::Error::new_spanned(
attribute,
"remove the other runtime entry attribute because application owns the runtime",
));
}
if segments
.last()
.is_some_and(|name| matches!(name.as_str(), "EnableScheduling" | "EnableAsync"))
{
return Err(syn::Error::new_spanned(
attribute,
"declare the scheduling component in application instead of using an entry attribute",
));
}
}
Ok(())
}
fn validate_application_parameter(argument: &FnArg) -> syn::Result<()> {
let FnArg::Typed(argument) = argument else {
return Err(syn::Error::new_spanned(
argument,
"application main cannot use a receiver parameter",
));
};
let Type::Path(path) = argument.ty.as_ref() else {
return Err(syn::Error::new_spanned(
&argument.ty,
"application main parameter must be Application",
));
};
if path
.path
.segments
.last()
.is_none_or(|segment| segment.ident != "Application")
{
return Err(syn::Error::new_spanned(
&argument.ty,
"application main parameter must be Application",
));
}
Ok(())
}
fn validate_return_type(output: &ReturnType) -> syn::Result<()> {
let ReturnType::Type(_, output_type) = output else {
return Err(syn::Error::new_spanned(
output,
"application main must return anyhow::Result<()>",
));
};
let Type::Path(path) = output_type.as_ref() else {
return Err(syn::Error::new_spanned(
output_type,
"application main must return anyhow::Result<()>",
));
};
let Some(result) = path.path.segments.last() else {
return Err(syn::Error::new_spanned(
output_type,
"application main must return anyhow::Result<()>",
));
};
let PathArguments::AngleBracketed(arguments) = &result.arguments else {
return Err(syn::Error::new_spanned(
output_type,
"application main must return anyhow::Result<()>",
));
};
let unit_success = matches!(
arguments.args.first(),
Some(GenericArgument::Type(Type::Tuple(tuple))) if tuple.elems.is_empty()
);
if result.ident != "Result" || arguments.args.len() != 1 || !unit_success {
return Err(syn::Error::new_spanned(
output_type,
"application main must return anyhow::Result<()>",
));
}
Ok(())
}
const CANONICAL_COMPONENT_ORDER: [&str; 17] = [
"log",
"nacos-config",
"telemetry",
"db",
"redis",
"cache",
"partition",
"saga",
"kafka",
"outbox",
"redis-job",
"grpc",
"auth",
"web",
"ws",
"nacos-discovery",
"scheduling",
];
fn validate_components(components: &[LitStr]) -> syn::Result<Vec<String>> {
let mut seen = HashSet::new();
let mut names = Vec::with_capacity(components.len());
for component in components.iter() {
let name = component.value();
if !CANONICAL_COMPONENT_ORDER.contains(&name.as_str()) {
return Err(syn::Error::new_spanned(
component,
format!("unknown application component `{name}`"),
));
}
if !seen.insert(name.clone()) {
return Err(syn::Error::new_spanned(
component,
format!("application component `{name}` is declared more than once"),
));
}
names.push(name);
}
if seen.contains("saga") && seen.insert("outbox".to_string()) {
names.push("outbox".to_string());
}
if seen.contains("outbox") && seen.insert("db".to_string()) {
names.push("db".to_string());
}
if seen.contains("redis-job") && seen.insert("redis".to_string()) {
names.push("redis".to_string());
}
names.sort_by_key(|name| {
CANONICAL_COMPONENT_ORDER
.iter()
.position(|canonical| canonical == name)
.expect("name validated against CANONICAL_COMPONENT_ORDER above")
});
Ok(names)
}
fn component_variant(name: &str) -> syn::Result<syn::Ident> {
let variant = match name {
"log" => "Log",
"nacos-config" => "NacosConfig",
"db" => "Db",
"redis" => "Redis",
"redis-job" => "RedisJob",
"telemetry" => "Telemetry",
"cache" => "Cache",
"partition" => "Partition",
"grpc" => "Grpc",
"saga" => "Saga",
"kafka" => "Kafka",
"outbox" => "Outbox",
"auth" => "Auth",
"web" => "Web",
"ws" => "Ws",
"nacos-discovery" => "NacosDiscovery",
"scheduling" => "Scheduling",
_ => {
return Err(syn::Error::new(
proc_macro2::Span::call_site(),
"component name was not validated",
));
}
};
Ok(format_ident!("{variant}"))
}
fn component_feature_module(name: &str) -> syn::Result<syn::Ident> {
match name {
"log" | "db" | "redis" | "telemetry" | "cache" | "partition" | "grpc" | "saga"
| "kafka" | "outbox" | "auth" | "web" | "ws" | "scheduling" => Ok(format_ident!("{name}")),
"redis-job" => Ok(format_ident!("redis_job")),
"nacos-config" => Ok(format_ident!("nacos_config")),
"nacos-discovery" => Ok(format_ident!("nacos_discovery")),
_ => Err(syn::Error::new(
proc_macro2::Span::call_site(),
"component name was not validated",
)),
}
}