use crate::mcp_common::{self, ParamSlot};
use darling::FromMeta;
use heck::ToUpperCamelCase;
use proc_macro2::{Ident, TokenStream};
use quote::{format_ident, quote};
use syn::parse::Parser;
use syn::ItemFn;
#[derive(Debug, FromMeta)]
pub struct McpPromptArgs {
pub(crate) description: String,
#[darling(default)]
pub(crate) name: Option<String>,
}
fn parse_prompt_attr_args(
args: TokenStream,
fn_ident: &syn::Ident,
) -> syn::Result<Vec<darling::ast::NestedMeta>> {
if args.is_empty() {
return Err(syn::Error::new_spanned(
fn_ident,
"mcp_prompt requires at least `description = \"...\"` attribute",
));
}
let parser =
syn::punctuated::Punctuated::<darling::ast::NestedMeta, syn::Token![,]>::parse_terminated;
Ok(parser
.parse2(args)
.map(|p| p.into_iter().collect::<Vec<_>>())
.unwrap_or_default())
}
struct PromptFnParams {
args_type: Option<syn::Type>,
state_inner_ty: Option<syn::Type>,
param_order: Vec<ParamSlot>,
has_extra: bool,
}
fn push_prompt_args_slot(
param: &syn::FnArg,
ty: syn::Type,
args_type: &mut Option<syn::Type>,
param_order: &mut Vec<ParamSlot>,
) -> syn::Result<()> {
if args_type.is_some() {
return Err(syn::Error::new_spanned(
param,
"mcp_prompt functions can have at most one args parameter",
));
}
*args_type = Some(ty);
param_order.push(ParamSlot::Args);
Ok(())
}
fn push_prompt_state_slot(
param: &syn::FnArg,
inner_ty: syn::Type,
state_inner_ty: &mut Option<syn::Type>,
param_order: &mut Vec<ParamSlot>,
) -> syn::Result<()> {
if state_inner_ty.is_some() {
return Err(syn::Error::new_spanned(
param,
"mcp_prompt functions can have at most one State<T> parameter",
));
}
*state_inner_ty = Some(inner_ty);
param_order.push(ParamSlot::State);
Ok(())
}
fn push_prompt_extra_slot(
param: &syn::FnArg,
has_extra: &mut bool,
param_order: &mut Vec<ParamSlot>,
) -> syn::Result<()> {
if *has_extra {
return Err(syn::Error::new_spanned(
param,
"mcp_prompt functions can have at most one RequestHandlerExtra parameter",
));
}
*has_extra = true;
param_order.push(ParamSlot::Extra);
Ok(())
}
fn classify_prompt_fn_params(input: &ItemFn) -> syn::Result<PromptFnParams> {
let mut args_type: Option<syn::Type> = None;
let mut state_inner_ty: Option<syn::Type> = None;
let mut has_extra = false;
let mut param_order: Vec<ParamSlot> = Vec::new();
for param in &input.sig.inputs {
let role = mcp_common::classify_param(param)?;
match role {
mcp_common::ParamRole::Args(ty) => {
push_prompt_args_slot(param, ty, &mut args_type, &mut param_order)?;
},
mcp_common::ParamRole::State { inner_ty, .. } => {
push_prompt_state_slot(param, inner_ty, &mut state_inner_ty, &mut param_order)?;
},
mcp_common::ParamRole::Extra => {
push_prompt_extra_slot(param, &mut has_extra, &mut param_order)?;
},
mcp_common::ParamRole::SelfRef => {
return Err(syn::Error::new_spanned(
param,
"standalone #[mcp_prompt] functions cannot have &self -- use #[mcp_server] for impl block prompts",
));
},
}
}
Ok(PromptFnParams {
args_type,
state_inner_ty,
param_order,
has_extra,
})
}
struct PromptStateCodegen {
struct_fields: TokenStream,
with_state_method: TokenStream,
constructor_default: TokenStream,
state_resolution: TokenStream,
}
fn generate_prompt_state_codegen(
state_inner_ty: Option<&syn::Type>,
struct_name: &syn::Ident,
prompt_name: &str,
) -> PromptStateCodegen {
let Some(inner) = state_inner_ty else {
return PromptStateCodegen {
struct_fields: quote! {},
with_state_method: quote! {},
constructor_default: quote! { #struct_name {} },
state_resolution: quote! {},
};
};
let inner_name = quote!(#inner).to_string();
PromptStateCodegen {
struct_fields: quote! { state: Option<std::sync::Arc<#inner>>, },
with_state_method: quote! {
pub fn with_state(mut self, state: impl Into<std::sync::Arc<#inner>>) -> Self {
self.state = Some(state.into());
self
}
},
constructor_default: quote! { #struct_name { state: None } },
state_resolution: quote! {
let state_val = pmcp::State(
self.state.as_ref()
.ok_or_else(|| pmcp::Error::internal(format!(
"State<{}> not provided for prompt '{}' -- call .with_state() during registration",
#inner_name, #prompt_name
)))?
.clone()
);
},
}
}
fn generate_prompt_args_deser(args_type: Option<&syn::Type>, prompt_name: &str) -> TokenStream {
let Some(at) = args_type else {
return quote! {};
};
quote! {
let typed_args: #at = pmcp::server::typed_prompt::deserialize_prompt_args(args, #prompt_name)?;
}
}
fn generate_prompt_metadata_body(
args_type: Option<&syn::Type>,
prompt_name: &str,
description: &str,
) -> TokenStream {
let Some(at) = args_type else {
return quote! {
fn metadata(&self) -> Option<pmcp::types::PromptInfo> {
Some(pmcp::types::PromptInfo::new(#prompt_name)
.with_description(#description))
}
};
};
quote! {
fn metadata(&self) -> Option<pmcp::types::PromptInfo> {
let mut info = pmcp::types::PromptInfo::new(#prompt_name)
.with_description(#description);
let schema = schemars::schema_for!(#at);
let json_schema = serde_json::to_value(&schema).unwrap_or_default();
let arguments = pmcp::server::typed_prompt::extract_prompt_arguments_from_schema(&json_schema);
if !arguments.is_empty() {
info = info.with_arguments(arguments);
}
Some(info)
}
}
}
pub fn expand_mcp_prompt(args: TokenStream, input: &ItemFn) -> syn::Result<TokenStream> {
let nested_metas = parse_prompt_attr_args(args, &input.sig.ident)?;
let macro_args = McpPromptArgs::from_list(&nested_metas)
.map_err(|e| syn::Error::new_spanned(&input.sig.ident, e.to_string()))?;
let fn_name = &input.sig.ident;
let fn_name_str = fn_name.to_string();
let prompt_name = macro_args.name.unwrap_or_else(|| fn_name_str.clone());
let is_async = input.sig.asyncness.is_some();
let struct_name = format_ident!("{}Prompt", fn_name_str.to_upper_camel_case());
let description = ¯o_args.description;
let impl_fn_name = format_ident!("__{}_impl", fn_name_str);
let mut impl_fn = input.clone();
impl_fn.sig.ident = impl_fn_name.clone();
let params = classify_prompt_fn_params(input)?;
let state_codegen =
generate_prompt_state_codegen(params.state_inner_ty.as_ref(), &struct_name, &prompt_name);
let args_deser = generate_prompt_args_deser(params.args_type.as_ref(), &prompt_name);
let struct_fields = state_codegen.struct_fields;
let with_state_method = state_codegen.with_state_method;
let constructor_default = state_codegen.constructor_default;
let state_resolution = state_codegen.state_resolution;
let extra_param_name: Ident = if params.has_extra {
format_ident!("extra")
} else {
format_ident!("_extra")
};
let call_args: Vec<TokenStream> = params
.param_order
.iter()
.map(|slot| match slot {
ParamSlot::Args => quote! { typed_args },
ParamSlot::State => quote! { state_val },
ParamSlot::Extra => quote! { #extra_param_name },
})
.collect();
let fn_call = if is_async {
quote! { #impl_fn_name(#(#call_args),*).await }
} else {
quote! { #impl_fn_name(#(#call_args),*) }
};
let handle_body = quote! {
#args_deser
#state_resolution
#fn_call
};
let metadata_body =
generate_prompt_metadata_body(params.args_type.as_ref(), &prompt_name, description);
let expanded = quote! {
#impl_fn
#[derive(Clone)]
pub struct #struct_name {
#struct_fields
}
impl std::fmt::Debug for #struct_name {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct(stringify!(#struct_name)).finish()
}
}
#[pmcp::async_trait]
impl pmcp::PromptHandler for #struct_name {
async fn handle(
&self,
args: std::collections::HashMap<String, String>,
#extra_param_name: pmcp::RequestHandlerExtra,
) -> pmcp::Result<pmcp::types::GetPromptResult> {
#handle_body
}
#metadata_body
}
impl #struct_name {
#with_state_method
}
pub fn #fn_name() -> #struct_name {
#constructor_default
}
};
Ok(expanded)
}