use crate::codegen::shared::is_optional_type;
use crate::codegen::validation;
use crate::model::ParamDef;
use crate::model::RouterDef;
use crate::model::ToolDef;
use proc_macro2::TokenStream;
use quote::quote;
pub fn generate_mcp_methods(router: &RouterDef) -> TokenStream {
let struct_type = &router.struct_type;
let dispatch_method = generate_mcp_dispatch_method(router);
let dispatch_method_mcp = generate_mcp_dispatch_method_mcp(router);
let tools_method = generate_mcp_tools_method(router);
let server_info_method = generate_mcp_server_info_method(router);
quote! {
impl #struct_type {
#dispatch_method
#dispatch_method_mcp
#tools_method
#server_info_method
}
}
}
fn generate_mcp_dispatch_method(router: &RouterDef) -> TokenStream {
let match_arms: Vec<_> = router.tools.iter().map(generate_tool_match_arm).collect();
quote! {
pub async fn handle_mcp_call(
&self,
method: &str,
params: ::serde_json::Value
) -> ::std::result::Result<::serde_json::Value, ::universal_tool_core::error::ToolError> {
match method {
#(#match_arms)*
_ => {
::std::result::Result::Err(
::universal_tool_core::error::ToolError::new(
::universal_tool_core::error::ErrorCode::NotFound,
::std::format!("Unknown method: {}", method)
)
)
}
}
}
}
}
fn generate_mcp_dispatch_method_mcp(router: &RouterDef) -> TokenStream {
let match_arms: Vec<_> = router
.tools
.iter()
.map(generate_tool_match_arm_mcp)
.collect();
quote! {
pub async fn handle_mcp_call_mcp(
&self,
method: &str,
params: ::serde_json::Value
) -> ::std::result::Result<::universal_tool_core::mcp::McpOutput, ::universal_tool_core::error::ToolError> {
match method {
#(#match_arms)*
_ => {
::std::result::Result::Err(
::universal_tool_core::error::ToolError::new(
::universal_tool_core::error::ErrorCode::NotFound,
::std::format!("Unknown method: {}", method)
)
)
}
}
}
}
}
fn generate_tool_match_arm(tool: &ToolDef) -> TokenStream {
let tool_name = &tool.tool_name;
let param_extractions = crate::codegen::validation::generate_params_extraction(tool, "mcp");
let has_includable_params = tool
.params
.iter()
.any(|p| crate::codegen::validation::should_include_param(p, "mcp"));
let object_assertion = if has_includable_params {
quote! {
let params = match params {
::serde_json::Value::Object(map) => map,
_ => {
return ::std::result::Result::Err(
::universal_tool_core::error::ToolError::new(
::universal_tool_core::error::ErrorCode::InvalidArgument,
"Parameters must be a JSON object"
)
);
}
};
}
} else {
quote! {}
};
let method_call = generate_method_call(tool);
quote! {
#tool_name => {
#object_assertion
#( #param_extractions )*
#method_call
}
}
}
fn generate_method_call(tool: &ToolDef) -> TokenStream {
let param_args: Vec<_> = tool
.params
.iter()
.map(|param| {
let name = ¶m.name;
let ty = ¶m.ty;
let ty_str = quote!(#ty).to_string();
if ty_str.contains("ProgressReporter") {
quote! { None }
} else if ty_str.contains("CancellationToken") {
quote! { ::universal_tool_core::mcp::CancellationToken::new() }
} else {
quote! { #name }
}
})
.collect();
let method_call =
crate::codegen::shared::generate_normalized_method_call(tool, quote! { self }, param_args);
quote! {
let result = #method_call?;
::std::result::Result::Ok(::serde_json::to_value(&result)?)
}
}
fn generate_tool_match_arm_mcp(tool: &ToolDef) -> TokenStream {
let tool_name = &tool.tool_name;
let param_extractions = crate::codegen::validation::generate_params_extraction(tool, "mcp");
let has_includable_params = tool
.params
.iter()
.any(|p| crate::codegen::validation::should_include_param(p, "mcp"));
let object_assertion = if has_includable_params {
quote! {
let params = match params {
::serde_json::Value::Object(map) => map,
_ => {
return ::std::result::Result::Err(
::universal_tool_core::error::ToolError::new(
::universal_tool_core::error::ErrorCode::InvalidArgument,
"Parameters must be a JSON object"
)
);
}
};
}
} else {
quote! {}
};
let param_args: Vec<_> = tool
.params
.iter()
.map(|param| {
let name = ¶m.name;
let ty = ¶m.ty;
let ty_str = quote!(#ty).to_string();
if ty_str.contains("ProgressReporter") {
quote! { None }
} else if ty_str.contains("CancellationToken") {
quote! { ::universal_tool_core::mcp::CancellationToken::new() }
} else {
quote! { #name }
}
})
.collect();
let method_call =
crate::codegen::shared::generate_normalized_method_call(tool, quote! { self }, param_args);
let output_mode_tokens = if let Some(mcp_config) = &tool.metadata.mcp_config {
if matches!(
mcp_config.output_mode,
Some(crate::model::McpOutputMode::Text)
) {
quote! {
let result = #method_call?;
let text = ::universal_tool_core::mcp::McpFormatter::mcp_format_text(&result);
::std::result::Result::Ok(::universal_tool_core::mcp::McpOutput::Text(text))
}
} else {
quote! {
let result = #method_call?;
let val = ::serde_json::to_value(&result)?;
::std::result::Result::Ok(::universal_tool_core::mcp::McpOutput::Json(val))
}
}
} else {
quote! {
let result = #method_call?;
let val = ::serde_json::to_value(&result)?;
::std::result::Result::Ok(::universal_tool_core::mcp::McpOutput::Json(val))
}
};
quote! {
#tool_name => {
#object_assertion
#( #param_extractions )*
#output_mode_tokens
}
}
}
fn generate_mcp_tools_method(router: &RouterDef) -> TokenStream {
let tool_definitions = router.tools.iter().map(generate_tool_definition);
quote! {
pub fn get_mcp_tools(&self) -> ::std::vec::Vec<::serde_json::Value> {
::std::vec![
#(#tool_definitions),*
]
}
}
}
fn generate_tool_definition(tool: &ToolDef) -> TokenStream {
let name = &tool.tool_name;
let description = &tool.metadata.description;
let schema = generate_tool_schema(tool);
if let Some(mcp_config) = &tool.metadata.mcp_config {
let mut has_annotations = false;
let mut annotation_fields = vec![];
if let Some(v) = mcp_config.annotations.read_only_hint {
has_annotations = true;
annotation_fields.push(quote! { "readOnlyHint": #v });
}
if let Some(v) = mcp_config.annotations.destructive_hint {
has_annotations = true;
annotation_fields.push(quote! { "destructiveHint": #v });
}
if let Some(v) = mcp_config.annotations.idempotent_hint {
has_annotations = true;
annotation_fields.push(quote! { "idempotentHint": #v });
}
if let Some(v) = mcp_config.annotations.open_world_hint {
has_annotations = true;
annotation_fields.push(quote! { "openWorldHint": #v });
}
if has_annotations {
quote! {
{
let schema = #schema;
let mut tool_def = ::serde_json::json!({
"name": #name,
"description": #description,
"inputSchema": schema
});
tool_def["annotations"] = ::serde_json::json!({
#(#annotation_fields),*
});
tool_def
}
}
} else {
quote! {
{
let schema = #schema;
::serde_json::json!({
"name": #name,
"description": #description,
"inputSchema": schema
})
}
}
}
} else {
quote! {
{
let schema = #schema;
::serde_json::json!({
"name": #name,
"description": #description,
"inputSchema": schema
})
}
}
}
}
fn generate_param_schema(param: &ParamDef) -> TokenStream {
let param_type = ¶m.ty;
let description = ¶m.metadata.description.as_deref().unwrap_or("");
quote! {
{
let mut settings = ::universal_tool_core::schemars::r#gen::SchemaSettings::draft07();
settings.inline_subschemas = true;
let mut schema_gen = settings.into_generator();
let schema = <#param_type as ::universal_tool_core::JsonSchema>::json_schema(&mut schema_gen);
let mut json_schema = ::serde_json::to_value(&schema).unwrap_or_else(|_| ::serde_json::json!({"type": "string"}));
if let ::serde_json::Value::Object(ref mut map) = json_schema {
if !#description.is_empty() {
map.insert("description".to_string(), ::serde_json::Value::String(#description.to_string()));
}
}
json_schema
}
}
}
fn generate_tool_schema(tool: &ToolDef) -> TokenStream {
let schema_params: Vec<_> = tool
.params
.iter()
.filter(|param| validation::should_include_param(param, "mcp"))
.collect();
if schema_params.is_empty() {
quote! {
::serde_json::json!({
"type": "object",
"properties": {},
"required": []
})
}
} else {
let param_schemas: Vec<_> = schema_params
.iter()
.map(|param| {
let name = ¶m.name.to_string();
let schema = generate_param_schema(param);
let required = !is_optional_type(¶m.ty);
quote! {
properties.insert(#name.to_string(), #schema);
if #required {
required.push(#name.to_string());
}
}
})
.collect();
quote! {
{
let mut properties = ::serde_json::Map::new();
let mut required = Vec::new();
#(#param_schemas)*
::serde_json::json!({
"type": "object",
"properties": properties,
"required": required
})
}
}
}
}
fn generate_mcp_server_info_method(router: &RouterDef) -> TokenStream {
let default_name = router
.struct_type
.segments
.last()
.map(|s| s.ident.to_string())
.unwrap_or_else(|| "tools".to_string());
let name_val: String = router
.metadata
.mcp_config
.as_ref()
.and_then(|c| c.name.as_ref())
.cloned()
.unwrap_or(default_name);
let version_tokens: TokenStream = if let Some(cfg) = &router.metadata.mcp_config {
if let Some(ver) = &cfg.version {
let ver_lit = ver.clone();
quote! { #ver_lit.to_string() }
} else {
quote! { env!("CARGO_PKG_VERSION").to_string() }
}
} else {
quote! { env!("CARGO_PKG_VERSION").to_string() }
};
quote! {
pub fn get_mcp_server_info(&self) -> (String, String) {
let name = #name_val.to_string();
let version = #version_tokens;
(name, version)
}
}
}