use std::collections::HashSet;
use nasa_macro_support::runtime_root;
use proc_macro::TokenStream;
use quote::{format_ident, quote};
use syn::{
parse_macro_input, punctuated::Punctuated, FnArg, GenericArgument, ItemFn, LitStr,
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(),
}
}
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; 12] = [
"log",
"nacos-config",
"telemetry",
"db",
"redis",
"cache",
"kafka",
"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);
}
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",
"telemetry" => "Telemetry",
"cache" => "Cache",
"kafka" => "Kafka",
"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" | "kafka" | "auth" | "web" | "ws"
| "scheduling" => Ok(format_ident!("{name}")),
"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",
)),
}
}