Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 13 additions & 13 deletions Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

26 changes: 13 additions & 13 deletions Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,7 @@ members = [
resolver = "2"

[workspace.package]
version = "0.10.1"
version = "0.11.0"
rust-version = "1.88"
edition = "2024"
license = "MIT OR Apache-2.0"
Expand Down Expand Up @@ -108,18 +108,18 @@ assert_matches = "1.5"
serde_yaml = "0.9"

# Framework internal dependencies (published versions)
pulseengine-mcp-protocol = { version = "0.10.1", path = "mcp-protocol" }
pulseengine-mcp-logging = { version = "0.10.1", path = "mcp-logging" }
pulseengine-mcp-auth = { version = "0.10.1", path = "mcp-auth" }
pulseengine-mcp-security = { version = "0.10.1", path = "mcp-security" }
pulseengine-mcp-security-middleware = { version = "0.10.1", path = "mcp-security-middleware" }
pulseengine-mcp-monitoring = { version = "0.10.1", path = "mcp-monitoring" }
pulseengine-mcp-transport = { version = "0.10.1", path = "mcp-transport" }
pulseengine-mcp-cli = { version = "0.10.1", path = "mcp-cli" }
pulseengine-mcp-cli-derive = { version = "0.10.1", path = "mcp-cli-derive" }
pulseengine-mcp-server = { version = "0.10.1", path = "mcp-server" }
pulseengine-mcp-macros = { version = "0.10.1", path = "mcp-macros" }
pulseengine-mcp-external-validation = { version = "0.10.1", path = "mcp-external-validation" }
pulseengine-mcp-protocol = { version = "0.11.0", path = "mcp-protocol" }
pulseengine-mcp-logging = { version = "0.11.0", path = "mcp-logging" }
pulseengine-mcp-auth = { version = "0.11.0", path = "mcp-auth" }
pulseengine-mcp-security = { version = "0.11.0", path = "mcp-security" }
pulseengine-mcp-security-middleware = { version = "0.11.0", path = "mcp-security-middleware" }
pulseengine-mcp-monitoring = { version = "0.11.0", path = "mcp-monitoring" }
pulseengine-mcp-transport = { version = "0.11.0", path = "mcp-transport" }
pulseengine-mcp-cli = { version = "0.11.0", path = "mcp-cli" }
pulseengine-mcp-cli-derive = { version = "0.11.0", path = "mcp-cli-derive" }
pulseengine-mcp-server = { version = "0.11.0", path = "mcp-server" }
pulseengine-mcp-macros = { version = "0.11.0", path = "mcp-macros" }
pulseengine-mcp-external-validation = { version = "0.11.0", path = "mcp-external-validation" }

[profile.release]
opt-level = "s"
Expand Down
6 changes: 6 additions & 0 deletions examples/hello-world/src/main.rs
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,12 @@ pub struct HelloWorld;
#[mcp_tools]
impl HelloWorld {
/// Say hello to someone
///
/// AI agents send flat arguments: `{"name": "Alice"}`
/// NOT nested: `{"params": {"name": "Alice"}}`
///
/// The parameter name "params" is just an internal variable -
/// AI agents see the struct's fields directly in the schema.
pub async fn say_hello(&self, params: SayHelloParams) -> anyhow::Result<String> {
let name = params.name.unwrap_or_else(|| "World".to_string());
Ok(format!("Hello, {name}!"))
Expand Down
187 changes: 154 additions & 33 deletions mcp-macros/src/mcp_tool.rs
Original file line number Diff line number Diff line change
Expand Up @@ -631,26 +631,69 @@ fn extract_parameters(

param_names.push(param_name.clone());
param_types.push(param_type.clone());
}
}
}
}

// Generate parameter extraction code with consistent error handling
if is_option_type(param_type) {
param_fields.push(quote! {
args.get(stringify!(#param_name))
.and_then(|v| serde_json::from_value(v.clone()).ok())
});
} else {
param_fields.push(quote! {
match args.get(stringify!(#param_name))
.and_then(|v| serde_json::from_value(v.clone()).ok()) {
Some(value) => value,
None => return Err(pulseengine_mcp_protocol::Error::invalid_params(
format!("Missing required parameter '{}' for tool '{}'. Expected type: {}",
stringify!(#param_name), #tool_name, stringify!(#param_type))
)),
}
});
// Detect single parameter case
if param_names.len() == 1 {
let param_name = &param_names[0];
let param_type = &param_types[0];

// Check if it's a custom struct (not primitive/std type)
if is_primitive_or_std_type(param_type) {
// Primitive - extract by name (standard behavior)
if is_option_type(param_type) {
param_fields.push(quote! {
args.get(stringify!(#param_name))
.and_then(|v| serde_json::from_value(v.clone()).ok())
});
} else {
param_fields.push(quote! {
match args.get(stringify!(#param_name))
.and_then(|v| serde_json::from_value(v.clone()).ok()) {
Some(value) => value,
None => return Err(pulseengine_mcp_protocol::Error::invalid_params(
format!("Missing required parameter '{}' for tool '{}'. Expected type: {}",
stringify!(#param_name), #tool_name, stringify!(#param_type))
)),
}
});
}
} else {
// Custom struct - deserialize entire args object (flattened)
param_fields.push(quote! {
match serde_json::from_value::<#param_type>(
serde_json::Value::Object(args.clone())
) {
Ok(value) => value,
Err(e) => return Err(pulseengine_mcp_protocol::Error::invalid_params(
format!("Failed to deserialize parameters for tool '{}': {}", #tool_name, e)
)),
}
});
}
} else {
// Multi-parameter - extract by name
for (param_name, param_type) in param_names.iter().zip(param_types.iter()) {
// Generate parameter extraction code with consistent error handling
if is_option_type(param_type) {
param_fields.push(quote! {
args.get(stringify!(#param_name))
.and_then(|v| serde_json::from_value(v.clone()).ok())
});
} else {
param_fields.push(quote! {
match args.get(stringify!(#param_name))
.and_then(|v| serde_json::from_value(v.clone()).ok()) {
Some(value) => value,
None => return Err(pulseengine_mcp_protocol::Error::invalid_params(
format!("Missing required parameter '{}' for tool '{}'. Expected type: {}",
stringify!(#param_name), #tool_name, stringify!(#param_type))
)),
}
});
}
}
}
Expand Down Expand Up @@ -879,6 +922,44 @@ fn extract_option_inner_type(ty: &syn::Type) -> (bool, &syn::Type) {
(false, ty)
}

/// Check if a type is a primitive or standard library type (not a custom struct)
fn is_primitive_or_std_type(ty: &syn::Type) -> bool {
match ty {
syn::Type::Path(type_path) => {
if let Some(segment) = type_path.path.segments.last() {
matches!(
segment.ident.to_string().as_str(),
"String"
| "str"
| "i8"
| "i16"
| "i32"
| "i64"
| "isize"
| "u8"
| "u16"
| "u32"
| "u64"
| "usize"
| "f32"
| "f64"
| "bool"
| "Vec"
| "HashMap"
| "BTreeMap"
| "HashSet"
| "BTreeSet"
| "Option"
| "Value" // serde_json::Value
)
} else {
false
}
}
_ => false,
}
}

/// Generate JSON schema for a specific type
fn generate_type_schema_for_type(ty: &syn::Type) -> TokenStream {
// Convert Rust type to JSON schema
Expand Down Expand Up @@ -946,7 +1027,9 @@ fn generate_method_call_with_params(
) -> syn::Result<TokenStream> {
let mut param_declarations = Vec::new();
let mut param_names = Vec::new();
let mut param_types = Vec::new();

// Collect all parameters (skip self)
for input in &sig.inputs {
match input {
syn::FnArg::Receiver(_) => continue, // Skip self
Expand All @@ -956,27 +1039,65 @@ fn generate_method_call_with_params(
let param_type = &*pat_type.ty;

param_names.push(param_name);

// Generate parameter extraction based on whether it's optional
if is_option_type(param_type) {
param_declarations.push(quote! {
let #param_name = args.get(stringify!(#param_name))
.and_then(|v| serde_json::from_value(v.clone()).ok());
});
} else {
param_declarations.push(quote! {
let #param_name = args.get(stringify!(#param_name))
.and_then(|v| serde_json::from_value(v.clone()).ok())
.ok_or_else(|| pulseengine_mcp_protocol::Error::invalid_params(
format!("Missing required parameter '{}'", stringify!(#param_name))
))?;
});
}
param_types.push(param_type);
}
}
}
}

// Detect single parameter case
if param_names.len() == 1 {
let param_name = param_names[0];
let param_type = param_types[0];

// Check if it's a custom struct (not primitive/std type)
if is_primitive_or_std_type(param_type) {
// Primitive - extract by name (standard behavior)
if is_option_type(param_type) {
param_declarations.push(quote! {
let #param_name = args.get(stringify!(#param_name))
.and_then(|v| serde_json::from_value(v.clone()).ok());
});
} else {
param_declarations.push(quote! {
let #param_name = args.get(stringify!(#param_name))
.and_then(|v| serde_json::from_value(v.clone()).ok())
.ok_or_else(|| pulseengine_mcp_protocol::Error::invalid_params(
format!("Missing required parameter '{}'", stringify!(#param_name))
))?;
});
}
} else {
// Custom struct - deserialize entire args object (flattened)
param_declarations.push(quote! {
let #param_name: #param_type = serde_json::from_value(
serde_json::Value::Object(args.clone())
).map_err(|e| pulseengine_mcp_protocol::Error::invalid_params(
format!("Failed to deserialize parameters: {}", e)
))?;
});
}
} else {
// Multi-parameter or no parameters - extract by name
for (param_name, param_type) in param_names.iter().zip(param_types.iter()) {
// Generate parameter extraction based on whether it's optional
if is_option_type(param_type) {
param_declarations.push(quote! {
let #param_name = args.get(stringify!(#param_name))
.and_then(|v| serde_json::from_value(v.clone()).ok());
});
} else {
param_declarations.push(quote! {
let #param_name = args.get(stringify!(#param_name))
.and_then(|v| serde_json::from_value(v.clone()).ok())
.ok_or_else(|| pulseengine_mcp_protocol::Error::invalid_params(
format!("Missing required parameter '{}'", stringify!(#param_name))
))?;
});
}
}
}

if param_declarations.is_empty() {
// No parameters - call method directly
if is_async {
Expand Down
Loading
Loading