"""文档入库测试.""" import tempfile from pathlib import Path import pytest from src.core.ingest import DocumentIngestor from src.core.splitters import MarkdownSplitter class TestMarkdownSplitter: """Markdown 分块器测试.""" @pytest.fixture def splitter(self): return MarkdownSplitter(max_size=1000, overlap=100) def test_split_simple_document(self, splitter): """简单文档按标题拆分.""" md = """# 标题一 这是第一段内容。 ## 标题二 这是第二段内容。 # 标题三 这是第三段内容。""" chunks = splitter.split(md, source_file="test.md") assert len(chunks) >= 3 # 每个 chunk 有元数据 for chunk in chunks: assert "content" in chunk assert chunk["source_file"] == "test.md" def test_chunk_has_heading_metadata(self, splitter): """chunk 附带标题元数据.""" md = "# 配置指南\n这里是配置说明。" chunks = splitter.split(md, source_file="config.md") assert len(chunks) >= 1 title = chunks[0]["section_title"] assert "配置指南" in title or title == "" def test_long_section_is_split(self, splitter): """超长章节被进一步拆分.""" # 创建一个超过 max_size 的段落 long_text = "这是很长的文本。" * 300 # ~3000 字符 md = f"# 长章节\n{long_text}" small_splitter = MarkdownSplitter(max_size=500, overlap=50) chunks = small_splitter.split(md, source_file="long.md") assert len(chunks) > 1 def test_empty_document(self, splitter): """空文档返回空列表.""" chunks = splitter.split("", source_file="empty.md") assert chunks == [] def test_code_blocks_preserved(self, splitter): """代码块不被拆分.""" md = """# 代码示例 ```python def hello(): print("world") ``` """ chunks = splitter.split(md, source_file="code.md") assert len(chunks) >= 1 # 代码块内容应在某个 chunk 中 all_content = " ".join(c["content"] for c in chunks) assert "def hello()" in all_content def test_separator_near_start_does_not_loop(self, splitter): """分隔符紧挨 start 时不会死循环 (回归测试, fix: start=max(start+1, next_start)). 场景: 超长段落中, 分隔符出现在距离 start 小于 overlap 的位置, _split_single_paragraph 的 start 会回退为负数, str.rfind 负索引绕回导致死循环. """ # 100 个句号 + 大量内容 → 句号密集在开头且很近 text = "。" * 80 + "内容文本" * 500 md = f"# 边界测试\n{text}" chunks = splitter.split(md, source_file="edge.md") # 不卡死即通过 assert len(chunks) > 0 # 验证内容完整 all_text = "".join(c["content"] for c in chunks) assert "内容文本" in all_text def test_dense_separators_in_long_para(self, splitter): """超长段落中分隔符密集分布也能正确分块.""" # ~2000 字符: 每段 20 个"内容文本" + "。",共 30 段 text = "" for i in range(30): text += "内容文本" * 20 + "。" * (3 if i % 5 == 0 else 1) + "\n" md = f"# 密集分隔符\n{text}" chunks = splitter.split(md, source_file="dense.md") assert len(chunks) > 1 # 超过 1000 字符应被拆分 class TestDocumentIngestor: """文档入库器测试.""" @pytest.fixture def temp_md_dir(self): with tempfile.TemporaryDirectory() as d: # 创建测试 Markdown 文件 md_path = Path(d) / "test.md" md_path.write_text("# 测试\n这是测试内容。", encoding="utf-8") yield d def test_read_markdown_file(self, temp_md_dir): """读取 Markdown 文件.""" content = Path(temp_md_dir + "/test.md").read_text(encoding="utf-8") assert "测试" in content assert "这是测试内容" in content class TestIngestorIntegration: """入库器集成测试 (使用真实 embedder).""" def test_ingest_content_real(self, tmp_path): """真实入库: 分块→嵌入→入库.""" from src.core.config import EmbedConfig from src.core.db import VectorDB from src.core.embedder import create_embedder db = VectorDB(persist_dir=str(tmp_path)) embedder = create_embedder(EmbedConfig(mode="local")) ingestor = DocumentIngestor(db, embedder, "test_integration") count = ingestor.ingest_content("# Hello\nWorld.", "hello.md") assert count > 0 assert ingestor.collection.count() == count def test_ingest_deduplicates(self, tmp_path): """重复入库同一文件会去重.""" from src.core.config import EmbedConfig from src.core.db import VectorDB from src.core.embedder import create_embedder db = VectorDB(persist_dir=str(tmp_path)) embedder = create_embedder(EmbedConfig(mode="local")) ingestor = DocumentIngestor(db, embedder, "test_dedup") ingestor.ingest_content("# A", "dup.md") c2 = ingestor.ingest_content("# B", "dup.md") assert ingestor.collection.count() == c2 class TestIngestFile: """ingest_file 方法测试.""" def test_ingest_file_markdown(self, tmp_path): """通过文件路径入库 .md 文件.""" from src.core.config import EmbedConfig from src.core.db import VectorDB from src.core.embedder import create_embedder md_file = tmp_path / "hello.md" md_file.write_text("# 测试\n这是测试内容。", encoding="utf-8") db = VectorDB(persist_dir=str(tmp_path / "db")) embedder = create_embedder(EmbedConfig(mode="local")) ingestor = DocumentIngestor(db, embedder, "test_file") count = ingestor.ingest_file(str(md_file)) assert count > 0 assert ingestor.collection.count() == count def test_ingest_file_text(self, tmp_path): """通过文件路径入库 .txt 文件.""" from src.core.config import EmbedConfig from src.core.db import VectorDB from src.core.embedder import create_embedder txt_file = tmp_path / "notes.txt" txt_file.write_text("这是一段纯文本内容。\n\n第二段内容在这里。", encoding="utf-8") db = VectorDB(persist_dir=str(tmp_path / "db")) embedder = create_embedder(EmbedConfig(mode="local")) ingestor = DocumentIngestor(db, embedder, "test_txt") count = ingestor.ingest_file(str(txt_file)) assert count > 0 class TestIngestDirectory: """ingest_directory 方法测试.""" def test_ingest_directory_mixed_formats(self, tmp_path): """入库包含多种格式的目录.""" from src.core.config import EmbedConfig from src.core.db import VectorDB from src.core.embedder import create_embedder (tmp_path / "a.md").write_text("# A\n内容 A", encoding="utf-8") (tmp_path / "b.txt").write_text("内容 B", encoding="utf-8") (tmp_path / "not_supported.xyz").write_text("不应被处理", encoding="utf-8") db = VectorDB(persist_dir=str(tmp_path / "db")) embedder = create_embedder(EmbedConfig(mode="local")) ingestor = DocumentIngestor(db, embedder, "test_dir") results = ingestor.ingest_directory(str(tmp_path)) assert len(results) >= 2 # a.md + b.txt, .xyz 被忽略 def test_ingest_directory_empty(self, tmp_path): """空目录返回空结果.""" from src.core.config import EmbedConfig from src.core.db import VectorDB from src.core.embedder import create_embedder db = VectorDB(persist_dir=str(tmp_path / "db")) embedder = create_embedder(EmbedConfig(mode="local")) ingestor = DocumentIngestor(db, embedder, "test_empty_dir") results = ingestor.ingest_directory(str(tmp_path)) assert results == {} class TestIncrementalIngest: """增量入库测试.""" def test_ingest_file_incremental_skips_unchanged(self, tmp_path): """增量模式: 未修改的文件跳过入库.""" from src.core.config import ChunkConfig, EmbedConfig from src.core.db import VectorDB from src.core.embedder import create_embedder from src.core.file_tracker import FileTracker file = tmp_path / "stable.md" file.write_text("# 稳定文档\n\n内容不变。", encoding="utf-8") tracker_path = str(tmp_path / "tracker.json") db = VectorDB(persist_dir=str(tmp_path / "db")) embedder = create_embedder(EmbedConfig(mode="local")) ingestor = DocumentIngestor( db, embedder, "test_incr", chunk_config=ChunkConfig(max_size=1000, overlap=100), file_tracker=FileTracker(tracker_path), ) count1 = ingestor.ingest_file(str(file), incremental=True) assert count1 > 0 count2 = ingestor.ingest_file(str(file), incremental=True) assert count2 == 0 # 跳过 def test_ingest_file_incremental_reingests_modified(self, tmp_path): """增量模式: 修改后的文件重新入库.""" from src.core.config import ChunkConfig, EmbedConfig from src.core.db import VectorDB from src.core.embedder import create_embedder from src.core.file_tracker import FileTracker file = tmp_path / "changing.md" file.write_text("# v1\n\n初始版本的内容段落。", encoding="utf-8") tracker_path = str(tmp_path / "tracker2.json") db = VectorDB(persist_dir=str(tmp_path / "db2")) embedder = create_embedder(EmbedConfig(mode="local")) ingestor = DocumentIngestor( db, embedder, "test_incr2", chunk_config=ChunkConfig(max_size=1000, overlap=100), file_tracker=FileTracker(tracker_path), ) count1 = ingestor.ingest_file(str(file), incremental=True) assert count1 > 0 file.write_text("# v2\n\n新增段落,内容完全不同了。", encoding="utf-8") count2 = ingestor.ingest_file(str(file), incremental=True) assert count2 > 0 def test_ingest_file_force_mode_always_reingests(self, tmp_path): """force=True 时始终重新入库(忽略 tracker).""" from src.core.config import ChunkConfig, EmbedConfig from src.core.db import VectorDB from src.core.embedder import create_embedder from src.core.file_tracker import FileTracker file = tmp_path / "force.md" file.write_text("# force test\n\n这是强制入库测试的内容。", encoding="utf-8") tracker_path = str(tmp_path / "tracker3.json") db = VectorDB(persist_dir=str(tmp_path / "db3")) embedder = create_embedder(EmbedConfig(mode="local")) ingestor = DocumentIngestor( db, embedder, "test_force", chunk_config=ChunkConfig(max_size=1000, overlap=100), file_tracker=FileTracker(tracker_path), ) count1 = ingestor.ingest_file(str(file), incremental=True) count2 = ingestor.ingest_file(str(file), incremental=True, force=True) assert count1 > 0 assert count2 > 0 # force 模式重新入库 def test_ingest_content_default_splitter(tmp_path): """未指定 splitter 时用 MarkdownSplitter.""" from src.core.config import EmbedConfig from src.core.db import VectorDB from src.core.embedder import create_embedder from src.core.ingest import DocumentIngestor db = VectorDB(persist_dir=str(tmp_path / "db")) embedder = create_embedder(EmbedConfig(mode="local")) ingestor = DocumentIngestor(db, embedder, "test_content") count = ingestor.ingest_content("# 测试\n\n一些内容。", "test.md") assert count >= 1 def test_ingest_directory_recursive(tmp_path): """ingest_directory 递归处理子目录.""" from src.core.config import EmbedConfig, ChunkConfig from src.core.db import VectorDB from src.core.embedder import create_embedder from src.core.ingest import DocumentIngestor (tmp_path / "sub").mkdir() (tmp_path / "a.md").write_text("# A\n\n内容A。", encoding="utf-8") (tmp_path / "sub" / "b.md").write_text("# B\n\n内容B。", encoding="utf-8") db = VectorDB(persist_dir=str(tmp_path / "db_r")) embedder = create_embedder(EmbedConfig(mode="local")) ingestor = DocumentIngestor( db, embedder, "test_recurse", chunk_config=ChunkConfig(max_size=1000, overlap=100), ) results = ingestor.ingest_directory(str(tmp_path)) assert len(results) >= 2 assert all(v > 0 for v in results.values())