use std::ops::ControlFlow;
use darling::{
FromAttributes,
util::{Override, SpannedValue},
};
use indexmap::IndexMap;
use itertools::Itertools as _;
use proc_macro2::{Delimiter, Group, Span, TokenStream, TokenTree};
use quote::ToTokens;
use syn::{Expr, Ident, Token, Type, Variant, punctuated::Punctuated, spanned::Spanned};
use crate::common::{
Description, FieldDefault, FlagTags, IdentString, compute_docs, compute_placeholder,
compute_tags, error_pair,
};
use super::create_non_colliding_ident;
#[derive(darling::FromAttributes, Debug)]
#[darling(attributes(debate))]
struct RawParsedFlagSetFlagAttr {
long: Option<SpannedValue<Override<String>>>,
short: Option<SpannedValue<Override<char>>>,
default: Option<SpannedValue<Override<Expr>>>,
placeholder: Option<SpannedValue<String>>,
#[darling(rename = "override")]
overridable: Option<SpannedValue<()>>,
}
#[derive(Clone, Copy)]
pub enum Family<Info> {
Solitary(Info),
Sibling,
}
impl<Info> Family<Info> {
pub fn merge(&mut self, other: Family<()>, info: Info) {
if matches!(other, Solitary(())) {
*self = Solitary(info)
}
}
}
use Family::*;
pub struct FlagSetFlagInfo<Info> {
pub docs: Description,
pub placeholder: SpannedValue<String>,
pub info: Info,
pub default: Option<FieldDefault>,
pub tags: FlagTags<SpannedValue<String>, SpannedValue<char>>,
pub overridable: Option<SpannedValue<()>>,
}
pub trait FlagSetFlagExtra<'a> {
fn as_flag_set_type(&self) -> FlagType<'a>;
fn field(&'a self) -> Option<&'a IdentString<'a>>;
fn site(&self) -> FlagSetFlagSite<'_>;
}
pub struct FlagFieldInfo<'a> {
pub ty: &'a Type,
pub ident: IdentString<'a>,
}
impl<'a> FlagSetFlagExtra<'a> for FlagFieldInfo<'a> {
fn as_flag_set_type(&self) -> FlagType<'a> {
FlagType::Typed(self.ty)
}
fn field(&'a self) -> Option<&'a IdentString<'a>> {
Some(&self.ident)
}
fn site(&self) -> FlagSetFlagSite<'_> {
FlagSetFlagSite::Field(&self.ident)
}
}
pub enum PlainFlagInfo<'a> {
Unit,
Newtype(&'a Type),
Struct(FlagFieldInfo<'a>),
}
impl<'a> PlainFlagInfo<'a> {
pub fn ty(&self) -> FlagType<'a> {
match *self {
PlainFlagInfo::Unit => FlagType::Unit,
PlainFlagInfo::Newtype(ty) => FlagType::Typed(ty),
PlainFlagInfo::Struct(ref info) => FlagType::Typed(info.ty),
}
}
}
impl<'a> FlagSetFlagExtra<'a> for PlainFlagInfo<'a> {
fn as_flag_set_type(&self) -> FlagType<'a> {
self.ty()
}
fn field(&self) -> Option<&IdentString<'a>> {
match *self {
Self::Unit => None,
Self::Newtype(_) => None,
Self::Struct(ref info) => Some(&info.ident),
}
}
fn site(&self) -> FlagSetFlagSite<'_> {
match *self {
Self::Unit => FlagSetFlagSite::Unit,
Self::Newtype(_) => FlagSetFlagSite::Newtype,
Self::Struct(ref field) => FlagSetFlagSite::Field(&field.ident),
}
}
}
#[expect(clippy::large_enum_variant)]
pub enum FlagSetVariant<'a> {
Plain(FlagSetFlagInfo<PlainFlagInfo<'a>>),
Struct(Vec<FlagSetFlagInfo<FlagFieldInfo<'a>>>),
}
#[derive(Clone, Copy)]
pub enum FlagType<'a> {
Unit,
Typed(&'a Type),
}
impl ToTokens for FlagType<'_> {
fn to_tokens(&self, tokens: &mut proc_macro2::TokenStream) {
match *self {
FlagType::Typed(ty) => ty.to_tokens(tokens),
FlagType::Unit => tokens.extend([TokenTree::Group(Group::new(
Delimiter::Parenthesis,
TokenStream::new(),
))]),
}
}
}
pub struct ParsedFlagSetInfo<'a> {
pub superposition: Ident,
pub variants: IndexMap<IdentString<'a>, FlagSetVariant<'a>>,
}
fn compute_flag_set_tags(
auto_long: Option<&Span>,
auto_short: Option<&Span>,
long: Option<SpannedValue<Override<String>>>,
short: Option<SpannedValue<Override<char>>>,
ident: &IdentString,
) -> syn::Result<FlagTags<SpannedValue<String>, SpannedValue<char>>> {
let long = long.or_else(|| auto_long.map(|long| SpannedValue::new(Override::Inherit, *long)));
let short =
short.or_else(|| auto_short.map(|short| SpannedValue::new(Override::Inherit, *short)));
compute_tags(long, short, ident)?.ok_or_else(|| {
syn::Error::new(
ident.span(),
"flag must have a short or long name. Add #[debate(short)] or \
#[debate(long)] to the top-level enum to automatically add it to \
all flags.",
)
})
}
macro_rules! create_flag {
(
$auto_long:ident,
$auto_short:ident,
source: $source:expr,
info: $info:expr,
attrs: $attrs:expr,
) => {
RawParsedFlagSetFlagAttr::from_attributes($attrs).and_then(|attr| {
Ok(FlagSetFlagInfo {
docs: compute_docs($attrs)?,
placeholder: compute_placeholder(attr.placeholder, $source)?,
default: FieldDefault::new(attr.default),
tags: compute_flag_set_tags(
$auto_long,
$auto_short,
attr.long,
attr.short,
$source,
)?,
overridable: attr.overridable,
info: $info,
})
})
};
}
impl<'a> ParsedFlagSetInfo<'a> {
pub fn from_variants(
variants: &'a Punctuated<Variant, Token![,]>,
auto_long: Option<&Span>,
auto_short: Option<&Span>,
) -> syn::Result<Self> {
if variants.is_empty() {
return Err(syn::Error::new(
variants.span(),
"must have at least one variant",
));
}
let variants: IndexMap<IdentString<'a>, FlagSetVariant<'a>> = variants
.into_iter()
.map(|variant| {
let variant_ident = IdentString::new(&variant.ident);
let mode = match variant.fields {
syn::Fields::Unit => create_flag!(
auto_long,
auto_short,
source: &variant_ident,
info: PlainFlagInfo::Unit,
attrs: &variant.attrs,
)
.map(FlagSetVariant::Plain),
syn::Fields::Unnamed(ref fields) => create_flag!(
auto_long,
auto_short,
source: &variant_ident,
info: fields
.unnamed
.iter()
.exactly_one()
.map(|field| &field.ty)
.map(PlainFlagInfo::Newtype)
.map_err(|_| {
syn::Error::new(
fields.span(),
"must have exactly one field (the flag argument)",
)
})?,
attrs: &variant.attrs,
)
.map(FlagSetVariant::Plain),
syn::Fields::Named(ref fields) => match fields
.named
.iter()
.map(|field| {
(
field
.ident
.as_ref()
.map(IdentString::new)
.expect("all fields in a struct variant have names"),
field,
)
})
.exactly_one()
{
Ok((ident, field)) => create_flag!(
auto_long,
auto_short,
source: &ident,
info: PlainFlagInfo::Struct(FlagFieldInfo {
ty: &field.ty, ident,
}),
attrs: &field.attrs,
)
.map(FlagSetVariant::Plain),
Err(fields) => fields
.map(|(ident, field)| {
create_flag!(
auto_long,
auto_short,
source: &ident,
info: FlagFieldInfo {
ident, ty: &field.ty,
},
attrs: &field.attrs,
)
})
.try_collect()
.map(FlagSetVariant::Struct),
},
};
mode.map(|mode| (variant_ident, mode))
})
.try_collect()?;
let superposition = create_non_colliding_ident("Superposition", variants.keys());
Ok(Self {
superposition,
variants,
})
}
}
pub enum FlagSetFlagSite<'a> {
Unit,
Newtype,
Field(&'a IdentString<'a>),
}
impl<'a> FlagSetFlagSite<'a> {
pub fn field(&self) -> Option<&'a IdentString<'a>> {
match self {
Self::Field(ident) => Some(ident),
Self::Unit | Self::Newtype => None,
}
}
}
pub struct FlagSetFlagVariant<'a> {
pub ident: &'a IdentString<'a>,
pub index: usize,
pub field_index: usize,
pub site: FlagSetFlagSite<'a>,
}
pub struct FlagSetFlag<'a> {
pub origin: &'a IdentString<'a>,
pub variants: Vec<FlagSetFlagVariant<'a>>,
pub docs: &'a Description,
pub placeholder: &'a str,
pub ty: FlagType<'a>,
pub tags: FlagTags<&'a str, char>,
pub overridable: bool,
pub ever_required: bool,
pub family: Family<&'a IdentString<'a>>,
}
impl<'a> FlagSetFlag<'a> {
#[inline(always)]
#[must_use]
pub fn unique_variant(&self) -> Option<&FlagSetFlagVariant<'a>> {
self.variants.iter().exactly_one().ok()
}
}
enum TagComparison {
Match,
Different,
LongMismatch,
ShortMismatch,
}
fn compare_tags(left: &FlagTags<&str, char>, right: &FlagTags<&str, char>) -> TagComparison {
use FlagTags::*;
match (*left, *right) {
(Long(left), Long(right)) if left == right => TagComparison::Match,
(Short(left), Short(right)) if left == right => TagComparison::Match,
(
LongShort {
long: l1,
short: s1,
},
LongShort {
long: l2,
short: s2,
},
) if l1 == l2 && s1 == s2 => TagComparison::Match,
(Long(left), LongShort { long: right, .. })
| (LongShort { long: left, .. }, Long(right))
| (LongShort { long: left, .. }, LongShort { long: right, .. })
if left == right =>
{
TagComparison::ShortMismatch
}
(Short(left), LongShort { short: right, .. })
| (LongShort { short: left, .. }, Short(right))
| (LongShort { short: left, .. }, LongShort { short: right, .. })
if left == right =>
{
TagComparison::LongMismatch
}
_ => TagComparison::Different,
}
}
fn try_find<I: IntoIterator, E>(
iterator: I,
mut filter: impl FnMut(&I::Item) -> Result<bool, E>,
) -> Result<Option<I::Item>, E> {
iterator
.into_iter()
.try_for_each(|item| match filter(&item) {
Err(err) => ControlFlow::Break(Err(err)),
Ok(true) => ControlFlow::Break(Ok(item)),
Ok(false) => ControlFlow::Continue(()),
})
.break_value()
.transpose()
}
fn add_or_update_flag<'a, Info>(
set: &mut Vec<FlagSetFlag<'a>>,
variant_flag: &'a FlagSetFlagInfo<Info>,
family: Family<()>,
variant: &'a IdentString<'a>,
variant_index: usize,
field_index: usize,
) -> syn::Result<usize>
where
Info: FlagSetFlagExtra<'a>,
{
let origin = variant_flag.info.field().unwrap_or(variant);
let tags = variant_flag.tags.simplify();
let mismatch_error = |existing_span, message| {
error_pair(
origin.span(),
"multiple instances of the same tag must be (mostly) identical",
existing_span,
message,
)
};
let existing_flag = try_find(set.iter_mut().enumerate(), |(_, existing)| {
let mismatch_error = |message| mismatch_error(existing.origin.span(), message);
match compare_tags(&tags, &existing.tags) {
TagComparison::LongMismatch => {
Err(mismatch_error("this instance has a different --long"))
}
TagComparison::ShortMismatch => {
Err(mismatch_error("this instance has a different -s short"))
}
TagComparison::Different => Ok(false),
TagComparison::Match => {
if let (Solitary(()), Solitary(existing_variant)) = (family, existing.family) {
Err(error_pair(
variant.span(),
"this variant will never be selected (another variant has the same flags)",
existing_variant.span(),
"this variant's flag is identical",
))
} else if variant_flag.placeholder.as_str() != existing.placeholder {
Err(mismatch_error("this instance has a different placeholder"))
} else if variant_flag.overridable.is_some() != existing.overridable {
Err(mismatch_error(
"this instance has a different `override` setting",
))
} else {
Ok(true)
}
}
}
})?;
let (flag_index, existing_flag) = match existing_flag {
Some(flag) => flag,
None => {
let index = set.len();
set.push(FlagSetFlag {
origin,
variants: Vec::new(),
docs: const { &Description::empty() },
placeholder: &variant_flag.placeholder,
ty: variant_flag.info.as_flag_set_type(),
tags,
overridable: variant_flag.overridable.is_some(),
family: Family::Sibling,
ever_required: false,
});
(
index,
set.last_mut().expect("we just pushed an item into the set"),
)
}
};
existing_flag.variants.push(FlagSetFlagVariant {
ident: variant,
index: variant_index,
field_index,
site: variant_flag.info.site(),
});
existing_flag.family.merge(family, variant);
existing_flag.ever_required = existing_flag.ever_required || variant_flag.default.is_none();
if existing_flag.docs.full.len() > variant_flag.docs.full.len() {
existing_flag.docs = &variant_flag.docs;
};
Ok(flag_index)
}
pub struct VariantFieldSource<'a, Info> {
pub index: usize,
pub unique: bool,
pub field: &'a FlagSetFlagInfo<Info>,
}
pub enum VariantFieldSourcesMode<'a> {
Plain(VariantFieldSource<'a, PlainFlagInfo<'a>>),
Struct(Vec<VariantFieldSource<'a, FlagFieldInfo<'a>>>),
}
pub struct VariantFieldSources<'a> {
pub mode: VariantFieldSourcesMode<'a>,
pub reachable: bool,
}
pub fn compute_grouped_flags<'a>(
variants: &'a IndexMap<IdentString<'a>, FlagSetVariant<'a>>,
) -> syn::Result<(
IndexMap<&'a IdentString<'a>, VariantFieldSources<'a>>,
Vec<FlagSetFlag<'a>>,
)> {
let mut flags = Vec::new();
let mut sources: IndexMap<&'a IdentString<'a>, VariantFieldSources<'a>> = variants
.iter()
.enumerate()
.map(|(variant_index, (variant_ident, variant))| {
match variant {
FlagSetVariant::Plain(flag) => add_or_update_flag(
&mut flags,
flag,
Family::Solitary(()),
variant_ident,
variant_index,
0,
)
.map(|flag_index| VariantFieldSource {
index: flag_index,
field: flag,
unique: false,
})
.map(VariantFieldSourcesMode::Plain),
FlagSetVariant::Struct(fields) => fields
.iter()
.enumerate()
.map(|(field_index, flag)| {
add_or_update_flag(
&mut flags,
flag,
Family::Sibling,
variant_ident,
variant_index,
field_index,
)
.map(|field_index| VariantFieldSource {
index: field_index,
field: flag,
unique: false,
})
})
.try_collect()
.map(VariantFieldSourcesMode::Struct),
}
.map(|mode| VariantFieldSources {
mode,
reachable: false,
})
.map(|sources| (variant_ident, sources))
})
.try_collect()?;
for (index, flag) in flags.iter().enumerate() {
if let Some(unique_variant) = flag.unique_variant() {
let source = sources
.get_mut(unique_variant.ident)
.expect("all source variants should exist");
source.reachable = true;
match source.mode {
VariantFieldSourcesMode::Plain(ref mut source) => source.unique = true,
VariantFieldSourcesMode::Struct(ref mut sources) => {
sources
.iter_mut()
.find(|source| source.index == index)
.expect("source flag must exist here")
.unique = true
}
}
}
}
Ok((sources, flags))
}