refactor: 将 MarkdownSplitter 移入 splitters 包,ingest.py 保留兼容 import

This commit is contained in:
2026-07-10 14:15:59 +08:00
parent 30c4d65cf0
commit 4dc51dccd2
5 changed files with 110 additions and 166 deletions
+2 -165
View File
@@ -1,179 +1,16 @@
"""Markdown 文档解析与入库模块.""" """Markdown 文档解析与入库模块."""
import logging import logging
import re import re
from dataclasses import dataclass
from pathlib import Path from pathlib import Path
from typing import Protocol
from src.core.db import VectorDB from src.core.db import VectorDB
from src.core.embedder import Embedder, batch_embed from src.core.embedder import Embedder, batch_embed
from src.core.splitters.markdown import MarkdownSplitter # 兼容旧 import 路径
from src.core.splitters.base import Splitter # 兼容旧 import 路径
logger = logging.getLogger("md-vector-db") logger = logging.getLogger("md-vector-db")
class Splitter(Protocol):
"""文档分块器接口 — 将文本拆分为带元数据的 chunk 列表.
每个 chunk 为 dict: {"content": str, "section_title": str, "heading_level": int, ...}
"""
def split(self, text: str, source_file: str = "") -> list[dict]: ...
class MarkdownSplitter:
"""Markdown 混合分块器:先按标题拆,超长再按段落拆."""
def __init__(self, max_size: int = 1000, overlap: int = 100):
self.max_size = max_size
self.overlap = overlap
def split(self, text: str, source_file: str = "") -> list[dict]:
"""将 Markdown 文本拆分为带元数据的 chunk 列表."""
if not text.strip():
return []
sections = self._split_by_headings(text)
chunks = []
for section in sections:
if len(section["content"]) <= self.max_size:
chunks.append(section)
else:
sub_chunks = self._split_by_paragraphs(
section["content"],
section["section_title"],
section["heading_level"],
)
chunks.extend(sub_chunks)
# 为所有 chunk 补充 source_file 和 chunk_index
for i, chunk in enumerate(chunks):
chunk["source_file"] = source_file or chunk.get("source_file", "")
chunk["chunk_index"] = i
return chunks
def _split_by_headings(self, text: str) -> list[dict]:
"""按 Markdown 标题拆分."""
heading_pattern = re.compile(r"^(#{1,6})\s+(.+)$", re.MULTILINE)
matches = list(heading_pattern.finditer(text))
if not matches:
return [{
"content": text.strip(),
"section_title": "",
"heading_level": 0,
}]
sections = []
for i, match in enumerate(matches):
level = len(match.group(1))
title = match.group(2).strip()
start = match.end()
end = matches[i + 1].start() if i + 1 < len(matches) else len(text)
content = text[start:end].strip()
if content:
sections.append({
"content": f"{match.group(0)}\n{content}",
"section_title": title,
"heading_level": level,
})
# 处理第一个标题之前的内容
if matches and matches[0].start() > 0:
preamble = text[:matches[0].start()].strip()
if preamble:
sections.insert(0, {
"content": preamble,
"section_title": "",
"heading_level": 0,
})
return sections
def _split_by_paragraphs(
self, text: str, section_title: str, heading_level: int
) -> list[dict]:
"""按段落边界拆分超长章节.
优先在段落边界拆分,若单个段落仍超长则按字符硬切。
"""
paragraphs = re.split(r"\n\n+", text)
chunks = []
current = ""
for para in paragraphs:
# 单一段落超出 max_size 时直接硬切
if len(para) > self.max_size:
# 先 flush 当前累积
if current.strip():
chunks.append({
"content": current.strip(),
"section_title": section_title,
"heading_level": heading_level,
})
current = ""
# 硬切该段落
for sub in self._split_single_paragraph(para):
chunks.append({
"content": sub,
"section_title": section_title,
"heading_level": heading_level,
})
continue
if len(current) + len(para) > self.max_size and current:
chunks.append({
"content": current.strip(),
"section_title": section_title,
"heading_level": heading_level,
})
# overlap: 保留上一块的末尾部分
if self.overlap > 0 and len(current) > self.overlap:
current = current[-self.overlap:] + "\n\n" + para
else:
current = para
else:
if current:
current += "\n\n" + para
else:
current = para
if current.strip():
chunks.append({
"content": current.strip(),
"section_title": section_title,
"heading_level": heading_level,
})
return chunks
def _split_single_paragraph(self, text: str) -> list[str]:
"""按字符边界拆分单个超长段落(带 overlap)。"""
parts = []
start = 0
while start < len(text):
end = start + self.max_size
if end >= len(text):
parts.append(text[start:].strip())
break
# 尝试在句号或空格处断开
break_point = end
for sep in ("", "", "", "\n", ". ", " "):
pos = text.rfind(sep, start, end)
if pos > start:
break_point = pos + len(sep)
break
part = text[start:break_point].strip()
if part:
parts.append(part)
# 确保 start 始终前进(避免分隔符距 start 小于 overlap 时 start 回退导致死循环)
next_start = break_point - self.overlap if self.overlap > 0 else break_point
start = max(start + 1, next_start)
return parts
class DocumentIngestor: class DocumentIngestor:
"""文档入库器: 读取 MD 文件 → 分块 → 嵌入 → 入库.""" """文档入库器: 读取 MD 文件 → 分块 → 嵌入 → 入库."""
+68
View File
@@ -0,0 +1,68 @@
"""Markdown 文档分块器."""
import re
from src.core.splitters.base import BaseTextSplitter
class MarkdownSplitter(BaseTextSplitter):
"""Markdown 混合分块器:先按标题拆,超长再按段落拆."""
def split(self, text: str, source_file: str = "") -> list[dict]:
if not text.strip():
return []
sections = self._split_by_headings(text)
chunks = []
for section in sections:
if len(section["content"]) <= self.max_size:
chunks.append(section)
else:
sub_chunks = self._split_by_paragraphs(
section["content"],
section["section_title"],
section["heading_level"],
)
chunks.extend(sub_chunks)
for i, chunk in enumerate(chunks):
chunk["source_file"] = source_file or chunk.get("source_file", "")
chunk["chunk_index"] = i
return chunks
def _split_by_headings(self, text: str) -> list[dict]:
heading_pattern = re.compile(r"^(#{1,6})\s+(.+)$", re.MULTILINE)
matches = list(heading_pattern.finditer(text))
if not matches:
return [{
"content": text.strip(),
"section_title": "",
"heading_level": 0,
}]
sections = []
for i, match in enumerate(matches):
level = len(match.group(1))
title = match.group(2).strip()
start = match.end()
end = matches[i + 1].start() if i + 1 < len(matches) else len(text)
content = text[start:end].strip()
if content:
sections.append({
"content": f"{match.group(0)}\n{content}",
"section_title": title,
"heading_level": level,
})
if matches and matches[0].start() > 0:
preamble = text[:matches[0].start()].strip()
if preamble:
sections.insert(0, {
"content": preamble,
"section_title": "",
"heading_level": 0,
})
return sections
+22
View File
@@ -0,0 +1,22 @@
"""分块器注册表 — 按文件后缀名路由到对应 Splitter."""
SUPPORTED_SUFFIXES: dict[str, str] = {
".md": "markdown",
".txt": "text",
".html": "html",
".pdf": "pdf",
}
_registry: dict[str, type] = {}
def register_splitter(name: str, splitter_cls: type) -> None:
"""注册一个分块器."""
_registry[name] = splitter_cls
def get_splitter(name: str):
"""获取已注册的分块器类."""
if name not in _registry:
raise ValueError(f"未注册的分块器: {name}")
return _registry[name]
+16
View File
@@ -0,0 +1,16 @@
"""纯文本文档分块器(待实现)."""
from src.core.splitters.base import BaseTextSplitter
class TextSplitter(BaseTextSplitter):
"""纯文本分块器 — 按段落和字符边界拆分."""
def split(self, text: str, source_file: str = "") -> list[dict]:
if not text.strip():
return []
chunks = self._split_by_paragraphs(text)
for i, chunk in enumerate(chunks):
chunk["source_file"] = source_file or chunk.get("source_file", "")
chunk["chunk_index"] = i
return chunks
+2 -1
View File
@@ -4,7 +4,8 @@ from pathlib import Path
import pytest import pytest
from src.core.ingest import MarkdownSplitter, DocumentIngestor from src.core.splitters import MarkdownSplitter
from src.core.ingest import DocumentIngestor
class TestMarkdownSplitter: class TestMarkdownSplitter: