use std::path::Path;
use std::sync::atomic::{AtomicUsize, Ordering};
use anyhow::Result;
use crate::config::schema::WikiConfig;
use crate::generate::chunk::Chunk;
use crate::generate::llm::{LlmProvider, Message};
use crate::generate::prompt;
use crate::generate::GenerationOutput;
use crate::model::{DocumentKind, EdgeKind, KnowledgeGraph, NodeId, Reference, WikiDocument};
pub struct WikiGenerator<'a, P: LlmProvider> {
provider: &'a P,
call_count: AtomicUsize,
failed: std::sync::Mutex<Vec<String>>,
semaphore: std::sync::Arc<tokio::sync::Semaphore>,
desc_cache: std::sync::Mutex<Option<ModuleDescCache>>,
head_short: std::sync::OnceLock<Option<String>>,
}
impl<P: LlmProvider> WikiGenerator<'_, P> {
fn head_short_for(&self, root: &crate::project::ProjectRoot) -> Option<String> {
self.head_short
.get_or_init(|| git_head_short(root))
.clone()
}
}
#[derive(serde::Serialize, serde::Deserialize)]
struct CacheEntry {
fingerprint: String,
description: String,
}
struct ModuleDescCache {
entries: std::collections::HashMap<String, CacheEntry>,
}
fn git_head_short(root: &crate::project::ProjectRoot) -> Option<String> {
let repo = git2::Repository::open(root.path()).ok()?;
let head = repo.head().ok()?;
let commit = head.peel_to_commit().ok()?;
let id = commit.id().to_string();
Some(id[..id.len().min(8)].to_string())
}
impl ModuleDescCache {
fn new() -> Self {
Self {
entries: std::collections::HashMap::new(),
}
}
fn load(path: &Path) -> Self {
let Ok(content) = std::fs::read_to_string(path) else {
return Self::new();
};
match serde_json::from_str::<std::collections::HashMap<String, CacheEntry>>(&content) {
Ok(entries) => Self { entries },
Err(e) => {
tracing::warn!(
"模块描述缓存解析失败(回退空缓存,按需重新生成): {} {}",
path.display(),
e
);
Self::new()
}
}
}
fn save(&self, path: &Path) {
let Ok(content) = serde_json::to_string(&self.entries) else {
return;
};
let Some(parent) = path.parent() else {
return;
};
let _ = std::fs::create_dir_all(parent);
let tmp = path.with_extension("json.tmp");
if std::fs::write(&tmp, content).is_ok() {
let _ = std::fs::rename(&tmp, path);
}
}
}
fn module_files_fingerprint(
module: &crate::model::ModuleCluster,
graph: &KnowledgeGraph,
root: &crate::project::ProjectRoot,
) -> String {
let mut files: Vec<&str> = module
.node_ids
.iter()
.filter_map(|nid| graph.graph.node_weight(*nid))
.filter_map(|n| n.file_path.as_deref())
.collect();
files.sort_unstable();
files.dedup();
files
.iter()
.map(|f| {
crate::incremental::state::GenerationState::compute_file_fingerprint(
&root.path().join(f),
)
.unwrap_or_else(|_| "missing".to_string())
})
.collect::<Vec<_>>()
.join("|")
}
pub const CITATION_RETRY_MAX: usize = 2;
pub const MERMAID_RETRY_MAX: usize = crate::output::mermaid_check::MERMAID_RETRY_MAX;
const CALL_RETRY_MAX: usize = 3;
async fn complete_with_retry<P: LlmProvider>(
provider: &P,
messages: &[Message],
module: &str,
) -> Result<String> {
let mut last_err = None;
for attempt in 0..CALL_RETRY_MAX {
match provider.complete(messages).await {
Ok(content) => return Ok(content),
Err(e) => {
tracing::warn!(
"LLM 调用失败(第 {} 次重试): {} {}",
attempt + 1,
module,
e
);
last_err = Some(e);
tokio::time::sleep(std::time::Duration::from_millis(500 << attempt)).await;
}
}
}
Err(last_err.expect("CALL_RETRY_MAX 至少为 1"))
}
impl<'a, P: LlmProvider> WikiGenerator<'a, P> {
pub fn new(provider: &'a P, max_concurrent: usize) -> Self {
let max = if max_concurrent == 0 { 1_000_000_000 } else { max_concurrent };
Self {
provider,
call_count: AtomicUsize::new(0),
failed: std::sync::Mutex::new(Vec::new()),
semaphore: std::sync::Arc::new(tokio::sync::Semaphore::new(max)),
desc_cache: std::sync::Mutex::new(None),
head_short: std::sync::OnceLock::new(),
}
}
pub fn llm_call_count(&self) -> usize {
self.call_count.load(Ordering::Relaxed)
}
pub fn failed_modules(&self) -> Vec<String> {
self.failed.lock().map(|g| g.clone()).unwrap_or_default()
}
pub(crate) fn record_failure(&self, module: String) {
if let Ok(mut failed) = self.failed.lock() {
failed.push(module);
}
}
pub async fn generate_wiki_page(
&self,
chunk: &Chunk,
card_summary: &str,
config: &WikiConfig,
root: &crate::project::ProjectRoot,
entity_ranges: Option<&crate::output::citation::EntityRanges>,
) -> Result<WikiDocument> {
if chunk.is_empty() {
anyhow::bail!("空块,跳过 Wiki 页面生成");
}
let language = &config.wiki.language;
let mut messages =
prompt::wiki_page_prompt(chunk, card_summary, language, &config.wiki.guide.notes);
let mut content = String::new();
let mut last_invalid = Vec::new();
let mut last_mermaid = Vec::new();
let retry_max = CITATION_RETRY_MAX.max(MERMAID_RETRY_MAX);
for attempt in 0..=retry_max {
self.call_count.fetch_add(1, Ordering::Relaxed);
content = complete_with_retry(self.provider, &messages, &chunk.module_path.join("::"))
.await?;
last_invalid = if content.trim().is_empty() {
vec![crate::output::citation::InvalidCitation {
citation: crate::output::citation::Citation {
path: String::new(),
start: 0,
end: 0,
},
reason: "输出为空(未生成任何内容)".into(),
}]
} else {
match entity_ranges {
Some(ranges) => {
crate::output::citation::validate_citations_against_entities(
root.path(),
&content,
ranges,
)
}
None => crate::output::citation::validate_citations(root.path(), &content),
}
};
last_mermaid = crate::output::mermaid_check::validate_mermaid_blocks(&content);
if last_invalid.is_empty() && last_mermaid.is_empty() {
break;
}
if !last_invalid.is_empty() {
tracing::warn!(
"Wiki 页面引用校验失败(第 {} 次,无效 {} 条): {}",
attempt + 1,
last_invalid.len(),
chunk.module_path.join("::")
);
}
if !last_mermaid.is_empty() {
tracing::warn!(
"Wiki 页面 Mermaid 校验失败(第 {} 次,坏块 {} 个): {}",
attempt + 1,
last_mermaid.len(),
chunk.module_path.join("::")
);
}
if attempt == retry_max {
break;
}
if !last_invalid.is_empty() {
messages.push(Message::user(
crate::output::citation::retry_feedback(&last_invalid),
));
}
if !last_mermaid.is_empty() {
messages.push(Message::user(
crate::output::mermaid_check::mermaid_retry_feedback(&last_mermaid),
));
}
}
if !last_invalid.is_empty() {
anyhow::bail!(
"Wiki 页面引用校验失败(重试 {} 次仍无效,共 {} 条无效引用): {}",
retry_max,
last_invalid.len(),
chunk.module_path.join("::")
);
}
if !last_mermaid.is_empty() {
tracing::warn!(
"Wiki 页面 Mermaid 重试耗尽({} 个坏块),降级为 text 块: {}",
last_mermaid.len(),
chunk.module_path.join("::")
);
content = crate::output::mermaid_check::degrade_mermaid_blocks(&content, &last_mermaid);
}
let now = chrono::Utc::now().to_rfc3339();
let based_on_commit = self.head_short_for(root);
Ok(WikiDocument {
title: chunk.module_path.join("::"),
kind: DocumentKind::WikiPage,
content,
language: config.wiki.language.clone(),
module_path: chunk.module_path.clone(),
references: build_references(chunk, &config.wiki.language),
last_updated: now,
based_on_commit,
fingerprint: None,
})
}
async fn complete_with_mermaid_guard(
&self,
messages: Vec<Message>,
label: &str,
) -> Result<String> {
complete_with_mermaid_guard_free(self.provider, messages, label, Some(&self.call_count))
.await
}
pub async fn generate_architecture(
&self,
output: &GenerationOutput,
graph: &KnowledgeGraph,
config: &WikiConfig,
root: &crate::project::ProjectRoot,
) -> Result<WikiDocument> {
let language = &config.wiki.language;
let modules = self.describe_modules(graph, language, config, root).await;
let messages =
prompt::architecture_overview_prompt(&modules, graph, language);
let content = self.complete_with_mermaid_guard(messages, "架构概览").await?;
let now = chrono::Utc::now().to_rfc3339();
let based_on_commit = self.head_short_for(root);
Ok(WikiDocument {
title: "架构概览".into(),
kind: DocumentKind::ArchitectureOverview,
content,
language: config.wiki.language.clone(),
module_path: vec![],
references: output
.cards
.iter()
.map(|c| Reference {
target_title: c.module_name.clone(),
target_path: format!(
"wiki/{}/{}.md",
config.wiki.language,
c.module_name.replace("::", "_")
),
relation: "module".into(),
})
.collect(),
last_updated: now,
based_on_commit,
fingerprint: None,
})
}
async fn describe_modules(
&self,
graph: &KnowledgeGraph,
language: &str,
config: &WikiConfig,
root: &crate::project::ProjectRoot,
) -> Vec<crate::model::ModuleCluster> {
let cache_path = config.output_dir().join(".state").join("module_descriptions.json");
{
let mut guard = self.desc_cache.lock().unwrap_or_else(|e| e.into_inner());
if guard.is_none() {
*guard = Some(ModuleDescCache::load(&cache_path));
}
}
let semaphore = self.semaphore.clone();
let futures: Vec<_> = graph
.modules
.iter()
.map(|module| {
let semaphore = semaphore.clone();
let cache_key = format!("{}@{}", module.name, language);
let fingerprint = module_files_fingerprint(module, graph, root);
async move {
if module.name == "src" || module.node_ids.is_empty() {
return module.clone();
}
{
let guard = self.desc_cache.lock().unwrap_or_else(|e| e.into_inner());
if let Some(cache) = guard.as_ref()
&& let Some(entry) = cache.entries.get(&cache_key)
&& entry.fingerprint == fingerprint
{
let mut enriched = module.clone();
enriched.description = Some(entry.description.clone());
return enriched;
}
}
let _permit = match semaphore.acquire().await {
Ok(p) => p,
Err(_) => return module.clone(),
};
let mut enriched = module.clone();
if let Ok(text) = self.describe_module(module, graph, language).await
&& !text.trim().is_empty()
{
let description = text.trim().to_string();
{
let mut guard = self.desc_cache.lock().unwrap_or_else(|e| e.into_inner());
if let Some(cache) = guard.as_mut() {
cache.entries.insert(
cache_key,
CacheEntry {
fingerprint,
description: description.clone(),
},
);
}
}
enriched.description = Some(description);
}
enriched
}
})
.collect();
let modules = futures::future::join_all(futures).await;
{
let guard = self.desc_cache.lock().unwrap_or_else(|e| e.into_inner());
if let Some(cache) = guard.as_ref() {
cache.save(&cache_path);
}
}
modules
}
async fn describe_module(
&self,
module: &crate::model::ModuleCluster,
graph: &KnowledgeGraph,
language: &str,
) -> Result<String> {
self.call_count.fetch_add(1, Ordering::Relaxed);
let entity_names = collect_module_entity_names(module, graph);
let messages = prompt::module_description_prompt(
&module.name,
&entity_names,
language,
);
self.provider.complete(&messages).await
}
pub async fn generate_overview(
&self,
output: &GenerationOutput,
graph: &KnowledgeGraph,
config: &WikiConfig,
root: &crate::project::ProjectRoot,
) -> Result<WikiDocument> {
let modules = self.describe_modules(graph, &config.wiki.language, config, root).await;
let messages = vec![Message::user(overview_prompt(&modules, &output.cards, graph, config))];
let content = self.complete_with_mermaid_guard(messages, "项目概览").await?;
let now = chrono::Utc::now().to_rfc3339();
let based_on_commit = self.head_short_for(root);
Ok(WikiDocument {
title: "项目概览".into(),
kind: DocumentKind::ProjectOverview,
content,
language: config.wiki.language.clone(),
module_path: vec![],
references: output
.cards
.iter()
.map(|c| Reference {
target_title: c.module_name.clone(),
target_path: format!(
"wiki/{}/{}.md",
config.wiki.language,
c.module_name.replace("::", "_")
),
relation: "module".into(),
})
.collect(),
last_updated: now,
based_on_commit,
fingerprint: None,
})
}
}
pub(crate) const DESCRIBE_ENTITY_CAP: usize = 30;
fn entity_priority(kind: &crate::model::NodeKind) -> u8 {
match kind {
crate::model::NodeKind::Function
| crate::model::NodeKind::Trait
| crate::model::NodeKind::Impl
| crate::model::NodeKind::Interface
| crate::model::NodeKind::Class
| crate::model::NodeKind::Macro => 0,
crate::model::NodeKind::Struct
| crate::model::NodeKind::Enum
| crate::model::NodeKind::Type => 1,
crate::model::NodeKind::Constant => 2,
_ => 3,
}
}
pub(crate) fn collect_module_entity_names(
module: &crate::model::ModuleCluster,
graph: &KnowledgeGraph,
) -> Vec<String> {
let mut names: Vec<(u8, String)> = module
.node_ids
.iter()
.filter_map(|nid| graph.graph.node_weight(*nid))
.filter(|n| {
!matches!(
n.kind,
crate::model::NodeKind::Project
| crate::model::NodeKind::Module
| crate::model::NodeKind::File
| crate::model::NodeKind::Variable
)
})
.map(|n| (entity_priority(&n.kind), n.name.clone()))
.collect();
names.sort_by(|a, b| a.0.cmp(&b.0).then_with(|| a.1.cmp(&b.1)));
names
.into_iter()
.map(|(_, name)| name)
.take(DESCRIBE_ENTITY_CAP)
.collect()
}
pub async fn complete_with_mermaid_guard_free<P: LlmProvider>(
provider: &P,
mut messages: Vec<Message>,
label: &str,
call_count: Option<&std::sync::atomic::AtomicUsize>,
) -> Result<String> {
use std::sync::atomic::Ordering;
let mut content = String::new();
let mut last_mermaid = Vec::new();
for attempt in 0..=MERMAID_RETRY_MAX {
if let Some(c) = call_count {
c.fetch_add(1, Ordering::Relaxed);
}
content = provider.complete(&messages).await?;
last_mermaid = crate::output::mermaid_check::validate_mermaid_blocks(&content);
if last_mermaid.is_empty() {
return Ok(content);
}
tracing::warn!(
"{label} Mermaid 校验失败(第 {} 次,坏块 {} 个)",
attempt + 1,
last_mermaid.len()
);
if attempt == MERMAID_RETRY_MAX {
break;
}
messages.push(Message::user(
crate::output::mermaid_check::mermaid_retry_feedback(&last_mermaid),
));
}
tracing::warn!("{label} Mermaid 重试耗尽({} 个坏块),降级为 text 块", last_mermaid.len());
Ok(crate::output::mermaid_check::degrade_mermaid_blocks(
&content,
&last_mermaid,
))
}
fn overview_prompt(
modules: &[crate::model::ModuleCluster],
cards: &[crate::model::KnowledgeCard],
graph: &KnowledgeGraph,
config: &WikiConfig,
) -> String {
let mut parts = Vec::new();
parts.push(format!(
"你是一个资深软件架构师,负责为整个项目生成人类可读的项目概览文档。\n\n\
请基于下面的模块聚类信息、各模块卡片摘要和模块间依赖摘要,输出以下结构:\n\n\
# 项目概览\n\n\
## 技术栈\n根据模块名称与依赖关系推断项目使用的技术栈。\n\n\
## 目录结构\n根据模块划分描述仓库的目录结构。\n\n\
## 核心模块\n列出核心模块及其职责。\n\n\
请用 {} 语言输出。保留 Markdown 格式。",
config.wiki.language
));
parts.push("## 模块列表".to_string());
for module in modules {
let desc = module.description.as_deref().unwrap_or("");
parts.push(format!(
"- {} (节点数: {}{})",
module.name,
module.node_ids.len(),
if desc.is_empty() {
String::new()
} else {
format!(", 职责: {}", desc)
}
));
}
if !cards.is_empty() {
parts.push("## 模块卡片摘要".to_string());
for card in cards {
let entities: Vec<&str> = card
.key_entities
.iter()
.map(|e| e.name.as_str())
.take(8)
.collect();
parts.push(format!(
"- {}: {}(关键实体: {})",
card.module_name,
card.summary,
if entities.is_empty() {
"无".to_string()
} else {
entities.join(", ")
}
));
}
}
let mut module_of: std::collections::HashMap<NodeId, &str> = Default::default();
for module in modules {
for nid in &module.node_ids {
module_of.insert(*nid, module.name.as_str());
}
}
let mut deps: std::collections::BTreeMap<(String, String), usize> = Default::default();
for edge in graph.graph.edge_weights() {
if edge.kind == EdgeKind::Contains {
continue;
}
let (Some(src), Some(dst)) = (module_of.get(&edge.source), module_of.get(&edge.target))
else {
continue;
};
*deps.entry((src.to_string(), dst.to_string())).or_default() += 1;
}
if deps.is_empty() {
parts.push("\n## 模块间依赖\n(图中未检测到模块间依赖边)".to_string());
} else {
parts.push("\n## 模块间依赖".to_string());
for ((src, dst), count) in deps {
parts.push(format!("- {} → {} ({} 条边)", src, dst, count));
}
}
parts.join("\n")
}
fn build_references(chunk: &Chunk, language: &str) -> Vec<Reference> {
chunk
.dependencies
.iter()
.map(|dep| Reference {
target_title: dep.clone(),
target_path: format!(
"wiki/{language}/{}.md",
dep.replace("::", "_")
),
relation: "depends_on".into(),
})
.collect()
}
pub fn fallback_architecture_doc(
graph: &KnowledgeGraph,
config: &WikiConfig,
kind: DocumentKind,
title: &str,
) -> WikiDocument {
use petgraph::visit::{EdgeRef, IntoEdgeReferences};
use std::collections::{BTreeMap, BTreeSet, HashMap};
let mut node_module: HashMap<NodeId, String> = HashMap::new();
for module in &graph.modules {
for nid in &module.node_ids {
node_module
.entry(*nid)
.or_insert_with(|| module.name.clone());
}
}
let mut deps: BTreeMap<String, BTreeSet<String>> = Default::default();
for edge in graph.graph.edge_references() {
if matches!(
graph.graph[edge.id()].kind,
EdgeKind::Calls | EdgeKind::Imports
) {
let (Some(src), Some(tgt)) = (
node_module.get(&edge.source()),
node_module.get(&edge.target()),
) else {
continue;
};
if src != tgt {
deps.entry(src.clone()).or_default().insert(tgt.clone());
}
}
}
let mut body = format!(
"# {title}\n\n> LLM 生成不可用,本页为确定性骨架:模块与依赖关系由知识图谱自动生成(无 LLM 摘要)。\n\n## 模块\n\n"
);
for module in &graph.modules {
body.push_str(&format!("- `{}`({} 个实体)", module.name, module.node_ids.len()));
if let Some(dl) = deps.get(&module.name)
&& !dl.is_empty()
{
body.push_str(&format!(" — 依赖 {}", dl.iter().cloned().collect::<Vec<_>>().join(", ")));
}
body.push('\n');
}
let mut refs: Vec<Reference> = graph
.modules
.iter()
.map(|m| Reference {
target_title: m.name.clone(),
target_path: format!(
"wiki/{}/{}.md",
config.wiki.language,
m.name.replace("::", "_")
),
relation: "module".into(),
})
.collect();
refs.sort_by(|a, b| a.target_title.cmp(&b.target_title));
WikiDocument {
title: title.to_string(),
kind,
content: body,
language: config.wiki.language.clone(),
module_path: vec![],
references: refs,
last_updated: chrono::Utc::now().to_rfc3339(),
based_on_commit: None,
fingerprint: None,
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::generate::chunk::chunk_by_file;
use crate::generate::llm::MockProvider;
use crate::ingest::parser::{Entity, FileInsight, ImportStmt};
use std::path::PathBuf;
struct ScriptedProvider {
responses: std::sync::Mutex<std::vec::IntoIter<String>>,
calls: std::sync::atomic::AtomicUsize,
}
impl ScriptedProvider {
fn new(responses: Vec<String>) -> Self {
Self {
responses: std::sync::Mutex::new(responses.into_iter()),
calls: std::sync::atomic::AtomicUsize::new(0),
}
}
}
impl LlmProvider for ScriptedProvider {
async fn complete(&self, _messages: &[Message]) -> Result<String> {
self.calls.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
self.responses
.lock()
.unwrap()
.next()
.ok_or_else(|| anyhow::anyhow!("预设响应耗尽"))
}
}
struct FlakyProvider {
fail_times: std::sync::atomic::AtomicUsize,
calls: std::sync::atomic::AtomicUsize,
}
impl FlakyProvider {
fn new(fail_times: usize) -> Self {
Self {
fail_times: std::sync::atomic::AtomicUsize::new(fail_times),
calls: std::sync::atomic::AtomicUsize::new(0),
}
}
}
impl LlmProvider for FlakyProvider {
async fn complete(&self, _messages: &[Message]) -> Result<String> {
self.calls.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let remaining = self.fail_times.fetch_sub(1, std::sync::atomic::Ordering::Relaxed);
if remaining > 0 {
Err(anyhow::anyhow!("模拟瞬时网络错误"))
} else {
Ok("重试成功".to_string())
}
}
}
#[tokio::test]
async fn test_complete_with_retry_recovers_after_transient_failure() {
let provider = FlakyProvider::new(1);
let content = complete_with_retry(&provider, &[], "src::test").await.unwrap();
assert_eq!(content, "重试成功");
assert_eq!(provider.calls.load(std::sync::atomic::Ordering::Relaxed), 2);
}
#[tokio::test]
async fn test_complete_with_retry_gives_up_after_max_attempts() {
let provider = FlakyProvider::new(10);
let err = complete_with_retry(&provider, &[], "src::test")
.await
.unwrap_err();
assert!(err.to_string().contains("模拟瞬时网络错误"));
assert_eq!(
provider.calls.load(std::sync::atomic::Ordering::Relaxed),
CALL_RETRY_MAX
);
}
fn make_test_chunk() -> Chunk {
let entity = Entity {
name: "Server".into(),
kind: "struct".into(),
line_start: 1,
line_end: 50,
doc_comment: Some("HTTP 服务".into()),
signature: None, visibility: None,
};
let insight = FileInsight {
path: PathBuf::from("src/server.rs"),
language: "rust".into(),
entities: vec![entity],
imports: vec![ImportStmt {
source: "tokio".into(),
alias: None,
line: 1,
}],
doc_comments: vec![],
source: String::new(),
};
chunk_by_file(&insight)
}
#[tokio::test]
async fn test_skip_empty_chunk() {
let provider = MockProvider::new();
let generator = WikiGenerator::new(&provider, 0);
let config = WikiConfig::default();
let root = crate::project::ProjectRoot::new(std::env::temp_dir());
let empty_chunk = Chunk {
module_path: vec![],
entities: vec![],
imports: vec![],
dependencies: vec![],
file_paths: vec![],
entity_sources: vec![],
};
let result = generator.generate_wiki_page(&empty_chunk, "", &config, &root, None).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_wiki_page_retries_on_invalid_citation() {
let dir = std::env::temp_dir().join(format!("code_repo_wiki_test_cite_retry_{}", std::process::id()));
let _ = std::fs::remove_dir_all(&dir);
let src = dir.join("src");
std::fs::create_dir_all(&src).unwrap();
std::fs::write(src.join("server.rs"), "pub struct Server;\n// comment\n").unwrap();
let root = crate::project::ProjectRoot::new(dir.clone());
let provider = ScriptedProvider::new(vec![
"模块职责是管理连接。核心实体 `Server` 定义见 nonexistent.rs:99。".to_string(),
"模块职责是管理连接。核心实体 `Server` 定义见 src/server.rs:1。".to_string(),
]);
let generator = WikiGenerator::new(&provider, 0);
let config = WikiConfig::default();
let chunk = make_test_chunk();
let doc = generator.generate_wiki_page(&chunk, "摘要", &config, &root, None).await.unwrap();
assert!(doc.content.contains("src/server.rs:1"), "重试后应使用有效引用");
assert_eq!(provider.calls.load(std::sync::atomic::Ordering::Relaxed), 2, "应调用 2 次(1 次失败 + 1 次重试)");
let _ = std::fs::remove_dir_all(&dir);
}
#[tokio::test]
async fn test_wiki_page_bails_when_citations_never_valid() {
let dir = std::env::temp_dir().join(format!("code_repo_wiki_test_cite_fail_{}", std::process::id()));
let _ = std::fs::remove_dir_all(&dir);
let root = crate::project::ProjectRoot::new(dir.clone());
let provider = ScriptedProvider::new(vec![
"引用 nonexistent.rs:99".to_string(),
"引用 nonexistent.rs:99".to_string(),
"引用 nonexistent.rs:99".to_string(),
]);
let generator = WikiGenerator::new(&provider, 0);
let config = WikiConfig::default();
let chunk = make_test_chunk();
let result = generator.generate_wiki_page(&chunk, "摘要", &config, &root, None).await;
assert!(result.is_err(), "重试耗尽后应报错");
let err = result.unwrap_err().to_string();
assert!(err.contains("引用校验失败"), "错误信息应说明引用校验失败: {err}");
assert_eq!(
provider.calls.load(std::sync::atomic::Ordering::Relaxed),
CITATION_RETRY_MAX + 1,
"应调用 CITATION_RETRY_MAX+1 次后放弃"
);
let _ = std::fs::remove_dir_all(&dir);
}
#[tokio::test]
async fn test_wiki_page_without_citations_passes() {
let dir = std::env::temp_dir().join(format!("code_repo_wiki_test_cite_ok_{}", std::process::id()));
let _ = std::fs::remove_dir_all(&dir);
let root = crate::project::ProjectRoot::new(dir.clone());
let provider = ScriptedProvider::new(vec!["模块职责是管理连接。".to_string()]);
let generator = WikiGenerator::new(&provider, 0);
let config = WikiConfig::default();
let chunk = make_test_chunk();
let doc = generator.generate_wiki_page(&chunk, "摘要", &config, &root, None).await.unwrap();
assert_eq!(doc.content, "模块职责是管理连接。");
assert_eq!(provider.calls.load(std::sync::atomic::Ordering::Relaxed), 1, "无引用无需重试");
let _ = std::fs::remove_dir_all(&dir);
}
#[tokio::test]
async fn test_wiki_page_retries_on_bad_mermaid() {
let dir = std::env::temp_dir().join(format!("code_repo_wiki_test_mermaid_retry_{}", std::process::id()));
let _ = std::fs::remove_dir_all(&dir);
let root = crate::project::ProjectRoot::new(dir.clone());
let provider = ScriptedProvider::new(vec![
"```mermaid\nflowchart LR\nA[hello world\nB --> C\n```\n".to_string(),
"```mermaid\nflowchart LR\nA[Start] --> B[End]\n```\n".to_string(),
]);
let generator = WikiGenerator::new(&provider, 0);
let config = WikiConfig::default();
let chunk = make_test_chunk();
let doc = generator.generate_wiki_page(&chunk, "摘要", &config, &root, None).await.unwrap();
assert!(doc.content.contains("A[Start] --> B[End]"), "重试后应保留好图");
assert_eq!(provider.calls.load(std::sync::atomic::Ordering::Relaxed), 2, "应调用 2 次(1 次坏图 + 1 次重试)");
let _ = std::fs::remove_dir_all(&dir);
}
#[tokio::test]
async fn test_wiki_page_retries_on_overlap_citation() {
let dir = std::env::temp_dir().join(format!("code_repo_wiki_test_cite_overlap_{}", std::process::id()));
let _ = std::fs::remove_dir_all(&dir);
let src = dir.join("src");
std::fs::create_dir_all(&src).unwrap();
let content: String = (1..=10).map(|i| format!("line{i}\n")).collect();
std::fs::write(src.join("server.rs"), content).unwrap();
let root = crate::project::ProjectRoot::new(dir.clone());
let mut ranges: crate::output::citation::EntityRanges =
crate::output::citation::EntityRanges::new();
ranges.insert("src/server.rs".to_string(), vec![(2, 4)]);
let provider = ScriptedProvider::new(vec![
"模块职责是管理连接。核心实体 `Server` 定义见 src/server.rs:8。".to_string(),
"模块职责是管理连接。核心实体 `Server` 定义见 src/server.rs:2。".to_string(),
]);
let generator = WikiGenerator::new(&provider, 0);
let config = WikiConfig::default();
let chunk = make_test_chunk();
let doc = generator
.generate_wiki_page(&chunk, "摘要", &config, &root, Some(&ranges))
.await
.unwrap();
assert!(doc.content.contains("src/server.rs:2"), "重试后应使用覆盖实体的引用");
assert_eq!(
provider.calls.load(std::sync::atomic::Ordering::Relaxed),
2,
"区间外引用应触发一次重试"
);
let _ = std::fs::remove_dir_all(&dir);
}
#[tokio::test]
async fn test_wiki_page_bails_when_overlap_never_valid() {
let dir = std::env::temp_dir().join(format!("code_repo_wiki_test_cite_overlap_fail_{}", std::process::id()));
let _ = std::fs::remove_dir_all(&dir);
let src = dir.join("src");
std::fs::create_dir_all(&src).unwrap();
let content: String = (1..=10).map(|i| format!("line{i}\n")).collect();
std::fs::write(src.join("server.rs"), content).unwrap();
let root = crate::project::ProjectRoot::new(dir.clone());
let mut ranges: crate::output::citation::EntityRanges =
crate::output::citation::EntityRanges::new();
ranges.insert("src/server.rs".to_string(), vec![(2, 4)]);
let provider = ScriptedProvider::new(vec![
"核心实体 `Server` 见 src/server.rs:8。".to_string(),
"核心实体 `Server` 见 src/server.rs:8。".to_string(),
"核心实体 `Server` 见 src/server.rs:8。".to_string(),
]);
let generator = WikiGenerator::new(&provider, 0);
let config = WikiConfig::default();
let chunk = make_test_chunk();
let result = generator
.generate_wiki_page(&chunk, "摘要", &config, &root, Some(&ranges))
.await;
assert!(result.is_err(), "区间重叠校验重试耗尽应报错");
let err = result.unwrap_err().to_string();
assert!(err.contains("引用校验失败"), "错误信息应说明引用校验失败: {err}");
let _ = std::fs::remove_dir_all(&dir);
}
#[tokio::test]
async fn test_wiki_page_passes_non_code_file_citation() {
let dir = std::env::temp_dir().join(format!("code_repo_wiki_test_cite_noncode_{}", std::process::id()));
let _ = std::fs::remove_dir_all(&dir);
std::fs::create_dir_all(&dir).unwrap();
std::fs::write(dir.join("README.md"), "docs\n").unwrap();
let root = crate::project::ProjectRoot::new(dir.clone());
let ranges: crate::output::citation::EntityRanges =
crate::output::citation::EntityRanges::new();
let provider = ScriptedProvider::new(vec![
"模块说明见 README.md:1。".to_string(),
]);
let generator = WikiGenerator::new(&provider, 0);
let config = WikiConfig::default();
let chunk = make_test_chunk();
let doc = generator
.generate_wiki_page(&chunk, "摘要", &config, &root, Some(&ranges))
.await
.unwrap();
assert!(doc.content.contains("README.md:1"), "无实体文件引用应放行");
assert_eq!(provider.calls.load(std::sync::atomic::Ordering::Relaxed), 1, "无需重试");
let _ = std::fs::remove_dir_all(&dir);
}
#[tokio::test]
async fn test_wiki_page_degrades_when_mermaid_never_valid() {
let dir = std::env::temp_dir().join(format!("code_repo_wiki_test_mermaid_degrade_{}", std::process::id()));
let _ = std::fs::remove_dir_all(&dir);
let root = crate::project::ProjectRoot::new(dir.clone());
let provider = ScriptedProvider::new(vec![
"```mermaid\nflowchart LR\nA[hello world\nB --> C\n```\n".to_string(),
"```mermaid\nflowchart LR\nA[hello world\nB --> C\n```\n".to_string(),
"```mermaid\nflowchart LR\nA[hello world\nB --> C\n```\n".to_string(),
]);
let generator = WikiGenerator::new(&provider, 0);
let config = WikiConfig::default();
let chunk = make_test_chunk();
let doc = generator.generate_wiki_page(&chunk, "摘要", &config, &root, None).await.unwrap();
assert!(!doc.content.contains("```mermaid"), "坏图不应再以 mermaid 块出现");
assert!(doc.content.contains("```text"), "坏块应降级为 text fence");
assert!(doc.content.contains("code-repo-wiki: mermaid parse failed"), "应含降级标记注释");
assert_eq!(
provider.calls.load(std::sync::atomic::Ordering::Relaxed),
MERMAID_RETRY_MAX + 1,
"应调用 MERMAID_RETRY_MAX+1 次后降级"
);
let _ = std::fs::remove_dir_all(&dir);
}
#[tokio::test]
async fn test_architecture_degrades_on_bad_mermaid() {
let provider = ScriptedProvider::new(vec![
"```mermaid\nflowchart LR\nA[hello world\n```\n".to_string(),
"```mermaid\nflowchart LR\nA[hello world\n```\n".to_string(),
"```mermaid\nflowchart LR\nA[hello world\n```\n".to_string(),
]);
let generator = WikiGenerator::new(&provider, 0);
let config = WikiConfig::default();
let graph = crate::model::KnowledgeGraph::default();
let output = crate::generate::GenerationOutput {
cards: vec![],
documents: vec![],
generation_stats: crate::generate::GenerationStats::default(),
timings: crate::GenerationTimings::default(),
};
let root = crate::project::ProjectRoot::new(
std::env::temp_dir().join(format!("rw_arch_mermaid_{}", std::process::id())),
);
let doc = generator
.generate_architecture(&output, &graph, &config, &root)
.await
.unwrap();
assert!(!doc.content.contains("```mermaid"), "坏图不应再以 mermaid 块出现");
assert!(doc.content.contains("code-repo-wiki: mermaid parse failed"), "应含降级标记注释");
let _ = std::fs::remove_dir_all(root.path());
}
#[test]
fn test_build_references() {
let chunk = Chunk {
module_path: vec!["crate".into(), "net".into()],
entities: vec![],
imports: vec![],
dependencies: vec!["tokio".into(), "serde".into()],
file_paths: vec![],
entity_sources: vec![],
};
let refs = build_references(&chunk, "zh");
assert_eq!(refs.len(), 2);
assert_eq!(refs[0].target_title, "tokio");
assert_eq!(refs[0].target_path, "wiki/zh/tokio.md");
}
#[test]
fn test_build_references_uses_underscore_like_write_path() {
let chunk = Chunk {
module_path: vec!["src".into(), "generate".into()],
entities: vec![],
imports: vec![],
dependencies: vec!["src::analysis".into(), "src::output".into()],
file_paths: vec![],
entity_sources: vec![],
};
let refs = build_references(&chunk, "zh");
assert_eq!(refs[0].target_path, "wiki/zh/src_analysis.md");
assert_eq!(refs[1].target_path, "wiki/zh/src_output.md");
}
#[test]
fn test_module_description_prompt_shape() {
let messages = prompt::module_description_prompt(
"src::net",
&["connect".into(), "listen".into()],
"zh",
);
let user = &messages[1].content;
assert!(user.contains("src::net"), "应含模块名");
assert!(user.contains("connect"), "应含实体名");
assert!(user.contains("listen"), "应含实体名");
assert!(messages[0].content.contains("30"), "zh 应约束 30 字内");
}
#[tokio::test]
async fn test_describe_modules_enriches_description() {
use crate::model::{CodeEdge, CodeNode, ModuleCluster, NodeId, NodeKind};
use petgraph::stable_graph::StableDiGraph;
let mut g = StableDiGraph::<CodeNode, CodeEdge>::new();
let f = g.add_node(CodeNode {
id: NodeId::new(0),
kind: NodeKind::File,
name: "net.rs".into(),
file_path: Some("src/net.rs".into()),
line_range: None,
doc_comment: None,
signature: None, visibility: None,
module_path: vec!["src".into(), "net".into()],
});
let e = g.add_node(CodeNode {
id: NodeId::new(1),
kind: NodeKind::Function,
name: "connect".into(),
file_path: Some("src/net.rs".into()),
line_range: None,
doc_comment: None,
signature: None, visibility: None,
module_path: vec!["src".into(), "net".into()],
});
let kg = KnowledgeGraph {
graph: g,
modules: vec![
ModuleCluster {
name: "src::net".into(),
node_ids: vec![f, e],
cohesion: 0.5,
coupling: 0.5,
description: None,
},
ModuleCluster {
name: "src".into(),
node_ids: vec![],
cohesion: 1.0,
coupling: 0.0,
description: None,
},
],
features: Vec::new(),
};
let provider = MockProvider::new();
let generator = WikiGenerator::new(&provider, 0);
let config = crate::config::schema::WikiConfig {
output_dir: Some(
std::env::temp_dir()
.join(format!("rw_desc_cache_test_{}", std::process::id())),
),
..Default::default()
};
let root = crate::project::ProjectRoot::new(
std::env::temp_dir().join(format!("rw_desc_root_test_{}", std::process::id())),
);
let enriched = generator.describe_modules(&kg, "zh", &config, &root).await;
assert_eq!(enriched.len(), 2);
assert!(enriched[0].description.is_some(), "带实体的模块应获得描述");
assert_eq!(enriched[1].name, "src");
assert!(enriched[1].description.is_none(), "src 兜底模块不描述");
}
#[tokio::test]
async fn test_describe_modules_cache_hit_skips_llm() {
use crate::model::{CodeEdge, CodeNode, ModuleCluster, NodeId, NodeKind};
use petgraph::stable_graph::StableDiGraph;
let mut g = StableDiGraph::<CodeNode, CodeEdge>::new();
g.add_node(CodeNode {
id: NodeId::new(0),
kind: NodeKind::File,
name: "net.rs".into(),
file_path: Some("src/net.rs".into()),
line_range: None,
doc_comment: None,
signature: None,
visibility: None,
module_path: vec!["src".into(), "net".into()],
});
let kg = KnowledgeGraph {
graph: g,
modules: vec![
ModuleCluster {
name: "src::net".into(),
node_ids: vec![NodeId::new(0)],
cohesion: 1.0,
coupling: 0.0,
description: None,
},
],
features: Vec::new(),
};
let provider = MockProvider::new();
let generator = WikiGenerator::new(&provider, 0);
let out_dir = std::env::temp_dir().join(format!("rw_desc_cache_hit_{}", std::process::id()));
let _ = std::fs::remove_dir_all(&out_dir);
let root_dir = std::env::temp_dir().join(format!("rw_desc_root_hit_{}", std::process::id()));
let _ = std::fs::remove_dir_all(&root_dir);
let config = crate::config::schema::WikiConfig {
output_dir: Some(out_dir),
..Default::default()
};
let root = crate::project::ProjectRoot::new(root_dir);
let first = generator.describe_modules(&kg, "zh", &config, &root).await;
assert!(first[0].description.is_some(), "首次应走 LLM 获得描述");
let calls_after_first = generator.llm_call_count();
assert!(calls_after_first > 0, "首次必须真实调用 LLM");
let second = generator.describe_modules(&kg, "zh", &config, &root).await;
assert_eq!(second[0].description, first[0].description, "缓存应返回相同描述");
assert_eq!(
generator.llm_call_count(),
calls_after_first,
"第二次调用必须命中缓存、不触发 LLM"
);
}
#[tokio::test]
async fn test_describe_modules_cache_invalidated_by_file_change() {
use crate::model::{CodeEdge, CodeNode, ModuleCluster, NodeId, NodeKind};
use petgraph::stable_graph::StableDiGraph;
let dir = std::env::temp_dir().join(format!("rw_desc_cache_chg_{}", std::process::id()));
let _ = std::fs::remove_dir_all(&dir);
std::fs::create_dir_all(dir.join("src")).unwrap();
std::fs::write(dir.join("src/net.rs"), "pub fn connect() {}\n").unwrap();
let mut g = StableDiGraph::<CodeNode, CodeEdge>::new();
g.add_node(CodeNode {
id: NodeId::new(0),
kind: NodeKind::File,
name: "net.rs".into(),
file_path: Some("src/net.rs".into()),
line_range: None,
doc_comment: None,
signature: None,
visibility: None,
module_path: vec!["src".into(), "net".into()],
});
let kg = KnowledgeGraph {
graph: g,
modules: vec![ModuleCluster {
name: "src::net".into(),
node_ids: vec![NodeId::new(0)],
cohesion: 1.0,
coupling: 0.0,
description: None,
}],
features: Vec::new(),
};
let provider = MockProvider::new();
let generator = WikiGenerator::new(&provider, 0);
let config = crate::config::schema::WikiConfig {
output_dir: Some(dir.join(".code-repo-wiki")),
..Default::default()
};
let root = crate::project::ProjectRoot::new(dir.clone());
generator.describe_modules(&kg, "zh", &config, &root).await;
let calls_after_first = generator.llm_call_count();
std::fs::write(dir.join("src/net.rs"), "pub fn connect() {}\npub fn listen() {}\n").unwrap();
generator.describe_modules(&kg, "zh", &config, &root).await;
assert!(
generator.llm_call_count() > calls_after_first,
"文件内容变化后必须重新调用 LLM"
);
let _ = std::fs::remove_dir_all(&dir);
}
#[tokio::test]
async fn test_describe_modules_recovers_from_corrupt_cache() {
use crate::model::{CodeEdge, CodeNode, ModuleCluster, NodeId, NodeKind};
use petgraph::stable_graph::StableDiGraph;
let dir = std::env::temp_dir().join(format!("rw_desc_cache_corrupt_{}", std::process::id()));
let _ = std::fs::remove_dir_all(&dir);
std::fs::create_dir_all(dir.join(".code-repo-wiki/.state")).unwrap();
std::fs::write(dir.join(".code-repo-wiki/.state/module_descriptions.json"), "{not-json").unwrap();
let mut g = StableDiGraph::<CodeNode, CodeEdge>::new();
g.add_node(CodeNode {
id: NodeId::new(0),
kind: NodeKind::File,
name: "net.rs".into(),
file_path: Some("src/net.rs".into()),
line_range: None,
doc_comment: None,
signature: None,
visibility: None,
module_path: vec!["src".into(), "net".into()],
});
let kg = KnowledgeGraph {
graph: g,
modules: vec![ModuleCluster {
name: "src::net".into(),
node_ids: vec![NodeId::new(0)],
cohesion: 1.0,
coupling: 0.0,
description: None,
}],
features: Vec::new(),
};
let provider = MockProvider::new();
let generator = WikiGenerator::new(&provider, 0);
let config = crate::config::schema::WikiConfig {
output_dir: Some(dir.join(".code-repo-wiki")),
..Default::default()
};
let root = crate::project::ProjectRoot::new(dir.clone());
let modules = generator.describe_modules(&kg, "zh", &config, &root).await;
assert!(modules[0].description.is_some(), "损坏缓存回退后应重新生成描述");
assert!(generator.llm_call_count() > 0, "损坏缓存必须触发 LLM 调用");
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn test_collect_module_entity_names_prioritizes_behavior_over_fields() {
use crate::model::{CodeEdge, CodeNode, ModuleCluster, NodeId, NodeKind};
use petgraph::stable_graph::StableDiGraph;
let mut g = StableDiGraph::<CodeNode, CodeEdge>::new();
let mut ids = Vec::new();
for i in 0..8 {
ids.push(g.add_node(CodeNode {
id: NodeId::new(i as usize),
kind: NodeKind::Variable,
name: format!("field_{i}"),
file_path: None,
line_range: None,
doc_comment: None,
signature: None,
visibility: None,
module_path: vec!["src".into(), "net".into()],
}));
}
for i in 0..5 {
ids.push(g.add_node(CodeNode {
id: NodeId::new(100 + i as usize),
kind: NodeKind::Constant,
name: format!("const_{i}"),
file_path: None,
line_range: None,
doc_comment: None,
signature: None,
visibility: None,
module_path: vec!["src".into(), "net".into()],
}));
}
for i in 0..10 {
ids.push(g.add_node(CodeNode {
id: NodeId::new(200 + i as usize),
kind: NodeKind::Function,
name: format!("fn_{i:02}"),
file_path: None,
line_range: None,
doc_comment: None,
signature: None,
visibility: None,
module_path: vec!["src".into(), "net".into()],
}));
}
for i in 0..5 {
ids.push(g.add_node(CodeNode {
id: NodeId::new(300 + i as usize),
kind: NodeKind::Struct,
name: format!("struct_{i}"),
file_path: None,
line_range: None,
doc_comment: None,
signature: None,
visibility: None,
module_path: vec!["src".into(), "net".into()],
}));
}
ids.push(g.add_node(CodeNode {
id: NodeId::new(999),
kind: NodeKind::File,
name: "net.rs".into(),
file_path: None,
line_range: None,
doc_comment: None,
signature: None,
visibility: None,
module_path: vec!["src".into(), "net".into()],
}));
let module = ModuleCluster {
name: "src::net".into(),
node_ids: ids,
cohesion: 0.5,
coupling: 0.5,
description: None,
};
let graph = KnowledgeGraph {
graph: g,
modules: vec![module.clone()],
features: Vec::new(),
};
let names = collect_module_entity_names(&module, &graph);
assert_eq!(names.len(), 20, "变量与容器不进入名额");
assert!(
!names.iter().any(|n| n.starts_with("field_")),
"字段级实体(variable)应被排除"
);
assert!(
!names.iter().any(|n| n == "net.rs"),
"容器节点(File)应被排除"
);
let fn_pos = names.iter().position(|n| n == "fn_00").unwrap();
let struct_pos = names.iter().position(|n| n == "struct_0").unwrap();
let const_pos = names.iter().position(|n| n == "const_0").unwrap();
assert!(fn_pos < struct_pos && struct_pos < const_pos, "优先级序: 函数 < 结构体 < 常量");
assert!(names.iter().position(|n| n == "fn_00").unwrap() < names.iter().position(|n| n == "fn_01").unwrap());
}
#[test]
fn test_collect_module_entity_names_caps_at_30() {
use crate::model::{CodeEdge, CodeNode, ModuleCluster, NodeId, NodeKind};
use petgraph::stable_graph::StableDiGraph;
let mut g = StableDiGraph::<CodeNode, CodeEdge>::new();
let mut ids = Vec::new();
for i in 0..40 {
ids.push(g.add_node(CodeNode {
id: NodeId::new(i as usize),
kind: NodeKind::Function,
name: format!("fn_{i:02}"),
file_path: None,
line_range: None,
doc_comment: None,
signature: None,
visibility: None,
module_path: vec!["src".into()],
}));
}
let module = ModuleCluster {
name: "src".into(),
node_ids: ids,
cohesion: 1.0,
coupling: 0.0,
description: None,
};
let graph = KnowledgeGraph {
graph: g,
modules: vec![module.clone()],
features: Vec::new(),
};
let names = collect_module_entity_names(&module, &graph);
assert_eq!(names.len(), DESCRIBE_ENTITY_CAP);
assert_eq!(names[0], "fn_00", "字典序稳定");
assert_eq!(names[29], "fn_29", "截断取前 30");
}
#[test]
fn test_overview_prompt_includes_card_summaries() {
use crate::model::KnowledgeCard;
let graph = KnowledgeGraph::default();
let config = WikiConfig::default();
let card = KnowledgeCard {
module_name: "src::net".into(),
module_type: "module".into(),
summary: "网络模块:连接管理与监听".into(),
key_entities: vec![crate::model::EntitySummary {
name: "connect".into(),
kind: "function".into(),
visibility: "public".into(),
doc: None,
source: None,
}],
dependencies: vec![],
dependents: vec![],
design_patterns: vec![],
todo_notes: vec![],
related_files: vec![],
coding_spec: None,
tech_stack: vec![],
architecture: None,
pending_manual_edits: vec![],
features: Vec::new(),
};
let prompt = overview_prompt(&[], &[card], &graph, &config);
assert!(prompt.contains("## 模块卡片摘要"), "应含卡片摘要节");
assert!(prompt.contains("src::net"), "应含模块名");
assert!(prompt.contains("网络模块"), "应含卡片摘要");
assert!(prompt.contains("connect"), "应含关键实体");
}
}