Skip to main content

code_repo_wiki/
mcp.rs

1//! MCP (Model Context Protocol) server(P1-3)
2//!
3//! 通过 stdio 暴露 code-repo-wiki 能力给任意 MCP 客户端(Claude Code、Cline、
4//! 自研 Agent 等):代码搜索、AST 符号查找、Wiki 页面/卡片读取、状态查询。
5//! 实现基于官方 Rust MCP SDK(rmcp 3.x,crates.io 维护)。
6//!
7//! 设计:server 无状态(每次工具调用现场加载配置与索引),与 CLI 共享
8//! 同一 lib 入口(execute_search / execute_ast_search),不复制业务逻辑。
9//! 项目根由 `--root` 参数指定(缺省解析 cwd 的 config.toml 项目级配置)。
10
11use std::path::{Path, PathBuf};
12
13use anyhow::Result;
14use rmcp::handler::server::router::tool::ToolRouter;
15use rmcp::handler::server::wrapper::Parameters;
16use rmcp::model::{ServerCapabilities, ServerInfo};
17use rmcp::service::{QuitReason, ServiceExt};
18use rmcp::{ServerHandler, schemars, tool, tool_handler, tool_router};
19
20use crate::project::ProjectRoot;
21
22/// MCP server:工具路由 + 配置/根注入
23#[derive(Debug, Clone)]
24pub struct RepoWikiMcp {
25    /// 工具路由(tool_handler 宏访问)
26    #[expect(dead_code, reason = "tool_handler 宏访问此路由字段")]
27    tool_router: ToolRouter<Self>,
28    /// 配置文件路径(config.toml 项目级、全局或 --config 指定;None=默认链)
29    config_path: Option<PathBuf>,
30    /// 项目根(代码扫描/git 定位基准)
31    root: ProjectRoot,
32}
33
34impl RepoWikiMcp {
35    /// 创建 server
36    pub fn new(config_path: Option<PathBuf>, root: ProjectRoot) -> Self {
37        Self {
38            tool_router: Self::tool_router(),
39            config_path,
40            root,
41        }
42    }
43}
44
45// ============ 工具定义(#[tool_router] 块内,方法可访问 self 配置) ============
46
47/// 搜索请求参数
48#[derive(Debug, serde::Deserialize, schemars::JsonSchema)]
49struct SearchRequest {
50    /// 搜索关键词,如 "config 加载"
51    query: String,
52    /// 返回结果数量(默认取配置 search.default_top_k)
53    top_k: Option<usize>,
54    /// 搜索引擎: text / semantic / hybrid(默认取配置文件 default_engine)
55    engine: Option<String>,
56}
57
58/// AST 查找请求参数
59#[derive(Debug, serde::Deserialize, schemars::JsonSchema)]
60struct AstSearchRequest {
61    /// 要查找的符号名(函数/结构体/trait/类等)
62    symbol: String,
63    /// 源语言(rust/python/go/...);省略时按文件扩展名自动推断
64    language: Option<String>,
65}
66
67/// 读页面请求参数
68#[derive(Debug, serde::Deserialize, schemars::JsonSchema)]
69struct ReadPageRequest {
70    /// 页面文件名(不含 .md),如 src_config、architecture、overview、api
71    page: String,
72    /// 语言目录(默认取配置 wiki.language)
73    lang: Option<String>,
74}
75
76/// 读卡片请求参数
77#[derive(Debug, serde::Deserialize, schemars::JsonSchema)]
78struct ReadCardRequest {
79    /// 卡片名(模块名),如 src_config、crate_net
80    card: String,
81    /// 语言目录(默认取配置 wiki.language)
82    lang: Option<String>,
83}
84
85#[tool_router(router = tool_router)]
86impl RepoWikiMcp {
87    /// 搜索代码实体:按关键词返回匹配的函数/结构体/类及文件位置(text/semantic/hybrid 引擎)
88    #[tool(description = "搜索代码实体:按关键词返回匹配的函数/结构体/类及文件位置(text/semantic/hybrid 引擎,与 CLI code-repo-wiki search 等价;需先运行 code-repo-wiki generate 构建搜索索引)")]
89    async fn search(&self, Parameters(SearchRequest { query, top_k, engine }): Parameters<SearchRequest>) -> String {
90        // 配置完整性检查:搜索前确认配置可加载(错误早暴露);v22 起
91        // 引擎/条数默认值硬编码,配置内容不再被本函数使用
92        let _config = match crate::config::resolve_mcp_config(self.config_path.as_deref(), &self.root) {
93            Ok(c) => c,
94            Err(e) => return format!("配置加载失败: {e}"),
95        };
96        let engine_type = match engine.as_deref() {
97            Some("text") => crate::config::schema::SearchEngineType::Text,
98            Some("semantic") => crate::config::schema::SearchEngineType::Semantic,
99            Some("hybrid") => crate::config::schema::SearchEngineType::Hybrid,
100            Some(other) => return format!("不支持的搜索引擎: {other}(可选: text/semantic/hybrid)"),
101            None => crate::config::schema::SEARCH_DEFAULT_ENGINE,
102        };
103        let top_k = clamp_top_k(top_k.unwrap_or(crate::config::schema::SEARCH_DEFAULT_TOP_K));
104        match crate::execute_search(self.config_path.as_deref(), &self.root, &query, top_k, &engine_type) {
105            Ok(hits) if hits.is_empty() => "未找到匹配结果".to_string(),
106            Ok(hits) => {
107                let mut out = format!("找到 {} 个结果:\n", hits.len());
108                for (i, hit) in hits.iter().enumerate() {
109                    let file = hit.node.file_path.as_deref().unwrap_or("-");
110                    let loc = match hit.node.line_range {
111                        Some((s, e)) => format!("{file}:{s}-{e}"),
112                        None => file.to_string(),
113                    };
114                    let sig = hit.node.signature.as_deref().unwrap_or(&hit.node.name);
115                    out.push_str(&format!("{}. `{sig}` — {loc}\n", i + 1));
116                }
117                out
118            }
119            Err(e) => format!("搜索失败: {e}"),
120        }
121    }
122
123    /// AST 精确符号查找:扫描源文件定位 函数/结构体/类 定义的 文件+行号+签名(不依赖搜索索引)
124    #[tool(description = "AST 精确符号查找:扫描源文件定位函数/结构体/类定义的 文件+行号+签名(与 CLI code-repo-wiki ast-search 等价)")]
125    async fn ast_search(&self, Parameters(AstSearchRequest { symbol, language }): Parameters<AstSearchRequest>) -> String {
126        match crate::execute_ast_search(self.config_path.as_deref(), &self.root, &symbol, language.as_deref()) {
127            Ok(hits) if hits.is_empty() => format!("未找到符号 \"{symbol}\" 的定义"),
128            Ok(hits) => {
129                let mut out = format!("找到 {} 个定义:\n", hits.len());
130                for (i, hit) in hits.iter().enumerate() {
131                    let sig = hit.node.signature.as_deref().unwrap_or(&hit.node.name);
132                    let loc = match (&hit.node.file_path, hit.node.line_range) {
133                        (Some(f), Some((s, e))) => format!("{f}:{s}-{e}"),
134                        (Some(f), None) => f.clone(),
135                        _ => "(unknown)".to_string(),
136                    };
137                    out.push_str(&format!("{}. {sig} — {loc}\n", i + 1));
138                }
139                out
140            }
141            Err(e) => format!("符号查找失败: {e}"),
142        }
143    }
144
145    /// 读取已生成的 Wiki 页面内容(模块页/架构概览/项目概览/api)
146    #[tool(description = "读取已生成的 Wiki 页面内容(wiki/{lang}/{page}.md,如 src_config、architecture、overview、api;需先运行 code-repo-wiki generate,未生成的页面报错)")]
147    async fn read_wiki_page(&self, Parameters(ReadPageRequest { page, lang }): Parameters<ReadPageRequest>) -> String {
148        let config = match crate::config::resolve_mcp_config(self.config_path.as_deref(), &self.root) {
149            Ok(c) => c,
150            Err(e) => return format!("配置加载失败: {e}"),
151        };
152        let lang = lang.unwrap_or_else(|| config.wiki.language.clone());
153        // 语言目录净化(S1,工具暴露给任意 Agent):lang 直接 join 进产物
154        // 路径,未净化时 `../..` 可穿越到 output_dir 之外读取任意 .md 文件
155        //(曾实测复现)。与 page 同规则但更严:语言目录名只允许
156        // [A-Za-z0-9_-] 单段(zh、en、zh-CN 等),拒绝一切路径分隔符与
157        // 绝对路径形态——校验失败明确报错,不读盘。
158        if let Err(e) = validate_lang_segment(&lang) {
159            return e;
160        }
161        // 参数净化(工具暴露给任意 Agent):拒绝路径穿越与绝对路径,
162        // 只允许单段文件名(页面名),防读取 output_dir 之外任意文件
163        if page.contains('/') || page.contains('\\') || page.contains("..") || Path::new(&page).is_absolute() {
164            return format!("非法的页面名: {page}(只允许单段文件名)");
165        }
166        let path = config.output_dir()
167            .join("wiki")
168            .join(&lang)
169            .join(format!("{page}.md"));
170        match std::fs::read_to_string(&path) {
171            Ok(content) => format!("{}\n\n{content}", path.display()),
172            Err(e) => format!(
173                "页面不存在或不可读(可先运行 code-repo-wiki generate 生成): {}: {e}",
174                path.display()
175            ),
176        }
177    }
178
179    /// 读取已生成的 Knowledge Card(AI 代理的结构化模块摘要)
180    #[tool(description = "读取已生成的 Knowledge Card 内容(cards/{lang}/{card}.md;需先运行 code-repo-wiki generate,未生成的卡片报错)")]
181    async fn read_card(&self, Parameters(ReadCardRequest { card, lang }): Parameters<ReadCardRequest>) -> String {
182        let config = match crate::config::resolve_mcp_config(self.config_path.as_deref(), &self.root) {
183            Ok(c) => c,
184            Err(e) => return format!("配置加载失败: {e}"),
185        };
186        let lang = lang.unwrap_or_else(|| config.wiki.language.clone());
187        // 语言目录净化(S1):同 read_wiki_page,lang 直接 join 进路径,
188        // 未净化可穿越读取 output_dir 之外任意 .md(实测复现)。
189        if let Err(e) = validate_lang_segment(&lang) {
190            return e;
191        }
192        // 同 read_wiki_page:净化路径穿越
193        if card.contains('/') || card.contains('\\') || card.contains("..") || Path::new(&card).is_absolute() {
194            return format!("非法的卡片名: {card}(只允许单段文件名)");
195        }
196        let path = config.output_dir()
197            .join("cards")
198            .join(&lang)
199            .join(format!("{card}.md"));
200        match std::fs::read_to_string(&path) {
201            Ok(content) => format!("{}\n\n{content}", path.display()),
202            Err(e) => format!(
203                "卡片不存在或不可读(可先运行 code-repo-wiki generate 生成): {}: {e}",
204                path.display()
205            ),
206        }
207    }
208
209    /// 查看 Wiki 生成状态:页面/卡片数量与 lint 健康检查结果
210    #[tool(description = "查看 Wiki 生成状态:页面/卡片数量与产物健康检查(孤儿页/断链/过时/引用)")]
211    async fn status(&self) -> String {
212        let config = match crate::config::resolve_mcp_config(self.config_path.as_deref(), &self.root) {
213            Ok(c) => c,
214            Err(e) => return format!("配置加载失败: {e}"),
215        };
216        // MCP server 由项目内启动,root = 当前工作目录;
217        // 源码根须相对 root 解析(见 commands::source_roots_from_include_rooted),
218        // 否则跨 cwd 调用时 lint 会扫到错误目录
219        let root = match crate::project::ProjectRoot::from_cwd() {
220            Ok(r) => r,
221            Err(e) => return format!("无法确定当前工作目录: {e}"),
222        };
223        let report = crate::commands::status_report(&config, &root);
224        if !report.ready {
225            return "Wiki 未生成(运行 code-repo-wiki generate 生成后可用)".to_string();
226        }
227        let mut out = format!("Wiki 就绪: {} 张页面, {} 张卡片\n", report.wiki_pages, report.cards);
228        if report.issues.is_empty() {
229            out.push_str("lint: 通过(无孤儿页/断链/过时/引用/覆盖问题)");
230        } else {
231            out.push_str(&format!("lint: 发现 {} 个问题:\n", report.issues.len()));
232            for issue in &report.issues {
233                out.push_str(&format!("- [{}] {}: {}\n", issue.kind, issue.path, issue.message));
234            }
235        }
236        out
237    }
238}
239
240#[tool_handler]
241impl ServerHandler for RepoWikiMcp {
242    fn get_info(&self) -> ServerInfo {
243        ServerInfo::new(ServerCapabilities::builder().enable_tools().build())
244    }
245}
246
247/// 启动 stdio MCP server(阻塞直到 stdin 关闭)
248///
249/// 客户端配置示例(opencode.json / claude_desktop_config.json):
250/// ```json
251/// { "command": "code-repo-wiki", "args": ["mcp", "--root", "."] }
252/// ```
253pub async fn serve_stdio(config_path: Option<&Path>, root: ProjectRoot) -> Result<QuitReason> {
254    let server = RepoWikiMcp::new(config_path.map(|p| p.to_path_buf()), root);
255    let service = server.serve(rmcp::transport::stdio()).await?;
256    Ok(service.waiting().await?)
257}
258
259/// 将 top_k 收敛到 1..=50
260///
261/// 为什么 clamp:MCP 工具响应是整条字符串直接回给 Agent 的,top_k 无上限时
262/// 一次 search 调用可以把全部命中(几千行签名+文件路径)塞进单条响应,
263/// 直接撑爆 Agent 上下文窗口。CLI 的 -k 无此限制(用户显式指定、自担后果),
264/// MCP 是面向任意客户端的公共接口,必须从工具侧兜底。1 下限避免
265/// top_k=0 时"返回 0 个结果"的无意义调用。
266fn clamp_top_k(top_k: usize) -> usize {
267    top_k.clamp(1, 50)
268}
269
270/// 语言目录名校验(S1:MCP lang 参数净化)
271///
272/// 只允许单段名([A-Za-z0-9_-],如 zh/en/zh-CN/zh_cn),拒绝一切其他字符——
273/// 路径分隔符(/ \)、点段(..)、空格、盘符、绝对路径形态全部落入拒绝集。
274/// 校验失败返回含"非法语言名"的错误串(工具直接回给 Agent),不读盘。
275fn validate_lang_segment(lang: &str) -> Result<(), String> {
276    if lang.is_empty() {
277        return Err("非法语言名: (只允许 [A-Za-z0-9_-] 单段名)".to_string());
278    }
279    if lang
280        .chars()
281        .all(|c| c.is_ascii_alphanumeric() || c == '-' || c == '_')
282    {
283        Ok(())
284    } else {
285        Err(format!("非法语言名: {lang}(只允许 [A-Za-z0-9_-] 单段名)"))
286    }
287}
288
289#[cfg(test)]
290mod tests {
291    use super::*;
292
293    /// top_k 边界收敛:0/超大值收敛到 1/50,区间内原样保留
294    #[test]
295    fn test_clamp_top_k() {
296        assert_eq!(clamp_top_k(0), 1);
297        assert_eq!(clamp_top_k(1), 1);
298        assert_eq!(clamp_top_k(50), 50);
299        assert_eq!(clamp_top_k(9999), 50);
300        assert_eq!(clamp_top_k(10), 10);
301    }
302
303    /// S1 语言目录名校验:合法单段名通过,穿越/分隔符/空串全部拒绝
304    #[test]
305    fn test_validate_lang_segment() {
306        // 合法:语言目录名形态
307        for ok in ["zh", "en", "zh-CN", "zh_cn", "EN", "pt-BR"] {
308            assert!(validate_lang_segment(ok).is_ok(), "{ok} 应通过校验");
309        }
310        // 非法:路径穿越与分隔符形态(lang 直接 join 进产物路径的攻击面)
311        for bad in [
312            "", "..", "../..", "/", "//a", "a/b", "a\\b", "C:/x", "/abs", "zh..", "zh zh",
313        ] {
314            let err = validate_lang_segment(bad).unwrap_err();
315            assert!(
316                err.contains("非法语言名"),
317                "{bad:?} 应被拒绝且报错可读: {err}"
318            );
319        }
320    }
321}