use crate::settings_schema::{self, SettingAttr};
use proc_macro2::TokenStream;
use quote::{format_ident, quote};
use syn::{ItemStruct, LitStr, Result};
pub(crate) fn settings_fragment_impl(args: TokenStream, input: ItemStruct) -> Result<TokenStream> {
let conf_crate = crate::crate_paths::get_reinhardt_conf_crate();
let mut section: Option<String> = None;
let mut default_policy: Option<String> = None;
let mut generate_validation: Option<bool> = None;
let parser = syn::meta::parser(|meta| {
if meta.path.is_ident("fragment") {
let _: syn::LitBool = meta.value()?.parse()?;
Ok(())
} else if meta.path.is_ident("section") {
let lit: LitStr = meta.value()?.parse()?;
section = Some(lit.value());
Ok(())
} else if meta.path.is_ident("default_policy") {
let lit: LitStr = meta.value()?.parse()?;
let val = lit.value();
if val != "required" && val != "optional" {
return Err(syn::Error::new(
lit.span(),
"invalid `default_policy` value, expected `\"required\"` or `\"optional\"`",
));
}
default_policy = Some(val);
Ok(())
} else if meta.path.is_ident("validate") {
let lit: syn::LitBool = meta.value()?.parse()?;
generate_validation = Some(lit.value());
Ok(())
} else {
Err(meta.error(
"expected `fragment = true`, `section = \"...\"`, `default_policy = \"...\"`, or `validate = true|false`",
))
}
});
syn::parse::Parser::parse2(parser, args)?;
let parsed_fields = settings_schema::parse_fields(&input)?;
let section = section.unwrap_or_else(|| {
settings_schema::infer_type_key(&input.ident.to_string())
.unwrap_or_else(|_| settings_schema::camel_to_snake(&input.ident.to_string()))
});
let default_policy_is_required = default_policy.as_deref() == Some("required");
let should_generate_validation = generate_validation.unwrap_or(true);
let struct_name = &input.ident;
let vis = &input.vis;
let trait_name = format_ident!("Has{}", struct_name);
let method_name: syn::Ident = syn::parse_str(§ion).map_err(|_| {
syn::Error::new(
struct_name.span(),
format!(
"`section` must be a valid Rust method identifier for the generated settings accessor; got `{section}`"
),
)
})?;
let has_derive = input.attrs.iter().any(|a| a.path().is_ident("derive"));
let derive_attr = if has_derive {
quote! {}
} else {
quote! { #[derive(Clone, Debug, serde::Serialize, serde::Deserialize)] }
};
let attrs = &input.attrs;
let semi_token = &input.semi_token;
let mut field_policy_entries = Vec::new();
let mut node_field_schema_entries = Vec::new();
let mut default_fn_defs = Vec::new();
let mut new_fields = Vec::new();
for field in &parsed_fields {
let field_name = &field.ident;
let field_name_str = &field.rust_name;
let field_key_str = &field.key;
let setting_attr = &field.setting_attr;
let already_has_serde_default = field.has_serde_default;
let cfg_attrs = &field.cfg_attrs;
let (requirement_tokens, has_default, serde_default_tokens) = match setting_attr {
Some(SettingAttr::Required) => (
quote! { #conf_crate::settings::policy::FieldRequirement::Required },
false,
quote! {},
),
Some(SettingAttr::Optional) => {
let serde_tokens = if already_has_serde_default {
quote! {}
} else {
quote! { #[serde(default)] }
};
(
quote! { #conf_crate::settings::policy::FieldRequirement::Optional },
true,
serde_tokens,
)
}
Some(SettingAttr::Default(expr)) => {
let fn_name = format_ident!("__default_{}_{}", struct_name, field_name);
let field_ty = &field.ty;
let expr_tokens: TokenStream = expr.parse().map_err(|e| {
syn::Error::new(
field.ident.span(),
format!("invalid default expression: {e}"),
)
})?;
default_fn_defs.push(quote! {
fn #fn_name() -> #field_ty {
#expr_tokens
}
});
let fn_name_str = fn_name.to_string();
let serde_tokens = if already_has_serde_default {
quote! {}
} else {
quote! { #[serde(default = #fn_name_str)] }
};
(
quote! { #conf_crate::settings::policy::FieldRequirement::Optional },
true,
serde_tokens,
)
}
None => {
if default_policy_is_required {
(
quote! { #conf_crate::settings::policy::FieldRequirement::Required },
false,
quote! {},
)
} else {
let serde_tokens = if already_has_serde_default {
quote! {}
} else {
quote! { #[serde(default)] }
};
(
quote! { #conf_crate::settings::policy::FieldRequirement::Optional },
true,
serde_tokens,
)
}
}
};
let value_schema = settings_schema::value_schema_tokens(&field.shape, &conf_crate);
field_policy_entries.push(quote! {
#(#cfg_attrs)*
#conf_crate::settings::policy::FieldPolicy {
name: #field_name_str,
requirement: #requirement_tokens,
has_default: #has_default,
}
});
node_field_schema_entries.push(quote! {
#(#cfg_attrs)*
#conf_crate::settings::schema::SettingsFieldSchema {
rust_name: #field_name_str,
key: #field_key_str,
policy: #conf_crate::settings::policy::FieldPolicy {
name: #field_name_str,
requirement: #requirement_tokens,
has_default: #has_default,
},
value: #value_schema,
}
});
let cleaned_attrs = &field.cleaned_attrs;
let field_vis = &field.vis;
let field_ty = &field.ty;
new_fields.push(quote! {
#(#cleaned_attrs)*
#serde_default_tokens
#field_vis #field_name: #field_ty
});
}
let schema_name = settings_schema::schema_type_name(struct_name);
let schema_fields = settings_schema::schema_struct_fields(&parsed_fields, &conf_crate);
let schema_inits = settings_schema::schema_struct_inits(&parsed_fields, &conf_crate);
let schema_root_marker_field = if schema_fields.is_empty() {
quote! {
__root: ::std::marker::PhantomData<fn() -> Root>,
}
} else {
quote! {}
};
let schema_root_marker_init = if schema_fields.is_empty() {
quote! {
__root: ::std::marker::PhantomData,
}
} else {
quote! {}
};
let struct_body = if semi_token.is_some() {
quote! { ; }
} else {
quote! {
{
#(#new_fields),*
}
}
};
let validation_impl = if should_generate_validation {
quote! {
impl #conf_crate::settings::fragment::SettingsValidation for #struct_name {}
}
} else {
quote! {}
};
let validate_override = if !should_generate_validation {
quote! {
fn validate(
&self,
profile: &#conf_crate::settings::profile::Profile,
) -> #conf_crate::settings::validation::ValidationResult {
<Self as #conf_crate::settings::fragment::SettingsValidation>::validate(self, profile)
}
}
} else {
quote! {}
};
Ok(quote! {
#derive_attr
#(#attrs)*
#vis struct #struct_name #struct_body
#(#default_fn_defs)*
#validation_impl
#[doc = "Typed schema references for this settings fragment."]
#[derive(Clone, Debug)]
#vis struct #schema_name<Root> {
__path: #conf_crate::settings::schema::SettingsPathBuf,
#schema_root_marker_field
#(#schema_fields,)*
}
impl<Root> #schema_name<Root> {
fn __from_path(path: #conf_crate::settings::schema::SettingsPathBuf) -> Self {
Self {
__path: path.clone(),
#schema_root_marker_init
#(#schema_inits,)*
}
}
#[must_use]
#[doc = "Return secret field references reachable from this settings fragment."]
pub fn secret_fields(&self) -> ::std::vec::Vec<#conf_crate::settings::schema::SecretFieldRef<Root, ()>> {
let mut paths = ::std::vec::Vec::new();
<#struct_name as #conf_crate::settings::schema::SettingsNode>::node_schema()
.collect_secret_paths(&mut paths);
paths
.into_iter()
.map(|path| #conf_crate::settings::schema::SecretFieldRef::<Root, ()>::new(self.__path.clone().extend(path)))
.collect()
}
}
impl #conf_crate::settings::schema::SettingsNode for #struct_name {
type Schema<Root> = #schema_name<Root>;
fn schema_at<Root>(
path: #conf_crate::settings::schema::SettingsPathBuf,
) -> Self::Schema<Root> {
#schema_name::__from_path(path)
}
fn node_schema() -> #conf_crate::settings::schema::SettingsNodeSchema {
#conf_crate::settings::schema::SettingsNodeSchema {
type_name: ::std::any::type_name::<#struct_name>(),
fields: ::std::vec![#(#node_field_schema_entries),*],
}
}
}
impl #conf_crate::settings::fragment::SettingsFragment for #struct_name {
type Accessor = dyn #trait_name;
fn section() -> &'static str {
#section
}
#validate_override
fn field_policies() -> &'static [#conf_crate::settings::policy::FieldPolicy] {
static POLICIES: &[#conf_crate::settings::policy::FieldPolicy] = &[
#(#field_policy_entries),*
];
POLICIES
}
}
impl #conf_crate::settings::fragment::HasSettings<#struct_name> for #struct_name {
fn get_settings(&self) -> &#struct_name {
self
}
}
#vis trait #trait_name {
fn #method_name(&self) -> &#struct_name;
}
impl<T: #conf_crate::settings::fragment::HasSettings<#struct_name>> #trait_name for T {
fn #method_name(&self) -> &#struct_name {
self.get_settings()
}
}
})
}