From 57b2e563bf851c1782c5707206b6eff8080e9df5 Mon Sep 17 00:00:00 2001 From: Ralf Anton Beier Date: Tue, 30 Sep 2025 17:27:55 +0200 Subject: [PATCH] fix(macros): flatten struct parameters to remove confusing wrapper Single struct parameters now deserialize from entire args object instead of requiring nested "params" key. AI agents send {"name": "Alice"} directly instead of {"params": {"name": "Alice"}}. - Add smart type detection (primitives vs custom structs) - Flatten custom structs, extract primitives by name - Add runtime tests validating flattening behavior - Update hello-world example docs BREAKING CHANGE: Single struct parameter handling changed Bumps version to 0.11.0 --- Cargo.lock | 26 ++-- Cargo.toml | 26 ++-- examples/hello-world/src/main.rs | 6 + mcp-macros/src/mcp_tool.rs | 187 +++++++++++++++++++++----- mcp-macros/tests/dual_pattern_test.rs | 106 +++++++++++++++ mcp-security-middleware/Cargo.toml | 2 +- 6 files changed, 293 insertions(+), 60 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 31ff18f9..77aac606 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2297,7 +2297,7 @@ dependencies = [ [[package]] name = "pulseengine-mcp-auth" -version = "0.10.1" +version = "0.11.0" dependencies = [ "aes-gcm", "anyhow", @@ -2336,7 +2336,7 @@ dependencies = [ [[package]] name = "pulseengine-mcp-cli" -version = "0.10.1" +version = "0.11.0" dependencies = [ "clap", "pulseengine-mcp-cli-derive", @@ -2355,7 +2355,7 @@ dependencies = [ [[package]] name = "pulseengine-mcp-cli-derive" -version = "0.10.1" +version = "0.11.0" dependencies = [ "async-trait", "clap", @@ -2373,7 +2373,7 @@ dependencies = [ [[package]] name = "pulseengine-mcp-external-validation" -version = "0.10.1" +version = "0.11.0" dependencies = [ "anyhow", "arbitrary", @@ -2411,7 +2411,7 @@ dependencies = [ [[package]] name = "pulseengine-mcp-integration-tests" -version = "0.10.1" +version = "0.11.0" dependencies = [ "anyhow", "assert_matches", @@ -2439,7 +2439,7 @@ dependencies = [ [[package]] name = "pulseengine-mcp-logging" -version = "0.10.1" +version = "0.11.0" dependencies = [ "chrono", "hex", @@ -2458,7 +2458,7 @@ dependencies = [ [[package]] name = "pulseengine-mcp-macros" -version = "0.10.1" +version = "0.11.0" dependencies = [ "anyhow", "async-trait", @@ -2484,7 +2484,7 @@ dependencies = [ [[package]] name = "pulseengine-mcp-monitoring" -version = "0.10.1" +version = "0.11.0" dependencies = [ "anyhow", "chrono", @@ -2504,7 +2504,7 @@ dependencies = [ [[package]] name = "pulseengine-mcp-protocol" -version = "0.10.1" +version = "0.11.0" dependencies = [ "async-trait", "chrono", @@ -2520,7 +2520,7 @@ dependencies = [ [[package]] name = "pulseengine-mcp-security" -version = "0.10.1" +version = "0.11.0" dependencies = [ "anyhow", "async-trait", @@ -2542,7 +2542,7 @@ dependencies = [ [[package]] name = "pulseengine-mcp-security-middleware" -version = "0.10.1" +version = "0.11.0" dependencies = [ "anyhow", "assert_matches", @@ -2574,7 +2574,7 @@ dependencies = [ [[package]] name = "pulseengine-mcp-server" -version = "0.10.1" +version = "0.11.0" dependencies = [ "anyhow", "async-trait", @@ -2602,7 +2602,7 @@ dependencies = [ [[package]] name = "pulseengine-mcp-transport" -version = "0.10.1" +version = "0.11.0" dependencies = [ "anyhow", "async-stream", diff --git a/Cargo.toml b/Cargo.toml index de92b81b..bf4fb783 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -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" @@ -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" diff --git a/examples/hello-world/src/main.rs b/examples/hello-world/src/main.rs index 359e1f2c..52889233 100644 --- a/examples/hello-world/src/main.rs +++ b/examples/hello-world/src/main.rs @@ -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 { let name = params.name.unwrap_or_else(|| "World".to_string()); Ok(format!("Hello, {name}!")) diff --git a/mcp-macros/src/mcp_tool.rs b/mcp-macros/src/mcp_tool.rs index 511cec6b..73e62c37 100644 --- a/mcp-macros/src/mcp_tool.rs +++ b/mcp-macros/src/mcp_tool.rs @@ -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 = ¶m_names[0]; + let param_type = ¶m_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)) + )), + } + }); } } } @@ -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 @@ -946,7 +1027,9 @@ fn generate_method_call_with_params( ) -> syn::Result { 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 @@ -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 { diff --git a/mcp-macros/tests/dual_pattern_test.rs b/mcp-macros/tests/dual_pattern_test.rs index 531d92b7..5ce26445 100644 --- a/mcp-macros/tests/dual_pattern_test.rs +++ b/mcp-macros/tests/dual_pattern_test.rs @@ -179,3 +179,109 @@ async fn test_all_three_patterns_generate_schemas() { panic!("Tools should be available"); } } + +#[tokio::test] +async fn test_struct_parameter_flattening_at_runtime() { + use pulseengine_mcp_protocol::CallToolRequestParam; + use pulseengine_mcp_server::McpToolsProvider; + + let server = DualPatternServer; + + // Test that struct params work with FLAT arguments + println!("Testing flat arguments (correct behavior)..."); + let request = CallToolRequestParam { + name: "rich_struct_tool".to_string(), + arguments: Some(serde_json::json!({ + "message": "Hello World", + "count": 42 + })), + }; + + let result = server.call_tool_impl(request).await; + assert!( + result.is_ok(), + "Flat arguments should work: {:?}", + result.err() + ); + println!("✅ Flat arguments work correctly!"); + + // Test that NESTED arguments (old broken behavior) now fail + println!("Testing nested 'params' wrapper (should fail)..."); + let nested_request = CallToolRequestParam { + name: "rich_struct_tool".to_string(), + arguments: Some(serde_json::json!({ + "params": { + "message": "Hello", + "count": 42 + } + })), + }; + + let result = server.call_tool_impl(nested_request).await; + assert!( + result.is_err(), + "Nested 'params' wrapper should not work - AI agents send flat args" + ); + println!("✅ Nested params correctly rejected!"); +} + +#[tokio::test] +async fn test_multi_param_still_works() { + use pulseengine_mcp_protocol::CallToolRequestParam; + use pulseengine_mcp_server::McpToolsProvider; + + let server = DualPatternServer; + + // Multi-parameter should continue working with named properties + println!("Testing multi-parameter tool..."); + let request = CallToolRequestParam { + name: "multi_param_tool".to_string(), + arguments: Some(serde_json::json!({ + "name": "Alice", + "age": 30, + "active": true + })), + }; + + let result = server.call_tool_impl(request).await; + assert!( + result.is_ok(), + "Multi-parameter should work: {:?}", + result.err() + ); + + if let Ok(call_result) = result { + if let Some(pulseengine_mcp_protocol::Content::Text { text }) = call_result.content.first() + { + assert!(text.contains("Alice")); + assert!(text.contains("30")); + assert!(text.contains("true")); + println!("✅ Multi-parameter tool works correctly!"); + } + } +} + +#[tokio::test] +async fn test_optional_struct_fields_work() { + use pulseengine_mcp_protocol::CallToolRequestParam; + use pulseengine_mcp_server::McpToolsProvider; + + let server = DualPatternServer; + + // Test struct with optional field - only send required field + println!("Testing optional field handling..."); + let request = CallToolRequestParam { + name: "rich_struct_tool".to_string(), + arguments: Some(serde_json::json!({ + "message": "Hello without count" + })), + }; + + let result = server.call_tool_impl(request).await; + assert!( + result.is_ok(), + "Optional fields should work when omitted: {:?}", + result.err() + ); + println!("✅ Optional field handling works correctly!"); +} diff --git a/mcp-security-middleware/Cargo.toml b/mcp-security-middleware/Cargo.toml index 1a4618b3..44c10223 100644 --- a/mcp-security-middleware/Cargo.toml +++ b/mcp-security-middleware/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "pulseengine-mcp-security-middleware" -version = "0.10.1" +version = "0.11.0" rust-version = "1.88" edition = "2024" license = "MIT OR Apache-2.0"