feat: 添加 TextSplitter 和 registry 自动选择机制
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user