use crate::codegen::error_handling;
use crate::codegen::shared::is_bool_type;
use crate::codegen::shared::is_custom_struct_type;
use crate::codegen::shared::is_hashmap_type;
use crate::codegen::shared::is_optional_type;
use crate::codegen::shared::is_vec_type;
use crate::model::ParamDef;
use crate::model::ToolDef;
use proc_macro2::TokenStream;
use quote::quote;
use syn::Type;
use syn::parse_str;
pub fn generate_param_extraction(param: &ParamDef, interface: &str) -> TokenStream {
let param_ident = ¶m.name;
match interface {
"cli" => generate_cli_param_extraction(param),
"rest" => {
quote! { let #param_ident = params.#param_ident; }
}
"mcp" => generate_mcp_param_extraction(param),
_ => quote! {},
}
}
fn generate_cli_param_extraction(param: &ParamDef) -> TokenStream {
let param_name = ¶m.name.to_string();
let param_ident = ¶m.name;
let param_type = ¶m.ty;
if is_custom_struct_type(¶m.ty) {
let missing_err = error_handling::generate_missing_param_error(param_name, "cli");
let parse_err = error_handling::generate_parse_error(param_name, "JSON object", "cli");
if is_optional_type(¶m.ty) {
quote! {
let #param_ident: #param_type = if let Some(json_str) = sub_matches.get_one::<String>(#param_name) {
Some(::serde_json::from_str(json_str)
.map_err(|e| #parse_err)?)
} else {
None
};
}
} else {
quote! {
let #param_ident: #param_type = {
let json_str = sub_matches.get_one::<String>(#param_name)
.ok_or_else(|| #missing_err)?;
::serde_json::from_str(json_str)
.map_err(|e| {
::universal_tool_core::prelude::ToolError::new(
::universal_tool_core::prelude::ErrorCode::InvalidArgument,
format!("Failed to parse JSON for {}: {}", #param_name, e)
)
})?
};
}
}
} else if is_vec_type(¶m.ty) {
quote! {
let #param_ident: #param_type = sub_matches.get_many::<String>(#param_name)
.map(|values| values.map(|s| s.clone()).collect())
.unwrap_or_else(Vec::new);
}
} else if is_hashmap_type(¶m.ty) {
let parse_err = error_handling::generate_parse_error(param_name, "key=value format", "cli");
let value_type = extract_hashmap_value_type(¶m.ty);
let value_type_str = quote!(#value_type).to_string();
let value_parse = match value_type_str.as_str() {
"String" => quote! { parts[1].to_string() },
"i32" => quote! {
parts[1].parse::<i32>().map_err(|_| {
::universal_tool_core::prelude::ToolError::new(
::universal_tool_core::prelude::ErrorCode::InvalidArgument,
format!("Invalid i32 value for {}: {}", #param_name, parts[1])
)
})?
},
"i64" => quote! {
parts[1].parse::<i64>().map_err(|_| {
::universal_tool_core::prelude::ToolError::new(
::universal_tool_core::prelude::ErrorCode::InvalidArgument,
format!("Invalid i64 value for {}: {}", #param_name, parts[1])
)
})?
},
"u32" => quote! {
parts[1].parse::<u32>().map_err(|_| {
::universal_tool_core::prelude::ToolError::new(
::universal_tool_core::prelude::ErrorCode::InvalidArgument,
format!("Invalid u32 value for {}: {}", #param_name, parts[1])
)
})?
},
"u64" => quote! {
parts[1].parse::<u64>().map_err(|_| {
::universal_tool_core::prelude::ToolError::new(
::universal_tool_core::prelude::ErrorCode::InvalidArgument,
format!("Invalid u64 value for {}: {}", #param_name, parts[1])
)
})?
},
"f32" => quote! {
parts[1].parse::<f32>().map_err(|_| {
::universal_tool_core::prelude::ToolError::new(
::universal_tool_core::prelude::ErrorCode::InvalidArgument,
format!("Invalid f32 value for {}: {}", #param_name, parts[1])
)
})?
},
"f64" => quote! {
parts[1].parse::<f64>().map_err(|_| {
::universal_tool_core::prelude::ToolError::new(
::universal_tool_core::prelude::ErrorCode::InvalidArgument,
format!("Invalid f64 value for {}: {}", #param_name, parts[1])
)
})?
},
"bool" => quote! {
parts[1].parse::<bool>().map_err(|_| {
::universal_tool_core::prelude::ToolError::new(
::universal_tool_core::prelude::ErrorCode::InvalidArgument,
format!("Invalid bool value for {}: {}", #param_name, parts[1])
)
})?
},
_ => quote! {
::serde_json::from_str(parts[1]).map_err(|_| {
::universal_tool_core::prelude::ToolError::new(
::universal_tool_core::prelude::ErrorCode::InvalidArgument,
format!("Invalid JSON value for {}: {}", #param_name, parts[1])
)
})?
},
};
quote! {
let #param_ident: #param_type = {
let mut map = std::collections::HashMap::new();
if let Some(values) = sub_matches.get_many::<String>(#param_name) {
for kv in values {
let parts: Vec<&str> = kv.splitn(2, '=').collect();
if parts.len() != 2 {
return Err(#parse_err);
}
let value = #value_parse;
map.insert(parts[0].to_string(), value);
}
}
map
};
}
} else {
let missing_err = error_handling::generate_missing_param_error(param_name, "cli");
if is_optional_type(¶m.ty) {
let inner_type = extract_option_inner_type(¶m.ty);
let type_str = quote!(#inner_type).to_string();
match type_str.as_str() {
"String" => quote! {
let #param_ident: #param_type = sub_matches.get_one::<String>(#param_name).cloned();
},
"i32" | "i64" | "u32" | "u64" | "f32" | "f64" => quote! {
let #param_ident: #param_type = sub_matches.get_one::<#inner_type>(#param_name).cloned();
},
_ => quote! {
let #param_ident: #param_type = sub_matches.get_one::<String>(#param_name)
.and_then(|s| ::serde_json::from_str(s).ok());
},
}
} else if is_bool_type(¶m.ty) {
quote! {
let #param_ident = sub_matches.get_flag(#param_name);
}
} else {
let type_str = quote!(#param_type).to_string();
match type_str.as_str() {
"String" => quote! {
let #param_ident = sub_matches.get_one::<String>(#param_name)
.ok_or_else(|| #missing_err)?
.clone();
},
"i32" => quote! {
let #param_ident = sub_matches.get_one::<i32>(#param_name)
.ok_or_else(|| #missing_err)?
.clone();
},
"i64" => quote! {
let #param_ident = sub_matches.get_one::<i64>(#param_name)
.ok_or_else(|| #missing_err)?
.clone();
},
"u32" => quote! {
let #param_ident = sub_matches.get_one::<u32>(#param_name)
.ok_or_else(|| #missing_err)?
.clone();
},
"u64" => quote! {
let #param_ident = sub_matches.get_one::<u64>(#param_name)
.ok_or_else(|| #missing_err)?
.clone();
},
"f32" => quote! {
let #param_ident = sub_matches.get_one::<f32>(#param_name)
.ok_or_else(|| #missing_err)?
.clone();
},
"f64" => quote! {
let #param_ident = sub_matches.get_one::<f64>(#param_name)
.ok_or_else(|| #missing_err)?
.clone();
},
_ => {
let parse_err =
error_handling::generate_parse_error(param_name, "value", "cli");
quote! {
let #param_ident: #param_type = sub_matches.get_one::<String>(#param_name)
.ok_or_else(|| #missing_err)?
.parse()
.map_err(|_| #parse_err)?;
}
}
}
}
}
}
fn generate_mcp_param_extraction(param: &ParamDef) -> TokenStream {
let param_name = ¶m.name.to_string();
let param_ident = ¶m.name;
let param_type = ¶m.ty;
let missing_err = error_handling::generate_missing_param_error(param_name, "mcp");
let parse_err =
error_handling::generate_parse_error(param_name, "e!(#param_type).to_string(), "mcp");
if is_optional_type(¶m.ty) {
quote! {
let #param_ident: #param_type = params.get(#param_name)
.map(|v| ::serde_json::from_value(v.clone())
.map_err(|_| #parse_err))
.transpose()?;
}
} else {
quote! {
let #param_ident: #param_type = params.get(#param_name)
.ok_or_else(|| #missing_err)
.and_then(|v| ::serde_json::from_value(v.clone())
.map_err(|_| #parse_err))?;
}
}
}
pub fn generate_params_extraction(tool: &ToolDef, interface: &str) -> Vec<TokenStream> {
tool.params
.iter()
.filter(|p| should_include_param(p, interface))
.map(|param| generate_param_extraction(param, interface))
.collect()
}
pub fn should_include_param(param: &ParamDef, interface: &str) -> bool {
match interface {
"mcp" => {
let param_ty = ¶m.ty;
let type_str = quote!(#param_ty).to_string();
!type_str.contains("ProgressReporter") && !type_str.contains("CancellationToken")
}
_ => true,
}
}
fn extract_option_inner_type(ty: &Type) -> Type {
if let Type::Path(type_path) = ty
&& let Some(segment) = type_path.path.segments.first()
&& segment.ident == "Option"
&& let syn::PathArguments::AngleBracketed(args) = &segment.arguments
&& let Some(syn::GenericArgument::Type(inner_ty)) = args.args.first()
{
return inner_ty.clone();
}
parse_str("String").unwrap()
}
fn extract_hashmap_value_type(ty: &Type) -> Type {
if let Type::Path(type_path) = ty
&& let Some(segment) = type_path.path.segments.last()
&& segment.ident == "HashMap"
&& let syn::PathArguments::AngleBracketed(args) = &segment.arguments
{
if let Some(syn::GenericArgument::Type(value_ty)) = args.args.iter().nth(1) {
return value_ty.clone();
}
}
parse_str("String").unwrap()
}
pub fn generate_cli_value_parser(param: &ParamDef) -> Option<TokenStream> {
let param_type = ¶m.ty;
if is_bool_type(¶m.ty)
|| is_vec_type(¶m.ty)
|| is_custom_struct_type(¶m.ty)
|| is_hashmap_type(¶m.ty)
{
None
} else if is_optional_type(¶m.ty) {
let inner_type = extract_option_inner_type(¶m.ty);
Some(quote! { ::universal_tool_core::cli::clap::value_parser!(#inner_type) })
} else {
Some(quote! { ::universal_tool_core::cli::clap::value_parser!(#param_type) })
}
}
pub fn generate_cli_arg_config(param: &ParamDef) -> TokenStream {
let param_name = ¶m.name.to_string();
let description = param
.metadata
.description
.as_deref()
.unwrap_or("Parameter value");
let value_parser = generate_cli_value_parser(param);
let is_required = !is_optional_type(¶m.ty)
&& !is_bool_type(¶m.ty)
&& !is_vec_type(¶m.ty)
&& !is_hashmap_type(¶m.ty);
let is_bool = is_bool_type(¶m.ty);
let is_multi = is_vec_type(¶m.ty) || is_hashmap_type(¶m.ty);
let base_arg = match (value_parser, is_required, is_bool, is_multi) {
(Some(vp), true, false, false) => quote! {
::universal_tool_core::cli::clap::Arg::new(#param_name)
.long(#param_name)
.help(#description)
.value_parser(#vp)
.required(true)
},
(Some(vp), false, false, false) => quote! {
::universal_tool_core::cli::clap::Arg::new(#param_name)
.long(#param_name)
.help(#description)
.value_parser(#vp)
},
(None, true, false, false) => quote! {
::universal_tool_core::cli::clap::Arg::new(#param_name)
.long(#param_name)
.help(#description)
.required(true)
},
(None, false, true, false) => quote! {
::universal_tool_core::cli::clap::Arg::new(#param_name)
.long(#param_name)
.help(#description)
.action(::universal_tool_core::cli::clap::ArgAction::SetTrue)
},
(None, false, false, true) => quote! {
::universal_tool_core::cli::clap::Arg::new(#param_name)
.long(#param_name)
.help(#description)
.action(::universal_tool_core::cli::clap::ArgAction::Append)
},
_ => quote! {
::universal_tool_core::cli::clap::Arg::new(#param_name)
.long(#param_name)
.help(#description)
},
};
let with_env = if let Some(env_var) = ¶m.metadata.env {
quote! { .env(#env_var) }
} else {
quote! {}
};
let with_default = if let Some(default_val) = ¶m.metadata.default {
quote! { .default_value(#default_val) }
} else {
quote! {}
};
quote! {
#base_arg
#with_env
#with_default
}
}
#[cfg(test)]
mod tests {
use super::*;
use syn::parse_quote;
#[test]
fn test_should_include_param() {
let normal_param = ParamDef {
name: parse_quote!(input),
ty: parse_quote!(String),
source: crate::model::ParamSource::Body,
is_optional: false,
metadata: Default::default(),
};
assert!(should_include_param(&normal_param, "cli"));
assert!(should_include_param(&normal_param, "rest"));
assert!(should_include_param(&normal_param, "mcp"));
let progress_param = ParamDef {
name: parse_quote!(progress),
ty: parse_quote!(ProgressReporter),
source: crate::model::ParamSource::Body,
is_optional: false,
metadata: Default::default(),
};
assert!(should_include_param(&progress_param, "cli"));
assert!(should_include_param(&progress_param, "rest"));
assert!(!should_include_param(&progress_param, "mcp"));
}
}