1use 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#[derive(Debug, Clone)]
24pub struct RepoWikiMcp {
25 #[expect(dead_code, reason = "tool_handler 宏访问此路由字段")]
27 tool_router: ToolRouter<Self>,
28 config_path: Option<PathBuf>,
30 root: ProjectRoot,
32}
33
34impl RepoWikiMcp {
35 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#[derive(Debug, serde::Deserialize, schemars::JsonSchema)]
49struct SearchRequest {
50 query: String,
52 top_k: Option<usize>,
54 engine: Option<String>,
56}
57
58#[derive(Debug, serde::Deserialize, schemars::JsonSchema)]
60struct AstSearchRequest {
61 symbol: String,
63 language: Option<String>,
65}
66
67#[derive(Debug, serde::Deserialize, schemars::JsonSchema)]
69struct ReadPageRequest {
70 page: String,
72 lang: Option<String>,
74}
75
76#[derive(Debug, serde::Deserialize, schemars::JsonSchema)]
78struct ReadCardRequest {
79 card: String,
81 lang: Option<String>,
83}
84
85#[tool_router(router = tool_router)]
86impl RepoWikiMcp {
87 #[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 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 #[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 #[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 if let Err(e) = validate_lang_segment(&lang) {
159 return e;
160 }
161 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 #[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 if let Err(e) = validate_lang_segment(&lang) {
190 return e;
191 }
192 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 #[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 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
247pub 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
259fn clamp_top_k(top_k: usize) -> usize {
267 top_k.clamp(1, 50)
268}
269
270fn 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 #[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 #[test]
305 fn test_validate_lang_segment() {
306 for ok in ["zh", "en", "zh-CN", "zh_cn", "EN", "pt-BR"] {
308 assert!(validate_lang_segment(ok).is_ok(), "{ok} 应通过校验");
309 }
310 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}