use syn::{Attribute, Expr, Lit, LitStr, Meta};
#[derive(Default)]
pub(crate) struct ContainerAttrs {
pub name: Option<String>,
pub description: Option<String>,
pub read_only: bool,
pub concurrency_safe: bool,
pub system_prompt: Option<String>,
pub handler: Option<String>,
pub allow_extra: bool,
}
#[derive(Default)]
pub(crate) struct FieldAttrs {
pub name: Option<String>,
pub description: Option<String>,
pub skip: bool,
pub default: bool,
}
const CONTAINER_KEYS: &str =
"name, description, read_only, concurrency_safe, system_prompt, handler, allow_extra";
const FIELD_KEYS: &str = "name, description, skip, default";
pub(crate) fn parse_container(attrs: &[Attribute]) -> syn::Result<ContainerAttrs> {
let mut out = ContainerAttrs::default();
for attr in attrs.iter().filter(|a| a.path().is_ident("tool")) {
attr.parse_nested_meta(|meta| {
if meta.path.is_ident("name") {
out.name = Some(string_value(&meta)?);
} else if meta.path.is_ident("description") {
out.description = Some(string_value(&meta)?);
} else if meta.path.is_ident("system_prompt") {
out.system_prompt = Some(string_value(&meta)?);
} else if meta.path.is_ident("handler") {
out.handler = Some(string_value(&meta)?);
} else if meta.path.is_ident("read_only") {
out.read_only = true;
} else if meta.path.is_ident("concurrency_safe") {
out.concurrency_safe = true;
} else if meta.path.is_ident("allow_extra") {
out.allow_extra = true;
} else {
return Err(meta.error(format!(
"unknown `tool` attribute; expected one of: {CONTAINER_KEYS}"
)));
}
Ok(())
})?;
}
Ok(out)
}
pub(crate) fn parse_field(attrs: &[Attribute]) -> syn::Result<FieldAttrs> {
let mut out = FieldAttrs::default();
for attr in attrs.iter().filter(|a| a.path().is_ident("tool")) {
attr.parse_nested_meta(|meta| {
if meta.path.is_ident("name") {
out.name = Some(string_value(&meta)?);
} else if meta.path.is_ident("description") {
out.description = Some(string_value(&meta)?);
} else if meta.path.is_ident("skip") {
out.skip = true;
} else if meta.path.is_ident("default") {
out.default = true;
} else {
return Err(meta.error(format!(
"unknown `tool` attribute; expected one of: {FIELD_KEYS}"
)));
}
Ok(())
})?;
}
Ok(out)
}
fn string_value(meta: &syn::meta::ParseNestedMeta<'_>) -> syn::Result<String> {
let value = meta.value()?;
let lit: LitStr = value.parse()?;
Ok(lit.value())
}
pub(crate) fn doc_string(attrs: &[Attribute]) -> Option<String> {
let mut lines = Vec::new();
for attr in attrs.iter().filter(|a| a.path().is_ident("doc")) {
if let Meta::NameValue(nv) = &attr.meta
&& let Expr::Lit(expr) = &nv.value
&& let Lit::Str(s) = &expr.lit
{
lines.push(s.value().trim().to_string());
}
}
if lines.is_empty() {
None
} else {
Some(lines.join(" "))
}
}
pub(crate) fn has_serde_default(attrs: &[Attribute]) -> bool {
let mut found = false;
for attr in attrs.iter().filter(|a| a.path().is_ident("serde")) {
let _ = attr.parse_nested_meta(|meta| {
if meta.path.is_ident("default") {
found = true;
}
Ok(())
});
}
found
}
pub(crate) fn serde_rename(attrs: &[Attribute]) -> Option<String> {
let mut plain = None;
let mut deserialize = None;
for attr in attrs.iter().filter(|a| a.path().is_ident("serde")) {
let Ok(list) = attr
.parse_args_with(syn::punctuated::Punctuated::<Meta, syn::Token![,]>::parse_terminated)
else {
continue;
};
for meta in &list {
match meta {
Meta::NameValue(nv) if nv.path.is_ident("rename") => {
if let Expr::Lit(expr) = &nv.value
&& let Lit::Str(lit) = &expr.lit
{
plain = Some(lit.value());
}
}
Meta::List(ml) if ml.path.is_ident("rename") => {
let Ok(inner) = syn::parse::Parser::parse2(
&syn::punctuated::Punctuated::<Meta, syn::Token![,]>::parse_terminated,
ml.tokens.clone(),
) else {
continue;
};
for meta in &inner {
if let Meta::NameValue(nv) = meta
&& nv.path.is_ident("deserialize")
&& let Expr::Lit(expr) = &nv.value
&& let Lit::Str(lit) = &expr.lit
{
deserialize = Some(lit.value());
}
}
}
_ => {}
}
}
}
deserialize.or(plain)
}
pub(crate) fn serde_rename_all(attrs: &[Attribute]) -> Option<RenameAll> {
let mut out = None;
for attr in attrs.iter().filter(|a| a.path().is_ident("serde")) {
let _ = attr.parse_nested_meta(|meta| {
if meta.path.is_ident("rename_all")
&& let Ok(lit) = meta.value()?.parse::<LitStr>()
{
out = RenameAll::from_str(&lit.value());
}
Ok(())
});
}
out
}
#[derive(Debug, Clone, Copy)]
pub(crate) enum RenameAll {
Lower,
Upper,
Pascal,
Camel,
Snake,
ScreamingSnake,
Kebab,
ScreamingKebab,
}
impl RenameAll {
pub(crate) fn from_str(s: &str) -> Option<Self> {
match s {
"lowercase" => Some(Self::Lower),
"UPPERCASE" => Some(Self::Upper),
"PascalCase" => Some(Self::Pascal),
"camelCase" => Some(Self::Camel),
"snake_case" => Some(Self::Snake),
"SCREAMING_SNAKE_CASE" => Some(Self::ScreamingSnake),
"kebab-case" => Some(Self::Kebab),
"SCREAMING_KEBAB_CASE" => Some(Self::ScreamingKebab),
_ => None,
}
}
pub(crate) fn apply(self, name: &str) -> String {
match self {
Self::Lower => name.to_lowercase(),
Self::Upper => name.to_uppercase(),
Self::Pascal => to_pascal_case(name),
Self::Camel => {
let pascal = to_pascal_case(name);
let mut chars = pascal.chars();
match chars.next() {
Some(first) => first.to_lowercase().collect::<String>() + chars.as_str(),
None => String::new(),
}
}
Self::Snake => serde_snake_case(name),
Self::ScreamingSnake => serde_snake_case(name).to_uppercase(),
Self::Kebab => serde_snake_case(name).replace('_', "-"),
Self::ScreamingKebab => serde_snake_case(name).replace('_', "-").to_uppercase(),
}
}
}
fn to_pascal_case(name: &str) -> String {
name.split('_')
.filter(|s| !s.is_empty())
.map(|word| {
let mut chars = word.chars();
match chars.next() {
Some(first) => first.to_uppercase().collect::<String>() + chars.as_str(),
None => String::new(),
}
})
.collect()
}
fn serde_snake_case(name: &str) -> String {
let chars: Vec<char> = name.chars().collect();
let mut out = String::new();
for (i, &ch) in chars.iter().enumerate() {
if ch.is_uppercase() {
let prev_lower = i > 0
&& chars
.get(i.wrapping_sub(1))
.is_some_and(|c| c.is_lowercase() || c.is_numeric());
let next_lower = chars
.get(i.wrapping_add(1))
.is_some_and(|c| c.is_lowercase());
if prev_lower || (i > 0 && next_lower) {
out.push('_');
}
out.extend(ch.to_lowercase());
} else {
out.push(ch);
}
}
out
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn rename_all_covers_the_full_serde_strategy_set() {
let cases: Vec<(&str, &str, &str)> = vec![
("lowercase", "FileName", "filename"),
("UPPERCASE", "FileName", "FILENAME"),
("PascalCase", "file_name", "FileName"),
("camelCase", "file_name", "fileName"),
("snake_case", "FileName", "file_name"),
("SCREAMING_SNAKE_CASE", "FileName", "FILE_NAME"),
("kebab-case", "file_name", "file-name"),
("SCREAMING_KEBAB_CASE", "file_name", "FILE-NAME"),
];
for (name, input, expected) in cases {
let strategy =
RenameAll::from_str(name).unwrap_or_else(|| panic!("unknown strategy: {name}"));
assert_eq!(
strategy.apply(input),
expected,
"strategy {name:?} on {input:?}"
);
}
}
#[test]
fn snake_case_preserves_consecutive_uppercase() {
assert_eq!(
RenameAll::from_str("snake_case").unwrap().apply("userID"),
"user_id"
);
assert_eq!(
RenameAll::from_str("snake_case")
.unwrap()
.apply("parseHTTPResponse"),
"parse_http_response"
);
assert_eq!(
RenameAll::from_str("snake_case").unwrap().apply("htmlID"),
"html_id"
);
}
#[test]
fn serde_rename_deserialize_form_is_preferred() {
use syn::parse_quote;
let attr: Attribute = parse_quote! {
#[serde(rename(deserialize = "from_wire", serialize = "to_wire"))]
};
assert_eq!(
serde_rename(&[attr]),
Some("from_wire".to_string()),
"the deserialize half wins over the serialize half"
);
let attr: Attribute = parse_quote! {
#[serde(rename = "simple")]
};
assert_eq!(serde_rename(&[attr]), Some("simple".to_string()));
}
#[test]
fn rename_all_from_str_rejects_unknown_names() {
assert!(RenameAll::from_str("NonsenseCase").is_none());
assert!(RenameAll::from_str("").is_none());
}
}