diff --git a/crates/context_server/src/protocol.rs b/crates/context_server/src/protocol.rs index b2b73890783..926b6bc177c 100644 --- a/crates/context_server/src/protocol.rs +++ b/crates/context_server/src/protocol.rs @@ -19,7 +19,7 @@ use std::sync::Arc; use std::time::Duration; -use anyhow::Result; +use anyhow::{Context as _, Result}; use collections::HashMap; use futures::{channel::oneshot, future::BoxFuture}; use gpui::AsyncApp; @@ -42,6 +42,18 @@ const MRTR_METHODS: &[&str] = &[ /// without input requests to fulfill there is no point retrying forever. const MAX_INPUT_REQUIRED_RETRIES: usize = 8; +fn result_type(result: &Value, require_result_type: bool) -> Result { + let Some(result_type) = result.get("resultType") else { + anyhow::ensure!( + !require_result_type, + "Modern MCP result is missing required resultType" + ); + return Ok(types::ResultType::Complete); + }; + + serde_json::from_value(result_type.clone()).context("invalid resultType in MCP result") +} + pub struct ModelContextProtocol { inner: Client, } @@ -94,21 +106,26 @@ impl ModelContextProtocol { .await; match discover_result { - Ok(result) => match serde_json::from_value::(result) { - Ok(response) => { - self.finish_discovery(client_info, request_meta, response) - .await - } - Err(error) => { - // A modern server must return a valid DiscoverResult; - // anything else is a legacy server answering an unknown - // method leniently. + Ok(result) => { + if result.get("resultType").is_none() { + // A legacy server may answer an unknown method leniently, + // including with a successful but empty result. log::debug!( - "treating malformed server/discover result as a legacy server: {error}" + "treating server/discover result without resultType as a legacy server" ); - self.legacy_initialize(client_info, None).await + return self.legacy_initialize(client_info, None).await; } - }, + + let result_type = result_type(&result, true)?; + anyhow::ensure!( + result_type == types::ResultType::Complete, + "server/discover returned an input_required result" + ); + let response = serde_json::from_value::(result) + .context("invalid server/discover result")?; + self.finish_discovery(client_info, request_meta, response) + .await + } Err(error) => { if let Some(rpc_error) = error.downcast_ref::() { // A well-formed UnsupportedProtocolVersionError is the @@ -331,16 +348,15 @@ impl InitializedContextServerProtocol { .request_with(T::METHOD, ¶ms, cancel_rx.as_mut(), timeout) .await?; - // Multi round-trip requests (MCP 2026-07-28): a modern server - // may answer with an interim `input_required` result instead of - // a final one. Legacy results have no `resultType` and are final - // by definition. - let input_required = self.is_modern() - && MRTR_METHODS.contains(&T::METHOD) - && result.get("resultType").and_then(|value| value.as_str()) - == Some("input_required"); - if !input_required { - return Ok(serde_json::from_value(result)?); + match result_type(&result, self.is_modern())? { + types::ResultType::Complete => return Ok(serde_json::from_value(result)?), + types::ResultType::InputRequired => { + anyhow::ensure!( + self.is_modern() && MRTR_METHODS.contains(&T::METHOD), + "Server returned an input_required result for unsupported method {}", + T::METHOD + ); + } } let interim: types::InputRequiredResult = serde_json::from_value(result)?; @@ -514,14 +530,13 @@ mod tests { }, cx.executor(), ) - .on_request::(|_params| async { - types::ListToolsResponse { - tools: Vec::new(), - next_cursor: None, - ttl_ms: Some(60_000), - cache_scope: Some(types::CacheScope::Private), - meta: None, - } + .on_raw_request(requests::ListTools::METHOD, |_params| async { + Ok(Some(serde_json::json!({ + "resultType": "complete", + "tools": [], + "ttlMs": 60_000, + "cacheScope": "private" + }))) }); let (protocol, received_messages) = connect(transport, cx).await; @@ -722,6 +737,40 @@ mod tests { assert!(!protocol.unwrap().is_modern()); } + #[gpui::test] + async fn test_unknown_discover_result_type_fails(cx: &mut TestAppContext) { + let transport = FakeTransport::new(cx.executor()) + .on_raw_request(requests::ServerDiscover::METHOD, |_params| async { + Ok(Some(serde_json::json!({ + "resultType": "future_result", + "supportedVersions": [types::VERSION_2026_07_28], + "capabilities": {} + }))) + }) + .on_request::(|_params| async { + types::InitializeResponse { + protocol_version: types::ProtocolVersion( + types::LATEST_LEGACY_PROTOCOL_VERSION.to_string(), + ), + server_info: client_info(), + capabilities: types::ServerCapabilities::default(), + meta: None, + } + }); + + let (protocol, received_messages) = connect(transport, cx).await; + let error = protocol.err().expect("discovery should fail"); + + assert!( + error.to_string().contains("invalid resultType"), + "unexpected error: {error}" + ); + assert!( + messages_with_method(&received_messages.lock(), requests::Initialize::METHOD) + .is_empty() + ); + } + #[gpui::test] async fn test_bare_unsupported_version_error_falls_back_to_legacy(cx: &mut TestAppContext) { // A -32022 without data.supported could be a legacy server's @@ -940,6 +989,74 @@ mod tests { ); } + #[gpui::test] + async fn test_modern_result_without_result_type_fails(cx: &mut TestAppContext) { + let transport = create_modern_fake_transport("modern-server", cx.executor()) + .on_raw_request(requests::ListTools::METHOD, |_params| async { + Ok(Some(serde_json::json!({ "tools": [] }))) + }); + + let (protocol, _) = connect(transport, cx).await; + let error = protocol + .unwrap() + .request::(()) + .await + .unwrap_err(); + + assert!( + error.to_string().contains("missing required resultType"), + "unexpected error: {error}" + ); + } + + #[gpui::test] + async fn test_unknown_result_type_fails(cx: &mut TestAppContext) { + let transport = create_modern_fake_transport("modern-server", cx.executor()) + .on_raw_request(requests::ListTools::METHOD, |_params| async { + Ok(Some(serde_json::json!({ + "resultType": "future_result", + "tools": [] + }))) + }); + + let (protocol, _) = connect(transport, cx).await; + let error = protocol + .unwrap() + .request::(()) + .await + .unwrap_err(); + + assert!( + error.to_string().contains("invalid resultType"), + "unexpected error: {error}" + ); + } + + #[gpui::test] + async fn test_input_required_result_for_unsupported_method_fails(cx: &mut TestAppContext) { + let transport = create_modern_fake_transport("modern-server", cx.executor()) + .on_raw_request(requests::ListTools::METHOD, |_params| async { + Ok(Some(serde_json::json!({ + "resultType": "input_required", + "tools": [] + }))) + }); + + let (protocol, _) = connect(transport, cx).await; + let error = protocol + .unwrap() + .request::(()) + .await + .unwrap_err(); + + assert!( + error + .to_string() + .contains("input_required result for unsupported method tools/list"), + "unexpected error: {error}" + ); + } + #[gpui::test] async fn test_legacy_results_without_result_type_are_final(cx: &mut TestAppContext) { let transport = create_fake_transport("legacy-server", cx.executor()) diff --git a/crates/context_server/src/test.rs b/crates/context_server/src/test.rs index 9094189b966..909c59f732d 100644 --- a/crates/context_server/src/test.rs +++ b/crates/context_server/src/test.rs @@ -9,8 +9,8 @@ use crate::{ client, transport::Transport, types::{ - DiscoverResponse, Implementation, InitializeResponse, ProtocolVersion, ServerCapabilities, - meta_keys, + DiscoverResponse, Implementation, InitializeResponse, ProtocolVersion, Request, + ServerCapabilities, meta_keys, }, }; @@ -45,11 +45,18 @@ pub fn create_modern_fake_transport_with_capabilities( executor: BackgroundExecutor, ) -> FakeTransport { let name = name.into(); - FakeTransport::new(executor).on_request::( + FakeTransport::new(executor).on_raw_request( + crate::types::requests::ServerDiscover::METHOD, move |_params| { let name = name.clone(); let capabilities = capabilities.clone(); - async move { create_discover_response(name, capabilities) } + async move { + let response = create_discover_response(name, capabilities); + let mut response = + serde_json::to_value(response).expect("discover response should serialize"); + response["resultType"] = serde_json::Value::String("complete".to_string()); + Ok(Some(response)) + } }, ) }