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
12 changes: 12 additions & 0 deletions crates/rmcp/src/model.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<D>(deserializer: D) -> Result<Self, D::Error>
where
Expand All @@ -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\"",
));
}
Comment thread
DaleSeo marked this conversation as resolved.

if helper.content.is_none()
&& helper.structured_content.is_none()
&& helper.is_error.is_none()
Expand Down
19 changes: 19 additions & 0 deletions crates/rmcp/src/model/mrtr.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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::<InputRequiredResult>(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!({
Expand Down
75 changes: 75 additions & 0 deletions crates/rmcp/tests/test_deserialization.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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::<CallToolResult>(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!({})));
Expand Down
Loading