feat(mcp): 验证有界工具目录并使陈旧契约失效

This commit is contained in:
2026-09-09 01:00:02 +08:00
parent 7c7a637e38
commit 6373829243
6 changed files with 578 additions and 21 deletions
+47 -3
View File
@@ -871,6 +871,9 @@ mod tests {
"mcp_cancel",
"mcp_remote_error",
"mcp_bad_version",
"mcp_pages",
"mcp_bad_result",
"mcp_idle_change",
];
if _mcp_deadline {
mcp_modes.push("mcp_deadline");
@@ -911,8 +914,26 @@ mod tests {
crate::extension_mcp::PROTOCOL_VERSION
);
assert_eq!(
session.list_tools(None, &cancel).unwrap()["tools"][0]["name"],
"echo"
session
.call_tool("echo", serde_json::json!({}), &cancel)
.unwrap_err()
.code,
"EXTENSION_MCP_CATALOG_REQUIRED"
);
assert_eq!(session.refresh_tools(&cancel).unwrap()[0].name, "echo");
assert_eq!(
session
.call_tool("missing", serde_json::json!({}), &cancel)
.unwrap_err()
.code,
"EXTENSION_MCP_TOOL_NOT_FOUND"
);
assert_eq!(
session
.call_tool("echo", serde_json::json!({"unexpected":1}), &cancel)
.unwrap_err()
.code,
"EXTENSION_MCP_ARGUMENTS_INVALID"
);
let cancellation = if mode == "mcp_cancel" {
let flag = Arc::clone(&cancel);
@@ -926,6 +947,22 @@ mod tests {
let tool_started = std::time::Instant::now();
let result = session.call_tool("echo", serde_json::json!({}), &cancel);
match mode {
"mcp_bad_result" => assert_eq!(
result.unwrap_err().code,
"EXTENSION_MCP_TOOL_RESULT_INVALID"
),
"mcp_idle_change" => {
assert_eq!(result.unwrap()["structuredContent"]["ok"], true);
std::thread::sleep(std::time::Duration::from_millis(200));
assert_eq!(
session
.call_tool("echo", serde_json::json!({}), &cancel)
.unwrap_err()
.code,
"EXTENSION_MCP_CATALOG_REQUIRED"
);
assert!(session.take_tools_changed());
}
"mcp_deadline" => {
assert_eq!(
result.unwrap_err().code,
@@ -977,11 +1014,18 @@ mod tests {
assert_eq!(result.unwrap()["content"][0]["text"], "native MCP success");
assert!(session.take_tools_changed());
assert!(!session.take_tools_changed());
assert_eq!(
session
.call_tool("echo", serde_json::json!({}), &cancel)
.unwrap_err()
.code,
"EXTENSION_MCP_CATALOG_REQUIRED"
);
}
}
if mode == "mcp_cancel" || mode == "mcp_wrong_id" {
assert_eq!(
session.list_tools(None, &cancel).unwrap_err().code,
session.refresh_tools(&cancel).err().unwrap().code,
"EXTENSION_MCP_SESSION_FAILED"
);
}
+90 -15
View File
@@ -92,6 +92,7 @@ pub struct Session<'a, 'p> {
tools: bool,
failed: bool,
tools_changed: bool,
catalog: Option<crate::extension_mcp_tools::Catalog>,
}
impl<'a, 'p> Session<'a, 'p> {
pub fn new(process: &'a Running<'p>, io: HostIo) -> Result<Self> {
@@ -103,6 +104,7 @@ impl<'a, 'p> Session<'a, 'p> {
tools: false,
failed: false,
tools_changed: false,
catalog: None,
})
}
pub fn initialize(&mut self, cancel: &AtomicBool) -> Result<String> {
@@ -140,7 +142,43 @@ impl<'a, 'p> Session<'a, 'p> {
self.ready = true;
Ok(version.unwrap().into())
}
pub fn list_tools(&mut self, cursor: Option<&str>, cancel: &AtomicBool) -> Result<Value> {
pub fn refresh_tools(
&mut self,
cancel: &AtomicBool,
) -> Result<Vec<crate::extension_mcp_tools::Description>> {
self.require_tools()?;
self.drain_pending()?;
self.catalog = None;
self.tools_changed = false;
let started = Instant::now();
let catalog = crate::extension_mcp_tools::Catalog::discover(|cursor| {
let remaining = Duration::from_secs(20).saturating_sub(started.elapsed());
if remaining.is_zero() {
return Err(HostError::new("EXTENSION_MCP_CATALOG_TIMEOUT"));
}
self.list_tools_page(cursor, cancel, remaining.min(Duration::from_secs(10)))
});
let catalog = match catalog {
Ok(catalog) => catalog,
Err(error) => {
self.abort();
return Err(error);
}
};
self.drain_pending()?;
if self.tools_changed {
return Err(HostError::new("EXTENSION_MCP_CATALOG_CHANGED"));
}
let descriptions = catalog.descriptions();
self.catalog = Some(catalog);
Ok(descriptions)
}
fn list_tools_page(
&mut self,
cursor: Option<&str>,
cancel: &AtomicBool,
budget: Duration,
) -> Result<Value> {
self.require_tools()?;
if cursor.is_some_and(|s| s.is_empty() || s.len() > 1024) {
return Err(invalid());
@@ -148,7 +186,7 @@ impl<'a, 'p> Session<'a, 'p> {
let value = self.request(
"tools/list",
cursor.map_or_else(|| json!({}), |c| json!({"cursor":c})),
Duration::from_secs(10),
budget,
cancel,
)?;
if !value["tools"]
@@ -170,6 +208,7 @@ impl<'a, 'p> Session<'a, 'p> {
cancel: &AtomicBool,
) -> Result<Value> {
self.require_tools()?;
self.drain_pending()?;
if name.is_empty()
|| name.len() > 256
|| name.chars().any(char::is_control)
@@ -177,18 +216,21 @@ impl<'a, 'p> Session<'a, 'p> {
{
return Err(invalid());
}
let tool = self
.catalog
.as_ref()
.ok_or_else(|| HostError::new("EXTENSION_MCP_CATALOG_REQUIRED"))?
.tool(name)?;
tool.validate_arguments(&arguments)?;
let value = self.request(
"tools/call",
json!({"name":name,"arguments":arguments}),
Duration::from_secs(60),
cancel,
)?;
if !value["content"].is_array()
|| value.get("isError").is_some_and(|v| !v.is_boolean())
|| serde_json::to_vec(&value).map_err(|_| invalid())?.len() > 256 * 1024
{
if let Err(error) = tool.validate_result(&value) {
self.abort();
return Err(HostError::new("EXTENSION_MCP_TOOL_RESULT_INVALID"));
return Err(error);
}
Ok(value)
}
@@ -204,6 +246,46 @@ impl<'a, 'p> Session<'a, 'p> {
}
Ok(())
}
fn server_message(&mut self, message: &Envelope) -> Result<bool> {
let Some(method) = &message.method else {
return Ok(false);
};
if let Some(id) = &message.id {
self.send(if method == "ping" { json!({"jsonrpc":"2.0","id":id,"result":{}}) }
else { json!({"jsonrpc":"2.0","id":id,"error":{"code":-32601,"message":"Method not supported"}}) })?;
} else if method == "notifications/tools/list_changed" {
self.tools_changed = true;
self.catalog = None;
}
Ok(true)
}
/// Apply already-received notifications before selecting a cached contract.
/// The eventual registry loop must also call this while the instance is idle.
pub fn drain_pending(&mut self) -> Result<()> {
let result = (|| {
for _ in 0..128 {
self.process.check_authorization()?;
match self.pump.receive(Duration::ZERO) {
Err(error) if error.code == "EXTENSION_IO_TIMEOUT" => return Ok(()),
Err(error) => return Err(error),
Ok(Event::Closed) => {
return Err(HostError::new("EXTENSION_MCP_CONNECTION_CLOSED"))
}
Ok(Event::Frame(bytes)) => {
let message = decode(&bytes)?;
if !self.server_message(&message)? {
return Err(HostError::new("EXTENSION_MCP_UNEXPECTED_RESPONSE"));
}
}
}
}
Err(HostError::new("EXTENSION_IO_RATE_LIMITED"))
})();
if result.is_err() {
self.abort();
}
result
}
fn send(&self, value: Value) -> Result<()> {
self.pump.send_wait(
serde_json::to_vec(&value).map_err(|_| invalid())?,
@@ -259,14 +341,7 @@ impl<'a, 'p> Session<'a, 'p> {
Err(error) => return Err(error),
};
let message = decode(&bytes)?;
if let Some(method) = message.method {
if let Some(id) = message.id {
// No sampling/roots/elicitation capability was advertised.
self.send(if method == "ping" { json!({"jsonrpc":"2.0","id":id,"result":{}}) }
else { json!({"jsonrpc":"2.0","id":id,"error":{"code":-32601,"message":"Method not supported"}}) })?;
} else if method == "notifications/tools/list_changed" {
self.tools_changed = true;
}
if self.server_message(&message)? {
continue;
}
if message.id != Some(json!(id)) {
@@ -0,0 +1,408 @@
//! Bounded, offline MCP tool contracts. Descriptions/annotations are untrusted
//! data and never confer permissions. URI content is validated, never fetched.
use crate::workspace::{HostError, Result};
use base64::Engine;
use serde::Serialize;
use serde_json::Value;
use std::{
collections::{BTreeMap, BTreeSet},
sync::Arc,
};
const MAX_TOOLS: usize = 500;
const MAX_PAGES: usize = 100;
const MAX_CATALOG_BYTES: usize = 4 * 1024 * 1024;
fn invalid() -> HostError {
HostError::new("EXTENSION_MCP_TOOL_SCHEMA_INVALID")
}
fn bounds(value: &Value, depth: usize, nodes: &mut usize, schema: bool) -> Result<()> {
*nodes += 1;
if depth > 24 || *nodes > 4096 {
return Err(HostError::new("EXTENSION_MCP_TOOL_LIMIT"));
}
match value {
Value::Object(map) => {
for (key, child) in map {
if schema && matches!(key.as_str(), "$ref" | "$dynamicRef" | "$recursiveRef") {
return Err(HostError::new("EXTENSION_MCP_SCHEMA_REFERENCE"));
}
bounds(child, depth + 1, nodes, schema)?;
}
}
Value::Array(items) => {
for child in items {
bounds(child, depth + 1, nodes, schema)?;
}
}
_ => {}
}
Ok(())
}
fn bounded(value: &Value, bytes: usize, schema: bool) -> Result<()> {
bounds(value, 0, &mut 0, schema)?;
if serde_json::to_vec(value).map_err(|_| invalid())?.len() > bytes {
return Err(HostError::new("EXTENSION_MCP_TOOL_LIMIT"));
}
Ok(())
}
fn compile(schema: &Value) -> Result<jsonschema::Validator> {
bounded(schema, 64 * 1024, true)?;
if !schema.is_object()
|| schema["type"] != "object"
|| schema
.get("$schema")
.is_some_and(|v| v != "https://json-schema.org/draft/2020-12/schema")
{
return Err(invalid());
}
jsonschema::options()
.offline()
.with_draft(jsonschema::Draft::Draft202012)
.with_pattern_options(jsonschema::PatternOptions::regex())
.should_validate_formats(true)
.should_ignore_unknown_formats(false)
.build(schema)
.map_err(|_| invalid())
}
#[derive(Clone, Serialize)]
pub struct Description {
pub name: String,
pub title: Option<String>,
pub description: Option<String>,
pub input_schema: Value,
pub output_schema: Option<Value>,
}
pub struct Tool {
description: Description,
input: jsonschema::Validator,
output: Option<jsonschema::Validator>,
}
impl Tool {
fn parse(value: &Value) -> Result<Self> {
bounded(value, 192 * 1024, false)?;
let name = value["name"]
.as_str()
.filter(|s| {
!s.is_empty()
&& s.len() <= 128
&& s.bytes()
.all(|b| b.is_ascii_alphanumeric() || b"_.-".contains(&b))
})
.ok_or_else(invalid)?;
fn optional(value: &Value, key: &str, max: usize) -> Result<Option<String>> {
value
.get(key)
.map(|v| {
v.as_str()
.filter(|s| s.len() <= max)
.map(str::to_owned)
.ok_or_else(invalid)
})
.transpose()
}
if let Some(execution) = value.get("execution") {
if !execution.is_object()
|| execution
.get("taskSupport")
.is_some_and(|v| v != "optional" && v != "forbidden")
{
return Err(HostError::new("EXTENSION_MCP_TASKS_UNSUPPORTED"));
}
}
let input = compile(&value["inputSchema"])?;
let output = value.get("outputSchema").map(compile).transpose()?;
Ok(Self {
description: Description {
name: name.into(),
title: optional(value, "title", 256)?,
description: optional(value, "description", 16 * 1024)?,
input_schema: value["inputSchema"].clone(),
output_schema: value.get("outputSchema").cloned(),
},
input,
output,
})
}
pub fn validate_arguments(&self, arguments: &Value) -> Result<()> {
bounded(arguments, 256 * 1024, false)?;
if !arguments.is_object() || !self.input.is_valid(arguments) {
return Err(HostError::new("EXTENSION_MCP_ARGUMENTS_INVALID"));
}
Ok(())
}
pub fn validate_result(&self, result: &Value) -> Result<()> {
bounded(result, 256 * 1024, false)?;
let bad = || HostError::new("EXTENSION_MCP_TOOL_RESULT_INVALID");
if !result.is_object() || result.get("isError").is_some_and(|v| !v.is_boolean()) {
return Err(bad());
}
let content = result["content"]
.as_array()
.filter(|v| v.len() <= 128)
.ok_or_else(bad)?;
fn uri(value: &Value) -> bool {
value
.as_str()
.is_some_and(|s| s.len() <= 2048 && reqwest::Url::parse(s).is_ok())
}
fn binary(value: &Value) -> bool {
value
.as_str()
.is_some_and(|s| base64::engine::general_purpose::STANDARD.decode(s).is_ok())
}
for item in content {
let valid = match item["type"].as_str() {
Some("text") => item["text"].is_string(),
Some("image" | "audio") => {
binary(&item["data"])
&& item["mimeType"].as_str().is_some_and(|s| {
s.len() <= 128
&& s.starts_with(if item["type"] == "image" {
"image/"
} else {
"audio/"
})
&& !s.chars().any(char::is_control)
})
}
Some("resource_link") => {
uri(&item["uri"])
&& item["name"]
.as_str()
.is_some_and(|s| !s.is_empty() && s.len() <= 256)
}
Some("resource") => {
let resource = &item["resource"];
uri(&resource["uri"])
&& match (resource.get("text"), resource.get("blob")) {
(Some(text), None) => text.is_string(),
(None, Some(blob)) => binary(blob),
_ => false,
}
}
_ => false,
};
if !valid {
return Err(bad());
}
}
if result
.get("structuredContent")
.is_some_and(|v| !v.is_object())
{
return Err(bad());
}
if result["isError"] != true {
if let Some(output) = &self.output {
if !result["structuredContent"].is_object()
|| !output.is_valid(&result["structuredContent"])
{
return Err(bad());
}
}
}
Ok(())
}
}
pub struct Catalog {
tools: BTreeMap<String, Arc<Tool>>,
}
impl Catalog {
pub fn discover(mut fetch: impl FnMut(Option<&str>) -> Result<Value>) -> Result<Self> {
let mut tools = BTreeMap::new();
let mut seen = BTreeSet::new();
let mut cursor: Option<String> = None;
let mut bytes = 0;
for _ in 0..MAX_PAGES {
let page = fetch(cursor.as_deref())?;
bytes += serde_json::to_vec(&page).map_err(|_| invalid())?.len();
if bytes > MAX_CATALOG_BYTES {
return Err(HostError::new("EXTENSION_MCP_CATALOG_LIMIT"));
}
let entries = page["tools"].as_array().ok_or_else(invalid)?;
if entries.len() > MAX_TOOLS - tools.len() {
return Err(HostError::new("EXTENSION_MCP_CATALOG_LIMIT"));
}
for entry in entries {
let tool = Tool::parse(entry)?;
if tools
.insert(tool.description.name.clone(), Arc::new(tool))
.is_some()
{
return Err(HostError::new("EXTENSION_MCP_TOOL_DUPLICATE"));
}
}
match page.get("nextCursor") {
None => return Ok(Self { tools }),
Some(next) => {
let next = next
.as_str()
.filter(|s| !s.is_empty() && s.len() <= 1024)
.ok_or_else(invalid)?;
if !seen.insert(next.to_owned()) {
return Err(HostError::new("EXTENSION_MCP_PAGINATION_CYCLE"));
}
cursor = Some(next.into());
}
}
}
Err(HostError::new("EXTENSION_MCP_CATALOG_LIMIT"))
}
pub fn descriptions(&self) -> Vec<Description> {
self.tools
.values()
.map(|tool| tool.description.clone())
.collect()
}
pub fn tool(&self, name: &str) -> Result<Arc<Tool>> {
self.tools
.get(name)
.cloned()
.ok_or_else(|| HostError::new("EXTENSION_MCP_TOOL_NOT_FOUND"))
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
fn tool() -> Value {
json!({"name":"Echo.V1","inputSchema":{"type":"object","required":["count"],"additionalProperties":false,"properties":{"count":{"type":"integer","minimum":1,"maximum":3}}},"outputSchema":{"type":"object","required":["ok"],"properties":{"ok":{"type":"boolean"}}},"annotations":{"readOnlyHint":true}})
}
#[test]
fn paginated_catalog_is_atomic_unique_and_bounded() {
let mut pages = 0;
let catalog = Catalog::discover(|cursor| {
pages += 1;
if pages == 1 {
assert!(cursor.is_none());
Ok(json!({"tools":[tool()],"nextCursor":"two"}))
} else {
assert_eq!(cursor, Some("two"));
let mut second = tool();
second["name"] = json!("second");
Ok(json!({"tools":[second]}))
}
})
.unwrap();
assert_eq!(catalog.descriptions().len(), 2);
assert_eq!(
catalog.tool("missing").err().unwrap().code,
"EXTENSION_MCP_TOOL_NOT_FOUND"
);
let mut pages = 0;
assert_eq!(
Catalog::discover(|_| {
pages += 1;
Ok(if pages == 1 {
json!({"tools":[tool()],"nextCursor":"two"})
} else {
json!({"tools":[tool()]})
})
})
.err()
.unwrap()
.code,
"EXTENSION_MCP_TOOL_DUPLICATE"
);
assert_eq!(
Catalog::discover(|_| Ok(json!({"tools":[],"nextCursor":"same"})))
.err()
.unwrap()
.code,
"EXTENSION_MCP_PAGINATION_CYCLE"
);
assert_eq!(
Catalog::discover(|_| Ok(json!({"tools":vec![tool();501]})))
.err()
.unwrap()
.code,
"EXTENSION_MCP_CATALOG_LIMIT"
);
let mut pages = 0;
assert_eq!(
Catalog::discover(|_| {
pages += 1;
Ok(json!({"tools":[],"nextCursor":pages.to_string()}))
})
.err()
.unwrap()
.code,
"EXTENSION_MCP_CATALOG_LIMIT"
);
assert_eq!(pages, 100);
assert_eq!(
Catalog::discover(|_| Ok(json!({"tools":[],"padding":"x".repeat(MAX_CATALOG_BYTES)})))
.err()
.unwrap()
.code,
"EXTENSION_MCP_CATALOG_LIMIT"
);
}
#[test]
fn contracts_validate_arguments_structured_results_and_never_resolve_references() {
let contract = Tool::parse(&tool()).unwrap();
contract.validate_arguments(&json!({"count":2})).unwrap();
for arguments in [
json!({}),
json!({"count":0}),
json!({"count":"2"}),
json!({"count":2,"extra":true}),
] {
assert_eq!(
contract.validate_arguments(&arguments).unwrap_err().code,
"EXTENSION_MCP_ARGUMENTS_INVALID"
);
}
let result =
json!({"content":[{"type":"text","text":"ok"}],"structuredContent":{"ok":true}});
contract.validate_result(&result).unwrap();
for result in [
json!({"content":[]}),
json!({"content":[],"structuredContent":{"ok":"wrong"}}),
json!({"content":[{"type":"image","mimeType":"image/png","data":"invalid base64 !"}],"structuredContent":{"ok":true}}),
json!({"content":[{"type":"text","text":42}],"structuredContent":{"ok":true}}),
] {
assert!(contract.validate_result(&result).is_err());
}
contract
.validate_result(
&json!({"content":[{"type":"text","text":"tool failed"}],"isError":true}),
)
.unwrap();
for target in [
"https://example.invalid/schema",
"file:///C:/private",
"#/$defs/recursive",
] {
let mut value = tool();
value["inputSchema"]["$ref"] = json!(target);
assert_eq!(
Tool::parse(&value).err().unwrap().code,
"EXTENSION_MCP_SCHEMA_REFERENCE"
);
}
for bad in [
json!({"type":"object","properties":{"value":{"pattern":"(?=x)"}}}),
json!({"type":"object","$schema":"unknown"}),
json!({"type":"array"}),
] {
let mut value = tool();
value["inputSchema"] = bad;
assert!(Tool::parse(&value).is_err());
}
let mut value = tool();
value["execution"] = json!({"taskSupport":"required"});
assert_eq!(
Tool::parse(&value).err().unwrap().code,
"EXTENSION_MCP_TASKS_UNSUPPORTED"
);
let mut deep = json!({});
for _ in 0..25 {
deep = json!({"nested":deep});
}
assert_eq!(
contract.validate_arguments(&deep).unwrap_err().code,
"EXTENSION_MCP_TOOL_LIMIT"
);
}
}
+3
View File
@@ -87,3 +87,6 @@ pub mod extension_io;
#[cfg(all(windows, feature = "desktop"))]
pub mod extension_mcp;
#[cfg(all(windows, feature = "desktop"))]
pub mod extension_mcp_tools;