1use rmcp::model::{ClientCapabilities, ElicitRequestParams, ErrorData, InputRequest, InputRequests, InputResponses};
2use serde::de::DeserializeOwned;
3
4pub const ELICITATION_UNSUPPORTED: &str = "This tool needs to ask the user for input, but the connected client does not support \
5 interactive input (MCP elicitation over protocol 2026-07-28 or newer).";
6
7pub fn input_requests_supported(capabilities: Option<&ClientCapabilities>, requests: &InputRequests) -> bool {
8 requests.values().all(|request| capabilities.is_some_and(|capabilities| supports_input(capabilities, request)))
9}
10
11pub fn parse_response<T: DeserializeOwned>(responses: &InputResponses, key: &str) -> Result<T, ErrorData> {
13 let response = responses
14 .get(key)
15 .ok_or_else(|| ErrorData::invalid_params(format!("missing input response for '{key}'"), None))?;
16 serde_json::from_value(response.clone())
17 .map_err(|e| ErrorData::invalid_params(format!("invalid input response for '{key}': {e}"), None))
18}
19
20fn supports_input(capabilities: &ClientCapabilities, request: &InputRequest) -> bool {
21 #[allow(unreachable_patterns)]
22 match request {
23 InputRequest::CreateMessage(_) => capabilities.sampling.is_some(),
24 InputRequest::Elicitation(request) => supports_elicitation(capabilities, &request.params),
25 InputRequest::ListRoots(_) => capabilities.roots.is_some(),
26 _ => false,
27 }
28}
29
30fn supports_elicitation(capabilities: &ClientCapabilities, params: &ElicitRequestParams) -> bool {
31 #[allow(unreachable_patterns)]
32 match params {
33 ElicitRequestParams::FormElicitationParams { .. } => capabilities
34 .elicitation
35 .as_ref()
36 .is_some_and(|elicitation| elicitation.form.is_some() || elicitation.url.is_none()),
37 ElicitRequestParams::UrlElicitationParams { .. } => {
38 capabilities.elicitation.as_ref().is_some_and(|elicitation| elicitation.url.is_some())
39 }
40 _ => false,
41 }
42}
43
44#[cfg(test)]
45mod tests {
46 use super::*;
47 use rmcp::model::{
48 ClientCapabilities, ElicitRequest, ElicitRequestParams, ElicitationCapability, ElicitationSchema,
49 FormElicitationCapability, InputRequest, UrlElicitationCapability,
50 };
51 use serde_json::json;
52
53 #[test]
54 fn validates_capabilities_against_every_request_in_the_batch() {
55 let requests = InputRequests::from([
56 ("form".to_string(), elicitation("Form")),
57 ("url".to_string(), url_elicitation("URL")),
58 ]);
59 let mut capabilities = ClientCapabilities::default();
60 capabilities.elicitation = Some(ElicitationCapability::new().with_form(FormElicitationCapability::new()));
61
62 assert!(!input_requests_supported(Some(&capabilities), &requests));
63 }
64
65 #[test]
66 fn validates_a_mixed_elicitation_batch_when_all_capabilities_are_present() {
67 let requests = InputRequests::from([
68 ("form".to_string(), elicitation("Form")),
69 ("url".to_string(), url_elicitation("URL")),
70 ]);
71 let mut capabilities = ClientCapabilities::default();
72 capabilities.elicitation = Some(
73 ElicitationCapability::new()
74 .with_form(FormElicitationCapability::new())
75 .with_url(UrlElicitationCapability::new()),
76 );
77
78 assert!(input_requests_supported(Some(&capabilities), &requests));
79 }
80
81 #[test]
82 fn parse_response_deserializes_a_key_from_the_batch() {
83 let responses = InputResponses::from([("answer".to_string(), json!(42))]);
84
85 assert_eq!(parse_response::<u64>(&responses, "answer").unwrap(), 42);
86 }
87
88 #[test]
89 fn parse_response_rejects_a_missing_key() {
90 let error = parse_response::<u64>(&InputResponses::new(), "answer").unwrap_err();
91
92 assert_eq!(error.message, "missing input response for 'answer'");
93 }
94
95 fn elicitation(message: &str) -> InputRequest {
96 InputRequest::Elicitation(ElicitRequest::new(ElicitRequestParams::FormElicitationParams {
97 meta: None,
98 message: message.to_string(),
99 requested_schema: ElicitationSchema::builder().build().unwrap(),
100 }))
101 }
102
103 fn url_elicitation(message: &str) -> InputRequest {
104 InputRequest::Elicitation(ElicitRequest::new(ElicitRequestParams::UrlElicitationParams {
105 meta: None,
106 message: message.to_string(),
107 url: "https://example.com/input".to_string(),
108 elicitation_id: "id".to_string(),
109 }))
110 }
111}