use proc_macro::{Delimiter, Group, Ident, Literal, Punct, Spacing, Span, TokenStream, TokenTree};
use std::{
collections::{HashMap, HashSet},
sync::{atomic::AtomicUsize, LazyLock},
};
#[proc_macro_attribute]
pub fn derive(attr: TokenStream, item: TokenStream) -> TokenStream {
expand_aliases(attr).into_iter().chain(item).collect()
}
fn add_doc(ts: &mut TokenStream, doc: impl AsRef<str>) {
ts.extend(TokenStream::from_iter([
TokenTree::Punct(Punct::new('#', Spacing::Joint)),
TokenTree::Group(Group::new(
Delimiter::Bracket,
TokenStream::from_iter([
TokenTree::Ident(Ident::new("doc", Span::call_site())),
TokenTree::Punct(Punct::new('=', Spacing::Alone)),
TokenTree::Literal(Literal::string(doc.as_ref())),
]),
)),
]));
}
fn compile_error(ts: &mut TokenStream, span: Span, msg: impl AsRef<str>) {
ts.extend([
TokenTree::Ident(Ident::new("compile_error", span)),
TokenTree::Punct({
let mut punct = Punct::new('!', Spacing::Alone);
punct.set_span(span);
punct
}),
TokenTree::Group({
let mut group = Group::new(Delimiter::Brace, {
TokenStream::from_iter(vec![TokenTree::Literal({
let mut string = Literal::string(msg.as_ref());
string.set_span(span);
string
})])
});
group.set_span(span);
group
}),
]);
}
fn expand_aliases(input: TokenStream) -> TokenStream {
let mut seen_expanded = HashSet::new();
let mut input = input.into_iter().peekable();
let mut output = TokenStream::new();
let mut documented_items = Vec::new();
let mut compile_errors = TokenStream::new();
while let Some(tt) = input.next() {
let TokenTree::Punct(ref p) = tt else {
output.extend([tt]);
continue;
};
if p.as_char() != '.' {
output.extend([tt]);
continue;
}
let Some(next) = input.next_if(|tt| {
if let TokenTree::Punct(dot) = tt {
dot.as_char() == '.'
} else {
false
}
}) else {
compile_error(
&mut compile_errors,
p.span(),
"after `.` we expect `.` followed by a derive alias, for example: `..Alias`",
);
continue;
};
let Some(TokenTree::Ident(alias)) = input.next() else {
compile_error(
&mut compile_errors,
next.span(),
"after `..` we expect a derive alias, for example: `..Alias`",
);
continue;
};
let alias_string = alias.to_string();
let Some(list_of_aliased) = DERIVE_ALIASES.get(&alias_string) else {
let available = format_list(DERIVE_ALIASES.keys());
let most_similar = most_similar_alias(&alias_string)
.map(|similar| {
format!(
"did you mean: `{similar}`? it expands to: {}\n\n",
format_list(
DERIVE_ALIASES
.get(similar)
.expect("it is an existing alias")
)
)
})
.unwrap_or_default();
compile_error(
&mut compile_errors,
alias.span(),
format!(
"The alias `{alias_string}` is undefined.\n\n{most_similar}All available aliases: {available}",
),
);
continue;
};
let mut alias_documentation = TokenStream::new();
alias_documentation.extend([
TokenTree::Punct(Punct::new('#', Spacing::Joint)),
TokenTree::Group(Group::new(
Delimiter::Bracket,
TokenStream::from_iter([
TokenTree::Ident(Ident::new("doc", Span::call_site())),
TokenTree::Group(Group::new(
Delimiter::Parenthesis,
TokenStream::from_iter([TokenTree::Ident(Ident::new(
"hidden",
Span::mixed_site(),
))]),
)),
]),
)),
TokenTree::Punct(Punct::new('#', Spacing::Joint)),
TokenTree::Group(Group::new(
Delimiter::Bracket,
TokenStream::from_iter([
TokenTree::Ident(Ident::new("allow", Span::call_site())),
TokenTree::Group(Group::new(
Delimiter::Parenthesis,
TokenStream::from_iter([TokenTree::Ident(Ident::new(
"non_camel_case_types",
Span::mixed_site(),
))]),
)),
]),
)),
]);
add_doc(
&mut alias_documentation,
format!("Derive alias `..{alias_string}` expands to the following derives:\n"),
);
let mut is_first_derive = true;
for derive in list_of_aliased {
let contains = seen_expanded.contains(derive);
seen_expanded.insert(derive);
if contains {
continue;
}
add_doc(&mut alias_documentation, format!("- [`{derive}`]"));
for (i, part) in derive.split("::").enumerate() {
if !is_first_derive && i == 0 {
output.extend([TokenTree::Punct(Punct::new(',', Spacing::Alone))]);
}
if i > 0 {
output.extend([
TokenTree::Punct(Punct::new(':', Spacing::Joint)),
TokenTree::Punct(Punct::new(':', Spacing::Joint)),
]);
}
output.extend([TokenTree::Ident(Ident::new(part, Span::call_site()))]);
is_first_derive = false;
}
}
#[cfg(not(feature = "workspace"))]
let whos = "crate's";
#[cfg(feature = "workspace")]
let whos = "workspace's";
add_doc(
&mut alias_documentation,
format!("\nDerive aliases are defined in file `derive_aliases.rs` next to the {whos} `Cargo.toml`"),
);
alias_documentation.extend([
TokenTree::Ident(Ident::new("struct", Span::call_site())),
TokenTree::Ident(Ident::new(
&format!(
"{alias_string}__derive__alias__{}",
ALIASES_OUTPUTTED.fetch_add(1, std::sync::atomic::Ordering::Acquire)
),
alias.span(),
)),
TokenTree::Punct(Punct::new(';', Spacing::Alone)),
]);
documented_items.push(alias_documentation);
}
let output = TokenStream::from_iter([
TokenTree::Punct(Punct::new('#', Spacing::Joint)),
TokenTree::Group(Group::new(
Delimiter::Bracket,
TokenStream::from_iter([
TokenTree::Punct(Punct::new(':', Spacing::Joint)),
TokenTree::Punct(Punct::new(':', Spacing::Joint)),
TokenTree::Ident(Ident::new("std", Span::call_site())),
TokenTree::Punct(Punct::new(':', Spacing::Joint)),
TokenTree::Punct(Punct::new(':', Spacing::Joint)),
TokenTree::Ident(Ident::new("prelude", Span::call_site())),
TokenTree::Punct(Punct::new(':', Spacing::Joint)),
TokenTree::Punct(Punct::new(':', Spacing::Joint)),
TokenTree::Ident(Ident::new("v1", Span::call_site())),
TokenTree::Punct(Punct::new(':', Spacing::Joint)),
TokenTree::Punct(Punct::new(':', Spacing::Joint)),
TokenTree::Ident(Ident::new("derive", Span::call_site())),
TokenTree::Group(Group::new(Delimiter::Parenthesis, output)),
]),
)),
]);
documented_items
.into_iter()
.flatten()
.chain(output)
.chain(compile_errors)
.collect()
}
fn format_list<'a>(list: impl IntoIterator<Item = &'a String>) -> String {
list.into_iter()
.enumerate()
.flat_map(|(i, key)| {
if i == 0 {
vec![key.as_str()]
} else {
vec![", ", key.as_str()]
}
})
.collect::<String>()
}
fn most_similar_alias(alias: impl AsRef<str>) -> Option<&'static String> {
DERIVE_ALIASES
.keys()
.map(|it| {
(
it,
strsim::normalized_damerau_levenshtein(it, alias.as_ref()),
)
})
.max_by(|a, b| a.1.total_cmp(&b.1))
.filter(|it| it.1 >= 0.70)
.map(|it| it.0)
}
static DERIVE_ALIASES: LazyLock<HashMap<String, Vec<String>>> = LazyLock::new(|| {
#[cfg(feature = "workspace")]
let Ok(dir) = std::env::var("CARGO_WORKSPACE_DIR") else {
panic!(concat!(
"\n\n`CARGO_WORKSPACE_DIR` environment variable must be set, which points to the directory containing the\n",
"workspace `Cargo.toml`. Since cargo currently doesn't set this variable, in your workspace create `.cargo/config.toml` with contents:\n",
"\n",
"[env]\n",
"CARGO_WORKSPACE_DIR = {{ value = \"\", relative = true }}\n",
))
};
#[cfg(not(feature = "workspace"))]
let Ok(dir) = std::env::var("CARGO_MANIFEST_DIR") else {
panic!("env variable `CARGO_MANIFEST_DIR` must be set, which points to the directory containing your crate's `Cargo.toml`. Cargo supplies this env variable by default")
};
let path = std::path::Path::new(&dir).join("derive_aliases.rs");
let content = std::fs::read_to_string(&path).unwrap_or_else(|err| {
panic!(
"expected {} to exist and contain derive aliases. error: {err}.\nhere's an example of syntax in `derive_aliases.rs` file:\n\n{EXAMPLE_DERIVE_ALIASES_RS}\n",
path.display()
)
});
let alias_map = parse_my_little_rust(&content);
fn resolve_derives(
alias: &str,
alias_map: &HashMap<String, Vec<String>>,
current: &mut Vec<String>,
) {
let Some(list_of_aliased) = alias_map.get(alias) else {
panic!("failed to parse aliases file: Alias `{alias}` does not exist");
};
for aliased in list_of_aliased {
if let Some(alias) = aliased.strip_prefix("..") {
resolve_derives(alias, alias_map, current);
} else {
current.push(aliased.clone());
}
}
}
let mut aliases_to_derives = HashMap::new();
for alias in alias_map.keys() {
let mut derives = vec![];
resolve_derives(alias, &alias_map, &mut derives);
aliases_to_derives.insert(alias.to_string(), derives);
}
aliases_to_derives
});
const EXAMPLE_DERIVE_ALIASES_RS: &str = "\
// Simple derive aliases
//
// `#[derive(..Copy, ..Eq)]` expands to `#[std::derive(Copy, Clone, PartialEq, Eq)]`
Copy = Copy, Clone;
Eq = PartialEq, Eq;
// You can nest them!
//
// `#[derive(..Ord, std::hash::Hash)]` expands to `#[std::derive(PartialOrd, Ord, PartialEq, Eq, std::hash::Hash)]`
Ord = PartialOrd, Ord, ..Eq;";
fn parse_my_little_rust(s: &str) -> HashMap<String, Vec<String>> {
let mut cleaned_from_block_comments = String::new();
let mut chars = s.chars().peekable();
while let Some(c) = chars.next() {
if c == '/' && chars.peek() == Some(&'*') {
chars.next();
while let Some(d) = chars.next() {
if d == '*' && chars.peek() == Some(&'/') {
chars.next();
break;
}
}
} else {
cleaned_from_block_comments.push(c);
}
}
let mut cleaned = String::new();
for line in cleaned_from_block_comments.lines() {
if let Some(idx) = line.find("//") {
cleaned.push_str(&line[..idx]);
cleaned.push('\n');
} else {
cleaned.push_str(line);
cleaned.push('\n');
}
}
let mut aliases = HashMap::new();
for stmt in cleaned.split(';').map(str::trim).filter(|s| !s.is_empty()) {
if let Some(path) = stmt.strip_prefix("use") {
let path = path.trim();
let path = path.strip_prefix('"').expect("`use` expects a string");
let path = path.strip_suffix('"').expect("`use` expects a string");
let other_contents = std::fs::read_to_string(path)
.unwrap_or_else(|err| panic!("failed to read derive aliases at {path}: {err}"));
aliases.extend(parse_my_little_rust(&other_contents));
} else if let Some((alias_name, alias_value)) = stmt.split_once('=') {
let alias_name = alias_name.trim().to_string();
let expansions: Vec<String> = alias_value
.split(',')
.map(|part| part.trim().to_string())
.filter(|s| !s.is_empty())
.collect();
aliases.insert(alias_name, expansions);
} else {
panic!("invalid derive alias file")
}
}
aliases
}
static ALIASES_OUTPUTTED: AtomicUsize = AtomicUsize::new(0);
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn my_little_rust() {
assert_eq!(
parse_my_little_rust(EXAMPLE_DERIVE_ALIASES_RS),
HashMap::from([
(
"Copy".to_string(),
vec!["Copy".to_string(), "Clone".to_string()]
),
(
"Eq".to_string(),
vec!["PartialEq".to_string(), "Eq".to_string(),]
),
(
"Ord".to_string(),
vec![
"PartialOrd".to_string(),
"Ord".to_string(),
"..Eq".to_string()
]
),
])
);
}
}