use std::path::{Path, PathBuf};
use anyhow::Result;
use rmcp::handler::server::router::tool::ToolRouter;
use rmcp::handler::server::wrapper::Parameters;
use rmcp::model::{ServerCapabilities, ServerInfo};
use rmcp::service::{QuitReason, ServiceExt};
use rmcp::{ServerHandler, schemars, tool, tool_handler, tool_router};
use crate::project::ProjectRoot;
#[derive(Debug, Clone)]
pub struct RepoWikiMcp {
#[expect(dead_code, reason = "tool_handler 宏访问此路由字段")]
tool_router: ToolRouter<Self>,
config_path: Option<PathBuf>,
root: ProjectRoot,
}
impl RepoWikiMcp {
pub fn new(config_path: Option<PathBuf>, root: ProjectRoot) -> Self {
Self {
tool_router: Self::tool_router(),
config_path,
root,
}
}
}
#[derive(Debug, serde::Deserialize, schemars::JsonSchema)]
struct SearchRequest {
query: String,
top_k: Option<usize>,
engine: Option<String>,
}
#[derive(Debug, serde::Deserialize, schemars::JsonSchema)]
struct AstSearchRequest {
symbol: String,
language: Option<String>,
}
#[derive(Debug, serde::Deserialize, schemars::JsonSchema)]
struct ReadPageRequest {
page: String,
lang: Option<String>,
}
#[derive(Debug, serde::Deserialize, schemars::JsonSchema)]
struct ReadCardRequest {
card: String,
lang: Option<String>,
}
#[tool_router(router = tool_router)]
impl RepoWikiMcp {
#[tool(description = "搜索代码实体:按关键词返回匹配的函数/结构体/类及文件位置(text/semantic/hybrid 引擎,与 CLI code-repo-wiki search 等价;需先运行 code-repo-wiki generate 构建搜索索引)")]
async fn search(&self, Parameters(SearchRequest { query, top_k, engine }): Parameters<SearchRequest>) -> String {
let _config = match crate::config::resolve_mcp_config(self.config_path.as_deref(), &self.root) {
Ok(c) => c,
Err(e) => return format!("配置加载失败: {e}"),
};
let engine_type = match engine.as_deref() {
Some("text") => crate::config::schema::SearchEngineType::Text,
Some("semantic") => crate::config::schema::SearchEngineType::Semantic,
Some("hybrid") => crate::config::schema::SearchEngineType::Hybrid,
Some(other) => return format!("不支持的搜索引擎: {other}(可选: text/semantic/hybrid)"),
None => crate::config::schema::SEARCH_DEFAULT_ENGINE,
};
let top_k = clamp_top_k(top_k.unwrap_or(crate::config::schema::SEARCH_DEFAULT_TOP_K));
match crate::execute_search(self.config_path.as_deref(), &self.root, &query, top_k, &engine_type) {
Ok(hits) if hits.is_empty() => "未找到匹配结果".to_string(),
Ok(hits) => {
let mut out = format!("找到 {} 个结果:\n", hits.len());
for (i, hit) in hits.iter().enumerate() {
let file = hit.node.file_path.as_deref().unwrap_or("-");
let loc = match hit.node.line_range {
Some((s, e)) => format!("{file}:{s}-{e}"),
None => file.to_string(),
};
let sig = hit.node.signature.as_deref().unwrap_or(&hit.node.name);
out.push_str(&format!("{}. `{sig}` — {loc}\n", i + 1));
}
out
}
Err(e) => format!("搜索失败: {e}"),
}
}
#[tool(description = "AST 精确符号查找:扫描源文件定位函数/结构体/类定义的 文件+行号+签名(与 CLI code-repo-wiki ast-search 等价)")]
async fn ast_search(&self, Parameters(AstSearchRequest { symbol, language }): Parameters<AstSearchRequest>) -> String {
match crate::execute_ast_search(self.config_path.as_deref(), &self.root, &symbol, language.as_deref()) {
Ok(hits) if hits.is_empty() => format!("未找到符号 \"{symbol}\" 的定义"),
Ok(hits) => {
let mut out = format!("找到 {} 个定义:\n", hits.len());
for (i, hit) in hits.iter().enumerate() {
let sig = hit.node.signature.as_deref().unwrap_or(&hit.node.name);
let loc = match (&hit.node.file_path, hit.node.line_range) {
(Some(f), Some((s, e))) => format!("{f}:{s}-{e}"),
(Some(f), None) => f.clone(),
_ => "(unknown)".to_string(),
};
out.push_str(&format!("{}. {sig} — {loc}\n", i + 1));
}
out
}
Err(e) => format!("符号查找失败: {e}"),
}
}
#[tool(description = "读取已生成的 Wiki 页面内容(wiki/{lang}/{page}.md,如 src_config、architecture、overview、api;需先运行 code-repo-wiki generate,未生成的页面报错)")]
async fn read_wiki_page(&self, Parameters(ReadPageRequest { page, lang }): Parameters<ReadPageRequest>) -> String {
let config = match crate::config::resolve_mcp_config(self.config_path.as_deref(), &self.root) {
Ok(c) => c,
Err(e) => return format!("配置加载失败: {e}"),
};
let lang = lang.unwrap_or_else(|| config.wiki.language.clone());
if let Err(e) = validate_lang_segment(&lang) {
return e;
}
if page.contains('/') || page.contains('\\') || page.contains("..") || Path::new(&page).is_absolute() {
return format!("非法的页面名: {page}(只允许单段文件名)");
}
let path = config.output_dir()
.join("wiki")
.join(&lang)
.join(format!("{page}.md"));
match std::fs::read_to_string(&path) {
Ok(content) => format!("{}\n\n{content}", path.display()),
Err(e) => format!(
"页面不存在或不可读(可先运行 code-repo-wiki generate 生成): {}: {e}",
path.display()
),
}
}
#[tool(description = "读取已生成的 Knowledge Card 内容(cards/{lang}/{card}.md;需先运行 code-repo-wiki generate,未生成的卡片报错)")]
async fn read_card(&self, Parameters(ReadCardRequest { card, lang }): Parameters<ReadCardRequest>) -> String {
let config = match crate::config::resolve_mcp_config(self.config_path.as_deref(), &self.root) {
Ok(c) => c,
Err(e) => return format!("配置加载失败: {e}"),
};
let lang = lang.unwrap_or_else(|| config.wiki.language.clone());
if let Err(e) = validate_lang_segment(&lang) {
return e;
}
if card.contains('/') || card.contains('\\') || card.contains("..") || Path::new(&card).is_absolute() {
return format!("非法的卡片名: {card}(只允许单段文件名)");
}
let path = config.output_dir()
.join("cards")
.join(&lang)
.join(format!("{card}.md"));
match std::fs::read_to_string(&path) {
Ok(content) => format!("{}\n\n{content}", path.display()),
Err(e) => format!(
"卡片不存在或不可读(可先运行 code-repo-wiki generate 生成): {}: {e}",
path.display()
),
}
}
#[tool(description = "查看 Wiki 生成状态:页面/卡片数量与产物健康检查(孤儿页/断链/过时/引用)")]
async fn status(&self) -> String {
let config = match crate::config::resolve_mcp_config(self.config_path.as_deref(), &self.root) {
Ok(c) => c,
Err(e) => return format!("配置加载失败: {e}"),
};
let root = match crate::project::ProjectRoot::from_cwd() {
Ok(r) => r,
Err(e) => return format!("无法确定当前工作目录: {e}"),
};
let report = crate::commands::status_report(&config, &root);
if !report.ready {
return "Wiki 未生成(运行 code-repo-wiki generate 生成后可用)".to_string();
}
let mut out = format!("Wiki 就绪: {} 张页面, {} 张卡片\n", report.wiki_pages, report.cards);
if report.issues.is_empty() {
out.push_str("lint: 通过(无孤儿页/断链/过时/引用/覆盖问题)");
} else {
out.push_str(&format!("lint: 发现 {} 个问题:\n", report.issues.len()));
for issue in &report.issues {
out.push_str(&format!("- [{}] {}: {}\n", issue.kind, issue.path, issue.message));
}
}
out
}
}
#[tool_handler]
impl ServerHandler for RepoWikiMcp {
fn get_info(&self) -> ServerInfo {
ServerInfo::new(ServerCapabilities::builder().enable_tools().build())
}
}
pub async fn serve_stdio(config_path: Option<&Path>, root: ProjectRoot) -> Result<QuitReason> {
let server = RepoWikiMcp::new(config_path.map(|p| p.to_path_buf()), root);
let service = server.serve(rmcp::transport::stdio()).await?;
Ok(service.waiting().await?)
}
fn clamp_top_k(top_k: usize) -> usize {
top_k.clamp(1, 50)
}
fn validate_lang_segment(lang: &str) -> Result<(), String> {
if lang.is_empty() {
return Err("非法语言名: (只允许 [A-Za-z0-9_-] 单段名)".to_string());
}
if lang
.chars()
.all(|c| c.is_ascii_alphanumeric() || c == '-' || c == '_')
{
Ok(())
} else {
Err(format!("非法语言名: {lang}(只允许 [A-Za-z0-9_-] 单段名)"))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_clamp_top_k() {
assert_eq!(clamp_top_k(0), 1);
assert_eq!(clamp_top_k(1), 1);
assert_eq!(clamp_top_k(50), 50);
assert_eq!(clamp_top_k(9999), 50);
assert_eq!(clamp_top_k(10), 10);
}
#[test]
fn test_validate_lang_segment() {
for ok in ["zh", "en", "zh-CN", "zh_cn", "EN", "pt-BR"] {
assert!(validate_lang_segment(ok).is_ok(), "{ok} 应通过校验");
}
for bad in [
"", "..", "../..", "/", "//a", "a/b", "a\\b", "C:/x", "/abs", "zh..", "zh zh",
] {
let err = validate_lang_segment(bad).unwrap_err();
assert!(
err.contains("非法语言名"),
"{bad:?} 应被拒绝且报错可读: {err}"
);
}
}
}