feat(sandbox): 使用逐调用截止时间驱动原生 MCP 会话
This commit is contained in:
@@ -620,6 +620,15 @@ mod tests {
|
|||||||
}
|
}
|
||||||
#[test]
|
#[test]
|
||||||
fn real_container_cannot_reach_ipv4_or_ipv6_loopback_listeners() {
|
fn real_container_cannot_reach_ipv4_or_ipv6_loopback_listeners() {
|
||||||
|
real_native_protocol_probes(false);
|
||||||
|
}
|
||||||
|
#[cfg(feature = "desktop")]
|
||||||
|
#[test]
|
||||||
|
#[ignore = "real MCP tools/call 60-second deadline acceptance; run explicitly"]
|
||||||
|
fn real_mcp_tool_deadline_terminates_hung_server_and_host_can_save() {
|
||||||
|
real_native_protocol_probes(true);
|
||||||
|
}
|
||||||
|
fn real_native_protocol_probes(_mcp_deadline: bool) {
|
||||||
use std::{
|
use std::{
|
||||||
net::{TcpListener, UdpSocket},
|
net::{TcpListener, UdpSocket},
|
||||||
os::windows::fs::OpenOptionsExt,
|
os::windows::fs::OpenOptionsExt,
|
||||||
@@ -856,6 +865,138 @@ mod tests {
|
|||||||
"from host broker"
|
"from host broker"
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
let mut mcp_modes = vec![
|
||||||
|
"mcp",
|
||||||
|
"mcp_wrong_id",
|
||||||
|
"mcp_cancel",
|
||||||
|
"mcp_remote_error",
|
||||||
|
"mcp_bad_version",
|
||||||
|
];
|
||||||
|
if _mcp_deadline {
|
||||||
|
mcp_modes.push("mcp_deadline");
|
||||||
|
}
|
||||||
|
for mode in mcp_modes {
|
||||||
|
use std::sync::{
|
||||||
|
atomic::{AtomicBool, Ordering},
|
||||||
|
Arc,
|
||||||
|
};
|
||||||
|
let mut mcp = claims.clone();
|
||||||
|
mcp.arguments = vec![mode.into()];
|
||||||
|
mcp.expires_at_ms = 120_000;
|
||||||
|
let permit = authority.issue(&mcp, 1).unwrap();
|
||||||
|
let prepared = context
|
||||||
|
.prepare(&authority, &permit, &mcp, &bound_entry, &broker, 2)
|
||||||
|
.unwrap();
|
||||||
|
let (suspended, io) = prepared
|
||||||
|
.create_suspended_with_stdio(&profile, &bound_entry)
|
||||||
|
.unwrap();
|
||||||
|
let running = unsafe { suspended.resume().unwrap() };
|
||||||
|
let mut session = crate::extension_mcp::Session::new(&running, io).unwrap();
|
||||||
|
let cancel = Arc::new(AtomicBool::new(false));
|
||||||
|
assert_eq!(
|
||||||
|
session
|
||||||
|
.call_tool("echo", serde_json::json!({}), &cancel)
|
||||||
|
.unwrap_err()
|
||||||
|
.code,
|
||||||
|
"EXTENSION_MCP_NOT_INITIALIZED"
|
||||||
|
);
|
||||||
|
if mode == "mcp_bad_version" {
|
||||||
|
assert_eq!(
|
||||||
|
session.initialize(&cancel).unwrap_err().code,
|
||||||
|
"EXTENSION_MCP_INITIALIZATION_INVALID"
|
||||||
|
);
|
||||||
|
} else {
|
||||||
|
assert_eq!(
|
||||||
|
session.initialize(&cancel).unwrap(),
|
||||||
|
crate::extension_mcp::PROTOCOL_VERSION
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
session.list_tools(None, &cancel).unwrap()["tools"][0]["name"],
|
||||||
|
"echo"
|
||||||
|
);
|
||||||
|
let cancellation = if mode == "mcp_cancel" {
|
||||||
|
let flag = Arc::clone(&cancel);
|
||||||
|
Some(std::thread::spawn(move || {
|
||||||
|
std::thread::sleep(std::time::Duration::from_millis(100));
|
||||||
|
flag.store(true, Ordering::Release);
|
||||||
|
}))
|
||||||
|
} else {
|
||||||
|
None
|
||||||
|
};
|
||||||
|
let tool_started = std::time::Instant::now();
|
||||||
|
let result = session.call_tool("echo", serde_json::json!({}), &cancel);
|
||||||
|
match mode {
|
||||||
|
"mcp_deadline" => {
|
||||||
|
assert_eq!(
|
||||||
|
result.unwrap_err().code,
|
||||||
|
"EXTENSION_TOOL_DEADLINE_EXCEEDED"
|
||||||
|
);
|
||||||
|
let elapsed = tool_started.elapsed();
|
||||||
|
assert!(
|
||||||
|
elapsed >= std::time::Duration::from_secs(59)
|
||||||
|
&& elapsed < std::time::Duration::from_secs(65)
|
||||||
|
);
|
||||||
|
eprintln!("real MCP tools/call deadline: {elapsed:?}");
|
||||||
|
let vault = tempfile::tempdir().unwrap();
|
||||||
|
let mut workspace =
|
||||||
|
crate::workspace::Workspace::open(vault.path()).unwrap();
|
||||||
|
workspace
|
||||||
|
.write(
|
||||||
|
"after-timeout.md",
|
||||||
|
"",
|
||||||
|
b"Host save after actual MCP deadline",
|
||||||
|
"local",
|
||||||
|
)
|
||||||
|
.unwrap();
|
||||||
|
drop(workspace);
|
||||||
|
let mut workspace =
|
||||||
|
crate::workspace::Workspace::open(vault.path()).unwrap();
|
||||||
|
assert_eq!(
|
||||||
|
workspace.read("after-timeout.md").unwrap().content,
|
||||||
|
"Host save after actual MCP deadline"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
"mcp_wrong_id" => assert_eq!(
|
||||||
|
result.unwrap_err().code,
|
||||||
|
"EXTENSION_MCP_RESPONSE_ID_MISMATCH"
|
||||||
|
),
|
||||||
|
"mcp_cancel" => {
|
||||||
|
assert_eq!(result.unwrap_err().code, "EXTENSION_MCP_CANCELLED");
|
||||||
|
cancellation.unwrap().join().unwrap();
|
||||||
|
}
|
||||||
|
"mcp_remote_error" => {
|
||||||
|
assert_eq!(result.unwrap_err().code, "EXTENSION_MCP_REMOTE_ERROR");
|
||||||
|
assert_eq!(
|
||||||
|
session
|
||||||
|
.call_tool("echo", serde_json::json!({}), &cancel)
|
||||||
|
.unwrap()["content"][0]["text"],
|
||||||
|
"native MCP success"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
_ => {
|
||||||
|
assert_eq!(result.unwrap()["content"][0]["text"], "native MCP success");
|
||||||
|
assert!(session.take_tools_changed());
|
||||||
|
assert!(!session.take_tools_changed());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if mode == "mcp_cancel" || mode == "mcp_wrong_id" {
|
||||||
|
assert_eq!(
|
||||||
|
session.list_tools(None, &cancel).unwrap_err().code,
|
||||||
|
"EXTENSION_MCP_SESSION_FAILED"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
session.shutdown().unwrap();
|
||||||
|
assert!(running
|
||||||
|
.wait(std::time::Duration::from_secs(5))
|
||||||
|
.unwrap()
|
||||||
|
.is_some());
|
||||||
|
let started = std::time::Instant::now();
|
||||||
|
while running.active_test_processes().unwrap() != 0 {
|
||||||
|
assert!(started.elapsed() < std::time::Duration::from_secs(5));
|
||||||
|
std::thread::sleep(std::time::Duration::from_millis(5));
|
||||||
|
}
|
||||||
|
}
|
||||||
for mode in ["wait_tree", "stderr_flood", "stdout_flood"] {
|
for mode in ["wait_tree", "stderr_flood", "stdout_flood"] {
|
||||||
let mut io_claims = claims.clone();
|
let mut io_claims = claims.clone();
|
||||||
io_claims.arguments = vec![mode.into()];
|
io_claims.arguments = vec![mode.into()];
|
||||||
|
|||||||
@@ -28,6 +28,8 @@ pub enum Event {
|
|||||||
}
|
}
|
||||||
struct State {
|
struct State {
|
||||||
stopped: AtomicBool,
|
stopped: AtomicBool,
|
||||||
|
#[cfg(test)]
|
||||||
|
writing: AtomicBool,
|
||||||
error: Mutex<Option<String>>,
|
error: Mutex<Option<String>>,
|
||||||
job: Job,
|
job: Job,
|
||||||
}
|
}
|
||||||
@@ -106,6 +108,8 @@ impl Pump {
|
|||||||
pub(crate) fn start(io: HostIo, job: Job) -> Result<Self> {
|
pub(crate) fn start(io: HostIo, job: Job) -> Result<Self> {
|
||||||
let state = Arc::new(State {
|
let state = Arc::new(State {
|
||||||
stopped: AtomicBool::new(false),
|
stopped: AtomicBool::new(false),
|
||||||
|
#[cfg(test)]
|
||||||
|
writing: AtomicBool::new(false),
|
||||||
error: Mutex::new(None),
|
error: Mutex::new(None),
|
||||||
job,
|
job,
|
||||||
});
|
});
|
||||||
@@ -130,7 +134,11 @@ impl Pump {
|
|||||||
match writes.recv_timeout(Duration::from_millis(20)) {
|
match writes.recv_timeout(Duration::from_millis(20)) {
|
||||||
Ok(frame) => {
|
Ok(frame) => {
|
||||||
state.check()?;
|
state.check()?;
|
||||||
|
#[cfg(test)]
|
||||||
|
state.writing.store(true, Ordering::Release);
|
||||||
write_frame(&mut input, &frame)?;
|
write_frame(&mut input, &frame)?;
|
||||||
|
#[cfg(test)]
|
||||||
|
state.writing.store(false, Ordering::Release);
|
||||||
}
|
}
|
||||||
Err(RecvTimeoutError::Timeout) => {}
|
Err(RecvTimeoutError::Timeout) => {}
|
||||||
Err(RecvTimeoutError::Disconnected) => return Ok(()),
|
Err(RecvTimeoutError::Disconnected) => return Ok(()),
|
||||||
@@ -201,7 +209,15 @@ impl Pump {
|
|||||||
}
|
}
|
||||||
/// Nonblocking admission; at most one pending write plus one in progress.
|
/// Nonblocking admission; at most one pending write plus one in progress.
|
||||||
pub fn send(&self, frame: Vec<u8>) -> Result<()> {
|
pub fn send(&self, frame: Vec<u8>) -> Result<()> {
|
||||||
|
self.send_wait(frame, Duration::ZERO)
|
||||||
|
}
|
||||||
|
/// Bounded admission for serial protocol notifications immediately followed
|
||||||
|
/// by a request; retries retain the same frame, without allocating copies.
|
||||||
|
pub(crate) fn send_wait(&self, mut frame: Vec<u8>, timeout: Duration) -> Result<()> {
|
||||||
self.state.check()?;
|
self.state.check()?;
|
||||||
|
if timeout > Duration::from_secs(1) {
|
||||||
|
return Err(HostError::new("EXTENSION_IO_TIMEOUT_INVALID"));
|
||||||
|
}
|
||||||
if frame.is_empty()
|
if frame.is_empty()
|
||||||
|| frame.len() > MAX_FRAME_BYTES
|
|| frame.len() > MAX_FRAME_BYTES
|
||||||
|| frame.contains(&b'\n')
|
|| frame.contains(&b'\n')
|
||||||
@@ -209,11 +225,25 @@ impl Pump {
|
|||||||
{
|
{
|
||||||
return Err(HostError::new("EXTENSION_PIPE_INVALID_FRAME"));
|
return Err(HostError::new("EXTENSION_PIPE_INVALID_FRAME"));
|
||||||
}
|
}
|
||||||
self.input
|
let sender = self
|
||||||
|
.input
|
||||||
.as_ref()
|
.as_ref()
|
||||||
.ok_or_else(|| HostError::new("EXTENSION_IO_INPUT_CLOSED"))?
|
.ok_or_else(|| HostError::new("EXTENSION_IO_INPUT_CLOSED"))?;
|
||||||
.try_send(frame)
|
let started = Instant::now();
|
||||||
.map_err(|_| HostError::new("EXTENSION_IO_INPUT_BACKPRESSURE"))
|
loop {
|
||||||
|
self.state.check()?;
|
||||||
|
match sender.try_send(frame) {
|
||||||
|
Ok(()) => return Ok(()),
|
||||||
|
Err(mpsc::TrySendError::Disconnected(_)) => {
|
||||||
|
return Err(HostError::new("EXTENSION_IO_INPUT_CLOSED"))
|
||||||
|
}
|
||||||
|
Err(mpsc::TrySendError::Full(value)) => frame = value,
|
||||||
|
}
|
||||||
|
if started.elapsed() >= timeout {
|
||||||
|
return Err(HostError::new("EXTENSION_IO_INPUT_BACKPRESSURE"));
|
||||||
|
}
|
||||||
|
std::thread::sleep(Duration::from_millis(1));
|
||||||
|
}
|
||||||
}
|
}
|
||||||
pub fn close_input(&mut self) {
|
pub fn close_input(&mut self) {
|
||||||
self.input.take();
|
self.input.take();
|
||||||
@@ -318,7 +348,7 @@ mod tests {
|
|||||||
);
|
);
|
||||||
pending != 0
|
pending != 0
|
||||||
});
|
});
|
||||||
if pending {
|
if pending && pump.state.writing.load(Ordering::Acquire) {
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
assert!(
|
assert!(
|
||||||
|
|||||||
@@ -0,0 +1,325 @@
|
|||||||
|
//! Serial MCP session over an already-authorized native instance. Package trust,
|
||||||
|
//! tool consent/schema validation and registry routing remain Host responsibilities.
|
||||||
|
use crate::{
|
||||||
|
extension_io::{Event, Pump},
|
||||||
|
extension_process::Running,
|
||||||
|
extension_stdio::HostIo,
|
||||||
|
workspace::{HostError, Result},
|
||||||
|
};
|
||||||
|
use serde::Deserialize;
|
||||||
|
use serde_json::{json, Value};
|
||||||
|
use std::{
|
||||||
|
sync::atomic::{AtomicBool, Ordering},
|
||||||
|
time::{Duration, Instant},
|
||||||
|
};
|
||||||
|
pub const PROTOCOL_VERSION: &str = "2025-11-25";
|
||||||
|
const VERSIONS: [&str; 4] = [PROTOCOL_VERSION, "2025-06-18", "2025-03-26", "2024-11-05"];
|
||||||
|
fn present<'de, D: serde::Deserializer<'de>, T: Deserialize<'de>>(
|
||||||
|
deserializer: D,
|
||||||
|
) -> std::result::Result<Option<T>, D::Error> {
|
||||||
|
T::deserialize(deserializer).map(Some)
|
||||||
|
}
|
||||||
|
#[derive(Deserialize)]
|
||||||
|
#[serde(deny_unknown_fields)]
|
||||||
|
struct Envelope {
|
||||||
|
jsonrpc: String,
|
||||||
|
#[serde(default, deserialize_with = "present")]
|
||||||
|
id: Option<Value>,
|
||||||
|
#[serde(default, deserialize_with = "present")]
|
||||||
|
method: Option<String>,
|
||||||
|
#[serde(default, deserialize_with = "present")]
|
||||||
|
params: Option<Value>,
|
||||||
|
#[serde(default, deserialize_with = "present")]
|
||||||
|
result: Option<Value>,
|
||||||
|
#[serde(default, deserialize_with = "present")]
|
||||||
|
error: Option<RemoteError>,
|
||||||
|
}
|
||||||
|
#[derive(Deserialize)]
|
||||||
|
#[serde(deny_unknown_fields)]
|
||||||
|
struct RemoteError {
|
||||||
|
code: i64,
|
||||||
|
message: String,
|
||||||
|
data: Option<Value>,
|
||||||
|
}
|
||||||
|
fn invalid() -> HostError {
|
||||||
|
HostError::new("EXTENSION_MCP_PROTOCOL_INVALID")
|
||||||
|
}
|
||||||
|
fn decode(bytes: &[u8]) -> Result<Envelope> {
|
||||||
|
let value: Envelope = serde_json::from_slice(bytes).map_err(|_| invalid())?;
|
||||||
|
if value.jsonrpc != "2.0" || value.params.as_ref().is_some_and(|v| !v.is_object()) {
|
||||||
|
return Err(invalid());
|
||||||
|
}
|
||||||
|
if let Some(id) = &value.id {
|
||||||
|
if id.as_i64().is_none() && !id.as_str().is_some_and(|s| !s.is_empty() && s.len() <= 128) {
|
||||||
|
return Err(invalid());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if let Some(method) = &value.method {
|
||||||
|
if method.is_empty()
|
||||||
|
|| method.len() > 256
|
||||||
|
|| method.chars().any(char::is_control)
|
||||||
|
|| value.result.is_some()
|
||||||
|
|| value.error.is_some()
|
||||||
|
{
|
||||||
|
return Err(invalid());
|
||||||
|
}
|
||||||
|
} else if value.id.is_none()
|
||||||
|
|| value.params.is_some()
|
||||||
|
|| value.result.is_some() == value.error.is_some()
|
||||||
|
{
|
||||||
|
return Err(invalid());
|
||||||
|
}
|
||||||
|
if value.result.as_ref().is_some_and(|v| !v.is_object()) {
|
||||||
|
return Err(invalid());
|
||||||
|
}
|
||||||
|
if let Some(error) = &value.error {
|
||||||
|
if error.message.len() > 4096
|
||||||
|
|| error.code < i32::MIN as i64
|
||||||
|
|| error.code > i32::MAX as i64
|
||||||
|
{
|
||||||
|
return Err(invalid());
|
||||||
|
}
|
||||||
|
// Error data is never propagated or logged; it may contain secrets.
|
||||||
|
let _ = &error.data;
|
||||||
|
}
|
||||||
|
Ok(value)
|
||||||
|
}
|
||||||
|
pub struct Session<'a, 'p> {
|
||||||
|
process: &'a Running<'p>,
|
||||||
|
pump: Pump,
|
||||||
|
next: u64,
|
||||||
|
ready: bool,
|
||||||
|
tools: bool,
|
||||||
|
failed: bool,
|
||||||
|
tools_changed: bool,
|
||||||
|
}
|
||||||
|
impl<'a, 'p> Session<'a, 'p> {
|
||||||
|
pub fn new(process: &'a Running<'p>, io: HostIo) -> Result<Self> {
|
||||||
|
Ok(Self {
|
||||||
|
process,
|
||||||
|
pump: process.start_io(io)?,
|
||||||
|
next: 1,
|
||||||
|
ready: false,
|
||||||
|
tools: false,
|
||||||
|
failed: false,
|
||||||
|
tools_changed: false,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
pub fn initialize(&mut self, cancel: &AtomicBool) -> Result<String> {
|
||||||
|
if self.ready {
|
||||||
|
return Err(HostError::new("EXTENSION_MCP_ALREADY_INITIALIZED"));
|
||||||
|
}
|
||||||
|
let value = self.request("initialize", json!({"protocolVersion":PROTOCOL_VERSION,"capabilities":{},"clientInfo":{"name":"OpenNexus","version":env!("CARGO_PKG_VERSION")}}), Duration::from_secs(10), cancel)?;
|
||||||
|
let version = value["protocolVersion"]
|
||||||
|
.as_str()
|
||||||
|
.filter(|v| VERSIONS.contains(v));
|
||||||
|
let info = value["serverInfo"].as_object();
|
||||||
|
let capabilities = value["capabilities"].as_object();
|
||||||
|
if version.is_none()
|
||||||
|
|| info.is_none()
|
||||||
|
|| capabilities.is_none()
|
||||||
|
|| !value["serverInfo"]["name"]
|
||||||
|
.as_str()
|
||||||
|
.is_some_and(|s| !s.is_empty() && s.len() <= 256)
|
||||||
|
|| !value["serverInfo"]["version"]
|
||||||
|
.as_str()
|
||||||
|
.is_some_and(|s| !s.is_empty() && s.len() <= 128)
|
||||||
|
|| value["capabilities"]
|
||||||
|
.get("tools")
|
||||||
|
.is_some_and(|v| !v.is_object())
|
||||||
|
{
|
||||||
|
self.abort();
|
||||||
|
return Err(HostError::new("EXTENSION_MCP_INITIALIZATION_INVALID"));
|
||||||
|
}
|
||||||
|
self.tools = value["capabilities"].get("tools").is_some();
|
||||||
|
if let Err(error) = self.send(json!({"jsonrpc":"2.0","method":"notifications/initialized"}))
|
||||||
|
{
|
||||||
|
self.abort();
|
||||||
|
return Err(error);
|
||||||
|
}
|
||||||
|
self.ready = true;
|
||||||
|
Ok(version.unwrap().into())
|
||||||
|
}
|
||||||
|
pub fn list_tools(&mut self, cursor: Option<&str>, cancel: &AtomicBool) -> Result<Value> {
|
||||||
|
self.require_tools()?;
|
||||||
|
if cursor.is_some_and(|s| s.is_empty() || s.len() > 1024) {
|
||||||
|
return Err(invalid());
|
||||||
|
}
|
||||||
|
let value = self.request(
|
||||||
|
"tools/list",
|
||||||
|
cursor.map_or_else(|| json!({}), |c| json!({"cursor":c})),
|
||||||
|
Duration::from_secs(10),
|
||||||
|
cancel,
|
||||||
|
)?;
|
||||||
|
if !value["tools"]
|
||||||
|
.as_array()
|
||||||
|
.is_some_and(|tools| tools.len() <= 500)
|
||||||
|
|| value
|
||||||
|
.get("nextCursor")
|
||||||
|
.is_some_and(|v| !v.as_str().is_some_and(|s| !s.is_empty() && s.len() <= 1024))
|
||||||
|
{
|
||||||
|
self.abort();
|
||||||
|
return Err(invalid());
|
||||||
|
}
|
||||||
|
Ok(value)
|
||||||
|
}
|
||||||
|
pub fn call_tool(
|
||||||
|
&mut self,
|
||||||
|
name: &str,
|
||||||
|
arguments: Value,
|
||||||
|
cancel: &AtomicBool,
|
||||||
|
) -> Result<Value> {
|
||||||
|
self.require_tools()?;
|
||||||
|
if name.is_empty()
|
||||||
|
|| name.len() > 256
|
||||||
|
|| name.chars().any(char::is_control)
|
||||||
|
|| !arguments.is_object()
|
||||||
|
{
|
||||||
|
return Err(invalid());
|
||||||
|
}
|
||||||
|
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
|
||||||
|
{
|
||||||
|
self.abort();
|
||||||
|
return Err(HostError::new("EXTENSION_MCP_TOOL_RESULT_INVALID"));
|
||||||
|
}
|
||||||
|
Ok(value)
|
||||||
|
}
|
||||||
|
fn require_tools(&self) -> Result<()> {
|
||||||
|
if self.failed {
|
||||||
|
return Err(HostError::new("EXTENSION_MCP_SESSION_FAILED"));
|
||||||
|
}
|
||||||
|
if !self.ready {
|
||||||
|
return Err(HostError::new("EXTENSION_MCP_NOT_INITIALIZED"));
|
||||||
|
}
|
||||||
|
if !self.tools {
|
||||||
|
return Err(HostError::new("EXTENSION_MCP_TOOLS_UNAVAILABLE"));
|
||||||
|
}
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
fn send(&self, value: Value) -> Result<()> {
|
||||||
|
self.pump.send_wait(
|
||||||
|
serde_json::to_vec(&value).map_err(|_| invalid())?,
|
||||||
|
Duration::from_millis(50),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
fn abort(&mut self) {
|
||||||
|
self.failed = true;
|
||||||
|
let _ = self.process.terminate();
|
||||||
|
}
|
||||||
|
fn request(
|
||||||
|
&mut self,
|
||||||
|
method: &str,
|
||||||
|
params: Value,
|
||||||
|
budget: Duration,
|
||||||
|
cancel: &AtomicBool,
|
||||||
|
) -> Result<Value> {
|
||||||
|
if self.failed {
|
||||||
|
return Err(HostError::new("EXTENSION_MCP_SESSION_FAILED"));
|
||||||
|
}
|
||||||
|
if cancel.load(Ordering::Acquire) {
|
||||||
|
return Err(HostError::new("EXTENSION_MCP_CANCELLED"));
|
||||||
|
}
|
||||||
|
let id = format!("opennexus.{}", self.next);
|
||||||
|
self.next = self.next.checked_add(1).ok_or_else(invalid)?;
|
||||||
|
let deadline = self.process.start_tool_call()?;
|
||||||
|
let started = Instant::now();
|
||||||
|
let result = (|| {
|
||||||
|
self.send(json!({"jsonrpc":"2.0","id":id,"method":method,"params":params}))?;
|
||||||
|
loop {
|
||||||
|
self.process.check_authorization()?;
|
||||||
|
deadline.check()?;
|
||||||
|
if cancel.load(Ordering::Acquire) || started.elapsed() >= budget {
|
||||||
|
// initialize cannot be cancelled at the protocol level.
|
||||||
|
if method != "initialize" {
|
||||||
|
let _ = self.send(json!({"jsonrpc":"2.0","method":"notifications/cancelled","params":{"requestId":id}}));
|
||||||
|
}
|
||||||
|
return Err(HostError::new(if cancel.load(Ordering::Acquire) {
|
||||||
|
"EXTENSION_MCP_CANCELLED"
|
||||||
|
} else {
|
||||||
|
"EXTENSION_MCP_TIMEOUT"
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
let received = self.pump.receive(Duration::from_millis(20));
|
||||||
|
self.process.check_authorization()?;
|
||||||
|
deadline.check()?;
|
||||||
|
let bytes = match received {
|
||||||
|
Ok(Event::Frame(bytes)) => bytes,
|
||||||
|
Ok(Event::Closed) => {
|
||||||
|
return Err(HostError::new("EXTENSION_MCP_CONNECTION_CLOSED"))
|
||||||
|
}
|
||||||
|
Err(error) if error.code == "EXTENSION_IO_TIMEOUT" => continue,
|
||||||
|
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;
|
||||||
|
}
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
if message.id != Some(json!(id)) {
|
||||||
|
return Err(HostError::new("EXTENSION_MCP_RESPONSE_ID_MISMATCH"));
|
||||||
|
}
|
||||||
|
return Ok(message.result);
|
||||||
|
}
|
||||||
|
})();
|
||||||
|
match result {
|
||||||
|
Ok(value) => {
|
||||||
|
if let Err(error) = deadline.finish() {
|
||||||
|
self.abort();
|
||||||
|
return Err(error);
|
||||||
|
}
|
||||||
|
value.ok_or_else(|| HostError::new("EXTENSION_MCP_REMOTE_ERROR"))
|
||||||
|
}
|
||||||
|
Err(error) => {
|
||||||
|
self.abort();
|
||||||
|
drop(deadline);
|
||||||
|
Err(error)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
pub fn take_tools_changed(&mut self) -> bool {
|
||||||
|
std::mem::take(&mut self.tools_changed)
|
||||||
|
}
|
||||||
|
pub fn shutdown(mut self) -> Result<()> {
|
||||||
|
self.pump.close_input();
|
||||||
|
let _ = self.process.wait(Duration::from_millis(200));
|
||||||
|
self.pump.shutdown()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
#[test]
|
||||||
|
fn malformed_and_ambiguous_envelopes_are_rejected() {
|
||||||
|
for input in [
|
||||||
|
r#"[]"#,
|
||||||
|
r#"{"jsonrpc":"2.0","id":null,"method":"ping"}"#,
|
||||||
|
r#"{"jsonrpc":"2.0","method":"ping","params":null}"#,
|
||||||
|
r#"{"jsonrpc":"2.0","id":1,"id":2,"result":{}}"#,
|
||||||
|
r#"{"jsonrpc":"1.0","id":1,"result":{}}"#,
|
||||||
|
r#"{"jsonrpc":"2.0","id":null,"result":{}}"#,
|
||||||
|
r#"{"jsonrpc":"2.0","id":1,"result":{},"error":{"code":1,"message":"bad"}}"#,
|
||||||
|
r#"{"jsonrpc":"2.0","id":1,"method":"ping","result":{}}"#,
|
||||||
|
] {
|
||||||
|
assert!(decode(input.as_bytes()).is_err(), "{input}");
|
||||||
|
}
|
||||||
|
assert!(decode(br#"{"jsonrpc":"2.0","id":"request","result":{}}"#).is_ok());
|
||||||
|
assert!(decode(br#"{"jsonrpc":"2.0","id":1,"method":"ping"}"#).is_ok());
|
||||||
|
assert!(
|
||||||
|
decode(br#"{"jsonrpc":"2.0","method":"notifications/tools/list_changed"}"#).is_ok()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -84,3 +84,6 @@ pub mod extension_stdio;
|
|||||||
|
|
||||||
#[cfg(all(windows, feature = "desktop"))]
|
#[cfg(all(windows, feature = "desktop"))]
|
||||||
pub mod extension_io;
|
pub mod extension_io;
|
||||||
|
|
||||||
|
#[cfg(all(windows, feature = "desktop"))]
|
||||||
|
pub mod extension_mcp;
|
||||||
|
|||||||
@@ -3,6 +3,42 @@ use std::net::{SocketAddr, TcpStream, UdpSocket};
|
|||||||
use std::time::Duration;
|
use std::time::Duration;
|
||||||
fn main() {
|
fn main() {
|
||||||
let args: Vec<_> = std::env::args().collect();
|
let args: Vec<_> = std::env::args().collect();
|
||||||
|
if args.get(1).is_some_and(|s| s.starts_with("mcp")) {
|
||||||
|
use std::io::{BufRead, Write};
|
||||||
|
fn read(reader: &mut impl BufRead) -> String { let mut line = String::new(); reader.read_line(&mut line).unwrap(); line }
|
||||||
|
fn id(line: &str) -> &str { line.split("\"id\":\"").nth(1).unwrap().split('"').next().unwrap() }
|
||||||
|
fn reply(id: &str, result: &str) { println!("{{\"jsonrpc\":\"2.0\",\"id\":\"{id}\",\"result\":{result}}}"); std::io::stdout().flush().unwrap(); }
|
||||||
|
let mut input = std::io::stdin().lock();
|
||||||
|
let request = read(&mut input);
|
||||||
|
assert!(request.contains("\"method\":\"initialize\""));
|
||||||
|
reply(id(&request), if args[1] == "mcp_bad_version" { r#"{"protocolVersion":"unknown","capabilities":{},"serverInfo":{"name":"fixture","version":"1"}}"# } else { r#"{"protocolVersion":"2025-11-25","capabilities":{"tools":{}},"serverInfo":{"name":"fixture","version":"1"}}"# });
|
||||||
|
assert!(read(&mut input).contains("notifications/initialized"));
|
||||||
|
let request = read(&mut input);
|
||||||
|
assert!(request.contains("tools/list"));
|
||||||
|
reply(id(&request), r#"{"tools":[{"name":"echo","inputSchema":{"type":"object"}}]}"#);
|
||||||
|
let mut request = read(&mut input);
|
||||||
|
assert!(request.contains("tools/call"));
|
||||||
|
if args[1] == "mcp_cancel" || args[1] == "mcp_deadline" {
|
||||||
|
let _child = std::process::Command::new(std::env::current_exe().unwrap()).arg("wait").spawn().unwrap();
|
||||||
|
std::thread::sleep(Duration::from_secs(120));
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
if args[1] == "mcp_remote_error" {
|
||||||
|
println!("{{\"jsonrpc\":\"2.0\",\"id\":\"{}\",\"error\":{{\"code\":-32602,\"message\":\"fixture error\"}}}}", id(&request));
|
||||||
|
std::io::stdout().flush().unwrap();
|
||||||
|
request = read(&mut input);
|
||||||
|
assert!(request.contains("tools/call"));
|
||||||
|
}
|
||||||
|
println!(r#"{{"jsonrpc":"2.0","id":"server-ping","method":"ping"}}"#);
|
||||||
|
std::io::stdout().flush().unwrap();
|
||||||
|
let ping = read(&mut input);
|
||||||
|
assert!(ping.contains("server-ping") && ping.contains("result"));
|
||||||
|
println!(r#"{{"jsonrpc":"2.0","method":"notifications/tools/list_changed"}}"#);
|
||||||
|
reply(if args[1] == "mcp_wrong_id" { "wrong-request" } else { id(&request) }, r#"{"content":[{"type":"text","text":"native MCP success"}]}"#);
|
||||||
|
// MCP server remains alive between calls until Host closes stdin.
|
||||||
|
while !read(&mut input).is_empty() {}
|
||||||
|
return;
|
||||||
|
}
|
||||||
if args.get(1).is_some_and(|s| s == "stderr_flood" || s == "stdout_flood") {
|
if args.get(1).is_some_and(|s| s == "stderr_flood" || s == "stdout_flood") {
|
||||||
use std::io::Write;
|
use std::io::Write;
|
||||||
let _child = std::process::Command::new(std::env::current_exe().unwrap()).arg("wait").spawn().unwrap();
|
let _child = std::process::Command::new(std::env::current_exe().unwrap()).arg("wait").spawn().unwrap();
|
||||||
|
|||||||
Reference in New Issue
Block a user