Files

324 lines
12 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""文档入库测试."""
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())