mirror of
https://github.com/zed-industries/zed.git
synced 2026-08-20 06:24:24 +00:00
context_server: Validate modern result types
This commit is contained in:
parent
56aa297fc7
commit
bb170a949e
2 changed files with 159 additions and 35 deletions
|
|
@ -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<types::ResultType> {
|
||||
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::<types::DiscoverResponse>(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::<types::DiscoverResponse>(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::<client::Error>() {
|
||||
// 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::<requests::ListTools, _>(|_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::<requests::Initialize, _>(|_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::<requests::ListTools>(())
|
||||
.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::<requests::ListTools>(())
|
||||
.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::<requests::ListTools>(())
|
||||
.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())
|
||||
|
|
|
|||
|
|
@ -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::<crate::types::requests::ServerDiscover, _>(
|
||||
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))
|
||||
}
|
||||
},
|
||||
)
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue