use std::path::Path;
use std::process::Stdio;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
mod common;
use common::{mock_config, unique_dir};
fn spawn_mcp(dir: &Path) -> tokio::process::Child {
tokio::process::Command::new(env!("CARGO_BIN_EXE_code-repo-wiki"))
.args(["mcp", "--config", "mcp-test.toml", "--root", "."])
.current_dir(dir)
.env("RUST_LOG", "off")
.env_remove("OPENAI_API_KEY")
.stdin(Stdio::piped())
.stdout(Stdio::piped())
.stderr(Stdio::null())
.kill_on_drop(true)
.spawn()
.expect("启动 code-repo-wiki mcp 失败")
}
async fn rpc_call(stdin: &mut tokio::process::ChildStdin, stdout: &mut tokio::process::ChildStdout, id: u64, method: &str, params: serde_json::Value) -> serde_json::Value {
let req = serde_json::json!({
"jsonrpc": "2.0",
"id": id,
"method": method,
"params": params,
});
stdin
.write_all(serde_json::to_string(&req).unwrap().as_bytes())
.await
.unwrap();
stdin.write_all(b"\n").await.unwrap();
loop {
let mut line = Vec::new();
let mut byte = [0u8; 1];
loop {
let n = stdout.read(&mut byte).await.unwrap();
if n == 0 {
panic!("MCP 进程提前退出(stdout EOF)");
}
line.push(byte[0]);
if line.ends_with(b"\n") {
break;
}
}
let trimmed = String::from_utf8(line).expect("MCP 响应应为合法 UTF-8");
let trimmed = trimmed.trim();
if !trimmed.is_empty() {
return serde_json::from_str(trimmed).unwrap_or_else(|e| panic!("响应非 JSON: {trimmed}: {e}"));
}
}
}
#[tokio::test]
async fn test_mcp_initialize_lists_tools_and_calls() {
let dir = unique_dir("server");
let _ = std::fs::remove_dir_all(&dir);
std::fs::create_dir_all(dir.join(".code-repo-wiki")).unwrap();
let config = mock_config();
std::fs::write(dir.join("mcp-test.toml"), &config).unwrap();
std::fs::create_dir_all(dir.join("src")).unwrap();
std::fs::write(dir.join("src").join("main.rs"), "pub fn hello_world() {}\n").unwrap();
let mut child = spawn_mcp(&dir);
let mut stdin = child.stdin.take().unwrap();
let mut stdout = child.stdout.take().unwrap();
let resp = rpc_call(
&mut stdin,
&mut stdout,
1,
"initialize",
serde_json::json!({
"protocolVersion": "2025-06-18",
"capabilities": {},
"clientInfo": {"name": "test", "version": "0.0.1"}
}),
)
.await;
assert!(resp["result"]["protocolVersion"].is_string(), "握手应返回协议版本: {resp}");
let _ = rpc_call(&mut stdin, &mut stdout, 2, "notifications/initialized", serde_json::json!({})).await;
let resp = rpc_call(&mut stdin, &mut stdout, 3, "tools/list", serde_json::json!({})).await;
let tools = resp["result"]["tools"].as_array().expect("tools 应为数组");
let names: Vec<&str> = tools
.iter()
.filter_map(|t| t["name"].as_str())
.collect();
for expected in ["search", "ast_search", "read_wiki_page", "read_card", "status"] {
assert!(names.contains(&expected), "工具 {expected} 未注册, 实际: {names:?}");
}
let resp = rpc_call(
&mut stdin,
&mut stdout,
4,
"tools/call",
serde_json::json!({
"name": "ast_search",
"arguments": {"symbol": "hello_world"}
}),
)
.await;
let text = resp["result"]["content"][0]["text"].as_str().expect("工具结果应有 text");
assert!(text.contains("hello_world"), "ast_search 应找到符号: {text}");
let resp = rpc_call(
&mut stdin,
&mut stdout,
5,
"tools/call",
serde_json::json!({"name": "status", "arguments": {}}),
)
.await;
let text = resp["result"]["content"][0]["text"].as_str().expect("工具结果应有 text");
assert!(text.contains("Wiki"), "status 应返回 wiki 状态: {text}");
let resp = rpc_call(
&mut stdin,
&mut stdout,
6,
"tools/call",
serde_json::json!({"name": "read_wiki_page", "arguments": {"page": "architecture"}}),
)
.await;
let text = resp["result"]["content"][0]["text"].as_str().expect("工具结果应有 text");
assert!(text.contains("code-repo-wiki generate"), "未生成时应有引导提示: {text}");
let resp = rpc_call(
&mut stdin,
&mut stdout,
7,
"tools/call",
serde_json::json!({"name": "search", "arguments": {"query": "hello", "engine": "text"}}),
)
.await;
assert!(
resp["result"]["isError"].as_bool() == Some(true) || resp["result"]["content"][0]["text"].is_string(),
"search 应有结果或错误信息: {resp}"
);
drop(stdin);
let _ = child.wait().await;
let _ = std::fs::remove_dir_all(&dir);
}
#[tokio::test]
async fn test_mcp_lang_traversal_rejected() {
let dir = unique_dir("traversal");
let _ = std::fs::remove_dir_all(&dir);
std::fs::create_dir_all(dir.join(".code-repo-wiki")).unwrap();
let config = mock_config();
std::fs::write(dir.join("mcp-test.toml"), &config).unwrap();
std::fs::create_dir_all(dir.join("src")).unwrap();
std::fs::write(dir.join("src").join("main.rs"), "pub fn hello_world() {}\n").unwrap();
let secret = dir.parent().unwrap().join(format!("secret_{}.md", std::process::id()));
std::fs::write(&secret, "SECRET-CONTENT").unwrap();
std::fs::create_dir_all(dir.join(".code-repo-wiki").join("wiki").join("zh")).unwrap();
std::fs::write(
dir.join(".code-repo-wiki").join("wiki").join("zh").join("architecture.md"),
"ok-content",
)
.unwrap();
let mut child = spawn_mcp(&dir);
let mut stdin = child.stdin.take().unwrap();
let mut stdout = child.stdout.take().unwrap();
let resp = rpc_call(
&mut stdin,
&mut stdout,
1,
"initialize",
serde_json::json!({
"protocolVersion": "2025-06-18",
"capabilities": {},
"clientInfo": {"name": "test", "version": "0.0.1"}
}),
)
.await;
assert!(resp["result"]["protocolVersion"].is_string());
let _ = rpc_call(&mut stdin, &mut stdout, 2, "notifications/initialized", serde_json::json!({})).await;
let resp = rpc_call(
&mut stdin,
&mut stdout,
3,
"tools/call",
serde_json::json!({"name": "read_wiki_page", "arguments": {"page": "secret", "lang": "../.."}}),
)
.await;
let text = resp["result"]["content"][0]["text"].as_str().expect("工具结果应有 text");
assert!(!text.contains("SECRET-CONTENT"), "穿越必须被拒绝, 泄漏: {text}");
assert!(text.contains("非法语言名"), "应返回明确的校验错误: {text}");
let resp = rpc_call(
&mut stdin,
&mut stdout,
4,
"tools/call",
serde_json::json!({"name": "read_card", "arguments": {"card": "secret", "lang": "../../x"}}),
)
.await;
let text = resp["result"]["content"][0]["text"].as_str().expect("工具结果应有 text");
assert!(!text.contains("SECRET-CONTENT"), "read_card 穿越必须被拒绝: {text}");
assert!(text.contains("非法语言名"), "read_card 应返回明确的校验错误: {text}");
let resp = rpc_call(
&mut stdin,
&mut stdout,
5,
"tools/call",
serde_json::json!({"name": "read_wiki_page", "arguments": {"page": "architecture", "lang": "zh"}}),
)
.await;
let text = resp["result"]["content"][0]["text"].as_str().expect("工具结果应有 text");
assert!(text.contains("ok-content"), "合法 lang 应正常读取: {text}");
drop(stdin);
let _ = child.wait().await;
let _ = std::fs::remove_file(&secret);
let _ = std::fs::remove_dir_all(&dir);
}