diff --git a/src/core/splitters/registry.py b/src/core/splitters/registry.py index 442ec4a..70a1006 100644 --- a/src/core/splitters/registry.py +++ b/src/core/splitters/registry.py @@ -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) diff --git a/src/core/splitters/text.py b/src/core/splitters/text.py index 4d351a0..41f4b31 100644 --- a/src/core/splitters/text.py +++ b/src/core/splitters/text.py @@ -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