diff --git a/crates/rmcp/src/model.rs b/crates/rmcp/src/model.rs index 06ed855e..7ee93d82 100644 --- a/crates/rmcp/src/model.rs +++ b/crates/rmcp/src/model.rs @@ -3814,6 +3814,8 @@ pub struct CallToolResult { // 2. Requires at least one known field to be present, so that `CallToolResult` doesn't // greedily match arbitrary JSON objects when used inside `#[serde(untagged)]` enums // (e.g. `ServerResult`), which would shadow `CustomResult`. +// 3. Rejects `resultType: "input_required"` so untagged `ServerResult` +// decoding selects `InputRequiredResult` and preserves input requests and state. impl<'de> Deserialize<'de> for CallToolResult { fn deserialize(deserializer: D) -> Result where @@ -3833,6 +3835,16 @@ impl<'de> Deserialize<'de> for CallToolResult { let helper = Helper::deserialize(deserializer)?; + if helper + .result_type + .as_ref() + .is_some_and(ResultType::is_input_required) + { + return Err(serde::de::Error::custom( + "CallToolResult cannot use resultType \"input_required\"", + )); + } + if helper.content.is_none() && helper.structured_content.is_none() && helper.is_error.is_none() diff --git a/crates/rmcp/src/model/mrtr.rs b/crates/rmcp/src/model/mrtr.rs index 40621a7c..8dac9a39 100644 --- a/crates/rmcp/src/model/mrtr.rs +++ b/crates/rmcp/src/model/mrtr.rs @@ -274,6 +274,12 @@ impl<'de> Deserialize<'de> for InputRequiredResult { } } + if helper.input_requests.is_none() && helper.request_state.is_none() { + return Err(serde::de::Error::custom( + "InputRequiredResult requires at least one of inputRequests or requestState", + )); + } + Ok(InputRequiredResult { result_type: ResultType::INPUT_REQUIRED, input_requests: helper.input_requests, @@ -451,6 +457,19 @@ mod tests { ); } + #[test] + fn rejects_missing_input_requests_and_request_state() { + let json = serde_json::json!({ + "resultType": "input_required", + "_meta": {} + }); + let err = serde_json::from_value::(json).unwrap_err(); + assert!( + err.to_string().contains("inputRequests or requestState"), + "error should mention the required continuation fields, got: {err}" + ); + } + #[test] fn rejects_missing_result_type() { let json = serde_json::json!({ diff --git a/crates/rmcp/tests/test_deserialization.rs b/crates/rmcp/tests/test_deserialization.rs index c3d08cd5..85707701 100644 --- a/crates/rmcp/tests/test_deserialization.rs +++ b/crates/rmcp/tests/test_deserialization.rs @@ -70,6 +70,81 @@ mod untagged_server_result { ); } + #[test] + fn input_required_result_with_meta_deserializes_to_correct_variant() { + let result = parse_result(wrap_response(json!({ + "resultType": "input_required", + "inputRequests": { + "username": { + "method": "elicitation/create", + "params": { + "message": "Please provide your username", + "requestedSchema": { + "type": "object", + "properties": { + "username": { "type": "string" } + }, + "required": ["username"] + } + } + } + }, + "requestState": "opaque-state", + "_meta": { + "io.modelcontextprotocol/serverInfo": { + "name": "test-server", + "version": "1.0.0" + } + } + }))); + + let ServerResult::InputRequiredResult(result) = result else { + panic!("expected InputRequiredResult, got {result:?}"); + }; + assert!( + result.input_requests.is_some_and(|requests| { + requests.len() == 1 && requests.contains_key("username") + }) + ); + assert_eq!(result.request_state.as_deref(), Some("opaque-state")); + assert_eq!( + result + .meta + .as_ref() + .and_then(|meta| meta.get("io.modelcontextprotocol/serverInfo")), + Some(&json!({ + "name": "test-server", + "version": "1.0.0" + })) + ); + } + + #[test] + fn call_tool_result_rejects_input_required_discriminator() { + assert!( + serde_json::from_value::(json!({ + "resultType": "input_required", + "requestState": "opaque-state", + "_meta": {} + })) + .is_err() + ); + } + + #[test] + fn invalid_input_required_result_falls_through_to_custom_result() { + let payload = json!({ + "resultType": "input_required", + "_meta": {} + }); + let result = parse_result(wrap_response(payload.clone())); + + let ServerResult::CustomResult(result) = result else { + panic!("expected CustomResult, got {result:?}"); + }; + assert_eq!(result.0, payload); + } + #[test] fn empty_object_deserializes_to_empty_result() { let result = parse_result(wrap_response(json!({})));