feat: 添加 TextSplitter 和 registry 自动选择机制

This commit is contained in:
2026-07-10 14:17:17 +08:00
parent 4dc51dccd2
commit 13d63ba6ff
2 changed files with 63 additions and 15 deletions
+59 -12
View File
@@ -1,22 +1,69 @@
"""分块器注册表 — 按文件后缀名路由到对应 Splitter."""
"""Splitter 注册表 — 按文件扩展名自动选择分块器."""
from pathlib import Path
from src.core.splitters.base import Splitter
SUPPORTED_SUFFIXES: dict[str, str] = {
# 扩展名 → Splitter 类名映射
_DEFAULT_MAP: dict[str, str] = {
".md": "markdown",
".markdown": "markdown",
".txt": "text",
".html": "html",
".pdf": "pdf",
".html": "html",
".htm": "html",
}
_registry: dict[str, type] = {}
# 所有支持的扩展名集合(供外部遍历文件使用)
SUPPORTED_SUFFIXES = frozenset(_DEFAULT_MAP.keys())
# 用户可注册自定义 Splitter
_custom_registry: dict[str, type[Splitter]] = {}
def register_splitter(name: str, splitter_cls: type) -> None:
"""注册一个分块器."""
_registry[name] = splitter_cls
def register_splitter(ext: str, splitter_cls: type[Splitter]) -> None:
"""注册自定义 Splitter 类."""
ext = ext.lower() if ext.startswith(".") else f".{ext}"
_custom_registry[ext] = splitter_cls
def get_splitter(name: str):
"""获取已注册的分块器类."""
if name not in _registry:
raise ValueError(f"未注册的分块器: {name}")
return _registry[name]
def get_splitter(
file_path: str,
max_size: int = 1000,
overlap: int = 100,
) -> Splitter:
"""根据文件扩展名自动选择 Splitter,未匹配回退到 TextSplitter.
Args:
file_path: 文件路径(用于提取扩展名)
max_size: 分块最大字符数
overlap: 相邻块重叠字符数
Returns:
对应格式的 Splitter 实例
"""
ext = Path(file_path).suffix.lower()
# 优先查用户自定义注册
if ext in _custom_registry:
return _custom_registry[ext](max_size=max_size, overlap=overlap)
kind = _DEFAULT_MAP.get(ext, "text")
if kind == "markdown":
from src.core.splitters.markdown import MarkdownSplitter
return MarkdownSplitter(max_size=max_size, overlap=overlap)
if kind == "text":
from src.core.splitters.text import TextSplitter
return TextSplitter(max_size=max_size, overlap=overlap)
if kind == "pdf":
from src.core.splitters.pdf import PDFSplitter
return PDFSplitter(max_size=max_size, overlap=overlap)
if kind == "html":
from src.core.splitters.html import HTMLSplitter
return HTMLSplitter(max_size=max_size, overlap=overlap)
# 回退
from src.core.splitters.text import TextSplitter
return TextSplitter(max_size=max_size, overlap=overlap)
+4 -3
View File
@@ -1,9 +1,9 @@
"""纯文本文档分块器(待实现)."""
"""纯文本分块器 — 按段落双换行切分."""
from src.core.splitters.base import BaseTextSplitter
class TextSplitter(BaseTextSplitter):
"""纯文本分块器 — 按段落和字符边界拆分."""
"""纯文本分块器:按 \n\n 切段落,超长按标点硬切."""
def split(self, text: str, source_file: str = "") -> list[dict]:
if not text.strip():
@@ -11,6 +11,7 @@ class TextSplitter(BaseTextSplitter):
chunks = self._split_by_paragraphs(text)
for i, chunk in enumerate(chunks):
chunk["source_file"] = source_file or chunk.get("source_file", "")
chunk["source_file"] = source_file
chunk["chunk_index"] = i
return chunks