use std::path::Path;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Citation {
pub path: String,
pub start: usize,
pub end: usize,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct InvalidCitation {
pub citation: Citation,
pub reason: String,
}
fn is_path_char(c: u8) -> bool {
c.is_ascii_alphanumeric() || matches!(c, b'.' | b'_' | b'/' | b'\\' | b'-')
}
fn has_valid_extension(path: &str) -> bool {
let last = path.rsplit(['/', '\\']).next().unwrap();
match last.rfind('.') {
Some(dot) => {
let ext = &last[dot + 1..];
!ext.is_empty() && ext.as_bytes()[0].is_ascii_alphabetic()
}
None => false,
}
}
fn fence_ranges(content: &str) -> Vec<(usize, usize)> {
let mut ranges = Vec::new();
let mut in_fence = false;
let mut fence_start = 0usize;
let mut offset = 0usize;
for line in content.split('\n') {
let trimmed = line.trim_start();
if in_fence {
if trimmed.starts_with("```") {
ranges.push((fence_start, offset + line.len() + 1));
in_fence = false;
}
} else if trimmed.starts_with("```") {
in_fence = true;
fence_start = offset;
}
offset += line.len() + 1;
}
if in_fence {
ranges.push((fence_start, offset));
}
ranges
}
pub fn extract_citations(content: &str) -> Vec<Citation> {
let fences = fence_ranges(content);
let bytes = content.as_bytes();
let mut out = Vec::new();
let mut i = 0;
let mut fence_idx = 0usize;
while i < bytes.len() {
while fence_idx < fences.len() && i >= fences[fence_idx].1 {
fence_idx += 1;
}
if fence_idx < fences.len() && i >= fences[fence_idx].0 {
i = fences[fence_idx].1;
continue;
}
if bytes[i] == b':' {
let mut p = i;
while p > 0 && is_path_char(bytes[p - 1]) {
p -= 1;
}
let mut path = &content[p..i];
path = path.trim_start_matches('-');
let has_drive_prefix = (p >= 2
&& bytes[p - 2].is_ascii_alphabetic()
&& bytes[p - 1] == b':'
&& matches!(bytes.get(p), Some(b'\\' | b'/')))
|| (path.len() >= 3
&& path.as_bytes()[0].is_ascii_alphabetic()
&& path.as_bytes()[1] == b':'
&& matches!(path.as_bytes()[2], b'\\' | b'/'));
let mut j = i + 1;
while j < bytes.len() && bytes[j].is_ascii_digit() {
j += 1;
}
let line_str = &content[i + 1..j];
let mut end: Option<usize> = None;
let mut consumed = j;
if j < bytes.len() && bytes[j] == b'-' {
let mut k = j + 1;
while k < bytes.len() && bytes[k].is_ascii_digit() {
k += 1;
}
if k > j + 1 {
end = content[j + 1..k].parse().ok();
consumed = k;
}
}
if !path.is_empty()
&& path.contains('.')
&& !path.contains("//")
&& !path.starts_with('/')
&& !has_drive_prefix
&& has_valid_extension(path)
&& !line_str.is_empty()
&& let Ok(start) = line_str.parse::<usize>()
&& start > 0
{
out.push(Citation {
path: path.to_string(),
start,
end: end.unwrap_or(start).max(start),
});
}
i = consumed;
} else {
i += 1;
}
}
out
}
fn check_citation_file_level(root: &Path, citation: &Citation) -> Option<String> {
if citation.path.split(['/', '\\']).any(|seg| seg == "..") {
return Some("路径含越界段 ..".to_string());
}
let abs = root.join(&citation.path);
let total_lines = std::fs::read_to_string(&abs)
.map(|s| s.lines().count())
.ok();
match total_lines {
None => Some("引用文件不存在或不可读".to_string()),
Some(n) if citation.end > n => {
Some(format!("行号越界: {}-{} 超出文件总行数 {}", citation.start, citation.end, n))
}
_ => None,
}
}
pub fn validate_citations(root: &Path, content: &str) -> Vec<InvalidCitation> {
extract_citations(content)
.into_iter()
.filter_map(|citation| {
check_citation_file_level(root, &citation)
.map(|reason| InvalidCitation { citation, reason })
})
.collect()
}
pub type EntityRanges = std::collections::HashMap<String, Vec<(usize, usize)>>;
pub fn citation_overlaps_entity(c: &Citation, ranges: &[(usize, usize)]) -> bool {
ranges
.iter()
.any(|&(entity_start, entity_end)| c.start <= entity_end && c.end >= entity_start)
}
pub fn validate_citations_against_entities(
root: &Path,
content: &str,
entity_ranges: &EntityRanges,
) -> Vec<InvalidCitation> {
let mut invalid = validate_citations(root, content);
let file_level_valid: Vec<Citation> = extract_citations(content)
.into_iter()
.filter(|c| check_citation_file_level(root, c).is_none())
.collect();
for citation in file_level_valid {
let key = crate::incremental::norm_sep(&citation.path);
if let Some(ranges) = entity_ranges.get(&key)
&& !citation_overlaps_entity(&citation, ranges)
{
invalid.push(InvalidCitation {
citation,
reason: "引用区间未覆盖任何实体(行号可能指向错误位置)".to_string(),
});
}
}
invalid
}
pub fn retry_feedback(invalid: &[InvalidCitation]) -> String {
let mut lines = String::from("上一版输出存在无效的源码引用,请修正后重新输出完整文档:\n");
for item in invalid {
lines.push_str(&format!(
"- `{}:{}` — {}\n",
item.citation.path, item.citation.start, item.reason
));
}
lines.push_str(
"要求:提及任何具体函数/结构体/文件时,必须携带真实存在的 `相对路径:行号` \
(如 `src/fs.rs:28`)或 `相对路径:起始行-结束行`(如 `src/fs.rs:28-45`)引用;\
不得编造不存在的文件或行号。",
);
lines
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_extract_single_line() {
let text = "见 src/fs.rs:28 的实现";
let cites = extract_citations(text);
assert_eq!(cites.len(), 1);
assert_eq!(cites[0].path, "src/fs.rs");
assert_eq!(cites[0].start, 28);
assert_eq!(cites[0].end, 28);
}
#[test]
fn test_extract_range() {
let text = "核心逻辑在 src/generate/mod.rs:164-179";
let cites = extract_citations(text);
assert_eq!(cites.len(), 1);
assert_eq!(cites[0].path, "src/generate/mod.rs");
assert_eq!(cites[0].start, 164);
assert_eq!(cites[0].end, 179);
}
#[test]
fn test_ignore_urls_and_times() {
let text = "见 https://example.com/a.rs:10,发生在 12:30 分";
let cites = extract_citations(text);
assert!(cites.is_empty(), "URL/时间格式不得误报: {:?}", cites);
}
#[test]
fn test_ignore_absolute_path_and_zero_line() {
assert!(extract_citations("绝对路径 /usr/bin/x.rs:5").is_empty());
assert!(extract_citations("零行号 src/a.rs:0").is_empty());
}
#[test]
fn test_multiple_citations() {
let text = "a.rs:3 与 src/b.rs:10-12 和 docs/c.md:1";
let cites = extract_citations(text);
assert_eq!(cites.len(), 3);
}
#[test]
fn test_windows_path_separator() {
let cites = extract_citations("见 src\\auth.rs:2 的实现");
assert_eq!(cites.len(), 1, "反斜杠路径应完整提取: {:?}", cites);
assert_eq!(cites[0].path, "src\\auth.rs");
assert_eq!(cites[0].start, 2);
assert!(extract_citations("在 C:\\repo\\x.rs:5 中").is_empty(), "盘符绝对路径应忽略");
}
#[test]
fn test_forward_slash_drive_prefix() {
assert!(extract_citations("见 C:/repo/x.rs:5 的实现").is_empty(), "正斜杠盘符应忽略");
assert!(extract_citations("-C:/repo/x.rs:5").is_empty(), "连字符+盘符应忽略");
assert!(extract_citations("在 C:\\repo\\x.rs:5 中").is_empty(), "反斜杠盘符应忽略");
}
#[test]
fn test_version_numbers_not_extracted() {
assert!(extract_citations("版本 v2.0:10 发布").is_empty(), "v2.0 是版本号");
assert!(extract_citations("时刻 1.2:30 记录").is_empty(), "1.2 是时刻");
let cites = extract_citations("见 src/v1.5.rs:3 的实现");
assert_eq!(cites.len(), 1);
assert_eq!(cites[0].path, "src/v1.5.rs");
assert_eq!(cites[0].start, 3);
}
#[test]
fn test_leading_hyphen_stripped() {
let cites = extract_citations("-src/fs.rs:10");
assert_eq!(cites.len(), 1);
assert_eq!(cites[0].path, "src/fs.rs");
let cites = extract_citations("见 my-file.rs:3 的实现");
assert_eq!(cites.len(), 1);
assert_eq!(cites[0].path, "my-file.rs");
}
#[test]
fn test_dotdot_paths_rejected_in_validate() {
let dir = std::env::temp_dir().join(format!("code_repo_wiki_dotdot_{}", std::process::id()));
let _ = std::fs::remove_dir_all(&dir);
std::fs::create_dir_all(dir.join("src")).unwrap();
std::fs::write(dir.join("lib.rs"), "line1\n").unwrap();
std::fs::write(dir.join("src").join("x.rs"), "line1\n").unwrap();
let cites = extract_citations("见 ../src/x.rs:5 与 src/../lib.rs:3");
assert_eq!(cites.len(), 2, "提取层应保留 .. 路径: {:?}", cites);
assert_eq!(cites[0].path, "../src/x.rs");
assert_eq!(cites[1].path, "src/../lib.rs");
let invalid = validate_citations(&dir, "见 ../src/x.rs:5 与 src/../lib.rs:3");
assert_eq!(invalid.len(), 2, "校验层应拒绝 .. 路径: {:?}", invalid);
for item in &invalid {
assert!(item.reason.contains("越界段 .."), "原因应说明越界: {:?}", item);
}
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn test_validate_ok_and_missing() {
let dir = std::env::temp_dir().join(format!("code_repo_wiki_cite_{}", 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").join("a.rs"), "line1\nline2\nline3\n").unwrap();
let invalid = validate_citations(&dir, "见 src/a.rs:3");
assert!(invalid.is_empty(), "存在的文件与合法行号应通过: {:?}", invalid);
let invalid = validate_citations(&dir, "见 src/a.rs:9");
assert_eq!(invalid.len(), 1);
assert_eq!(invalid[0].citation.path, "src/a.rs");
assert!(invalid[0].reason.contains("越界"));
let invalid = validate_citations(&dir, "见 src/missing.rs:1");
assert_eq!(invalid.len(), 1);
assert!(invalid[0].reason.contains("不存在"));
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn test_retry_feedback_lists_all() {
let invalid = vec![InvalidCitation {
citation: Citation { path: "src/a.rs".into(), start: 99, end: 99 },
reason: "行号越界".into(),
}];
let feedback = retry_feedback(&invalid);
assert!(feedback.contains("src/a.rs:99"));
assert!(feedback.contains("行号越界"));
assert!(feedback.contains("重新输出完整文档"));
}
}
#[test]
fn test_extract_skips_fenced_code_blocks() {
let text = "见 src/a.rs:1 的实现。\n\n```rust\nlet cfg = load(\"src/config.rs:99\");\n```\n";
let cites = extract_citations(text);
assert_eq!(cites.len(), 1, "围栏外应提取 1 条, 实际: {cites:?}");
assert_eq!(cites[0].path, "src/a.rs");
}
#[test]
fn test_extract_skips_mermaid_blocks() {
let text = "```mermaid\nflowchart LR\nA --> |src/fs.rs:28| B\n```\n正文 src/b.rs:1\n";
let cites = extract_citations(text);
assert_eq!(cites.len(), 1, "mermaid 块内不应提取, 实际: {cites:?}");
assert_eq!(cites[0].path, "src/b.rs");
}
#[test]
fn test_extract_unclosed_fence_skips_to_end() {
let text = "```rust\nlet x = load(\"src/a.rs:1\");\n正文 src/b.rs:2\n";
let cites = extract_citations(text);
assert!(cites.is_empty(), "未闭合围栏后不应提取: {cites:?}");
}
#[test]
fn test_extract_alternating_fences() {
let text = "正文一 src/a.rs:1\n```rust\nsrc/x.rs:2\n```\n正文二 src/b.rs:3\n```text\nsrc/y.rs:4\n```\n正文三 src/c.rs:5\n";
let cites = extract_citations(text);
let paths: Vec<&str> = cites.iter().map(|c| c.path.as_str()).collect();
assert_eq!(paths, vec!["src/a.rs", "src/b.rs", "src/c.rs"], "只应提取正文引用: {paths:?}");
}
#[test]
fn test_extract_indented_fence() {
let text = "正文 src/a.rs:1\n ```rust\n src/b.rs:2\n ```\n";
let cites = extract_citations(text);
assert_eq!(cites.len(), 1, "缩进围栏内部不应提取, 实际: {cites:?}");
assert_eq!(cites[0].path, "src/a.rs");
}
#[test]
fn test_citation_overlaps_entity() {
let ranges = vec![(10usize, 20usize), (30, 40)];
let c = |start: usize, end: usize| Citation { path: "x.rs".into(), start, end };
assert!(citation_overlaps_entity(&c(12, 18), &ranges));
assert!(citation_overlaps_entity(&c(8, 15), &ranges), "跨入区间应算覆盖");
assert!(citation_overlaps_entity(&c(15, 25), &ranges), "跨出区间应算覆盖");
assert!(citation_overlaps_entity(&c(10, 10), &ranges), "起点即实体起点应覆盖");
assert!(citation_overlaps_entity(&c(20, 20), &ranges), "终点即实体终点应覆盖");
assert!(citation_overlaps_entity(&c(35, 35), &ranges), "第二个实体单行");
assert!(!citation_overlaps_entity(&c(21, 29), &ranges), "实体间隙不应覆盖");
assert!(!citation_overlaps_entity(&c(1, 5), &ranges), "实体之前不应覆盖");
assert!(!citation_overlaps_entity(&c(41, 50), &ranges), "实体之后不应覆盖");
assert!(!citation_overlaps_entity(&c(21, 21), &ranges), "紧邻实体终点外一行不覆盖");
assert!(!citation_overlaps_entity(&c(10, 10), &[]));
}
#[test]
fn test_validate_against_entities_overlap() {
let dir = std::env::temp_dir().join(format!("code_repo_wiki_cite_ovl_{}", std::process::id()));
let _ = std::fs::remove_dir_all(&dir);
std::fs::create_dir_all(dir.join("src")).unwrap();
let src = dir.join("src").join("a.rs");
let content: String = (1..=10).map(|i| format!("line{i}\n")).collect();
std::fs::write(&src, content).unwrap();
std::fs::write(dir.join("README.md"), "docs\n").unwrap();
let mut ranges: EntityRanges = EntityRanges::new();
ranges.insert("src/a.rs".to_string(), vec![(2, 4), (7, 9)]);
let invalid = validate_citations_against_entities(&dir, "见 src/a.rs:3", &ranges);
assert!(invalid.is_empty(), "覆盖实体应通过: {invalid:?}");
let invalid = validate_citations_against_entities(&dir, "见 src/a.rs:7-9", &ranges);
assert!(invalid.is_empty(), "区间覆盖第二个实体应通过: {invalid:?}");
let invalid = validate_citations_against_entities(&dir, "见 src/a.rs:5-6", &ranges);
assert_eq!(invalid.len(), 1, "实体间隙引用应无效: {invalid:?}");
assert!(invalid[0].reason.contains("未覆盖任何实体"), "原因应说明: {:?}", invalid[0]);
let invalid = validate_citations_against_entities(&dir, "见 src/a.rs:1", &ranges);
assert_eq!(invalid.len(), 1, "实体之前引用应无效: {invalid:?}");
let invalid = validate_citations_against_entities(&dir, "见 README.md:1", &ranges);
assert!(invalid.is_empty(), "无实体文件引用应放行: {invalid:?}");
let invalid = validate_citations_against_entities(&dir, "见 src/a.rs:99", &ranges);
assert_eq!(invalid.len(), 1);
assert!(invalid[0].reason.contains("越界"));
let mut win_ranges: EntityRanges = EntityRanges::new();
win_ranges.insert("src\\a.rs".to_string(), vec![(2, 4)]);
let invalid = validate_citations_against_entities(&dir, "见 src/a.rs:3", &win_ranges);
assert!(invalid.is_empty(), "反斜杠表键经 norm_sep 应命中: {invalid:?}");
let _ = std::fs::remove_dir_all(&dir);
}