Skip to content

工具

lazyllm.tools.agent.code_interpreter

code_interpreter(code, language='python')

Interpret the code and return the code interpreter result (include stdout, stderr, returncode, etc.).

Parameters:

  • code (str) –

    The code to interpret.

  • language (str, default: 'python' ) –

    The language of the code. Default is 'python'.

Source code in lazyllm/tools/agent/code_interpreter.py
@fc_register('tool', execute_in_sandbox=False)
def code_interpreter(code: str, language: str = 'python') -> str:
    """
    Interpret the code and return the code interpreter result (include stdout, stderr, returncode, etc.).

    Args:
        code (str): The code to interpret.
        language (str): The language of the code. Default is 'python'.
    """
    call_once(_sandbox_once, _create_sandbox)
    return _sandbox(code=code, language=language)

lazyllm.tools.sandbox.LazyLLMSandboxBase

Bases: ModuleBase

沙箱执行基类,定义统一的代码执行接口与语言检查逻辑。

Parameters:

  • output_dir_path (str | None, default: None ) –

    输出文件保存目录,默认当前工作目录,可能会覆盖当前工作目录下的文件。

  • return_trace (bool, default: False ) –

    是否返回中间执行信息(由 ModuleBase 控制)。

Notes

子类需实现 _is_available_execute 方法。

Source code in lazyllm/tools/sandbox/sandbox_base.py
class LazyLLMSandboxBase(ModuleBase, metaclass=LazyLLMRegisterMetaClass):
    """沙箱执行基类,定义统一的代码执行接口与语言检查逻辑。

Args:
    output_dir_path (str | None): 输出文件保存目录,默认当前工作目录,可能会覆盖当前工作目录下的文件。
    return_trace (bool): 是否返回中间执行信息(由 ModuleBase 控制)。

Notes:
    子类需实现 `_is_available` 与 `_execute` 方法。
"""
    SUPPORTED_LANGUAGES: List[str] = []

    def __init__(self, output_dir_path: Optional[str] = None, return_trace: bool = False,
                 project_dir: Optional[str] = None, return_sandbox_result: bool = False):
        super().__init__(return_trace=return_trace)
        self._output_dir_path = output_dir_path or os.getcwd()
        self._project_dir = project_dir
        if self._project_dir and not os.path.isdir(self._project_dir):
            raise FileNotFoundError(f'Project directory not found: {self._project_dir}')
        self._available_checked = False
        self._return_sandbox_result = return_sandbox_result

    def _check_available(self) -> None:
        raise NotImplementedError

    def _create_context(self) -> dict:
        raise NotImplementedError

    def _execute(self, code: str, language: str, context: dict,
                 output_files: Optional[List[str]] = None) -> _SandboxResult:
        raise NotImplementedError

    def _process_input_files(self, input_files: List[str], context: dict) -> None:
        raise NotImplementedError

    def _process_output_files(self, result: _SandboxResult, output_files: List[str], context: dict) -> List[str]:
        raise NotImplementedError

    def _process_project_dir(self, context: dict) -> None:
        raise NotImplementedError

    def _cleanup_context(self, context: dict) -> None:
        pass

    def _validate_input_files(self, input_files: Optional[List[str]]) -> Optional[List[str]]:
        if not input_files:
            return input_files
        for f in input_files:
            if not os.path.isfile(f):
                raise FileNotFoundError(f'Input file not found: {f}')

    def _collect_project_py_files(self) -> Generator[Tuple[str, str], None, None]:
        if not self._project_dir:
            return
        abs_dir = os.path.abspath(self._project_dir)
        for root, _, files in os.walk(abs_dir):
            for name in files:
                if name.endswith('.py'):
                    abs_path = os.path.join(root, name)
                    yield abs_path, os.path.relpath(abs_path, abs_dir)

    def _ensure_output_dir(self) -> None:
        os.makedirs(self._output_dir_path, exist_ok=True)

    def forward(self, code: str, language: str = 'python', input_files: Optional[List[str]] = None,
                output_files: Optional[List[str]] = None) -> dict:
        """统一执行入口,负责语言校验并调用具体实现。

Args:
    code (str): 待执行的代码。
    language (str): 代码语言,默认 'python'。
    input_files (list[str] | None): 输入文件路径列表,可选。
    output_files (list[str] | None): 需要回传的输出文件列表,可选。

**Returns:**

    由具体沙箱实现返回的结果(通常为 dict 或错误信息字符串)。
"""
        if not self._available_checked:
            self._check_available()
            self._available_checked = True

        if language not in self.SUPPORTED_LANGUAGES:
            raise ValueError(f'Language {language} not supported by {self.__class__.__name__}')
        self._validate_input_files(input_files)

        context = self._create_context()
        try:
            if self._project_dir:
                self._process_project_dir(context)
            if input_files:
                self._process_input_files(input_files, context)

            result = self._execute(code, language, context, output_files)

            if output_files and result.success:
                result.output_files = self._process_output_files(result, output_files, context)

            if self._return_sandbox_result:
                return result.to_dict()
            else:
                if result.success:
                    match = re.search(rf'^{SANDBOX_TOOL_RESULT_PREFIX}(.*)', result.stdout, re.MULTILINE)
                    return match.group(1) if match else result.stdout
                return result.stderr
        finally:
            self._cleanup_context(context)

forward(code, language='python', input_files=None, output_files=None)

统一执行入口,负责语言校验并调用具体实现。

Parameters:

  • code (str) –

    待执行的代码。

  • language (str, default: 'python' ) –

    代码语言,默认 'python'。

  • input_files (list[str] | None, default: None ) –

    输入文件路径列表,可选。

  • output_files (list[str] | None, default: None ) –

    需要回传的输出文件列表,可选。

Returns:

由具体沙箱实现返回的结果(通常为 dict 或错误信息字符串)。
Source code in lazyllm/tools/sandbox/sandbox_base.py
    def forward(self, code: str, language: str = 'python', input_files: Optional[List[str]] = None,
                output_files: Optional[List[str]] = None) -> dict:
        """统一执行入口,负责语言校验并调用具体实现。

Args:
    code (str): 待执行的代码。
    language (str): 代码语言,默认 'python'。
    input_files (list[str] | None): 输入文件路径列表,可选。
    output_files (list[str] | None): 需要回传的输出文件列表,可选。

**Returns:**

    由具体沙箱实现返回的结果(通常为 dict 或错误信息字符串)。
"""
        if not self._available_checked:
            self._check_available()
            self._available_checked = True

        if language not in self.SUPPORTED_LANGUAGES:
            raise ValueError(f'Language {language} not supported by {self.__class__.__name__}')
        self._validate_input_files(input_files)

        context = self._create_context()
        try:
            if self._project_dir:
                self._process_project_dir(context)
            if input_files:
                self._process_input_files(input_files, context)

            result = self._execute(code, language, context, output_files)

            if output_files and result.success:
                result.output_files = self._process_output_files(result, output_files, context)

            if self._return_sandbox_result:
                return result.to_dict()
            else:
                if result.success:
                    match = re.search(rf'^{SANDBOX_TOOL_RESULT_PREFIX}(.*)', result.stdout, re.MULTILINE)
                    return match.group(1) if match else result.stdout
                return result.stderr
        finally:
            self._cleanup_context(context)

lazyllm.tools.sandbox.DummySandbox

Bases: LazyLLMSandboxBase

本地沙箱实现(python-only),用于在受限环境中执行代码。

特点: - 通过 AST + SecurityVisitor 做基础安全检查。 - 在临时目录中运行代码,执行完毕后清理。 - 返回 stdout/stderr/returncode 的字典结果。

Parameters:

  • timeout (int, default: 30 ) –

    超时时间(秒),默认 30。

  • project_dir (str | None, default: None ) –

    若指定,将项目内 .py 文件复制到沙箱执行目录,便于引用。

  • return_trace (bool, default: False ) –

    是否返回中间执行信息。

Examples:

>>> from lazyllm.tools.sandbox import DummySandbox
>>> sandbox = DummySandbox(timeout=10)
>>> result = sandbox(code="print(1 + 1)")
>>> print(result['stdout'].strip())
2
Source code in lazyllm/tools/sandbox/dummy_sandbox.py
class DummySandbox(LazyLLMSandboxBase):
    """本地沙箱实现(python-only),用于在受限环境中执行代码。

特点:
- 通过 AST + SecurityVisitor 做基础安全检查。
- 在临时目录中运行代码,执行完毕后清理。
- 返回 stdout/stderr/returncode 的字典结果。

Args:
    timeout (int): 超时时间(秒),默认 30。
    project_dir (str | None): 若指定,将项目内 .py 文件复制到沙箱执行目录,便于引用。
    return_trace (bool): 是否返回中间执行信息。


Examples:
    >>> from lazyllm.tools.sandbox import DummySandbox
    >>> sandbox = DummySandbox(timeout=10)
    >>> result = sandbox(code="print(1 + 1)")
    >>> print(result['stdout'].strip())
    2
    """
    SUPPORTED_LANGUAGES: List[str] = ['python']

    def __init__(self, timeout: int = 30, return_trace: bool = False, project_dir: Optional[str] = None,
                 return_sandbox_result: bool = False):
        super().__init__(return_trace=return_trace, project_dir=project_dir,
                         return_sandbox_result=return_sandbox_result)
        self._timeout = timeout

    def _check_available(self) -> bool:
        return True

    def _check_code_safety(self, code: str) -> Tuple[bool, Optional[str]]:
        try:
            tree = ast.parse(code)
        except SyntaxError as e:
            return False, f'Syntax error: {e}'
        try:
            SecurityVisitor().visit(tree)
        except ValueError as e:
            return False, str(e)
        return True, None

    def _run_in_subprocess(self, script_path: str, cwd: str,
                           env: Optional[Dict[str, str]] = None) -> Dict[str, Any]:
        proc = subprocess.Popen(
            [sys.executable, '-u', script_path],
            stdout=subprocess.PIPE, stderr=subprocess.PIPE,
            cwd=cwd, env=env or os.environ.copy(), text=True, bufsize=1,
        )
        try:
            stdout, stderr = proc.communicate(timeout=self._timeout)
        except subprocess.TimeoutExpired:
            proc.kill()
            proc.wait()
            raise
        return {'returncode': proc.returncode, 'stdout': stdout, 'stderr': stderr}

    def execute_script(self, source_dir: str, rel_path: str, args: Optional[List[str]] = None,
                       cwd: str = '.', allow_unsafe: bool = False) -> Dict[str, Any]:
        """在临时执行目录中运行一个已物化的 Skill 脚本。

该方法会将 `source_dir` 的完整目录树复制到临时目录,校验脚本路径和工作目录均未逃逸出临时目录,
再根据扩展名选择解释器执行。`.py` 文件使用当前 Python 解释器,`.sh` 和 `.bash` 文件使用 Bash,
其他扩展名使用 `sh`。执行结束后会清理临时目录。

Args:
    source_dir (str): 已物化的 Skill 包根目录。
    rel_path (str): 相对于 `source_dir` 的脚本路径。
    args (list[str] | None): 传递给脚本的参数。
    cwd (str): 相对于 `source_dir` 的工作目录,默认为 `.`。
    allow_unsafe (bool): 预留的审批参数;DummySandbox 当前不提供审批边界,因此会忽略该参数。

**Returns:**

    dict:包含 `status`、`stdout`、`stderr`、`exit_code` 和 `cwd`。脚本不存在时返回
    `status='missing'`;非零退出码返回 `status='failed'`。

Notes:
    DummySandbox 只提供临时目录和子进程执行边界,并非强安全隔离。它不会限制脚本读取宿主机文件、
    访问网络或继承当前进程环境。不要用它执行未经信任的代码。


Examples:
    >>> import tempfile
    >>> from pathlib import Path
    >>> from lazyllm.tools.sandbox import DummySandbox
    >>> with tempfile.TemporaryDirectory() as root:
    ...     script = Path(root) / "scripts" / "check.py"
    ...     script.parent.mkdir()
    ...     _ = script.write_text("print('ok')\\n", encoding="utf-8")
    ...     result = DummySandbox(timeout=10).execute_script(
    ...         source_dir=root,
    ...         rel_path="scripts/check.py",
    ...         args=[],
    ...     )
    ...     print(result['stdout'].strip())
    ok
    """
        del allow_unsafe  # DummySandbox currently has no approval boundary.
        context = self._create_context()
        try:
            sandbox_root = context['temp_dir']
            shutil.copytree(source_dir, sandbox_root, dirs_exist_ok=True)
            script_path = self._resolve_child(sandbox_root, rel_path, 'rel_path')
            run_cwd = self._resolve_child(sandbox_root, cwd or '.', 'cwd')
            if not os.path.isfile(script_path):
                return {
                    'status': 'missing',
                    'path': script_path,
                    'rel_path': rel_path,
                    'cwd': run_cwd,
                }
            if not os.path.isdir(run_cwd):
                raise FileNotFoundError(f'cwd not found: {run_cwd}')
            ext = os.path.splitext(script_path)[1].lower()
            runner = sys.executable if ext == '.py' else 'bash' if ext in ('.sh', '.bash') else 'sh'
            completed = subprocess.run(
                [runner, script_path, *(args or [])],
                cwd=run_cwd,
                env=os.environ.copy(),
                text=True,
                capture_output=True,
                timeout=self._timeout,
            )
            return {
                'status': 'ok' if completed.returncode == 0 else 'failed',
                'stdout': completed.stdout,
                'stderr': completed.stderr,
                'exit_code': completed.returncode,
                'cwd': run_cwd,
            }
        finally:
            self._cleanup_context(context)

    @staticmethod
    def _resolve_child(root: str, rel_path: str, label: str) -> str:
        root_real = os.path.realpath(os.path.abspath(root))
        target = os.path.realpath(os.path.abspath(os.path.join(root_real, rel_path)))
        if os.path.commonpath([root_real, target]) != root_real:
            raise ValueError(f'{label} must stay inside the sandbox directory.')
        return target

    def _create_context(self) -> dict:
        return {'temp_dir': tempfile.mkdtemp(prefix='lazyllm_sandbox_')}

    def _cleanup_context(self, context: dict) -> None:
        temp_dir = context.get('temp_dir')
        if temp_dir:
            shutil.rmtree(temp_dir, ignore_errors=True)

    def _process_input_files(self, input_files: List[str], context: dict) -> None:
        for f in input_files:
            try:
                shutil.copy(f, context['temp_dir'])
            except Exception as e:
                LOG.warning(f'DummySandbox: failed to copy input file {f!r}: {e}')

    def _process_project_dir(self, context: dict) -> None:
        temp_dir = context['temp_dir']
        for abs_path, rel_path in self._collect_project_py_files():
            dst = os.path.join(temp_dir, rel_path)
            os.makedirs(os.path.dirname(dst), exist_ok=True)
            shutil.copy(abs_path, dst)

    def _process_output_files(self, result: _SandboxResult, output_files: List[str], context: dict) -> List[str]:
        self._ensure_output_dir()
        collected = []
        for name in output_files:
            src = os.path.join(context['temp_dir'], name)
            dst = os.path.join(self._output_dir_path, name)
            try:
                if os.path.exists(src):
                    shutil.move(src, dst)
                    collected.append(dst)
            except Exception as e:
                LOG.warning(f'DummySandbox: failed to move output file {src!r}: {e}')
        return collected

    def _execute(self, code: str, language: str, context: dict,
                 output_files: Optional[List[str]] = None) -> _SandboxResult:
        is_safe, msg = self._check_code_safety(code)
        if not is_safe:
            return _SandboxResult(success=False, error_message=msg)

        temp_dir = context['temp_dir']
        try:
            script_path = os.path.join(temp_dir, '_script.py')
            with open(script_path, 'w', encoding='utf-8') as f:
                f.write(code)
            proc_result = self._run_in_subprocess(script_path, cwd=temp_dir)
            return _SandboxResult(
                success=(proc_result['returncode'] == 0),
                stdout=proc_result['stdout'],
                stderr=proc_result['stderr'],
                returncode=proc_result['returncode'],
            )
        except subprocess.TimeoutExpired:
            return _SandboxResult(success=False, error_message=f'Execution timed out after {self._timeout} seconds')
        except Exception as e:
            return _SandboxResult(success=False, error_message=str(e))

execute_script(source_dir, rel_path, args=None, cwd='.', allow_unsafe=False)

在临时执行目录中运行一个已物化的 Skill 脚本。

该方法会将 source_dir 的完整目录树复制到临时目录,校验脚本路径和工作目录均未逃逸出临时目录, 再根据扩展名选择解释器执行。.py 文件使用当前 Python 解释器,.sh.bash 文件使用 Bash, 其他扩展名使用 sh。执行结束后会清理临时目录。

Parameters:

  • source_dir (str) –

    已物化的 Skill 包根目录。

  • rel_path (str) –

    相对于 source_dir 的脚本路径。

  • args (list[str] | None, default: None ) –

    传递给脚本的参数。

  • cwd (str, default: '.' ) –

    相对于 source_dir 的工作目录,默认为 .

  • allow_unsafe (bool, default: False ) –

    预留的审批参数;DummySandbox 当前不提供审批边界,因此会忽略该参数。

Returns:

dict:包含 `status`、`stdout`、`stderr`、`exit_code` 和 `cwd`。脚本不存在时返回
`status='missing'`;非零退出码返回 `status='failed'`。
Notes

DummySandbox 只提供临时目录和子进程执行边界,并非强安全隔离。它不会限制脚本读取宿主机文件、 访问网络或继承当前进程环境。不要用它执行未经信任的代码。

Examples:

>>> import tempfile
>>> from pathlib import Path
>>> from lazyllm.tools.sandbox import DummySandbox
>>> with tempfile.TemporaryDirectory() as root:
...     script = Path(root) / "scripts" / "check.py"
...     script.parent.mkdir()
...     _ = script.write_text("print('ok')\n", encoding="utf-8")
...     result = DummySandbox(timeout=10).execute_script(
...         source_dir=root,
...         rel_path="scripts/check.py",
...         args=[],
...     )
...     print(result['stdout'].strip())
ok
Source code in lazyllm/tools/sandbox/dummy_sandbox.py
    def execute_script(self, source_dir: str, rel_path: str, args: Optional[List[str]] = None,
                       cwd: str = '.', allow_unsafe: bool = False) -> Dict[str, Any]:
        """在临时执行目录中运行一个已物化的 Skill 脚本。

该方法会将 `source_dir` 的完整目录树复制到临时目录,校验脚本路径和工作目录均未逃逸出临时目录,
再根据扩展名选择解释器执行。`.py` 文件使用当前 Python 解释器,`.sh` 和 `.bash` 文件使用 Bash,
其他扩展名使用 `sh`。执行结束后会清理临时目录。

Args:
    source_dir (str): 已物化的 Skill 包根目录。
    rel_path (str): 相对于 `source_dir` 的脚本路径。
    args (list[str] | None): 传递给脚本的参数。
    cwd (str): 相对于 `source_dir` 的工作目录,默认为 `.`。
    allow_unsafe (bool): 预留的审批参数;DummySandbox 当前不提供审批边界,因此会忽略该参数。

**Returns:**

    dict:包含 `status`、`stdout`、`stderr`、`exit_code` 和 `cwd`。脚本不存在时返回
    `status='missing'`;非零退出码返回 `status='failed'`。

Notes:
    DummySandbox 只提供临时目录和子进程执行边界,并非强安全隔离。它不会限制脚本读取宿主机文件、
    访问网络或继承当前进程环境。不要用它执行未经信任的代码。


Examples:
    >>> import tempfile
    >>> from pathlib import Path
    >>> from lazyllm.tools.sandbox import DummySandbox
    >>> with tempfile.TemporaryDirectory() as root:
    ...     script = Path(root) / "scripts" / "check.py"
    ...     script.parent.mkdir()
    ...     _ = script.write_text("print('ok')\\n", encoding="utf-8")
    ...     result = DummySandbox(timeout=10).execute_script(
    ...         source_dir=root,
    ...         rel_path="scripts/check.py",
    ...         args=[],
    ...     )
    ...     print(result['stdout'].strip())
    ok
    """
        del allow_unsafe  # DummySandbox currently has no approval boundary.
        context = self._create_context()
        try:
            sandbox_root = context['temp_dir']
            shutil.copytree(source_dir, sandbox_root, dirs_exist_ok=True)
            script_path = self._resolve_child(sandbox_root, rel_path, 'rel_path')
            run_cwd = self._resolve_child(sandbox_root, cwd or '.', 'cwd')
            if not os.path.isfile(script_path):
                return {
                    'status': 'missing',
                    'path': script_path,
                    'rel_path': rel_path,
                    'cwd': run_cwd,
                }
            if not os.path.isdir(run_cwd):
                raise FileNotFoundError(f'cwd not found: {run_cwd}')
            ext = os.path.splitext(script_path)[1].lower()
            runner = sys.executable if ext == '.py' else 'bash' if ext in ('.sh', '.bash') else 'sh'
            completed = subprocess.run(
                [runner, script_path, *(args or [])],
                cwd=run_cwd,
                env=os.environ.copy(),
                text=True,
                capture_output=True,
                timeout=self._timeout,
            )
            return {
                'status': 'ok' if completed.returncode == 0 else 'failed',
                'stdout': completed.stdout,
                'stderr': completed.stderr,
                'exit_code': completed.returncode,
                'cwd': run_cwd,
            }
        finally:
            self._cleanup_context(context)

lazyllm.tools.sandbox.SandboxFusion

Bases: LazyLLMSandboxBase

远程沙箱实现,通过 HTTP API 执行代码并获取结果。

支持语言:python / bash。可配置编译超时、运行超时、内存限制,并支持上传工程文件与拉取输出文件。

Parameters:

  • base_url (str, default: config['sandbox_fusion_base_url'] ) –

    远程沙箱服务地址,默认来自 config['sandbox_fusion_base_url']。

  • compile_timeout (int, default: 10 ) –

    编译超时(秒),默认 10。

  • run_timeout (int, default: 10 ) –

    运行超时(秒),默认 10。

  • memory_limit_mb (int, default: -1 ) –

    内存限制(MB),-1 表示不限制。

  • project_dir (str | None, default: None ) –

    若指定,将工程目录下的 .py 文件上传到沙箱。

Notes

需要配置 LAZYLLM_SANDBOX_FUSION_BASE_URL 或显式传入 base_url。

Examples:

>>> from lazyllm import config
>>> from lazyllm.tools.sandbox import SandboxFusion
>>> config['sandbox_fusion_base_url'] = "http://localhost:8000"
>>> sandbox = SandboxFusion(run_timeout=5)
>>> result = sandbox(code="print('ok')")
>>> print(result['stdout'].strip())
ok
Source code in lazyllm/tools/sandbox/sandbox_fusion.py
class SandboxFusion(LazyLLMSandboxBase):
    """远程沙箱实现,通过 HTTP API 执行代码并获取结果。

支持语言:python / bash。可配置编译超时、运行超时、内存限制,并支持上传工程文件与拉取输出文件。

Args:
    base_url (str): 远程沙箱服务地址,默认来自 config['sandbox_fusion_base_url']。
    compile_timeout (int): 编译超时(秒),默认 10。
    run_timeout (int): 运行超时(秒),默认 10。
    memory_limit_mb (int): 内存限制(MB),-1 表示不限制。
    project_dir (str | None): 若指定,将工程目录下的 .py 文件上传到沙箱。

Notes:
    需要配置 LAZYLLM_SANDBOX_FUSION_BASE_URL 或显式传入 base_url。


Examples:
    >>> from lazyllm import config
    >>> from lazyllm.tools.sandbox import SandboxFusion
    >>> config['sandbox_fusion_base_url'] = "http://localhost:8000"
    >>> sandbox = SandboxFusion(run_timeout=5)
    >>> result = sandbox(code="print('ok')")
    >>> print(result['stdout'].strip())
    ok
    """
    __lazyllm_registry_key__ = 'sandbox_fusion'
    SUPPORTED_LANGUAGES: List[str] = ['python', 'bash']

    def __init__(self, base_url: str = config['sandbox_fusion_base_url'], compile_timeout: int = 10,
                 run_timeout: int = 10, memory_limit_mb: int = -1, project_dir: str = None,
                 return_sandbox_result: bool = False, return_trace: bool = False):
        super().__init__(return_trace=return_trace, project_dir=project_dir, return_sandbox_result=return_sandbox_result)
        self._base_url = base_url
        self._compile_timeout = compile_timeout
        self._run_timeout = run_timeout
        self._memory_limit_mb = memory_limit_mb
        self._project_files_cache = None

    @property
    def url(self) -> str:
        return f'{self._base_url}/run_code'

    def _check_available(self) -> None:
        try:
            resp = requests.get(f'{self._base_url}/v1/ping', timeout=2)
            if resp.status_code != 200:
                raise ValueError(f'SandboxFusion ping failed: status={resp.status_code}, text={resp.text}')
        except Exception as e:
            raise ValueError(f'SandboxFusion _check_available error: {e}')

    def _call_api(self, call_params: dict[str, Any]):
        headers = {'Content-Type': 'application/json', 'Accept': 'application/json'}
        try:
            resp = requests.post(self.url, headers=headers, json=call_params)
            resp.raise_for_status()
            return resp.json()
        except req_exc.RequestException as e:
            LOG.error(f'API Request Error: {e}')
        except JSONDecodeError as e:
            LOG.error(f'API Response JSON Decode Error: {e}')
        except Exception as e:
            LOG.exception(f'Unexpected Error: {e}')
        return 'Sandbox API Call Failed'

    def _create_context(self) -> dict:
        return {'files': {}}

    def _process_input_files(self, input_files: List[str], context: dict) -> None:
        for f in input_files:
            context['files'][f] = self._encode_file_base64(f)

    def _process_project_dir(self, context: dict) -> None:
        if self._project_files_cache is None:
            self._project_files_cache = {
                rel: self._encode_file_base64(abs_p)
                for abs_p, rel in self._collect_project_py_files()
            }
        context['files'].update(self._project_files_cache)

    def _process_output_files(self, result: _SandboxResult, output_files: List[str], context: dict) -> List[str]:
        self._ensure_output_dir()
        response_files = context.get('response_files') or {}
        collected = []
        for name in output_files:
            b64 = response_files.get(name)
            if b64 is None:
                LOG.warning(f'SandboxFusion: requested output file {name!r} not found in response')
                continue
            path = os.path.join(self._output_dir_path, name)
            with open(path, 'wb') as f:
                f.write(base64.b64decode(b64))
            collected.append(path)
        return collected

    def _execute(self, code: str, language: str, context: dict,
                 output_files: Optional[List[str]] = None) -> _SandboxResult:
        call_params = {
            'code': code,
            'compile_timeout': self._compile_timeout,
            'run_timeout': self._run_timeout,
            'memory_limit_mb': self._memory_limit_mb,
            'language': language,
            'files': context['files'],
        }
        if output_files:
            call_params['fetch_files'] = output_files

        response = self._call_api(call_params)
        if isinstance(response, str):
            return _SandboxResult(success=False, error_message=response)

        context['response_files'] = response.get('files') or {}
        run_result = response.get('run_result') or {}
        returncode = run_result.get('return_code', -1)
        return _SandboxResult(
            success=(response.get('status') == 'Success' and returncode == 0),
            stdout=run_result.get('stdout', ''),
            stderr=run_result.get('stderr', ''),
            returncode=returncode,
        )

    @staticmethod
    def _encode_file_base64(path: str) -> str:
        encoded = file_to_base64(path)
        if encoded is None:
            raise ValueError(f'Failed to encode file to base64: {path}')
        return encoded[0]

lazyllm.tools.IntentClassifier

Bases: ModuleBase

意图分类模块,用于根据输入文本在给定的意图列表中进行分类。 支持中英文自动选择提示模板,并可通过示例、提示、约束和注意事项增强分类效果。

Parameters:

  • llm

    用于意图分类的大语言模型实例。

  • intent_list (list, default: None ) –

    可选,意图类别列表,例如 ["聊天", "天气", "问答"]。

  • prompt (str, default: '' ) –

    可选,自定义提示语,插入到系统提示模板中。

  • constrain (str, default: '' ) –

    可选,分类约束条件说明。

  • attention (str, default: '' ) –

    可选,提示注意事项。

  • examples (list[list[str, str]], default: None ) –

    可选,分类示例列表,每个元素为 [输入文本, 标签]。

  • return_trace (bool, default: False ) –

    是否返回执行过程的 trace,默认为 False。

Examples:

>>> import lazyllm
>>> from lazyllm.tools import IntentClassifier
>>> classifier_llm = lazyllm.OnlineChatModule(source="openai")
>>> chatflow_intent_list = ["Chat", "Financial Knowledge Q&A", "Employee Information Query", "Weather Query"]
>>> classifier = IntentClassifier(classifier_llm, intent_list=chatflow_intent_list)
>>> classifier.start()
>>> print(classifier('What is the weather today'))
Weather Query
>>>
>>> with IntentClassifier(classifier_llm) as ic:
>>>     ic.case['Weather Query', lambda x: '38.5°C']
>>>     ic.case['Chat', lambda x: 'permission denied']
>>>     ic.case['Financial Knowledge Q&A', lambda x: 'Calling Financial RAG']
>>>     ic.case['Employee Information Query', lambda x: 'Beijing']
...
>>> ic.start()
>>> print(ic('What is the weather today'))
38.5°C
Source code in lazyllm/tools/classifier/intent_classifier.py
class IntentClassifier(ModuleBase):
    """意图分类模块,用于根据输入文本在给定的意图列表中进行分类。
支持中英文自动选择提示模板,并可通过示例、提示、约束和注意事项增强分类效果。

Args:
    llm: 用于意图分类的大语言模型实例。
    intent_list (list): 可选,意图类别列表,例如 ["聊天", "天气", "问答"]。
    prompt (str): 可选,自定义提示语,插入到系统提示模板中。
    constrain (str): 可选,分类约束条件说明。
    attention (str): 可选,提示注意事项。
    examples (list[list[str, str]]): 可选,分类示例列表,每个元素为 [输入文本, 标签]。
    return_trace (bool): 是否返回执行过程的 trace,默认为 False。


Examples:
        >>> import lazyllm
        >>> from lazyllm.tools import IntentClassifier
        >>> classifier_llm = lazyllm.OnlineChatModule(source="openai")
        >>> chatflow_intent_list = ["Chat", "Financial Knowledge Q&A", "Employee Information Query", "Weather Query"]
        >>> classifier = IntentClassifier(classifier_llm, intent_list=chatflow_intent_list)
        >>> classifier.start()
        >>> print(classifier('What is the weather today'))
        Weather Query
        >>>
        >>> with IntentClassifier(classifier_llm) as ic:
        >>>     ic.case['Weather Query', lambda x: '38.5°C']
        >>>     ic.case['Chat', lambda x: 'permission denied']
        >>>     ic.case['Financial Knowledge Q&A', lambda x: 'Calling Financial RAG']
        >>>     ic.case['Employee Information Query', lambda x: 'Beijing']
        ...
        >>> ic.start()
        >>> print(ic('What is the weather today'))
        38.5°C
    """
    def __init__(self, llm, intent_list: list = None,
                 *, prompt: str = '', constrain: str = '', attention: str = '',
                 examples: Optional[list[list[str, str]]] = None, return_trace: bool = False) -> None:
        super().__init__(return_trace=return_trace)
        self._intent_list = intent_list or []
        self._llm = llm
        self._prompt, self._constrain, self._attention, self._examples = prompt, constrain, attention, examples or []
        if self._intent_list:
            self._init()

    def _init(self):
        def choose_prompt():
            # Use chinese prompt if intent elements have chinese character, otherwise use english version
            for ele in self._intent_list:
                for ch in ele:
                    # chinese unicode range
                    if '\u4e00' <= ch <= '\u9fff':
                        return ch_prompt_classifier_template
            return en_prompt_classifier_template

        example_template = '\nUser: {{{{"human_input": "{inp}", "intent_list": {intent}}}}}\nAssistant: {label}\n'
        examples = ''.join([example_template.format(
            inp=input, intent=self._intent_list, label=label) for input, label in self._examples])
        prompt = choose_prompt().replace(
            '{user_prompt}', f' {self._prompt}').replace('{attention}', self._attention).replace(
            '{user_constrains}', f' {self._constrain}').replace('{user_examples}', f' {examples}')
        self._llm = self._llm.share(prompt=AlpacaPrompter(dict(system=prompt, user='${input}')
                                                          ).pre_hook(self.intent_promt_hook)).used_by(self._module_id)
        self._impl = pipeline(self._llm, self.post_process_result)

    def intent_promt_hook(
        self,
        input: Union[str, List, Dict[str, str], None] = None,
        history: List[Union[List[str], Dict[str, Any]]] = [],  # noqa B006
        tools: Union[List[Dict[str, Any]], None] = None,
        label: Union[str, None] = None,
    ):
        """意图分类的预处理 Hook。
将输入文本与意图列表打包为 JSON,并生成历史对话信息字符串。

Args:
    input (str | List | Dict | None): 输入文本,仅支持字符串类型。
    history (List): 历史对话记录,默认为空列表。
    tools (List[Dict] | None): 工具信息,可选。
    label (str | None): 标签,可选。

**Returns:**

- tuple: 输入数据字典, 历史记录列表, 工具信息, 标签
"""
        input_json = {}
        if isinstance(input, str):
            input_json = {'human_input': input, 'intent_list': self._intent_list}
        else:
            raise ValueError(f'Unexpected type for input: {type(input)}')

        history_info = chat_history_to_str(history)
        history = []
        input_text = json.dumps(input_json, ensure_ascii=False)
        return dict(history_info=history_info, input=input_text), history, tools, label

    def post_process_result(self, input):
        """意图分类结果的后处理。
如果结果在意图列表中则直接返回,否则返回意图列表的第一个元素。

Args:
    input (str): 分类模型输出结果。

**Returns:**

- str: 最终的分类标签。
"""
        input = input.strip()
        return input if input in self._intent_list else self._intent_list[0]

    def forward(self, input: str, llm_chat_history: List[Dict[str, Any]] = None):
        if llm_chat_history is not None and self._llm._module_id not in globals['chat_history']:
            globals['chat_history'][self._llm._module_id] = llm_chat_history
        return self._impl(input)

    def __enter__(self):
        assert not self._intent_list, 'Intent list is already set'
        self._sw = switch()
        self._sw.__enter__()
        return self

    @property
    def case(self):
        return switch.Case(self)

    @property
    def submodules(self):
        submodule = []
        if isinstance(self._impl, switch):
            self._impl.for_each(lambda x: isinstance(x, ModuleBase), lambda x: submodule.append(x))
        return super().submodules + submodule

    # used by switch.Case
    def _add_case(self, cond, func):
        assert isinstance(cond, str), 'intent must be string'
        self._intent_list.append(cond)
        self._sw.case[cond, func]

    def __exit__(self, exc_type, exc_val, exc_tb):
        self._sw.__exit__(exc_type, exc_val, exc_tb)
        self._init()
        self._sw._set_conversion(self._impl)
        self._impl = self._sw

intent_promt_hook(input=None, history=[], tools=None, label=None)

意图分类的预处理 Hook。 将输入文本与意图列表打包为 JSON,并生成历史对话信息字符串。

Parameters:

  • input (str | List | Dict | None, default: None ) –

    输入文本,仅支持字符串类型。

  • history (List, default: [] ) –

    历史对话记录,默认为空列表。

  • tools (List[Dict] | None, default: None ) –

    工具信息,可选。

  • label (str | None, default: None ) –

    标签,可选。

Returns:

  • tuple: 输入数据字典, 历史记录列表, 工具信息, 标签
Source code in lazyllm/tools/classifier/intent_classifier.py
    def intent_promt_hook(
        self,
        input: Union[str, List, Dict[str, str], None] = None,
        history: List[Union[List[str], Dict[str, Any]]] = [],  # noqa B006
        tools: Union[List[Dict[str, Any]], None] = None,
        label: Union[str, None] = None,
    ):
        """意图分类的预处理 Hook。
将输入文本与意图列表打包为 JSON,并生成历史对话信息字符串。

Args:
    input (str | List | Dict | None): 输入文本,仅支持字符串类型。
    history (List): 历史对话记录,默认为空列表。
    tools (List[Dict] | None): 工具信息,可选。
    label (str | None): 标签,可选。

**Returns:**

- tuple: 输入数据字典, 历史记录列表, 工具信息, 标签
"""
        input_json = {}
        if isinstance(input, str):
            input_json = {'human_input': input, 'intent_list': self._intent_list}
        else:
            raise ValueError(f'Unexpected type for input: {type(input)}')

        history_info = chat_history_to_str(history)
        history = []
        input_text = json.dumps(input_json, ensure_ascii=False)
        return dict(history_info=history_info, input=input_text), history, tools, label

post_process_result(input)

意图分类结果的后处理。 如果结果在意图列表中则直接返回,否则返回意图列表的第一个元素。

Parameters:

  • input (str) –

    分类模型输出结果。

Returns:

  • str: 最终的分类标签。
Source code in lazyllm/tools/classifier/intent_classifier.py
    def post_process_result(self, input):
        """意图分类结果的后处理。
如果结果在意图列表中则直接返回,否则返回意图列表的第一个元素。

Args:
    input (str): 分类模型输出结果。

**Returns:**

- str: 最终的分类标签。
"""
        input = input.strip()
        return input if input in self._intent_list else self._intent_list[0]

lazyllm.tools.Document

Bases: ModuleBase, BuiltinGroups

初始化一个文档管理模块,支持可选的向量化、存储和用户界面。

Document 模块提供了统一的文档数据集管理接口,支持本地文件、云端文件或临时文档文件。它可以选择运行文档管理服务或 Web UI,并支持多种向量化模型和自定义存储后端。

Parameters:

  • dataset_path (Optional[str], default: None ) –

    数据集目录路径。如果路径不存在,系统会尝试在 lazyllm.config["data_path"] 中查找。

  • embed (Optional[Union[Callable, Dict[str, Callable]]], default: None ) –

    文档向量化函数或函数字典。若为字典,键为 embedding 名称,值为对应的模型。

  • create_ui (bool, default: False ) –

    是否创建文档管理 UI。该能力要求当前存在可用的 DocServer,可与 manager=Truemanager=DocServer(...) 组合使用。

  • manager (Union[bool, str, DocServer, DocumentProcessor], default: False ) –

    文档管理模式。True 表示启动本地 DocServer 及其 parsing service;DocServer(...) 表示连接已有文档管理服务;DocumentProcessor(...) 表示仅连接解析服务,此时必须提供非 map 的 store_conf'ui' 仅作为 manager=True, create_ui=True 的兼容写法保留。

  • server (Union[bool, int], default: False ) –

    是否为知识库运行服务接口。True 表示启动默认服务;整型数值表示自定义端口;False 表示关闭。默认为 False

  • name (Optional[str], default: None ) –

    文档集合的名称标识符。默认为系统默认名称。

  • launcher (Optional[LazyLLMLaunchersBase], default: None ) –

    启动器实例,用于管理服务进程。默认使用远程异步启动器。

  • doc_files (Optional[List[str]], default: None ) –

    临时文档文件列表。当使用此参数时,dataset_path 必须为 None,且仅支持 MapStore。

  • doc_fields (Optional[Dict[str, GlobalMetadataDesc]], default: None ) –

    元数据字段配置,用于存储和检索文档属性。

  • store_conf (Optional[Dict], default: None ) –

    存储配置。默认使用内存中的 MapStore。

  • display_name (Optional[str], default: '' ) –

    文档模块的可读显示名称。默认为集合名称。

  • description (Optional[str], default: 'algorithm description' ) –

    文档集合的描述。默认为 "algorithm description"

  • schema_extractor (Optional[Union[LLMBase, SchemaExtractor]], default: None ) –

    可选 schema extractor,用于元数据 schema 分析与注册。

  • enable_path_monitoring (Optional[bool], default: None ) –

    是否监控本地数据目录的文件新增和删除。仅在未接入 DocServer / DocumentProcessor 的本地模式下默认开启。

Examples:

>>> import lazyllm
>>> from lazyllm.tools import Document
>>> m = lazyllm.OnlineEmbeddingModule(source="glm")
>>> documents = Document(dataset_path='your_doc_path', embed=m, manager=False)  # or documents = Document(dataset_path='your_doc_path', embed={"key": m}, manager=False)
>>> m1 = lazyllm.TrainableModule("bge-large-zh-v1.5").start()
>>> document1 = Document(dataset_path='your_doc_path', embed={"online": m, "local": m1}, manager=False)
>>> store_conf = {
>>>     "segment_store": {
>>>         "type": "map",
>>>         "kwargs": {
>>>             "uri": "/tmp/tmp_segments.db",
>>>         },
>>>     },
>>>     "vector_store": {
>>>         "type": "milvus",
>>>         "kwargs": {
>>>             "uri": "/tmp/tmp_milvus.db",
>>>             "index_kwargs": {
>>>                 "index_type": "FLAT",
>>>                 "metric_type": "COSINE",
>>>             },
>>>         },
>>>     },
>>> }
>>> doc_fields = {
>>>     'author': DocField(data_type=DataType.VARCHAR, max_size=128, default_value=' '),
>>>     'public_year': DocField(data_type=DataType.INT32),
>>> }
>>> document2 = Document(dataset_path='your_doc_path', embed={"online": m, "local": m1}, store_conf=store_conf, doc_fields=doc_fields)
Source code in lazyllm/tools/rag/document.py
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
789
790
791
792
793
794
795
796
797
798
799
800
801
802
803
804
805
806
807
808
809
810
811
812
813
814
815
816
817
818
819
820
821
822
823
824
825
826
827
828
829
830
831
832
833
834
835
836
837
838
839
840
841
842
843
844
845
846
847
848
849
850
851
852
853
854
855
856
857
858
859
860
861
862
863
864
865
866
867
868
869
870
871
872
873
874
875
876
877
878
879
880
881
882
883
884
885
886
887
888
889
890
891
892
893
894
895
896
897
898
899
900
901
902
903
904
905
906
907
908
909
910
911
912
913
914
915
916
917
918
919
class Document(ModuleBase, BuiltinGroups, metaclass=_MetaDocument):
    """初始化一个文档管理模块,支持可选的向量化、存储和用户界面。

``Document`` 模块提供了统一的文档数据集管理接口,支持本地文件、云端文件或临时文档文件。它可以选择运行文档管理服务或 Web UI,并支持多种向量化模型和自定义存储后端。

Args:
    dataset_path (Optional[str]): 数据集目录路径。如果路径不存在,系统会尝试在 ``lazyllm.config["data_path"]`` 中查找。
    embed (Optional[Union[Callable, Dict[str, Callable]]]): 文档向量化函数或函数字典。若为字典,键为 embedding 名称,值为对应的模型。
    create_ui (bool, optional): 是否创建文档管理 UI。该能力要求当前存在可用的 ``DocServer``,可与 ``manager=True`` 或 ``manager=DocServer(...)`` 组合使用。
    manager (Union[bool, str, DocServer, DocumentProcessor], optional): 文档管理模式。``True`` 表示启动本地 ``DocServer`` 及其 parsing service;``DocServer(...)`` 表示连接已有文档管理服务;``DocumentProcessor(...)`` 表示仅连接解析服务,此时必须提供非 map 的 ``store_conf``;``'ui'`` 仅作为 ``manager=True, create_ui=True`` 的兼容写法保留。
    server (Union[bool, int], optional): 是否为知识库运行服务接口。``True`` 表示启动默认服务;整型数值表示自定义端口;``False`` 表示关闭。默认为 ``False``。
    name (Optional[str]): 文档集合的名称标识符。默认为系统默认名称。
    launcher (Optional[Launcher]): 启动器实例,用于管理服务进程。默认使用远程异步启动器。
    doc_files (Optional[List[str]]): 临时文档文件列表。当使用此参数时,``dataset_path`` 必须为 ``None``,且仅支持 MapStore。
    doc_fields (Optional[Dict[str, DocField]]): 元数据字段配置,用于存储和检索文档属性。
    store_conf (Optional[Dict]): 存储配置。默认使用内存中的 MapStore。
    display_name (Optional[str]): 文档模块的可读显示名称。默认为集合名称。
    description (Optional[str]): 文档集合的描述。默认为 ``"algorithm description"``。
    schema_extractor (Optional[Union[LLMBase, SchemaExtractor]]): 可选 schema extractor,用于元数据 schema 分析与注册。
    enable_path_monitoring (Optional[bool]): 是否监控本地数据目录的文件新增和删除。仅在未接入 ``DocServer`` / ``DocumentProcessor`` 的本地模式下默认开启。


Examples:
    >>> import lazyllm
    >>> from lazyllm.tools import Document
    >>> m = lazyllm.OnlineEmbeddingModule(source="glm")
    >>> documents = Document(dataset_path='your_doc_path', embed=m, manager=False)  # or documents = Document(dataset_path='your_doc_path', embed={"key": m}, manager=False)
    >>> m1 = lazyllm.TrainableModule("bge-large-zh-v1.5").start()
    >>> document1 = Document(dataset_path='your_doc_path', embed={"online": m, "local": m1}, manager=False)

    >>> store_conf = {
    >>>     "segment_store": {
    >>>         "type": "map",
    >>>         "kwargs": {
    >>>             "uri": "/tmp/tmp_segments.db",
    >>>         },
    >>>     },
    >>>     "vector_store": {
    >>>         "type": "milvus",
    >>>         "kwargs": {
    >>>             "uri": "/tmp/tmp_milvus.db",
    >>>             "index_kwargs": {
    >>>                 "index_type": "FLAT",
    >>>                 "metric_type": "COSINE",
    >>>             },
    >>>         },
    >>>     },
    >>> }
    >>> doc_fields = {
    >>>     'author': DocField(data_type=DataType.VARCHAR, max_size=128, default_value=' '),
    >>>     'public_year': DocField(data_type=DataType.INT32),
    >>> }
    >>> document2 = Document(dataset_path='your_doc_path', embed={"online": m, "local": m1}, store_conf=store_conf, doc_fields=doc_fields)
    """
    class _Manager(ModuleBase):
        @staticmethod
        def _resolve_dataset_path(dataset_path: Optional[str]) -> Optional[str]:
            if not dataset_path:
                return dataset_path
            if os.path.exists(dataset_path):
                return os.path.join(os.getcwd(), dataset_path)
            default_path = os.path.join(lazyllm.config['data_path'], dataset_path)
            return default_path if os.path.exists(default_path) else dataset_path

        @staticmethod
        def _decide_service_mode(manager, store_conf, processor, dataset_path) -> Tuple[bool, bool]:
            """Returns ``(spawn_doc_server, connect_doc_server)``."""
            if isinstance(manager, str) and manager != 'ui':
                raise ValueError(f'Unsupported manager value: {manager}')
            spawn = bool(manager) and not isinstance(manager, DocServer)
            connect = isinstance(manager, DocServer)
            if (not spawn and not connect and not processor and is_persistent_store(store_conf)
                    and dataset_path and not os.path.isfile(dataset_path)
                    and not any(iter_embedded_store_endpoints(store_conf))):
                lazyllm.LOG.info(f'Persistent store detected (type={store_conf.get("type")}),'
                                 f' auto-enabling DocServer for production-grade file tracking and scan.')
                spawn = True
            return spawn, connect

        @staticmethod
        def _reject_embedded_store_with_service_mode(store_conf, *, spawn, connect, processor):
            """Service-mode RAG + an embedded single-process backend races the subprocesses on
            shared on-disk state; force the user to point at a networked endpoint instead."""
            if not (spawn or connect or processor is not None):
                return
            embedded = list(iter_embedded_store_endpoints(store_conf))
            if not embedded:
                return
            raise ValueError(
                'Document with `manager=True` / `manager=DocServer(...)` / `manager=DocumentProcessor(...)`'
                ' does not support embedded (filesystem-bound) vector stores. Point the store config at a'
                ' remote service (http/https/tcp/grpc/unix scheme), e.g.'
                " Milvus: {'type': 'milvus', 'kwargs': {'uri': os.getenv('MILVUS_URI', 'http://<host>:19530')}}"
                " or Chroma: {'type': 'chroma', 'kwargs': {'uri': 'http://<host>:8000'}}."
                f' Offending store(s): {embedded!r}.')

        def _iter_kbs(self):
            return self._kbs._impl._m if isinstance(self._kbs, ServerModule) else self._kbs

        def __init__(self, dataset_path: Optional[str],
                     embed: Optional[Union[Callable, Dict[str, Callable]]] = None,
                     manager: Union[bool, str, DocServer] = False, server: Union[bool, int] = False,
                     name: Optional[str] = None, launcher: Optional[Launcher] = None,
                     store_conf: Optional[Dict] = None, doc_fields: Optional[Dict[str, DocField]] = None,
                     cloud: bool = False, doc_files: Optional[List[str]] = None,
                     processor: Optional[DocumentProcessor] = None, display_name: Optional[str] = '',
                     description: Optional[str] = 'algorithm description',
                     schema_extractor: Optional[Union[LLMBase, SchemaExtractor]] = None,
                     create_ui: bool = False):
            super().__init__()
            self._origin_path, self._doc_files, self._cloud = dataset_path, doc_files, cloud
            self._dataset_path = self._resolve_dataset_path(dataset_path)
            self._embed = self._get_embeds(embed)
            self._processor = processor
            self._create_ui = create_ui
            self._spawn_doc_server = False
            self._doc_processor_started = False

            spawn_doc_server, connect_doc_server = self._decide_service_mode(
                manager, store_conf, processor, self._dataset_path)
            self._reject_embedded_store_with_service_mode(
                store_conf, spawn=spawn_doc_server, connect=connect_doc_server, processor=processor)

            self._launcher: Launcher = launcher if launcher else (
                lazyllm.launchers.empty(sync=False) if spawn_doc_server else lazyllm.launchers.remote(sync=False))
            self._doc_impl_dataset_path = self._dataset_path if not (spawn_doc_server or connect_doc_server) else None
            self._doc_processor = None
            if spawn_doc_server:
                self._spawn_doc_server = True
                self._doc_processor = DocumentProcessor(launcher=self._launcher, pythonpath=_LOCAL_PYTHONPATH)
                self._submodules.remove(self._doc_processor)
            elif connect_doc_server:
                self._manager = manager
                parser_url = getattr(getattr(manager, '_raw_impl', None), '_parser_url', None) or manager.parser_url
                if parser_url:
                    self._doc_processor = DocumentProcessor(url=parser_url)
            self._schema_extractor = schema_extractor
            self._store_conf = store_conf
            self._display_name = display_name
            self._description = description
            name = name or RAG_DEFAULT_GROUP_NAME
            if not display_name: display_name = name
            doc_processor = self._doc_processor or processor
            self._kbs = CallableDict({name: DocImpl(
                embed=self._embed, dataset_path=self._doc_impl_dataset_path, doc_files=doc_files,
                global_metadata_desc=doc_fields, store=store_conf, processor=doc_processor,
                algo_name=name, display_name=display_name, description=description,
                schema_extractor=schema_extractor)})

            if create_ui and not self._spawn_doc_server:
                self.ensure_doc_web()
            if server:
                self._kbs = ServerModule(self._kbs, port=(None if isinstance(server, bool) else int(server)))
            self._global_metadata_desc = doc_fields

        @property
        def url(self):
            if hasattr(self, '_manager'): return self._manager._url
            return None

        @property
        @deprecated('Document.manager.url')
        def _url(self):
            return self.url

        @property
        def web_url(self):
            if hasattr(self, '_docweb'): return self._docweb.url
            return None

        def ensure_doc_web(self):
            if hasattr(self, '_docweb'):
                return self._docweb
            if self._spawn_doc_server and not hasattr(self, '_manager'):
                raise ValueError('`create_ui=True` with `manager=True` requires `Document.start()` before using the UI')
            if not hasattr(self, '_manager') or not isinstance(self._manager, DocServer):
                raise ValueError(
                    '`create_ui=True` requires an available DocServer. '
                    'Set `manager=True` or pass `manager=DocServer(...)`.'
                )
            self._docweb = DocWebModule(doc_server=self._manager)
            return self._docweb

        def _ensure_doc_processor_started(self):
            if self._doc_processor and not self._doc_processor_started:
                self._doc_processor.start()
                self._doc_processor_started = True

        def _ensure_managed_services_started(self):
            if self._spawn_doc_server:
                self._ensure_doc_processor_started()
                if not hasattr(self, '_manager'):
                    # Start DocServer with scanning disabled; enable only after
                    # all KBs + parser algorithms are registered so the first
                    # scan sees a consistent routing table.
                    self._manager = DocServer(
                        launcher=self._launcher,
                        storage_dir=self._dataset_path,
                        parser_url=self._doc_processor.url,
                        pythonpath=_LOCAL_PYTHONPATH,
                        enable_scan=bool(self._dataset_path),
                    )
                    self._manager.start()
                    kbs = self._iter_kbs()
                    for kb_name in kbs:
                        self._manager.ensure_kb_registered(kb_name)
                    for impl in kbs.values():
                        impl._lazy_init()
                    self._manager.enable_scanning()
                if self._create_ui and not hasattr(self, '_docweb'):
                    self.ensure_doc_web()
                    self._docweb.start()

        def _get_deploy_tasks(self):
            if self._spawn_doc_server and not hasattr(self, '_manager'):
                return lazyllm.pipeline(self._ensure_managed_services_started)
            return None

        def _get_embeds(self, embed):
            embeds = embed if isinstance(embed, dict) else {EMBED_DEFAULT_KEY: embed} if embed else {}
            for index, module in enumerate(embeds.values()):
                if isinstance(module, ModuleBase):
                    setattr(self, f'_embed_module_{index}', module)
            return embeds

        def add_kb_group(self, name, doc_fields: Optional[Dict[str, DocField]] = None,
                         store_conf: Optional[Dict] = None, embed: Optional[Union[Callable, Dict[str, Callable]]] = None,
                         schema_extractor: Optional[Union[LLMBase, SchemaExtractor]] = None):
            embed = self._get_embeds(embed) if embed else self._embed
            schema_extractor = schema_extractor or self._schema_extractor
            if isinstance(schema_extractor, ModuleBase):
                setattr(self, f'_schema_extractor_{name}', schema_extractor)
            impl = DocImpl(
                dataset_path=self._doc_impl_dataset_path, embed=embed, kb_group_name=name,
                global_metadata_desc=doc_fields,
                store=store_conf or (None if (self._doc_processor or self._processor) else self._store_conf),
                processor=self._doc_processor or self._processor,
                algo_name=name, display_name=name, description='',
                schema_extractor=schema_extractor,
            )
            self._iter_kbs()[name] = impl
            # Register KB with DocServer if it's already running so the next scan cycle picks it up.
            if hasattr(self, '_manager') and isinstance(self._manager, DocServer):
                self._manager.ensure_kb_registered(name)
                impl._lazy_init()

        def get_doc_by_kb_group(self, name):
            return self._iter_kbs()[name]

        def stop(self):
            if hasattr(self, '_docweb'):
                self._docweb.stop()
            self._launcher.cleanup()

        def __call__(self, *args, **kw):
            return self._kbs(*args, **kw)

    def __new__(cls, *args, **kw):
        if url := kw.pop('url', None):
            name = kw.pop('name', None)
            if args or kw:
                raise TypeError(
                    f"When 'url' is provided, only 'name' is allowed. "
                    f'Got args={args}, extra kwargs={kw}'
                )
            return UrlDocument(url, name)
        else:
            return super().__new__(cls)

    @staticmethod
    def _coerce_document_processor_manager(manager, store_conf, dataset_path):
        """Validate the ``manager=DocumentProcessor(...)`` combination.

        Returns ``(processor, manager)``: when ``manager`` is a ``DocumentProcessor``
        the returned ``processor`` is the original instance and ``manager`` becomes
        ``False``; otherwise ``processor`` is ``None`` and ``manager`` passes through.
        """
        if not isinstance(manager, DocumentProcessor):
            return None, manager
        if store_conf is not None:
            raise ValueError(
                '`store_conf` must not be passed to `Document` when `manager` is a DocumentProcessor; '
                'set `store_conf` on the DocumentProcessor instance instead.'
            )
        if getattr(manager, '_store_conf', None) is None:
            raise ValueError(
                '`manager=DocumentProcessor(...)` requires the DocumentProcessor to have `store_conf` set; '
                'pass `store_conf=...` when constructing the DocumentProcessor.'
            )
        if is_local_map_store(manager._store_conf):
            raise ValueError('`manager=DocumentProcessor(...)` does not support pure local map store')
        if dataset_path is not None:
            raise ValueError(
                '`manager=DocumentProcessor(...)` does not accept a local `dataset_path`: the external'
                ' parsing service does not own directory scanning / lifecycle management. Use'
                ' `manager=True` or `manager=DocServer(...)` for scan-based ingestion, or drop'
                ' `dataset_path` and upload documents via explicit API calls.')
        return manager, False

    def __init__(self, dataset_path: Optional[str] = None, embed: Optional[Union[Callable, Dict[str, Callable]]] = None,
                 create_ui: bool = False,
                 manager: Union[bool, str, DocServer, 'Document._Manager', DocumentProcessor] = False,
                 server: Union[bool, int] = False, name: Optional[str] = None,
                 launcher: Optional[Launcher] = None, doc_files: Optional[List[str]] = None,
                 doc_fields: Dict[str, DocField] = None,
                 store_conf: Optional[Dict] = None, display_name: Optional[str] = '',
                 description: Optional[str] = 'algorithm description',
                 schema_extractor: Optional[Union[LLMBase, SchemaExtractor]] = None,
                 enable_path_monitoring: Optional[bool] = None):
        super().__init__()
        if create_ui:
            warnings.warn('`create_ui=True` (and the legacy `manager="ui"` alias) is deprecated and will be removed'
                          ' in a future release. Prefer `manager=True` and interact with DocServer via its HTTP API'
                          ' / SDK instead.', DeprecationWarning, stacklevel=2)
        if isinstance(manager, str):
            if manager != 'ui': raise ValueError(f'Unsupported manager value: {manager}')
            create_ui = manager = True
        if enable_path_monitoring is not None:
            warnings.warn('`enable_path_monitoring` is deprecated: DocImpl no longer polls the dataset '
                          'directory. Persistent-store setups auto-upgrade to DocServer which owns scanning; '
                          'map-store setups get a one-time ingest at `_lazy_init`. The parameter is accepted '
                          'for backward compatibility but has no effect.', DeprecationWarning, stacklevel=2)
        if isinstance(dataset_path, (tuple, list)):
            doc_fields = dataset_path
            dataset_path = None
        if doc_files is not None:
            assert dataset_path is None and not manager, (
                'Manager and dataset_path are not supported for Document with temp-files')
            assert store_conf is None or store_conf['type'] == 'map', (
                'Only map store is supported for Document with temp-files')

        name = name or RAG_DEFAULT_GROUP_NAME

        if isinstance(manager, Document._Manager):
            assert not server, 'Server information is already set by manager'
            assert not launcher, 'Launcher information is already set by manager'
            assert not manager._cloud, 'manager is not allowed to share in cloud mode'
            assert manager._doc_files is None, 'manager is not allowed to share with temp files'
            if dataset_path != manager._dataset_path and dataset_path != manager._origin_path:
                raise RuntimeError(f'Document path mismatch, expected `{manager._dataset_path}`'
                                   f'while received `{dataset_path}`')
            manager.add_kb_group(name=name, doc_fields=doc_fields, store_conf=store_conf, embed=embed,
                                 schema_extractor=schema_extractor)
            if create_ui:
                manager.ensure_doc_web()
            self._manager = manager
            self._curr_group = name
        else:
            processor, manager = self._coerce_document_processor_manager(manager, store_conf, dataset_path)
            cloud = processor is not None
            self._manager = Document._Manager(dataset_path, embed, manager, server, name, launcher, store_conf,
                                              doc_fields, cloud=cloud, doc_files=doc_files, processor=processor,
                                              display_name=display_name, description=description,
                                              schema_extractor=schema_extractor, create_ui=create_ui)
            self._curr_group = name
        self._graph_document: weakref.ref = None

    @staticmethod
    def list_all_files_in_directory(dataset_path: str, skip_hidden_path: bool = True,
                                    recursive: bool = True) -> List[str]:
        """列出指定目录路径中的所有文件。

该方法会以递归或非递归方式遍历目录并收集所有文件路径。可以选择跳过隐藏文件和目录(以 “.” 开头的)。如果传入的路径本身是文件,则返回仅包含该文件路径的列表。

Args:
    dataset_path (str): 要列出文件列表的目录。
    skip_hidden_path (bool, optional): 是否跳过隐藏文件和目录(以 “.” 开头)。默认值为 True
    recursive (bool, optional): 是否递归搜索子目录。如果为 False,则只返回当前目录下的文件。默认值为 True。

**Returns:**

- List[str]: 绝对文件路径列表。如果路径不存在或不是目录,则返回空列表。
"""
        if not os.path.exists(dataset_path):
            return []
        if not os.path.isdir(dataset_path):
            return [dataset_path] if os.path.isfile(dataset_path) else []
        files_list = []
        if recursive:
            for root, dirs, files in os.walk(os.path.abspath(dataset_path)):
                if skip_hidden_path:
                    if any(part.startswith('.') for part in root.split(os.sep) if part):
                        continue
                    dirs[:] = [d for d in dirs if not d.startswith('.')]
                    files = [f for f in files if not f.startswith('.')]
                files_list.extend(os.path.join(root, f) for f in files)
        else:
            for item in os.listdir(dataset_path):
                if skip_hidden_path and item.startswith('.'):
                    continue
                item_path = os.path.join(dataset_path, item)
                if os.path.isfile(item_path):
                    files_list.append(item_path)
        return files_list

    def _list_all_files_in_dataset(self, skip_hidden_path: bool = True) -> List[str]:
        return self.list_all_files_in_directory(self._manager._dataset_path, skip_hidden_path)

    @property
    def url(self):
        assert isinstance(self._manager._kbs, ServerModule), 'Document is not a service, please set `manager` to `True`'
        return self._manager._kbs._url

    @deprecated('Use SchemaExtractor directly')
    def connect_sql_manager(self, sql_manager: SqlManager, schma=None,
                            force_refresh: bool = True):
        """.. deprecated:: 已废弃,请直接使用 SchemaExtractor。

此方法已移除,请使用 ``SchemaExtractor`` 配合 ``register_schema_set`` 替代。
"""
        raise NotImplementedError(
            'connect_sql_manager is removed. Use SchemaExtractor with register_schema_set instead.'
        )

    def get_sql_manager(self):
        """获取当前文档模块绑定的 SchemaExtractor 的 NL2SQL 管理器实例,可用于构建 SqlCall。

**Returns:**\\n
- SqlManager: SQL 管理器实例。
"""
        ext = self._schema_extractor
        if ext is None:
            raise ValueError('No schema extractor configured for this Document')
        return ext.sql_manager_for_nl2sql()

    def extract_db_schema(
        self, llm: Union[OnlineChatModule, TrainableModule] = None, print_schema: bool = False
    ):
        """基于文档数据集和大语言模型自动提取数据库表模式(schema)并注册。

Args:
    llm (Union[OnlineChatModule, TrainableModule], optional): 用于 schema 分析的 LLM,默认使用 SchemaExtractor 自带的 LLM。
    print_schema (bool, optional): 是否在日志中打印提取的 schema。默认为 ``False``。
"""
        ext = self._schema_extractor
        if ext is None:
            raise ValueError('No schema extractor configured for this Document')
        file_paths = self._list_all_files_in_dataset()
        result = ext.analyze_schema_and_register(data=file_paths)
        if print_schema:
            lazyllm.LOG.info(f'Extracted Schema:\n\t{result}\n')
        return result

    def update_database(self, llm: Union[OnlineChatModule, TrainableModule] = None):
        """使用 SchemaExtractor 解析文档并将提取的信息更新到数据库。

Args:
    llm (Union[OnlineChatModule, TrainableModule], optional): 用于信息抽取的 LLM,默认使用 SchemaExtractor 自带的 LLM。
"""
        ext = self._schema_extractor
        if ext is None:
            raise ValueError('No schema extractor configured for this Document')
        file_paths = self._list_all_files_in_dataset()
        for fp in file_paths:
            ext.extract_and_store(data=fp)

    @deprecated('Document(dataset_path, manager=doc.manager, name=xx, doc_fields=xx, store_conf=xx)')
    def create_kb_group(self, name: str, doc_fields: Optional[Dict[str, DocField]] = None,
                        store_conf: Optional[Dict] = None) -> 'Document':
        """创建一个新的知识库分组(KB Group),并返回绑定到该分组的文档对象。

知识库分组用于在同一个文档模块中划分不同的文档集合,每个分组可以有独立的字段定义和存储配置。

Args:
    name (str): 知识库分组的名称。
    doc_fields (Optional[Dict[str, DocField]]): 文档字段定义。指定每个字段的名称、类型和描述。
    store_conf (Optional[Dict]): 存储配置,用于定义存储后端及其参数。

**Returns:**

- Document: 一个绑定到新建知识库分组的文档对象副本。
"""
        self._manager.add_kb_group(name=name, doc_fields=doc_fields, store_conf=store_conf)
        doc = copy.copy(self)
        doc._curr_group = name
        return doc

    @property
    @deprecated('Document._manager')
    def _impls(self): return self._manager

    @property
    def _impl(self) -> DocImpl: return self._manager.get_doc_by_kb_group(self._curr_group)

    @property
    def _schema_extractor(self):
        # Compat shim: read through the active DocImpl so shared-manager KBs keep per-group values.
        # read through the active DocImpl so shared-manager KBs keep per-group values.
        impl = self._manager.get_doc_by_kb_group(self._curr_group)
        return getattr(impl, '_schema_extractor', None)

    @property
    def manager(self): return self._manager._processor or self._manager

    def activate_group(self, group_name: str, embed_keys: Optional[Union[str, List[str]]] = None,
                       enable_embed: bool = True):
        """激活指定的知识库分组,并可选择指定要启用的 embedding key。

激活后,文档模块会在该分组下执行检索和存储操作。如果未指定 embedding key,则默认启用所有可用的 embedding。

Args:
    group_name (str): 要激活的知识库分组名称。
    embed_keys (Optional[Union[str, List[str]]]): 需要启用的 embedding key,可以是单个字符串或字符串列表。默认为空列表,表示启用全部 embedding。
"""
        if embed_keys and not enable_embed:
            raise ValueError('`enable_embed` must be set to True when `embed_keys` is provided')
        # if embed_keys is None, use default embed keys
        if (enable_embed and not embed_keys) and self._manager._embed:
            embed_keys = self._manager._embed.keys()
        if isinstance(embed_keys, str): embed_keys = [embed_keys]
        self._impl.activate_group(group_name, embed_keys, enable_embed)

    def activate_groups(self, groups: Union[str, List[str]], **kwargs):
        """批量激活多个知识库分组。

该方法会依次调用 `activate_group` 来激活传入的所有分组。

Args:
    groups (Union[str, List[str]]): 要激活的分组名称或分组名称列表。
"""
        if isinstance(groups, str): groups = [groups]
        for group in groups:
            self.activate_group(group, **kwargs)

    @DynamicDescriptor
    def create_node_group(self, name: str = None, *, transform: Callable, parent: str = LAZY_ROOT_NAME,
                          trans_node: bool = None, num_workers: int = 0, display_name: str = None,
                          ref: str = None, group_type: NodeGroupType = NodeGroupType.CHUNK,
                          lazy_mode: str = None, **kwargs) -> None:
        """
创建一个由指定规则生成的 node group。

Args:
    name (str): node group 的名称。
    transform (Callable): 将 node 转换成 node group 的转换规则,函数原型是 `(DocNode, group_name, **kwargs) -> List[DocNode]`。目前内置的有 [SentenceSplitter][lazyllm.tools.SentenceSplitter]。用户也可以自定义转换规则。
    trans_node (bool): 决定了transform的输入和输出是 `DocNode` 还是 `str` ,默认为None。只有在 `transform` 为 `Callable` 时才可以设置为true。
    num_workers (int): Transform时所用的新线程数量,默认为0
    parent (str): 需要进一步转换的节点。转换之后得到的一系列新的节点将会作为该父节点的子节点。如果不指定则从根节点开始转换。
    ref (str): 当前节点组引用的其他节点组名称。引用的节点组必须是父节点组的后代。在转换时,ref 指定的节点组中的相关节点会作为参数传递给 transform 函数(如果 transform 函数支持 ref 参数)。
    kwargs: 和具体实现相关的参数。


Examples:

    >>> import lazyllm
    >>> from lazyllm.tools import Document, SentenceSplitter
    >>> m = lazyllm.OnlineEmbeddingModule(source="glm")
    >>> documents = Document(dataset_path='your_doc_path', embed=m, manager=False)
    >>> documents.create_node_group(name="sentences", transform=SentenceSplitter, chunk_size=1024, chunk_overlap=100)
    >>> # Example with ref parameter: create a node group that references another group
    >>> documents.create_node_group(name="fine_chunks", parent="sentences",
    ...                             transform=SentenceSplitter, chunk_size=128, chunk_overlap=12)
    >>> def transform_with_ref(text, ref):
    ...     # ref contains nodes from the referenced group
    ...     return "
    ".join(ref)
    >>> documents.create_node_group(name="summary_chunks", parent="sentences",
    ...                             transform=transform_with_ref, ref="fine_chunks")
    """
        assert ref is None or parent != ref, 'parent and ref must be different'
        if isinstance(self, type):
            DocImpl.create_global_node_group(name, transform=transform, parent=parent, trans_node=trans_node,
                                             num_workers=num_workers, display_name=display_name,
                                             group_type=group_type, ref=ref, lazy_mode=lazy_mode, **kwargs)
        else:
            self._impl.create_node_group(name, transform=transform, parent=parent, trans_node=trans_node,
                                         num_workers=num_workers, display_name=display_name, group_type=group_type,
                                         ref=ref, lazy_mode=lazy_mode, **kwargs)

    @DynamicDescriptor
    def add_reader(self, pattern: str, func: Optional[Callable] = None):
        """
用于实例指定文件读取器,作用范围仅对注册的 Document 对象可见。注册的文件读取器必须是 Callable 对象。只能通过函数调用的方式进行注册。并且通过实例注册的文件读取器的优先级高于通过类注册的文件读取器,并且实例和类注册的文件读取器的优先级高于系统默认的文件读取器。即优先级的顺序是:实例文件读取器 > 类文件读取器 > 系统默认文件读取器。

Args:
    pattern (str): 文件读取器适用的匹配规则
    func (Callable): 文件读取器,必须是Callable的对象


Examples:

    >>> from lazyllm.tools.rag import Document, DocNode
    >>> from lazyllm.tools.rag.readers import ReaderBase
    >>> class YmlReader(ReaderBase):
    ...     def _load_data(self, file, fs=None):
    ...         try:
    ...             import yaml
    ...         except ImportError:
    ...             raise ImportError("yaml is required to read YAML file: `pip install pyyaml`")
    ...         with open(file, 'r') as f:
    ...             data = yaml.safe_load(f)
    ...         print("Call the class YmlReader.")
    ...         return [DocNode(text=data)]
    ...
    >>> def processYml(file):
    ...     with open(file, 'r') as f:
    ...         data = f.read()
    ...     print("Call the function processYml.")
    ...     return [DocNode(text=data)]
    ...
    >>> doc1 = Document(dataset_path="your_files_path")
    >>> doc2 = Document(dataset_path="your_files_path")
    >>> doc1.add_reader("**/*.yml", YmlReader)
    >>> print(doc1._impl._local_file_reader)
    {'**/*.yml': <class '__main__.YmlReader'>}
    >>> print(doc2._impl._local_file_reader)
    {}
    >>> files = ["your_yml_files"]
    >>> Document.register_global_reader("**/*.yml", processYml)
    >>> doc1._impl._reader.load_data(input_files=files)
    Call the class YmlReader.
    >>> doc2._impl._reader.load_data(input_files=files)
    Call the function processYml.
    """
        if isinstance(self, type):
            return DocImpl.register_global_reader(pattern=pattern, func=func)
        else:
            self._impl.add_reader(pattern, func)

    @classmethod
    def register_global_reader(cls, pattern: str, func: Optional[Callable] = None):
        """
用于指定文件读取器,作用范围对于所有的 Document 对象都可见。注册的文件读取器必须是 Callable 对象。可以使用装饰器的方式进行注册,也可以通过函数调用的方式进行注册。

Args:
    pattern (str): 文件读取器适用的匹配规则
    func (Callable): 文件读取器,必须是Callable的对象


Examples:

    >>> from lazyllm.tools.rag import Document, DocNode
    >>> @Document.register_global_reader("**/*.yml")
    >>> def processYml(file):
    ...     with open(file, 'r') as f:
    ...         data = f.read()
    ...     return [DocNode(text=data)]
    ...
    >>> doc1 = Document(dataset_path="your_files_path")
    >>> doc2 = Document(dataset_path="your_files_path")
    >>> files = ["your_yml_files"]
    >>> docs1 = doc1._impl._reader.load_data(input_files=files)
    >>> docs2 = doc2._impl._reader.load_data(input_files=files)
    >>> print(docs1[0].text == docs2[0].text)
    # True
    """
        return cls.add_reader(pattern, func)

    def get_store(self):
        """获取存储占位符对象。

该方法返回一个存储层的占位符,用于延迟绑定具体的存储实现。调用者可以基于此对象进行存储相关的配置或扩展。

**Returns:**

- StorePlaceholder: 存储占位符对象。
"""
        return StorePlaceholder()

    def get_embed(self):
        """获取 embedding 占位符对象。

该方法返回一个 embedding 层的占位符,用于延迟绑定具体的 embedding 实现。调用者可以基于此对象进行 embedding 相关的配置或扩展。

**Returns:**

- EmbedPlaceholder: embedding 占位符对象。
"""
        return EmbedPlaceholder()

    def register_index(self, index_type: str, index_cls: IndexBase, *args, **kwargs) -> None:
        """注册索引类型。

该方法允许用户为文档模块注册新的索引类型,以便扩展检索能力。注册后,可以通过索引类型来调用对应的索引实现。

Args:
    index_type (str): 索引类型的名称。
    index_cls (IndexBase): 索引类,需继承自 ``IndexBase``。
    *args: 初始化索引类时的可变参数。
    **kwargs: 初始化索引类时的关键字参数。
"""
        self._impl.register_index(index_type, index_cls, *args, **kwargs)

    def _forward(self, func_name: str, *args, **kw):
        return self._manager(self._curr_group, func_name, *args, **kw)

    def start(self):
        return super().start()

    def find_parent(self, target) -> Callable:
        """查找目标的父节点。

该方法返回一个可调用对象,用于执行父节点查找操作。它会延迟调用底层实现以获取指定目标的父节点。

Args:
    target: 需要查找父节点的目标。

**Returns:**

- Callable: 可调用对象,用于执行父节点查找。


Examples:

    >>> import lazyllm
    >>> from lazyllm.tools import Document, SentenceSplitter
    >>> m = lazyllm.OnlineEmbeddingModule(source="glm")
    >>> documents = Document(dataset_path='your_doc_path', embed=m, manager=False)
    >>> documents.create_node_group(name="parent", transform=SentenceSplitter, chunk_size=1024, chunk_overlap=100)
    >>> documents.create_node_group(name="children", transform=SentenceSplitter, parent="parent", chunk_size=1024, chunk_overlap=100)
    >>> documents.find_parent('children')
    """
        return functools.partial(self._forward, 'find_parent', group=target)

    def find_children(self, target) -> Callable:
        """查找目标的子节点。

该方法返回一个可调用对象,用于执行子节点查找操作。它会延迟调用底层实现以获取指定目标的所有子节点。

Args:
    target: 需要查找子节点的目标。

**Returns:**

- Callable: 可调用对象,用于执行子节点查找。


Examples:

    >>> import lazyllm
    >>> from lazyllm.tools import Document, SentenceSplitter
    >>> m = lazyllm.OnlineEmbeddingModule(source="glm")
    >>> documents = Document(dataset_path='your_doc_path', embed=m, manager=False)
    >>> documents.create_node_group(name="parent", transform=SentenceSplitter, chunk_size=1024, chunk_overlap=100)
    >>> documents.create_node_group(name="children", transform=SentenceSplitter, parent="parent", chunk_size=1024, chunk_overlap=100)
    >>> documents.find_children('parent')
    """
        return functools.partial(self._forward, 'find_children', group=target)

    def find(self, target) -> Callable:
        """查找目标。

该方法返回一个可调用对象,用于执行目标查找操作。它会延迟调用底层实现以获取指定的目标对象。

Args:
    target: 需要查找的目标。

**Returns:**

- Callable: 可调用对象,用于执行目标查找。
"""
        return functools.partial(self._forward, 'find', group=target)

    def forward(self, *args, **kw) -> List[DocNode]:
        return self._forward('retrieve', *args, **kw)

    def clear_cache(self, group_names: Optional[List[str]] = None) -> None:
        """清理缓存。

该方法用于清理文档模块的缓存,可以指定要清理的分组名称列表。如果未指定分组名称,则默认清理所有分组的缓存。

Args:
    group_names (Optional[List[str]]): 需要清理缓存的分组名称列表。默认为 ``None``,表示清理全部缓存。
"""
        return self._forward('clear_cache', group_names)

    def drop_algorithm(self):
        """
用于删除当前文档集合的在文档解析服务中注册的算法信息。
"""
        return self._forward('drop_algorithm')

    def analyze_schema_by_llm(self, kb_id: Optional[str] = None, doc_ids: Optional[List[str]] = None):
        """
用于使用大模型为文档管理模块中特定的知识库或文档集合自动抽取字段集合,返回自动生成的Pydantic Model。
支持传入特定知识库id和文档id列表。

Args:
    kb_id: 目标知识库id
    doc_ids: 目标文档id列表
"""
        return self._forward('_analyze_schema_by_llm', kb_id, doc_ids)

    def register_schema_set(self, schema_set: Type[BaseModel], kb_id: Optional[str] = DEFAULT_KB_ID,
                            force_refresh: bool = False) -> str:
        """
手动注册一个 Pydantic Model 作为当前算法的字段集合(schema),并绑定到指定知识库。
如果该知识库已绑定其他 schema,默认会报错;传入 ``force_refresh=True`` 则会替换旧绑定并清理旧数据。

Args:
    schema_set (Type[BaseModel]): 要注册的 Pydantic 模型,用作 schema 定义。
    kb_id (Optional[str]): 目标知识库 ID,默认为 ``DEFAULT_KB_ID``。
    force_refresh (bool): 若已有绑定,是否强制刷新并覆盖。默认 ``False``。

Returns:
    str: 生成的 schema_set_id。
"""
        return self._forward('_register_schema_set', schema_set, kb_id, force_refresh)

    def get_nodes(self, uids: Optional[List[str]] = None, doc_ids: Optional[Set] = None,
                  group: Optional[str] = None, kb_id: Optional[str] = None, numbers: Optional[Set] = None,
                  limit: Optional[int] = None, offset: int = 0, return_total: bool = False,
                  sort_by_number: bool = False) -> Union[List[DocNode], Tuple[List[DocNode], int]]:
        """按条件获取节点列表。

Args:
    uids (Optional[List[str]]): 指定节点 uid 列表。
    doc_ids (Optional[Set]): 指定文档 id 集合。
    group (Optional[str]): 节点组名。
    kb_id (Optional[str]): 知识库 id。
    numbers (Optional[Set]): 节点编号集合。

**Returns:**

- List[DocNode]: 命中的节点列表。


Examples:
    >>> import lazyllm
    >>> from lazyllm.tools import Document
    >>> doc = Document()
    >>> nodes = doc.get_nodes(doc_ids={'doc_1'}, group='CoarseChunk', kb_id='kb_1', numbers={1, 2})
    """
        return self._forward(
            '_get_nodes', uids, doc_ids, group, kb_id, numbers, limit, offset, return_total, sort_by_number,
        )

    def get_window_nodes(self, node: DocNode, span: tuple[int, int] = (-5, 5),
                         merge: bool = False) -> Union[List[DocNode], DocNode]:
        """获取指定节点在同一文档内的窗口节点。

Args:
    node (DocNode): 目标节点。
    span (tuple[int, int]): 窗口范围,基于 node.number 的相对偏移。
    merge (bool): 是否将窗口节点合并为一个节点返回。

**Returns:**

- Union[List[DocNode], DocNode]: 窗口节点列表,或合并后的单节点。


Examples:
    >>> import lazyllm
    >>> from lazyllm.tools import Document
    >>> doc = Document()
    >>> node = doc.get_nodes(doc_ids={'doc_1'}, group='CoarseChunk', kb_id='kb_1', numbers={10})[0]
    >>> window_nodes = doc.get_window_nodes(node, span=(-2, 2), merge=False)
    """
        return self._forward('_get_window_nodes', node, span, merge)

    def keyword_search(self, group, keyword, doc_id='', kb_id=None,
                       phrase=True, sort_by='score', size=10, file_name=None):
        """在指定文档内做关键词精准匹配,与全库检索的 :meth:`find` 互补。

通过 ``doc_id`` 或 ``file_name`` 定位目标文档(二选一,``file_name`` 优先),支持精确短语匹配或单词级匹配,可控制排序方式与返回数量。

Args:
    group (str): 节点组名(如 ``"block"`` 或 ``"line"``)。
    keyword (str): 待匹配的关键词或短语。
    doc_id (str): 目标文档 ID,默认为空字符串。与 ``file_name`` 二选一,若同时提供则 ``file_name`` 优先。
    kb_id (Optional[str]): 知识库过滤条件(可选)。
    phrase (bool): True 为精确子串匹配,False 要求所有单词均出现。
    sort_by (str): ``"score"`` 按相关性排序,``"number"`` 按文档原始顺序排序。
    size (int): 最大返回条数。
    file_name (Optional[str]): 按文件名过滤,与 ``doc_id`` 二选一。提供此参数时 ``doc_id`` 被忽略。

Returns:
    List[dict]: 命中的切片列表。
"""
        return self._forward('_keyword_search', group, keyword, doc_id, kb_id, phrase, sort_by, size, file_name)

    def _get_post_process_tasks(self):
        return lazyllm.pipeline(lambda *a: self._forward('_lazy_init'))

    def __repr__(self):
        return lazyllm.make_repr('Module', 'Document', manager=hasattr(self._manager, '_manager'),
                                 server=isinstance(self._manager._kbs, ServerModule))

activate_group(group_name, embed_keys=None, enable_embed=True)

激活指定的知识库分组,并可选择指定要启用的 embedding key。

激活后,文档模块会在该分组下执行检索和存储操作。如果未指定 embedding key,则默认启用所有可用的 embedding。

Parameters:

  • group_name (str) –

    要激活的知识库分组名称。

  • embed_keys (Optional[Union[str, List[str]]], default: None ) –

    需要启用的 embedding key,可以是单个字符串或字符串列表。默认为空列表,表示启用全部 embedding。

Source code in lazyllm/tools/rag/document.py
    def activate_group(self, group_name: str, embed_keys: Optional[Union[str, List[str]]] = None,
                       enable_embed: bool = True):
        """激活指定的知识库分组,并可选择指定要启用的 embedding key。

激活后,文档模块会在该分组下执行检索和存储操作。如果未指定 embedding key,则默认启用所有可用的 embedding。

Args:
    group_name (str): 要激活的知识库分组名称。
    embed_keys (Optional[Union[str, List[str]]]): 需要启用的 embedding key,可以是单个字符串或字符串列表。默认为空列表,表示启用全部 embedding。
"""
        if embed_keys and not enable_embed:
            raise ValueError('`enable_embed` must be set to True when `embed_keys` is provided')
        # if embed_keys is None, use default embed keys
        if (enable_embed and not embed_keys) and self._manager._embed:
            embed_keys = self._manager._embed.keys()
        if isinstance(embed_keys, str): embed_keys = [embed_keys]
        self._impl.activate_group(group_name, embed_keys, enable_embed)

activate_groups(groups, **kwargs)

批量激活多个知识库分组。

该方法会依次调用 activate_group 来激活传入的所有分组。

Parameters:

  • groups (Union[str, List[str]]) –

    要激活的分组名称或分组名称列表。

Source code in lazyllm/tools/rag/document.py
    def activate_groups(self, groups: Union[str, List[str]], **kwargs):
        """批量激活多个知识库分组。

该方法会依次调用 `activate_group` 来激活传入的所有分组。

Args:
    groups (Union[str, List[str]]): 要激活的分组名称或分组名称列表。
"""
        if isinstance(groups, str): groups = [groups]
        for group in groups:
            self.activate_group(group, **kwargs)

add_reader(pattern, func=None)

用于实例指定文件读取器,作用范围仅对注册的 Document 对象可见。注册的文件读取器必须是 Callable 对象。只能通过函数调用的方式进行注册。并且通过实例注册的文件读取器的优先级高于通过类注册的文件读取器,并且实例和类注册的文件读取器的优先级高于系统默认的文件读取器。即优先级的顺序是:实例文件读取器 > 类文件读取器 > 系统默认文件读取器。

Parameters:

  • pattern (str) –

    文件读取器适用的匹配规则

  • func (Callable, default: None ) –

    文件读取器,必须是Callable的对象

Examples:

>>> from lazyllm.tools.rag import Document, DocNode
>>> from lazyllm.tools.rag.readers import ReaderBase
>>> class YmlReader(ReaderBase):
...     def _load_data(self, file, fs=None):
...         try:
...             import yaml
...         except ImportError:
...             raise ImportError("yaml is required to read YAML file: `pip install pyyaml`")
...         with open(file, 'r') as f:
...             data = yaml.safe_load(f)
...         print("Call the class YmlReader.")
...         return [DocNode(text=data)]
...
>>> def processYml(file):
...     with open(file, 'r') as f:
...         data = f.read()
...     print("Call the function processYml.")
...     return [DocNode(text=data)]
...
>>> doc1 = Document(dataset_path="your_files_path")
>>> doc2 = Document(dataset_path="your_files_path")
>>> doc1.add_reader("**/*.yml", YmlReader)
>>> print(doc1._impl._local_file_reader)
{'**/*.yml': <class '__main__.YmlReader'>}
>>> print(doc2._impl._local_file_reader)
{}
>>> files = ["your_yml_files"]
>>> Document.register_global_reader("**/*.yml", processYml)
>>> doc1._impl._reader.load_data(input_files=files)
Call the class YmlReader.
>>> doc2._impl._reader.load_data(input_files=files)
Call the function processYml.
Source code in lazyllm/tools/rag/document.py
    @DynamicDescriptor
    def add_reader(self, pattern: str, func: Optional[Callable] = None):
        """
用于实例指定文件读取器,作用范围仅对注册的 Document 对象可见。注册的文件读取器必须是 Callable 对象。只能通过函数调用的方式进行注册。并且通过实例注册的文件读取器的优先级高于通过类注册的文件读取器,并且实例和类注册的文件读取器的优先级高于系统默认的文件读取器。即优先级的顺序是:实例文件读取器 > 类文件读取器 > 系统默认文件读取器。

Args:
    pattern (str): 文件读取器适用的匹配规则
    func (Callable): 文件读取器,必须是Callable的对象


Examples:

    >>> from lazyllm.tools.rag import Document, DocNode
    >>> from lazyllm.tools.rag.readers import ReaderBase
    >>> class YmlReader(ReaderBase):
    ...     def _load_data(self, file, fs=None):
    ...         try:
    ...             import yaml
    ...         except ImportError:
    ...             raise ImportError("yaml is required to read YAML file: `pip install pyyaml`")
    ...         with open(file, 'r') as f:
    ...             data = yaml.safe_load(f)
    ...         print("Call the class YmlReader.")
    ...         return [DocNode(text=data)]
    ...
    >>> def processYml(file):
    ...     with open(file, 'r') as f:
    ...         data = f.read()
    ...     print("Call the function processYml.")
    ...     return [DocNode(text=data)]
    ...
    >>> doc1 = Document(dataset_path="your_files_path")
    >>> doc2 = Document(dataset_path="your_files_path")
    >>> doc1.add_reader("**/*.yml", YmlReader)
    >>> print(doc1._impl._local_file_reader)
    {'**/*.yml': <class '__main__.YmlReader'>}
    >>> print(doc2._impl._local_file_reader)
    {}
    >>> files = ["your_yml_files"]
    >>> Document.register_global_reader("**/*.yml", processYml)
    >>> doc1._impl._reader.load_data(input_files=files)
    Call the class YmlReader.
    >>> doc2._impl._reader.load_data(input_files=files)
    Call the function processYml.
    """
        if isinstance(self, type):
            return DocImpl.register_global_reader(pattern=pattern, func=func)
        else:
            self._impl.add_reader(pattern, func)

analyze_schema_by_llm(kb_id=None, doc_ids=None)

用于使用大模型为文档管理模块中特定的知识库或文档集合自动抽取字段集合,返回自动生成的Pydantic Model。 支持传入特定知识库id和文档id列表。

Parameters:

  • kb_id (Optional[str], default: None ) –

    目标知识库id

  • doc_ids (Optional[List[str]], default: None ) –

    目标文档id列表

Source code in lazyllm/tools/rag/document.py
    def analyze_schema_by_llm(self, kb_id: Optional[str] = None, doc_ids: Optional[List[str]] = None):
        """
用于使用大模型为文档管理模块中特定的知识库或文档集合自动抽取字段集合,返回自动生成的Pydantic Model。
支持传入特定知识库id和文档id列表。

Args:
    kb_id: 目标知识库id
    doc_ids: 目标文档id列表
"""
        return self._forward('_analyze_schema_by_llm', kb_id, doc_ids)

clear_cache(group_names=None)

清理缓存。

该方法用于清理文档模块的缓存,可以指定要清理的分组名称列表。如果未指定分组名称,则默认清理所有分组的缓存。

Parameters:

  • group_names (Optional[List[str]], default: None ) –

    需要清理缓存的分组名称列表。默认为 None,表示清理全部缓存。

Source code in lazyllm/tools/rag/document.py
    def clear_cache(self, group_names: Optional[List[str]] = None) -> None:
        """清理缓存。

该方法用于清理文档模块的缓存,可以指定要清理的分组名称列表。如果未指定分组名称,则默认清理所有分组的缓存。

Args:
    group_names (Optional[List[str]]): 需要清理缓存的分组名称列表。默认为 ``None``,表示清理全部缓存。
"""
        return self._forward('clear_cache', group_names)

connect_sql_manager(sql_manager, schma=None, force_refresh=True)

.. deprecated:: 已废弃,请直接使用 SchemaExtractor。

此方法已移除,请使用 SchemaExtractor 配合 register_schema_set 替代。

Source code in lazyllm/tools/rag/document.py
    @deprecated('Use SchemaExtractor directly')
    def connect_sql_manager(self, sql_manager: SqlManager, schma=None,
                            force_refresh: bool = True):
        """.. deprecated:: 已废弃,请直接使用 SchemaExtractor。

此方法已移除,请使用 ``SchemaExtractor`` 配合 ``register_schema_set`` 替代。
"""
        raise NotImplementedError(
            'connect_sql_manager is removed. Use SchemaExtractor with register_schema_set instead.'
        )

create_kb_group(name, doc_fields=None, store_conf=None)

创建一个新的知识库分组(KB Group),并返回绑定到该分组的文档对象。

知识库分组用于在同一个文档模块中划分不同的文档集合,每个分组可以有独立的字段定义和存储配置。

Parameters:

  • name (str) –

    知识库分组的名称。

  • doc_fields (Optional[Dict[str, GlobalMetadataDesc]], default: None ) –

    文档字段定义。指定每个字段的名称、类型和描述。

  • store_conf (Optional[Dict], default: None ) –

    存储配置,用于定义存储后端及其参数。

Returns:

  • Document: 一个绑定到新建知识库分组的文档对象副本。
Source code in lazyllm/tools/rag/document.py
    @deprecated('Document(dataset_path, manager=doc.manager, name=xx, doc_fields=xx, store_conf=xx)')
    def create_kb_group(self, name: str, doc_fields: Optional[Dict[str, DocField]] = None,
                        store_conf: Optional[Dict] = None) -> 'Document':
        """创建一个新的知识库分组(KB Group),并返回绑定到该分组的文档对象。

知识库分组用于在同一个文档模块中划分不同的文档集合,每个分组可以有独立的字段定义和存储配置。

Args:
    name (str): 知识库分组的名称。
    doc_fields (Optional[Dict[str, DocField]]): 文档字段定义。指定每个字段的名称、类型和描述。
    store_conf (Optional[Dict]): 存储配置,用于定义存储后端及其参数。

**Returns:**

- Document: 一个绑定到新建知识库分组的文档对象副本。
"""
        self._manager.add_kb_group(name=name, doc_fields=doc_fields, store_conf=store_conf)
        doc = copy.copy(self)
        doc._curr_group = name
        return doc

create_node_group(name=None, *, transform, parent=LAZY_ROOT_NAME, trans_node=None, num_workers=0, display_name=None, ref=None, group_type=NodeGroupType.CHUNK, lazy_mode=None, **kwargs)

创建一个由指定规则生成的 node group。

Parameters:

  • name (str, default: None ) –

    node group 的名称。

  • transform (Callable) –

    将 node 转换成 node group 的转换规则,函数原型是 (DocNode, group_name, **kwargs) -> List[DocNode]。目前内置的有 SentenceSplitter。用户也可以自定义转换规则。

  • trans_node (bool, default: None ) –

    决定了transform的输入和输出是 DocNode 还是 str ,默认为None。只有在 transformCallable 时才可以设置为true。

  • num_workers (int, default: 0 ) –

    Transform时所用的新线程数量,默认为0

  • parent (str, default: LAZY_ROOT_NAME ) –

    需要进一步转换的节点。转换之后得到的一系列新的节点将会作为该父节点的子节点。如果不指定则从根节点开始转换。

  • ref (str, default: None ) –

    当前节点组引用的其他节点组名称。引用的节点组必须是父节点组的后代。在转换时,ref 指定的节点组中的相关节点会作为参数传递给 transform 函数(如果 transform 函数支持 ref 参数)。

  • kwargs

    和具体实现相关的参数。

Examples:

>>> import lazyllm
>>> from lazyllm.tools import Document, SentenceSplitter
>>> m = lazyllm.OnlineEmbeddingModule(source="glm")
>>> documents = Document(dataset_path='your_doc_path', embed=m, manager=False)
>>> documents.create_node_group(name="sentences", transform=SentenceSplitter, chunk_size=1024, chunk_overlap=100)
>>> # Example with ref parameter: create a node group that references another group
>>> documents.create_node_group(name="fine_chunks", parent="sentences",
...                             transform=SentenceSplitter, chunk_size=128, chunk_overlap=12)
>>> def transform_with_ref(text, ref):
...     # ref contains nodes from the referenced group
...     return "
".join(ref)
>>> documents.create_node_group(name="summary_chunks", parent="sentences",
...                             transform=transform_with_ref, ref="fine_chunks")
Source code in lazyllm/tools/rag/document.py
    @DynamicDescriptor
    def create_node_group(self, name: str = None, *, transform: Callable, parent: str = LAZY_ROOT_NAME,
                          trans_node: bool = None, num_workers: int = 0, display_name: str = None,
                          ref: str = None, group_type: NodeGroupType = NodeGroupType.CHUNK,
                          lazy_mode: str = None, **kwargs) -> None:
        """
创建一个由指定规则生成的 node group。

Args:
    name (str): node group 的名称。
    transform (Callable): 将 node 转换成 node group 的转换规则,函数原型是 `(DocNode, group_name, **kwargs) -> List[DocNode]`。目前内置的有 [SentenceSplitter][lazyllm.tools.SentenceSplitter]。用户也可以自定义转换规则。
    trans_node (bool): 决定了transform的输入和输出是 `DocNode` 还是 `str` ,默认为None。只有在 `transform` 为 `Callable` 时才可以设置为true。
    num_workers (int): Transform时所用的新线程数量,默认为0
    parent (str): 需要进一步转换的节点。转换之后得到的一系列新的节点将会作为该父节点的子节点。如果不指定则从根节点开始转换。
    ref (str): 当前节点组引用的其他节点组名称。引用的节点组必须是父节点组的后代。在转换时,ref 指定的节点组中的相关节点会作为参数传递给 transform 函数(如果 transform 函数支持 ref 参数)。
    kwargs: 和具体实现相关的参数。


Examples:

    >>> import lazyllm
    >>> from lazyllm.tools import Document, SentenceSplitter
    >>> m = lazyllm.OnlineEmbeddingModule(source="glm")
    >>> documents = Document(dataset_path='your_doc_path', embed=m, manager=False)
    >>> documents.create_node_group(name="sentences", transform=SentenceSplitter, chunk_size=1024, chunk_overlap=100)
    >>> # Example with ref parameter: create a node group that references another group
    >>> documents.create_node_group(name="fine_chunks", parent="sentences",
    ...                             transform=SentenceSplitter, chunk_size=128, chunk_overlap=12)
    >>> def transform_with_ref(text, ref):
    ...     # ref contains nodes from the referenced group
    ...     return "
    ".join(ref)
    >>> documents.create_node_group(name="summary_chunks", parent="sentences",
    ...                             transform=transform_with_ref, ref="fine_chunks")
    """
        assert ref is None or parent != ref, 'parent and ref must be different'
        if isinstance(self, type):
            DocImpl.create_global_node_group(name, transform=transform, parent=parent, trans_node=trans_node,
                                             num_workers=num_workers, display_name=display_name,
                                             group_type=group_type, ref=ref, lazy_mode=lazy_mode, **kwargs)
        else:
            self._impl.create_node_group(name, transform=transform, parent=parent, trans_node=trans_node,
                                         num_workers=num_workers, display_name=display_name, group_type=group_type,
                                         ref=ref, lazy_mode=lazy_mode, **kwargs)

drop_algorithm()

用于删除当前文档集合的在文档解析服务中注册的算法信息。

Source code in lazyllm/tools/rag/document.py
    def drop_algorithm(self):
        """
用于删除当前文档集合的在文档解析服务中注册的算法信息。
"""
        return self._forward('drop_algorithm')

extract_db_schema(llm=None, print_schema=False)

基于文档数据集和大语言模型自动提取数据库表模式(schema)并注册。

Parameters:

  • llm (Union[OnlineChatModule, TrainableModule], default: None ) –

    用于 schema 分析的 LLM,默认使用 SchemaExtractor 自带的 LLM。

  • print_schema (bool, default: False ) –

    是否在日志中打印提取的 schema。默认为 False

Source code in lazyllm/tools/rag/document.py
    def extract_db_schema(
        self, llm: Union[OnlineChatModule, TrainableModule] = None, print_schema: bool = False
    ):
        """基于文档数据集和大语言模型自动提取数据库表模式(schema)并注册。

Args:
    llm (Union[OnlineChatModule, TrainableModule], optional): 用于 schema 分析的 LLM,默认使用 SchemaExtractor 自带的 LLM。
    print_schema (bool, optional): 是否在日志中打印提取的 schema。默认为 ``False``。
"""
        ext = self._schema_extractor
        if ext is None:
            raise ValueError('No schema extractor configured for this Document')
        file_paths = self._list_all_files_in_dataset()
        result = ext.analyze_schema_and_register(data=file_paths)
        if print_schema:
            lazyllm.LOG.info(f'Extracted Schema:\n\t{result}\n')
        return result

find(target)

查找目标。

该方法返回一个可调用对象,用于执行目标查找操作。它会延迟调用底层实现以获取指定的目标对象。

Parameters:

  • target

    需要查找的目标。

Returns:

  • Callable: 可调用对象,用于执行目标查找。
Source code in lazyllm/tools/rag/document.py
    def find(self, target) -> Callable:
        """查找目标。

该方法返回一个可调用对象,用于执行目标查找操作。它会延迟调用底层实现以获取指定的目标对象。

Args:
    target: 需要查找的目标。

**Returns:**

- Callable: 可调用对象,用于执行目标查找。
"""
        return functools.partial(self._forward, 'find', group=target)

find_children(target)

查找目标的子节点。

该方法返回一个可调用对象,用于执行子节点查找操作。它会延迟调用底层实现以获取指定目标的所有子节点。

Parameters:

  • target

    需要查找子节点的目标。

Returns:

  • Callable: 可调用对象,用于执行子节点查找。

Examples:

>>> import lazyllm
>>> from lazyllm.tools import Document, SentenceSplitter
>>> m = lazyllm.OnlineEmbeddingModule(source="glm")
>>> documents = Document(dataset_path='your_doc_path', embed=m, manager=False)
>>> documents.create_node_group(name="parent", transform=SentenceSplitter, chunk_size=1024, chunk_overlap=100)
>>> documents.create_node_group(name="children", transform=SentenceSplitter, parent="parent", chunk_size=1024, chunk_overlap=100)
>>> documents.find_children('parent')
Source code in lazyllm/tools/rag/document.py
    def find_children(self, target) -> Callable:
        """查找目标的子节点。

该方法返回一个可调用对象,用于执行子节点查找操作。它会延迟调用底层实现以获取指定目标的所有子节点。

Args:
    target: 需要查找子节点的目标。

**Returns:**

- Callable: 可调用对象,用于执行子节点查找。


Examples:

    >>> import lazyllm
    >>> from lazyllm.tools import Document, SentenceSplitter
    >>> m = lazyllm.OnlineEmbeddingModule(source="glm")
    >>> documents = Document(dataset_path='your_doc_path', embed=m, manager=False)
    >>> documents.create_node_group(name="parent", transform=SentenceSplitter, chunk_size=1024, chunk_overlap=100)
    >>> documents.create_node_group(name="children", transform=SentenceSplitter, parent="parent", chunk_size=1024, chunk_overlap=100)
    >>> documents.find_children('parent')
    """
        return functools.partial(self._forward, 'find_children', group=target)

find_parent(target)

查找目标的父节点。

该方法返回一个可调用对象,用于执行父节点查找操作。它会延迟调用底层实现以获取指定目标的父节点。

Parameters:

  • target

    需要查找父节点的目标。

Returns:

  • Callable: 可调用对象,用于执行父节点查找。

Examples:

>>> import lazyllm
>>> from lazyllm.tools import Document, SentenceSplitter
>>> m = lazyllm.OnlineEmbeddingModule(source="glm")
>>> documents = Document(dataset_path='your_doc_path', embed=m, manager=False)
>>> documents.create_node_group(name="parent", transform=SentenceSplitter, chunk_size=1024, chunk_overlap=100)
>>> documents.create_node_group(name="children", transform=SentenceSplitter, parent="parent", chunk_size=1024, chunk_overlap=100)
>>> documents.find_parent('children')
Source code in lazyllm/tools/rag/document.py
    def find_parent(self, target) -> Callable:
        """查找目标的父节点。

该方法返回一个可调用对象,用于执行父节点查找操作。它会延迟调用底层实现以获取指定目标的父节点。

Args:
    target: 需要查找父节点的目标。

**Returns:**

- Callable: 可调用对象,用于执行父节点查找。


Examples:

    >>> import lazyllm
    >>> from lazyllm.tools import Document, SentenceSplitter
    >>> m = lazyllm.OnlineEmbeddingModule(source="glm")
    >>> documents = Document(dataset_path='your_doc_path', embed=m, manager=False)
    >>> documents.create_node_group(name="parent", transform=SentenceSplitter, chunk_size=1024, chunk_overlap=100)
    >>> documents.create_node_group(name="children", transform=SentenceSplitter, parent="parent", chunk_size=1024, chunk_overlap=100)
    >>> documents.find_parent('children')
    """
        return functools.partial(self._forward, 'find_parent', group=target)

get_embed()

获取 embedding 占位符对象。

该方法返回一个 embedding 层的占位符,用于延迟绑定具体的 embedding 实现。调用者可以基于此对象进行 embedding 相关的配置或扩展。

Returns:

  • EmbedPlaceholder: embedding 占位符对象。
Source code in lazyllm/tools/rag/document.py
    def get_embed(self):
        """获取 embedding 占位符对象。

该方法返回一个 embedding 层的占位符,用于延迟绑定具体的 embedding 实现。调用者可以基于此对象进行 embedding 相关的配置或扩展。

**Returns:**

- EmbedPlaceholder: embedding 占位符对象。
"""
        return EmbedPlaceholder()

get_nodes(uids=None, doc_ids=None, group=None, kb_id=None, numbers=None, limit=None, offset=0, return_total=False, sort_by_number=False)

按条件获取节点列表。

Parameters:

  • uids (Optional[List[str]], default: None ) –

    指定节点 uid 列表。

  • doc_ids (Optional[Set], default: None ) –

    指定文档 id 集合。

  • group (Optional[str], default: None ) –

    节点组名。

  • kb_id (Optional[str], default: None ) –

    知识库 id。

  • numbers (Optional[Set], default: None ) –

    节点编号集合。

Returns:

  • List[DocNode]: 命中的节点列表。

Examples:

>>> import lazyllm
>>> from lazyllm.tools import Document
>>> doc = Document()
>>> nodes = doc.get_nodes(doc_ids={'doc_1'}, group='CoarseChunk', kb_id='kb_1', numbers={1, 2})
Source code in lazyllm/tools/rag/document.py
    def get_nodes(self, uids: Optional[List[str]] = None, doc_ids: Optional[Set] = None,
                  group: Optional[str] = None, kb_id: Optional[str] = None, numbers: Optional[Set] = None,
                  limit: Optional[int] = None, offset: int = 0, return_total: bool = False,
                  sort_by_number: bool = False) -> Union[List[DocNode], Tuple[List[DocNode], int]]:
        """按条件获取节点列表。

Args:
    uids (Optional[List[str]]): 指定节点 uid 列表。
    doc_ids (Optional[Set]): 指定文档 id 集合。
    group (Optional[str]): 节点组名。
    kb_id (Optional[str]): 知识库 id。
    numbers (Optional[Set]): 节点编号集合。

**Returns:**

- List[DocNode]: 命中的节点列表。


Examples:
    >>> import lazyllm
    >>> from lazyllm.tools import Document
    >>> doc = Document()
    >>> nodes = doc.get_nodes(doc_ids={'doc_1'}, group='CoarseChunk', kb_id='kb_1', numbers={1, 2})
    """
        return self._forward(
            '_get_nodes', uids, doc_ids, group, kb_id, numbers, limit, offset, return_total, sort_by_number,
        )

get_sql_manager()

获取当前文档模块绑定的 SchemaExtractor 的 NL2SQL 管理器实例,可用于构建 SqlCall。

Returns:\n - SqlManager: SQL 管理器实例。

Source code in lazyllm/tools/rag/document.py
    def get_sql_manager(self):
        """获取当前文档模块绑定的 SchemaExtractor 的 NL2SQL 管理器实例,可用于构建 SqlCall。

**Returns:**\\n
- SqlManager: SQL 管理器实例。
"""
        ext = self._schema_extractor
        if ext is None:
            raise ValueError('No schema extractor configured for this Document')
        return ext.sql_manager_for_nl2sql()

get_store()

获取存储占位符对象。

该方法返回一个存储层的占位符,用于延迟绑定具体的存储实现。调用者可以基于此对象进行存储相关的配置或扩展。

Returns:

  • StorePlaceholder: 存储占位符对象。
Source code in lazyllm/tools/rag/document.py
    def get_store(self):
        """获取存储占位符对象。

该方法返回一个存储层的占位符,用于延迟绑定具体的存储实现。调用者可以基于此对象进行存储相关的配置或扩展。

**Returns:**

- StorePlaceholder: 存储占位符对象。
"""
        return StorePlaceholder()

get_window_nodes(node, span=(-5, 5), merge=False)

获取指定节点在同一文档内的窗口节点。

Parameters:

  • node (DocNode) –

    目标节点。

  • span (tuple[int, int], default: (-5, 5) ) –

    窗口范围,基于 node.number 的相对偏移。

  • merge (bool, default: False ) –

    是否将窗口节点合并为一个节点返回。

Returns:

  • Union[List[DocNode], DocNode]: 窗口节点列表,或合并后的单节点。

Examples:

>>> import lazyllm
>>> from lazyllm.tools import Document
>>> doc = Document()
>>> node = doc.get_nodes(doc_ids={'doc_1'}, group='CoarseChunk', kb_id='kb_1', numbers={10})[0]
>>> window_nodes = doc.get_window_nodes(node, span=(-2, 2), merge=False)
Source code in lazyllm/tools/rag/document.py
    def get_window_nodes(self, node: DocNode, span: tuple[int, int] = (-5, 5),
                         merge: bool = False) -> Union[List[DocNode], DocNode]:
        """获取指定节点在同一文档内的窗口节点。

Args:
    node (DocNode): 目标节点。
    span (tuple[int, int]): 窗口范围,基于 node.number 的相对偏移。
    merge (bool): 是否将窗口节点合并为一个节点返回。

**Returns:**

- Union[List[DocNode], DocNode]: 窗口节点列表,或合并后的单节点。


Examples:
    >>> import lazyllm
    >>> from lazyllm.tools import Document
    >>> doc = Document()
    >>> node = doc.get_nodes(doc_ids={'doc_1'}, group='CoarseChunk', kb_id='kb_1', numbers={10})[0]
    >>> window_nodes = doc.get_window_nodes(node, span=(-2, 2), merge=False)
    """
        return self._forward('_get_window_nodes', node, span, merge)

在指定文档内做关键词精准匹配,与全库检索的 :meth:find 互补。

通过 doc_idfile_name 定位目标文档(二选一,file_name 优先),支持精确短语匹配或单词级匹配,可控制排序方式与返回数量。

Parameters:

  • group (str) –

    节点组名(如 "block""line")。

  • keyword (str) –

    待匹配的关键词或短语。

  • doc_id (str, default: '' ) –

    目标文档 ID,默认为空字符串。与 file_name 二选一,若同时提供则 file_name 优先。

  • kb_id (Optional[str], default: None ) –

    知识库过滤条件(可选)。

  • phrase (bool, default: True ) –

    True 为精确子串匹配,False 要求所有单词均出现。

  • sort_by (str, default: 'score' ) –

    "score" 按相关性排序,"number" 按文档原始顺序排序。

  • size (int, default: 10 ) –

    最大返回条数。

  • file_name (Optional[str], default: None ) –

    按文件名过滤,与 doc_id 二选一。提供此参数时 doc_id 被忽略。

Returns:

  • List[dict]: 命中的切片列表。

Source code in lazyllm/tools/rag/document.py
    def keyword_search(self, group, keyword, doc_id='', kb_id=None,
                       phrase=True, sort_by='score', size=10, file_name=None):
        """在指定文档内做关键词精准匹配,与全库检索的 :meth:`find` 互补。

通过 ``doc_id`` 或 ``file_name`` 定位目标文档(二选一,``file_name`` 优先),支持精确短语匹配或单词级匹配,可控制排序方式与返回数量。

Args:
    group (str): 节点组名(如 ``"block"`` 或 ``"line"``)。
    keyword (str): 待匹配的关键词或短语。
    doc_id (str): 目标文档 ID,默认为空字符串。与 ``file_name`` 二选一,若同时提供则 ``file_name`` 优先。
    kb_id (Optional[str]): 知识库过滤条件(可选)。
    phrase (bool): True 为精确子串匹配,False 要求所有单词均出现。
    sort_by (str): ``"score"`` 按相关性排序,``"number"`` 按文档原始顺序排序。
    size (int): 最大返回条数。
    file_name (Optional[str]): 按文件名过滤,与 ``doc_id`` 二选一。提供此参数时 ``doc_id`` 被忽略。

Returns:
    List[dict]: 命中的切片列表。
"""
        return self._forward('_keyword_search', group, keyword, doc_id, kb_id, phrase, sort_by, size, file_name)

list_all_files_in_directory(dataset_path, skip_hidden_path=True, recursive=True) staticmethod

列出指定目录路径中的所有文件。

该方法会以递归或非递归方式遍历目录并收集所有文件路径。可以选择跳过隐藏文件和目录(以 “.” 开头的)。如果传入的路径本身是文件,则返回仅包含该文件路径的列表。

Parameters:

  • dataset_path (str) –

    要列出文件列表的目录。

  • skip_hidden_path (bool, default: True ) –

    是否跳过隐藏文件和目录(以 “.” 开头)。默认值为 True

  • recursive (bool, default: True ) –

    是否递归搜索子目录。如果为 False,则只返回当前目录下的文件。默认值为 True。

Returns:

  • List[str]: 绝对文件路径列表。如果路径不存在或不是目录,则返回空列表。
Source code in lazyllm/tools/rag/document.py
    @staticmethod
    def list_all_files_in_directory(dataset_path: str, skip_hidden_path: bool = True,
                                    recursive: bool = True) -> List[str]:
        """列出指定目录路径中的所有文件。

该方法会以递归或非递归方式遍历目录并收集所有文件路径。可以选择跳过隐藏文件和目录(以 “.” 开头的)。如果传入的路径本身是文件,则返回仅包含该文件路径的列表。

Args:
    dataset_path (str): 要列出文件列表的目录。
    skip_hidden_path (bool, optional): 是否跳过隐藏文件和目录(以 “.” 开头)。默认值为 True
    recursive (bool, optional): 是否递归搜索子目录。如果为 False,则只返回当前目录下的文件。默认值为 True。

**Returns:**

- List[str]: 绝对文件路径列表。如果路径不存在或不是目录,则返回空列表。
"""
        if not os.path.exists(dataset_path):
            return []
        if not os.path.isdir(dataset_path):
            return [dataset_path] if os.path.isfile(dataset_path) else []
        files_list = []
        if recursive:
            for root, dirs, files in os.walk(os.path.abspath(dataset_path)):
                if skip_hidden_path:
                    if any(part.startswith('.') for part in root.split(os.sep) if part):
                        continue
                    dirs[:] = [d for d in dirs if not d.startswith('.')]
                    files = [f for f in files if not f.startswith('.')]
                files_list.extend(os.path.join(root, f) for f in files)
        else:
            for item in os.listdir(dataset_path):
                if skip_hidden_path and item.startswith('.'):
                    continue
                item_path = os.path.join(dataset_path, item)
                if os.path.isfile(item_path):
                    files_list.append(item_path)
        return files_list

register_global_reader(pattern, func=None) classmethod

用于指定文件读取器,作用范围对于所有的 Document 对象都可见。注册的文件读取器必须是 Callable 对象。可以使用装饰器的方式进行注册,也可以通过函数调用的方式进行注册。

Parameters:

  • pattern (str) –

    文件读取器适用的匹配规则

  • func (Callable, default: None ) –

    文件读取器,必须是Callable的对象

Examples:

>>> from lazyllm.tools.rag import Document, DocNode
>>> @Document.register_global_reader("**/*.yml")
>>> def processYml(file):
...     with open(file, 'r') as f:
...         data = f.read()
...     return [DocNode(text=data)]
...
>>> doc1 = Document(dataset_path="your_files_path")
>>> doc2 = Document(dataset_path="your_files_path")
>>> files = ["your_yml_files"]
>>> docs1 = doc1._impl._reader.load_data(input_files=files)
>>> docs2 = doc2._impl._reader.load_data(input_files=files)
>>> print(docs1[0].text == docs2[0].text)
# True
Source code in lazyllm/tools/rag/document.py
    @classmethod
    def register_global_reader(cls, pattern: str, func: Optional[Callable] = None):
        """
用于指定文件读取器,作用范围对于所有的 Document 对象都可见。注册的文件读取器必须是 Callable 对象。可以使用装饰器的方式进行注册,也可以通过函数调用的方式进行注册。

Args:
    pattern (str): 文件读取器适用的匹配规则
    func (Callable): 文件读取器,必须是Callable的对象


Examples:

    >>> from lazyllm.tools.rag import Document, DocNode
    >>> @Document.register_global_reader("**/*.yml")
    >>> def processYml(file):
    ...     with open(file, 'r') as f:
    ...         data = f.read()
    ...     return [DocNode(text=data)]
    ...
    >>> doc1 = Document(dataset_path="your_files_path")
    >>> doc2 = Document(dataset_path="your_files_path")
    >>> files = ["your_yml_files"]
    >>> docs1 = doc1._impl._reader.load_data(input_files=files)
    >>> docs2 = doc2._impl._reader.load_data(input_files=files)
    >>> print(docs1[0].text == docs2[0].text)
    # True
    """
        return cls.add_reader(pattern, func)

register_index(index_type, index_cls, *args, **kwargs)

注册索引类型。

该方法允许用户为文档模块注册新的索引类型,以便扩展检索能力。注册后,可以通过索引类型来调用对应的索引实现。

Parameters:

  • index_type (str) –

    索引类型的名称。

  • index_cls (IndexBase) –

    索引类,需继承自 IndexBase

  • *args

    初始化索引类时的可变参数。

  • **kwargs

    初始化索引类时的关键字参数。

Source code in lazyllm/tools/rag/document.py
    def register_index(self, index_type: str, index_cls: IndexBase, *args, **kwargs) -> None:
        """注册索引类型。

该方法允许用户为文档模块注册新的索引类型,以便扩展检索能力。注册后,可以通过索引类型来调用对应的索引实现。

Args:
    index_type (str): 索引类型的名称。
    index_cls (IndexBase): 索引类,需继承自 ``IndexBase``。
    *args: 初始化索引类时的可变参数。
    **kwargs: 初始化索引类时的关键字参数。
"""
        self._impl.register_index(index_type, index_cls, *args, **kwargs)

register_schema_set(schema_set, kb_id=DEFAULT_KB_ID, force_refresh=False)

手动注册一个 Pydantic Model 作为当前算法的字段集合(schema),并绑定到指定知识库。 如果该知识库已绑定其他 schema,默认会报错;传入 force_refresh=True 则会替换旧绑定并清理旧数据。

Parameters:

  • schema_set (Type[BaseModel]) –

    要注册的 Pydantic 模型,用作 schema 定义。

  • kb_id (Optional[str], default: DEFAULT_KB_ID ) –

    目标知识库 ID,默认为 DEFAULT_KB_ID

  • force_refresh (bool, default: False ) –

    若已有绑定,是否强制刷新并覆盖。默认 False

Returns:

  • str ( str ) –

    生成的 schema_set_id。

Source code in lazyllm/tools/rag/document.py
    def register_schema_set(self, schema_set: Type[BaseModel], kb_id: Optional[str] = DEFAULT_KB_ID,
                            force_refresh: bool = False) -> str:
        """
手动注册一个 Pydantic Model 作为当前算法的字段集合(schema),并绑定到指定知识库。
如果该知识库已绑定其他 schema,默认会报错;传入 ``force_refresh=True`` 则会替换旧绑定并清理旧数据。

Args:
    schema_set (Type[BaseModel]): 要注册的 Pydantic 模型,用作 schema 定义。
    kb_id (Optional[str]): 目标知识库 ID,默认为 ``DEFAULT_KB_ID``。
    force_refresh (bool): 若已有绑定,是否强制刷新并覆盖。默认 ``False``。

Returns:
    str: 生成的 schema_set_id。
"""
        return self._forward('_register_schema_set', schema_set, kb_id, force_refresh)

update_database(llm=None)

使用 SchemaExtractor 解析文档并将提取的信息更新到数据库。

Parameters:

Source code in lazyllm/tools/rag/document.py
    def update_database(self, llm: Union[OnlineChatModule, TrainableModule] = None):
        """使用 SchemaExtractor 解析文档并将提取的信息更新到数据库。

Args:
    llm (Union[OnlineChatModule, TrainableModule], optional): 用于信息抽取的 LLM,默认使用 SchemaExtractor 自带的 LLM。
"""
        ext = self._schema_extractor
        if ext is None:
            raise ValueError('No schema extractor configured for this Document')
        file_paths = self._list_all_files_in_dataset()
        for fp in file_paths:
            ext.extract_and_store(data=fp)

lazyllm.tools.rag.store.ChromaStore

Bases: EmbedResolveMixin, LazyLLMStoreBase

ChromaStore 是基于 Chroma 的向量存储实现,继承自 LazyLLMStoreBase,支持向量写入、检索与持久化。

Parameters:

  • uri (Optional[str], default: None ) –

    Chroma 连接 URI,当未指定 dir 时必填。

  • dir (Optional[str], default: None ) –

    本地持久化存储路径,提供时使用 PersistentClient 模式。

  • index_kwargs (Optional[Union[Dict, List]], default: None ) –

    Collection 配置参数,如索引类型、距离度量方式等。

  • client_kwargs (Optional[Dict], default: None ) –

    传递给 Chroma 客户端的额外参数。

  • **kwargs

    预留扩展参数。

Source code in lazyllm/tools/rag/store/vector/chroma_store.py
class ChromaStore(EmbedResolveMixin, LazyLLMStoreBase):
    """
ChromaStore 是基于 Chroma 的向量存储实现,继承自 LazyLLMStoreBase,支持向量写入、检索与持久化。

Args:
    uri (Optional[str]): Chroma 连接 URI,当未指定 `dir` 时必填。
    dir (Optional[str]): 本地持久化存储路径,提供时使用 PersistentClient 模式。
    index_kwargs (Optional[Union[Dict, List]]): Collection 配置参数,如索引类型、距离度量方式等。
    client_kwargs (Optional[Dict]): 传递给 Chroma 客户端的额外参数。
    **kwargs: 预留扩展参数。
"""
    capability = StoreCapability.VECTOR
    need_embedding = True
    supports_index_registration = False

    def __init__(self, uri: Optional[str] = None, dir: Optional[str] = None,
                 index_kwargs: Optional[Union[Dict, List]] = None, client_kwargs: Optional[Dict] = None,
                 **kwargs) -> None:
        assert uri or (dir), 'uri or dir must be provided'
        self._index_kwargs = index_kwargs or DEFAULT_INDEX_CONFIG
        self._client_kwargs = client_kwargs or {}
        if dir:
            self._dir = dir
        else:
            self._dir, self._host, self._port = self._parse_uri(uri)
        self._primary_key = 'uid'

    @property
    def dir(self):
        """
存储目录属性。

**Returns:**

- Optional[str]: 以斜杠结尾的目录路径,若未配置则返回 None。
"""
        if not self._dir: return None
        p = Path(self._dir)
        p = p if p.suffix else (p / 'chroma.sqlite3')
        return str(p.resolve(strict=False))

    def _parse_uri(self, uri: str):
        windows_drive = re.match(r'^[a-zA-Z]:[\\/]', uri or '')
        if ('://' not in uri) and (windows_drive or os.path.isabs(uri)):
            return os.path.abspath(uri), None, None

        p = urlparse(uri)

        if p.scheme == '':
            return os.path.abspath(uri), None, None

        if p.scheme == 'file':
            path = p.path
            if os.name == 'nt' and path.startswith('/') and re.match(r'^/[a-zA-Z]:', path):
                path = path.lstrip('/')  # file:///C:/... -> C:/...
            return os.path.abspath(path), None, None

        scheme = p.scheme
        if scheme.startswith('chroma+'):
            scheme = scheme.split('+', 1)[1]  # http or https

        if scheme in ('http', 'https'):
            host = p.hostname or '127.0.0.1'
            port = p.port or (443 if scheme == 'https' else 80)
            return None, host, port

        raise ValueError(f'Unsupported URI scheme in "{uri}". '
                         'Use file:///path or plain path for local; http(s)://host:port for remote.')

    @override
    def connect(self, embed_dims: Optional[Dict[str, int]] = None,
                embed_datatypes: Optional[Dict[str, DataType]] = None,
                embed: Optional[Dict[str, Callable]] = None,
                global_metadata_desc: Optional[Dict[str, GlobalMetadataDesc]] = None, **kwargs):
        """
初始化 Chroma 客户端并配置向量化及元数据相关设定。

Args:
    embed_dims (Optional[Dict[str, int]]): 每个嵌入键对应的向量维度,未提供时默认为空字典。
    embed_datatypes (Optional[Dict[str, DataType]]): 每个嵌入键的数据类型,仅支持 FLOAT_VECTOR 或 SPARSE_FLOAT_VECTOR。
    global_metadata_desc (Optional[Dict[str, GlobalMetadataDesc]]): 全局元数据字段的描述,支持类型:字符串、整型、浮点型、布尔型。
    **kwargs: 预留扩展参数。
"""
        self._global_metadata_desc = global_metadata_desc or {}
        self._embed_dims = embed_dims or {}
        self._embed_datatypes = embed_datatypes or {}
        self._embed = embed or {}
        self._ddl_lock = threading.Lock()
        for k, v in self._global_metadata_desc.items():
            if v.data_type not in [DataType.VARCHAR, DataType.INT32, DataType.FLOAT, DataType.BOOLEAN]:
                raise ValueError(f'[Chroma Store] Unsupported data type {v.data_type} for global metadata {k}'
                                 ' (only string, int, float, bool are supported)')
        for k, v in self._embed_datatypes.items():
            if v not in [DataType.FLOAT_VECTOR, DataType.SPARSE_FLOAT_VECTOR]:
                raise ValueError(f'[Chroma Store] Unsupported data type {v} for embed key {k}'
                                 ' (only float vector and sparse float vector are supported)')
        if self._dir:
            self._client = chromadb.PersistentClient(path=self._dir, **self._client_kwargs)
            LOG.success(f'Initialzed chroma in path: {self._dir}')
        else:
            self._client = chromadb.HttpClient(host=self._host, port=self._port, **self._client_kwargs)
            LOG.success(f'Initialzed chroma in host: {self._host}, port: {self._port}')

    @override
    def upsert(self, collection_name: str, data: List[dict]) -> bool:
        """
批量写入或更新记录(切片的id及向量数据)到 Chroma。

Args:
    collection_name (str): 集合名称。
    data (List[dict]): 文档切片数据列表。

**Returns:**

- bool: 操作成功返回 True,否则 False。
"""
        try:
            # NOTE chroma only support single embedding for each collection
            if not data:
                LOG.warning(f'[Chroma Store - upsert] No data to upsert for collection {collection_name}')
                return
            data_embeddings = data[0].get('embedding', {})
            if not data_embeddings: return
            embed_keys = list(data_embeddings.keys())
            for embed_key in embed_keys:
                with self._ddl_lock:
                    self._resolve_missing_embed_specs({embed_key})
                if embed_key not in self._embed_datatypes:
                    raise ValueError(f'Embed key {embed_key} not found in embed_datatypes')
                collection = self._client.get_or_create_collection(
                    name=self._gen_collection_name(collection_name, embed_key), configuration=self._index_kwargs)
                for i in range(0, len(data), INSERT_BATCH_SIZE):
                    collection.upsert(**self._serialize_data(data[i: i + INSERT_BATCH_SIZE], embed_key))
            return True
        except Exception as e:
            LOG.error(f'[Chroma Store - upsert] Failed to create collection {collection_name}: {e}')
            LOG.error(traceback.format_exc())
            raise e

    def _serialize_data(self, data: List[dict], embed_key: str) -> List[dict]:
        res = {'ids': [], 'embeddings': [], 'metadatas': []}
        for d in data:
            res['ids'].append(d.get('uid'))
            res['embeddings'].append(d.get('embedding', {}).get(embed_key))
            res['metadatas'].append({self._gen_global_meta_key(k): v for k, v in d.get('global_meta', {}).items()
                                     if k in self._global_metadata_desc})
        return res

    @override
    def delete(self, collection_name: str, criteria: Optional[dict] = None, **kwargs) -> bool:
        """
删除整个集合或指定记录。

Args:
    collection_name (str): 要删除的集合名称。
    criteria (Optional[dict]): 若为 None,则删除整个集合;否则按字典条件删除匹配的记录(例如按 doc_id、uid、kb_id 删除)。
    **kwargs: 预留扩展参数。

**Returns:**

- bool: 删除成功返回 True,否则返回 False。
"""
        try:
            if not criteria:
                for embed_key in self._embed_datatypes.keys():
                    try:
                        self._client.delete_collection(name=self._gen_collection_name(collection_name, embed_key))
                    except Exception:
                        continue
                return True
            else:
                filters = self._construct_criteria(criteria)
                for embed_key in self._embed_datatypes.keys():
                    try:
                        collection_name = self._gen_collection_name(collection_name, embed_key)
                        collection = self._client.get_collection(name=collection_name)
                        collection.delete(**filters)
                    except chromadb.errors.NotFoundError:
                        continue
                    except Exception as e:
                        LOG.error(f'[Chroma Store - delete] Failed to delete collection {collection_name}: {e}')
                        LOG.error(traceback.format_exc())
                        raise e
                return True
        except Exception as e:
            LOG.error(f'[Chroma Store - delete] Failed to delete collection {collection_name}: {e}')
            LOG.error(traceback.format_exc())
            return False

    @override
    def get(self, collection_name: str, criteria: Optional[dict] = None, **kwargs) -> List[dict]:
        """
根据条件检索记录。

Args:
    collection_name (str): 要查询的集合名称。
    criteria (Optional[dict]): 过滤条件,如主键或元数据(例如 doc_id、kb_id)。若为 None,则返回集合中所有记录。

**Returns:**

- List[dict]: 记录列表,每条记录包含:
    - 'uid': 记录的唯一标识符。
    - 'global_meta': 全局元数据字段的字典。
    - 'embedding': 嵌入键到对应向量的映射。
"""
        try:
            filters = self._construct_criteria(criteria) if criteria else {}
            all_data = []
            for key in self._embed_datatypes:
                try:
                    coll = self._client.get_collection(
                        name=self._gen_collection_name(collection_name, key)
                    )
                    data = coll.get(include=['metadatas', 'embeddings'], **filters)
                    all_data.append((key, data))
                except chromadb.errors.NotFoundError:
                    LOG.error(f'[ChromaStore - get] Collection {collection_name} not found')
                    continue
                except Exception as e:
                    LOG.error(f'[ChromaStore - get] Failed to get collection {collection_name}: {e}')
                    LOG.error(traceback.format_exc())
                    raise e

            res: Dict[str, Dict[str, Any]] = defaultdict(lambda: {
                'uid': None, 'global_meta': {}, 'embedding': {}})
            for embed_key, data in all_data:
                ids = data['ids']
                metas = data['metadatas']
                embs = data['embeddings']

                for uid, meta, emb in zip(ids, metas, embs):
                    entry = res[uid]
                    entry['uid'] = uid
                    if not entry['global_meta']:
                        entry['global_meta'] = {
                            k[len(GLOBAL_META_KEY_PREFIX):]: v
                            for k, v in meta.items()
                        }
                    entry['embedding'][embed_key] = list(emb)
            return list(res.values())
        except Exception as e:
            LOG.error(f'[ChromaStore - get] task fail: {e}')
            LOG.error(traceback.format_exc())

    @override
    def collection_exists(self, collection_name: str) -> bool:
        # Chroma uses one sub-collection per embed_key; the group collection
        # is considered to exist when at least one embed_key sub-collection exists.
        for embed_key in self._embed_datatypes:
            try:
                self._client.get_collection(name=self._gen_collection_name(collection_name, embed_key))
                return True
            except Exception:
                continue
        return False

    @override
    def search(self, collection_name: str, query_embedding: List[float], embed_key: str, topk: Optional[int] = 10,
               filters: Optional[Dict[str, Union[str, int, List, Set]]] = None,
               **kwargs) -> List[dict]:
        """
执行向量相似度检索。

Args:
    collection_name (str): 要查询的集合名称。
    query_embedding (List[float]): 用于检索的向量。
    embed_key (str): 指定使用的向量空间 key。
    topk (int, optional): 返回的结果数量,默认为 10。
    filters (Optional[Dict[str, Union[str, int, List, Set]]]): 可选的元数据过滤条件,用于限制检索结果。

**Returns:**

- List[dict]: 匹配结果列表,每条记录包含:
    - 'uid': 匹配记录的唯一标识符。
    - 'score': 相似度分数(1 - 距离)。
"""
        try:
            collection = self._client.get_collection(name=self._gen_collection_name(collection_name, embed_key))

            filters = self._construct_filter_expr(filters) if filters else {}
            query_results = collection.query(query_embeddings=[query_embedding], n_results=topk, **filters)
            res = []
            for i, r_list in enumerate(query_results['ids']):
                for j, uid in enumerate(r_list):
                    dis = query_results['distances'][i][j]
                    res.append({'uid': uid, 'score': 1 - dis})
            return res
        except chromadb.errors.NotFoundError:
            LOG.error(f'[ChromaStore - search] Collection {collection_name} not found')
            return []
        except Exception as e:
            LOG.error(f'[ChromaStore - search] task fail: {e}')
            LOG.error(traceback.format_exc())

    def _construct_criteria(self, criteria: dict) -> dict:
        res = {}
        if self._primary_key in criteria:
            res['ids'] = criteria[self._primary_key]
        else:
            where_conditions = []
            for key, vaule in criteria.items():
                if key not in self._global_metadata_desc:
                    continue
                field_key = self._gen_global_meta_key(key)
                if isinstance(vaule, list):
                    where_conditions.append({field_key: {'$in': vaule}})
                elif isinstance(vaule, str):
                    where_conditions.append({field_key: {'$eq': vaule}})
                else:
                    raise ValueError(f'invalid criteria type: {type(vaule)}')

            if where_conditions:
                if len(where_conditions) == 1:
                    res['where'] = where_conditions[0]
                else:
                    res['where'] = {'$and': where_conditions}
        return res

    def _construct_filter_expr(self, filters: Dict[str, Union[str, int, List, Set]]) -> str:
        where_conditions = []
        for name, candidates in filters.items():
            desc = self._global_metadata_desc.get(name)
            if not desc:
                raise ValueError(f'cannot find desc of field [{name}]')
            key = self._gen_global_meta_key(name)
            if isinstance(candidates, str):
                candidates = [candidates]
            elif (not isinstance(candidates, List)) and (not isinstance(candidates, Set)):
                candidates = list(candidates)
            where_conditions.append({key: {'$in': candidates}})

        if not where_conditions:
            return {}
        elif len(where_conditions) == 1:
            return {'where': where_conditions[0]}
        else:
            return {'where': {'$and': where_conditions}}

    def _gen_global_meta_key(self, k: str) -> str:
        return GLOBAL_META_KEY_PREFIX + k

    def _gen_collection_name(self, collection_name: str, embed_key: str) -> str:
        return collection_name + '_' + embed_key + '_embed'

dir property

存储目录属性。

Returns:

  • Optional[str]: 以斜杠结尾的目录路径,若未配置则返回 None。

connect(embed_dims=None, embed_datatypes=None, embed=None, global_metadata_desc=None, **kwargs)

初始化 Chroma 客户端并配置向量化及元数据相关设定。

Parameters:

  • embed_dims (Optional[Dict[str, int]], default: None ) –

    每个嵌入键对应的向量维度,未提供时默认为空字典。

  • embed_datatypes (Optional[Dict[str, DataType]], default: None ) –

    每个嵌入键的数据类型,仅支持 FLOAT_VECTOR 或 SPARSE_FLOAT_VECTOR。

  • global_metadata_desc (Optional[Dict[str, GlobalMetadataDesc]], default: None ) –

    全局元数据字段的描述,支持类型:字符串、整型、浮点型、布尔型。

  • **kwargs

    预留扩展参数。

Source code in lazyllm/tools/rag/store/vector/chroma_store.py
    @override
    def connect(self, embed_dims: Optional[Dict[str, int]] = None,
                embed_datatypes: Optional[Dict[str, DataType]] = None,
                embed: Optional[Dict[str, Callable]] = None,
                global_metadata_desc: Optional[Dict[str, GlobalMetadataDesc]] = None, **kwargs):
        """
初始化 Chroma 客户端并配置向量化及元数据相关设定。

Args:
    embed_dims (Optional[Dict[str, int]]): 每个嵌入键对应的向量维度,未提供时默认为空字典。
    embed_datatypes (Optional[Dict[str, DataType]]): 每个嵌入键的数据类型,仅支持 FLOAT_VECTOR 或 SPARSE_FLOAT_VECTOR。
    global_metadata_desc (Optional[Dict[str, GlobalMetadataDesc]]): 全局元数据字段的描述,支持类型:字符串、整型、浮点型、布尔型。
    **kwargs: 预留扩展参数。
"""
        self._global_metadata_desc = global_metadata_desc or {}
        self._embed_dims = embed_dims or {}
        self._embed_datatypes = embed_datatypes or {}
        self._embed = embed or {}
        self._ddl_lock = threading.Lock()
        for k, v in self._global_metadata_desc.items():
            if v.data_type not in [DataType.VARCHAR, DataType.INT32, DataType.FLOAT, DataType.BOOLEAN]:
                raise ValueError(f'[Chroma Store] Unsupported data type {v.data_type} for global metadata {k}'
                                 ' (only string, int, float, bool are supported)')
        for k, v in self._embed_datatypes.items():
            if v not in [DataType.FLOAT_VECTOR, DataType.SPARSE_FLOAT_VECTOR]:
                raise ValueError(f'[Chroma Store] Unsupported data type {v} for embed key {k}'
                                 ' (only float vector and sparse float vector are supported)')
        if self._dir:
            self._client = chromadb.PersistentClient(path=self._dir, **self._client_kwargs)
            LOG.success(f'Initialzed chroma in path: {self._dir}')
        else:
            self._client = chromadb.HttpClient(host=self._host, port=self._port, **self._client_kwargs)
            LOG.success(f'Initialzed chroma in host: {self._host}, port: {self._port}')

delete(collection_name, criteria=None, **kwargs)

删除整个集合或指定记录。

Parameters:

  • collection_name (str) –

    要删除的集合名称。

  • criteria (Optional[dict], default: None ) –

    若为 None,则删除整个集合;否则按字典条件删除匹配的记录(例如按 doc_id、uid、kb_id 删除)。

  • **kwargs

    预留扩展参数。

Returns:

  • bool: 删除成功返回 True,否则返回 False。
Source code in lazyllm/tools/rag/store/vector/chroma_store.py
    @override
    def delete(self, collection_name: str, criteria: Optional[dict] = None, **kwargs) -> bool:
        """
删除整个集合或指定记录。

Args:
    collection_name (str): 要删除的集合名称。
    criteria (Optional[dict]): 若为 None,则删除整个集合;否则按字典条件删除匹配的记录(例如按 doc_id、uid、kb_id 删除)。
    **kwargs: 预留扩展参数。

**Returns:**

- bool: 删除成功返回 True,否则返回 False。
"""
        try:
            if not criteria:
                for embed_key in self._embed_datatypes.keys():
                    try:
                        self._client.delete_collection(name=self._gen_collection_name(collection_name, embed_key))
                    except Exception:
                        continue
                return True
            else:
                filters = self._construct_criteria(criteria)
                for embed_key in self._embed_datatypes.keys():
                    try:
                        collection_name = self._gen_collection_name(collection_name, embed_key)
                        collection = self._client.get_collection(name=collection_name)
                        collection.delete(**filters)
                    except chromadb.errors.NotFoundError:
                        continue
                    except Exception as e:
                        LOG.error(f'[Chroma Store - delete] Failed to delete collection {collection_name}: {e}')
                        LOG.error(traceback.format_exc())
                        raise e
                return True
        except Exception as e:
            LOG.error(f'[Chroma Store - delete] Failed to delete collection {collection_name}: {e}')
            LOG.error(traceback.format_exc())
            return False

get(collection_name, criteria=None, **kwargs)

根据条件检索记录。

Parameters:

  • collection_name (str) –

    要查询的集合名称。

  • criteria (Optional[dict], default: None ) –

    过滤条件,如主键或元数据(例如 doc_id、kb_id)。若为 None,则返回集合中所有记录。

Returns:

  • List[dict]: 记录列表,每条记录包含:
    • 'uid': 记录的唯一标识符。
    • 'global_meta': 全局元数据字段的字典。
    • 'embedding': 嵌入键到对应向量的映射。
Source code in lazyllm/tools/rag/store/vector/chroma_store.py
    @override
    def get(self, collection_name: str, criteria: Optional[dict] = None, **kwargs) -> List[dict]:
        """
根据条件检索记录。

Args:
    collection_name (str): 要查询的集合名称。
    criteria (Optional[dict]): 过滤条件,如主键或元数据(例如 doc_id、kb_id)。若为 None,则返回集合中所有记录。

**Returns:**

- List[dict]: 记录列表,每条记录包含:
    - 'uid': 记录的唯一标识符。
    - 'global_meta': 全局元数据字段的字典。
    - 'embedding': 嵌入键到对应向量的映射。
"""
        try:
            filters = self._construct_criteria(criteria) if criteria else {}
            all_data = []
            for key in self._embed_datatypes:
                try:
                    coll = self._client.get_collection(
                        name=self._gen_collection_name(collection_name, key)
                    )
                    data = coll.get(include=['metadatas', 'embeddings'], **filters)
                    all_data.append((key, data))
                except chromadb.errors.NotFoundError:
                    LOG.error(f'[ChromaStore - get] Collection {collection_name} not found')
                    continue
                except Exception as e:
                    LOG.error(f'[ChromaStore - get] Failed to get collection {collection_name}: {e}')
                    LOG.error(traceback.format_exc())
                    raise e

            res: Dict[str, Dict[str, Any]] = defaultdict(lambda: {
                'uid': None, 'global_meta': {}, 'embedding': {}})
            for embed_key, data in all_data:
                ids = data['ids']
                metas = data['metadatas']
                embs = data['embeddings']

                for uid, meta, emb in zip(ids, metas, embs):
                    entry = res[uid]
                    entry['uid'] = uid
                    if not entry['global_meta']:
                        entry['global_meta'] = {
                            k[len(GLOBAL_META_KEY_PREFIX):]: v
                            for k, v in meta.items()
                        }
                    entry['embedding'][embed_key] = list(emb)
            return list(res.values())
        except Exception as e:
            LOG.error(f'[ChromaStore - get] task fail: {e}')
            LOG.error(traceback.format_exc())

search(collection_name, query_embedding, embed_key, topk=10, filters=None, **kwargs)

执行向量相似度检索。

Parameters:

  • collection_name (str) –

    要查询的集合名称。

  • query_embedding (List[float]) –

    用于检索的向量。

  • embed_key (str) –

    指定使用的向量空间 key。

  • topk (int, default: 10 ) –

    返回的结果数量,默认为 10。

  • filters (Optional[Dict[str, Union[str, int, List, Set]]], default: None ) –

    可选的元数据过滤条件,用于限制检索结果。

Returns:

  • List[dict]: 匹配结果列表,每条记录包含:
    • 'uid': 匹配记录的唯一标识符。
    • 'score': 相似度分数(1 - 距离)。
Source code in lazyllm/tools/rag/store/vector/chroma_store.py
    @override
    def search(self, collection_name: str, query_embedding: List[float], embed_key: str, topk: Optional[int] = 10,
               filters: Optional[Dict[str, Union[str, int, List, Set]]] = None,
               **kwargs) -> List[dict]:
        """
执行向量相似度检索。

Args:
    collection_name (str): 要查询的集合名称。
    query_embedding (List[float]): 用于检索的向量。
    embed_key (str): 指定使用的向量空间 key。
    topk (int, optional): 返回的结果数量,默认为 10。
    filters (Optional[Dict[str, Union[str, int, List, Set]]]): 可选的元数据过滤条件,用于限制检索结果。

**Returns:**

- List[dict]: 匹配结果列表,每条记录包含:
    - 'uid': 匹配记录的唯一标识符。
    - 'score': 相似度分数(1 - 距离)。
"""
        try:
            collection = self._client.get_collection(name=self._gen_collection_name(collection_name, embed_key))

            filters = self._construct_filter_expr(filters) if filters else {}
            query_results = collection.query(query_embeddings=[query_embedding], n_results=topk, **filters)
            res = []
            for i, r_list in enumerate(query_results['ids']):
                for j, uid in enumerate(r_list):
                    dis = query_results['distances'][i][j]
                    res.append({'uid': uid, 'score': 1 - dis})
            return res
        except chromadb.errors.NotFoundError:
            LOG.error(f'[ChromaStore - search] Collection {collection_name} not found')
            return []
        except Exception as e:
            LOG.error(f'[ChromaStore - search] task fail: {e}')
            LOG.error(traceback.format_exc())

upsert(collection_name, data)

批量写入或更新记录(切片的id及向量数据)到 Chroma。

Parameters:

  • collection_name (str) –

    集合名称。

  • data (List[dict]) –

    文档切片数据列表。

Returns:

  • bool: 操作成功返回 True,否则 False。
Source code in lazyllm/tools/rag/store/vector/chroma_store.py
    @override
    def upsert(self, collection_name: str, data: List[dict]) -> bool:
        """
批量写入或更新记录(切片的id及向量数据)到 Chroma。

Args:
    collection_name (str): 集合名称。
    data (List[dict]): 文档切片数据列表。

**Returns:**

- bool: 操作成功返回 True,否则 False。
"""
        try:
            # NOTE chroma only support single embedding for each collection
            if not data:
                LOG.warning(f'[Chroma Store - upsert] No data to upsert for collection {collection_name}')
                return
            data_embeddings = data[0].get('embedding', {})
            if not data_embeddings: return
            embed_keys = list(data_embeddings.keys())
            for embed_key in embed_keys:
                with self._ddl_lock:
                    self._resolve_missing_embed_specs({embed_key})
                if embed_key not in self._embed_datatypes:
                    raise ValueError(f'Embed key {embed_key} not found in embed_datatypes')
                collection = self._client.get_or_create_collection(
                    name=self._gen_collection_name(collection_name, embed_key), configuration=self._index_kwargs)
                for i in range(0, len(data), INSERT_BATCH_SIZE):
                    collection.upsert(**self._serialize_data(data[i: i + INSERT_BATCH_SIZE], embed_key))
            return True
        except Exception as e:
            LOG.error(f'[Chroma Store - upsert] Failed to create collection {collection_name}: {e}')
            LOG.error(traceback.format_exc())
            raise e

lazyllm.tools.rag.store.MilvusStore

Bases: EmbedResolveMixin, LazyLLMStoreBase

基于 Milvus 的向量存储实现,继承自 StoreBase。支持向量写入、删除、相似度检索,兼容标量过滤。

Parameters:

  • uri (str, default: '' ) –

    Milvus 连接 URI(如 "tcp://localhost:19530")。如果为本地路径则使用milvus-lite,否则为远程模式(需要独立部署milvus服务,例如standalone/distributed版本)。

  • db_name (str, default: 'lazyllm' ) –

    Milvus 中使用的数据库名称,默认为 "lazyllm"。

  • index_kwargs (Optional[Union[Dict, List]], default: None ) –

    索引创建参数(例如 {"index_type": "IVF_FLAT", "metric_type": "CONSINE"} ,支持按向量模型的key配置列表)。

  • client_kwargs (Optional[Dict], default: None ) –

    传递给 milvus 客户端的额外参数。

Source code in lazyllm/tools/rag/store/vector/milvus_store.py
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
class MilvusStore(EmbedResolveMixin, LazyLLMStoreBase):
    """
基于 Milvus 的向量存储实现,继承自 StoreBase。支持向量写入、删除、相似度检索,兼容标量过滤。

Args:
    uri (str): Milvus 连接 URI(如 "tcp://localhost:19530")。如果为本地路径则使用milvus-lite,否则为远程模式(需要独立部署milvus服务,例如standalone/distributed版本)。
    db_name (str): Milvus 中使用的数据库名称,默认为 "lazyllm"。
    index_kwargs (Optional[Union[Dict, List]]): 索引创建参数(例如 {"index_type": "IVF_FLAT", "metric_type": "CONSINE"} ,支持按向量模型的key配置列表)。
    client_kwargs (Optional[Dict]): 传递给 milvus 客户端的额外参数。
"""
    capability = StoreCapability.VECTOR
    need_embedding = True
    supports_index_registration = False

    def __init__(self, uri: str = '', db_name: str = 'lazyllm', index_kwargs: Optional[Union[Dict, List]] = None,
                 client_kwargs: Optional[Dict] = None):
        # one database, different collection for each group (for standalone, add prefix to collection name)
        # when there's data need upsert, collection creation happen.
        self._uri = uri
        self._db_name = db_name
        self._index_kwargs = index_kwargs
        self._client_kwargs = client_kwargs or {}
        self._primary_key = 'uid'
        if self._uri and parse.urlparse(self._uri).scheme.lower() in ['unix', 'http', 'https', 'tcp', 'grpc']:
            self._is_remote = True
        else:
            self._is_remote = False

    @classmethod
    def rebuild(cls, uri, db_name, index_kwargs, client_kwargs):
        return cls(uri=uri, db_name=db_name, index_kwargs=index_kwargs, client_kwargs=client_kwargs)

    def __reduce__(self):
        return self.rebuild, (self._uri, self._db_name, self._index_kwargs, self._client_kwargs)

    @property
    def dir(self):
        """
存储目录属性,基于 URI 推断。远程模式返回 None。

**Returns:**

- Optional[str]: 本地 milvus.db 文件的目录路径,或 None。
"""
        if self._is_remote: return None
        p = Path(self._uri)
        p = p if p.suffix else (p / 'milvus.db')
        return str(p.resolve(strict=False))

    @override
    def try_read_dims_from_schema(self, collections: List[str]) -> Tuple[Dict[str, int], Dict[str, DataType]]:
        embed_dims, embed_datatypes = {}, {}
        if not self._uri:
            return embed_dims, embed_datatypes
        try:
            tmp = pymilvus.MilvusClient(uri=self._uri, **self._client_kwargs)
            if self._is_remote and self._db_name:
                tmp.using_database(self._db_name)
            try:
                sparse_dt = pymilvus.DataType.SPARSE_FLOAT_VECTOR
                sparse_ids = {sparse_dt, getattr(sparse_dt, 'value', None)}
                for collection_name in collections:
                    if not tmp.has_collection(collection_name):
                        continue
                    desc = tmp.describe_collection(collection_name=collection_name)
                    for field in desc.get('fields', []):
                        name = field.get('name', '') or ''
                        if not name.startswith(EMBED_PREFIX):
                            continue
                        key = name[len(EMBED_PREFIX):]
                        raw_dtype = field.get('type', field.get('dtype'))
                        params = field.get('params') or {}
                        dim = params.get('dim')
                        if raw_dtype in sparse_ids:
                            embed_datatypes[key] = DataType.SPARSE_FLOAT_VECTOR
                        elif dim is not None:
                            embed_dims[key] = int(dim)
                            embed_datatypes[key] = DataType.FLOAT_VECTOR
            finally:
                tmp.close()
        except Exception as e:
            LOG.warning(f'[Milvus] Could not read embed dims from schema: {e}')
        return embed_dims, embed_datatypes

    @override
    def connect(self, embed_dims: Optional[Dict[str, int]] = None,
                embed_datatypes: Optional[Dict[str, DataType]] = None,
                embed: Optional[Dict[str, Callable]] = None,
                global_metadata_desc: Optional[Dict[str, GlobalMetadataDesc]] = None, **kwargs):
        """
初始化 Milvus 客户端,传入向量化模型参数和全局元数据描述。

Args:
    embed_dims (Dict[str, int]): 每个嵌入键对应的向量维度。
    embed_datatypes (Dict[str, DataType]): 每个嵌入键的数据类型。
    global_metadata_desc (Dict[str, GlobalMetadataDesc]): 全局元数据字段的描述。
    kwargs: 其他连接参数
"""
        self._embed_dims = embed_dims or {}
        self._embed_datatypes = embed_datatypes or {}
        self._embed = embed or {}
        self._global_metadata_desc = global_metadata_desc or {}
        if self._index_kwargs is not None and self._embed_datatypes:
            self._index_kwargs = self.validate_milvus_embed_keys(self._index_kwargs)
        self._set_constants()

        self._ddl_lock = threading.Lock()
        self._db_ready = False
        self._ensure_database()

        max_pool_size = int(self._client_kwargs.pop('max_pool_size', 8))
        self._client_pool = _ClientPool(self._new_client, max_size=max_pool_size)
        LOG.info('[Milvus Vector Store] init success!')

    def _new_client(self):
        kwargs = dict(self._client_kwargs)
        try:
            c = pymilvus.MilvusClient(uri=self._uri, **kwargs)
            if self._is_remote and self._db_name:
                c.using_database(self._db_name)
            return c
        except Exception as e:
            LOG.error(f'[Milvus Store - _new_client] error: {e}')
            raise e

    def _ensure_database(self):
        if not (self._is_remote and self._db_name) or self._db_ready:
            return
        tmp = pymilvus.MilvusClient(uri=self._uri, **self._client_kwargs)
        try:
            with self._ddl_lock:
                if self._db_ready:
                    return
                need_create = True
                try:
                    db_list = tmp.list_databases()
                    need_create = self._db_name not in db_list
                except Exception:
                    pass
                if need_create:
                    try:
                        tmp.create_database(self._db_name)
                    except Exception as e:
                        if 'already exist' not in str(e).lower():
                            raise
                self._db_ready = True
        finally:
            tmp.close()

    @contextmanager
    def _client_context(self):
        c = self._client_pool.acquire()
        try:
            yield c
        finally:
            self._client_pool.release(c)

    def _row_has_valid_embedding(self, d: dict, required_embed_keys: Set[str]) -> bool:
        """True if row has every collection-required embed key with a non-empty value."""
        emb = d.get('embedding')
        if not emb or not isinstance(emb, dict):
            return False
        for k in required_embed_keys:
            if _is_empty_embedding_value(emb.get(k)):
                return False
        return True

    def _collection_embed_keys(self, client, collection_name: str) -> Set[str]:
        desc = client.describe_collection(collection_name=collection_name)
        return {
            field.get('name')[len(EMBED_PREFIX):]
            for field in desc.get('fields', [])
            if field.get('name', '').startswith(EMBED_PREFIX)
        }

    def _data_embed_keys(self, data: List[dict]) -> Set[str]:
        keys = set()
        for row in data:
            emb = row.get('embedding')
            if not isinstance(emb, dict):
                continue
            keys.update(k for k, v in emb.items() if not _is_empty_embedding_value(v))
        return keys

    @override
    def upsert(self, collection_name: str, data: List[dict]) -> bool:  # noqa: C901
        """
批量写入或更新切片数据到 Milvus 集合。

Args:
    collection_name (str): 集合名称,通常为 "group_embedKey" 格式。
    data (List[dict]): 切片数据列表。

**Returns:**

- bool: 操作成功返回 True,否则 False。
"""
        try:
            if not data: return True
            with self._client_context() as client:
                collection_exists = client.has_collection(collection_name)
                required_embed_keys = (
                    self._collection_embed_keys(client, collection_name)
                    if collection_exists else self._data_embed_keys(data)
                )
                if not required_embed_keys:
                    return True

                # Only require embeddings that belong to this collection. Different node groups may use
                # different embedding models, e.g. text groups use BGE while image groups use SigLIP.
                valid_data = [d for d in data if self._row_has_valid_embedding(d, required_embed_keys)]
                dropped = len(data) - len(valid_data)
                if dropped:
                    LOG.warning(f'[Milvus Store - upsert] Dropping {dropped} rows with missing/empty embedding for '
                                f'collection {collection_name}, required embeddings: {sorted(required_embed_keys)}.')
                data = valid_data
                if not data:
                    return True

                if not collection_exists:
                    with self._ddl_lock:
                        if not client.has_collection(collection_name):
                            self._resolve_missing_embed_specs(required_embed_keys)
                            if self._index_kwargs is not None:
                                self._index_kwargs = self.validate_milvus_embed_keys(self._index_kwargs)
                            embed_kwargs = {}
                            for embed_key in required_embed_keys:
                                assert self._embed_datatypes.get(embed_key), \
                                    f'cannot find embedding params for embed [{embed_key}]'
                                if embed_key not in embed_kwargs:
                                    embed_kwargs[embed_key] = {
                                        'dtype': self._type2milvus[self._embed_datatypes[embed_key]]
                                    }
                                if self._embed_dims.get(embed_key):
                                    embed_kwargs[embed_key]['dim'] = self._embed_dims[embed_key]
                            self._create_collection(client, collection_name, embed_kwargs)

                for i in range(0, len(data), MILVUS_UPSERT_BATCH_SIZE):
                    client.upsert(collection_name=collection_name,
                                  data=[self._serialize_data(d, required_embed_keys)
                                        for d in data[i:i + MILVUS_UPSERT_BATCH_SIZE]])
            return True
        except Exception as e:
            LOG.error(f'[Milvus Store - upsert] error: {e}')
            LOG.error(traceback.format_exc())
            return False

    @override
    def delete(self, collection_name: str, criteria: Optional[dict] = None, **kwargs) -> bool:
        """
删除整个集合或按条件删除指定记录。

Args:
    collection_name (str): 目标集合名称。
    criteria (Optional[dict]): 若为 None 则删除整个集合;否则按 uid 列表或元数据条件过滤。
    kwargs: 其他查询参数

**Returns:**

- bool: 如果删除成功返回True,否则返回False。
"""
        try:
            with self._client_context() as client:
                if not client.has_collection(collection_name):
                    return True
                client.load_collection(collection_name)
                if not criteria:
                    with self._ddl_lock:
                        if client.has_collection(collection_name):
                            client.drop_collection(collection_name=collection_name)
                else:
                    client.delete(collection_name=collection_name, **self._construct_criteria(criteria))
            return True
        except Exception as e:
            LOG.error(f'[Milvus Store - delete] error: {e}')
            LOG.error(traceback.format_exc())
            return False

    @override
    def get(self, collection_name: str, criteria: Optional[dict] = None, **kwargs) -> List[dict]:  # noqa: C901
        """
检索匹配主键或元数据过滤条件的记录。

Args:
    collection_name (str): 待查询集合。
    criteria (Optional[dict]): 包含 'uid' 列表或元数据字段过滤条件。
    kwargs: 其他查询参数

**Returns:**

- List[dict]: 每项包含 'uid' 及 'embedding' 映射。
"""
        try:
            with self._client_context() as client:
                if not client.has_collection(collection_name):
                    return []
                client.load_collection(collection_name)
                col_desc = client.describe_collection(collection_name=collection_name)
                field_names = [field.get('name') for field in col_desc.get('fields', [])
                               if field.get('name').startswith(EMBED_PREFIX)]
                query_kwargs = self._construct_criteria(criteria) if criteria else {}
                if version.parse(pymilvus.__version__) < version.parse('2.4.11'):
                    # For older versions, batch query manually
                    res = self._batch_query_legacy(client, collection_name, field_names, query_kwargs)
                else:
                    if criteria and self._primary_key in criteria:
                        ids = criteria[self._primary_key]
                        if isinstance(ids, str):
                            ids = [ids]
                        query_kwargs = {'filter': f'{self._primary_key} in {ids}'}
                        # return all fields
                        field_names = None
                    else:
                        query_kwargs.update(**kwargs)

                    iterator = client.query_iterator(collection_name=collection_name,
                                                     batch_size=MILVUS_PAGINATION_OFFSET,
                                                     output_fields=field_names, **query_kwargs)
                    res = []
                    while True:
                        result = iterator.next()
                        if not result:
                            iterator.close()
                            break
                        res += result
            return [self._deserialize_data(r) for r in res]
        except Exception as e:
            LOG.error(f'[Milvus Store - get] error: {e}')
            LOG.error(traceback.format_exc())
            return []

    def _batch_query_legacy(self, client, collection_name: str, field_names: List[str], kwargs: dict) -> List[dict]:
        res = []
        offset = 0
        batch_size = MILVUS_PAGINATION_OFFSET

        while True:
            try:
                # Add offset and limit to filters for pagination
                batch_kwargs = dict(kwargs)
                batch_kwargs['offset'] = offset
                batch_kwargs['limit'] = batch_size

                batch_res = client.query(collection_name=collection_name, output_fields=field_names, **batch_kwargs)
                if not batch_res:
                    break

                res.extend(batch_res)
                if len(batch_res) < batch_size:
                    break
                offset += batch_size
            except Exception as e:
                LOG.error(f'[Milvus Store - _batch_query_legacy] error: {e}')
                raise
        return res

    def _set_constants(self):
        self._type2milvus = {
            DataType.VARCHAR: pymilvus.DataType.VARCHAR,
            DataType.ARRAY: pymilvus.DataType.ARRAY,
            DataType.FLOAT_VECTOR: pymilvus.DataType.FLOAT_VECTOR,
            DataType.INT32: pymilvus.DataType.INT32,
            DataType.INT64: pymilvus.DataType.INT64,
            DataType.SPARSE_FLOAT_VECTOR: pymilvus.DataType.SPARSE_FLOAT_VECTOR,
            DataType.STRING: pymilvus.DataType.STRING,
        }
        self._builtin_keys = {
            'uid': {'dtype': pymilvus.DataType.VARCHAR, 'max_length': 256, 'is_primary': True}
        }
        self._constant_fields = self._get_constant_fields()

    def _get_constant_fields(self) -> list:
        field_list = []
        for k, kws in self._builtin_keys.items():
            field_list.append(pymilvus.FieldSchema(name=k, **kws))
        for k, desc in self._global_metadata_desc.items():
            field_name = self._gen_global_meta_key(k)
            if desc.data_type == DataType.ARRAY:
                if desc.element_type is None:
                    raise ValueError(f'Milvus field [{field_name}]: '
                                     '`element_type` is required when `data_type` is ARRAY.')
                field_args = {'element_type': self._type2milvus[desc.element_type], 'max_capacity': desc.max_size}
                if desc.element_type == DataType.VARCHAR: field_args['max_length'] = 65535
            elif desc.data_type == DataType.VARCHAR:
                field_args = {'max_length': desc.max_size}
            else:
                field_args = {}
            field_list.append(pymilvus.FieldSchema(name=field_name, dtype=self._type2milvus[desc.data_type],
                                                   default_value=desc.default_value, **field_args))
        return field_list

    def _create_collection(self, client, collection_name: str, embed_kwargs: Dict[str, Dict],  # noqa: C901
                           retry: int = 0):
        field_list = copy.deepcopy(self._constant_fields)
        index_params = client.prepare_index_params()
        original_index_kwargs = copy.deepcopy(self._index_kwargs)

        # Pre-process index_kwargs to create a lookup dictionary for O(1) access
        index_kwargs_lookup = {}
        if isinstance(original_index_kwargs, dict):
            original_index_kwargs = [original_index_kwargs]
        for item in original_index_kwargs:
            embed_key = item.get('embed_key', None)
            if not embed_key:
                raise ValueError(f'cannot find `embed_key` in `index_kwargs` of `{item}`')
            # add default values to the params of each index item with no overrides
            self._ensure_params_defaults(item)
            index_kwargs_lookup[embed_key] = item.copy()
            index_kwargs_lookup[embed_key].pop('embed_key', None)
        for k, kws in embed_kwargs.items():
            embed_field_name = self._gen_embed_key(k)
            field_list.append(pymilvus.FieldSchema(name=embed_field_name, **kws))

            if k in index_kwargs_lookup:
                index_params.add_index(field_name=embed_field_name, **index_kwargs_lookup[k])

        schema = pymilvus.CollectionSchema(fields=field_list, auto_id=False, enable_dynamic_field=False)
        try:
            client.create_collection(collection_name=collection_name, schema=schema, index_params=index_params)
        except pymilvus.MilvusException as e:
            msg = getattr(e, 'message', str(e))
            if 'invalid index type' in msg.lower():
                if retry >= MILVUS_INDEX_MAX_RETRY:
                    LOG.error(f'[Milvus Store] index fallback exceeded max retries ({MILVUS_INDEX_MAX_RETRY}),'
                              f' last error: {msg}')
                    raise
                try:
                    wrong_index_type = msg.split('invalid index type: ')[1]
                    if ',' in wrong_index_type:
                        wrong_index_type = wrong_index_type.split(',')[0].strip()
                except Exception:
                    LOG.error(f'[Milvus Store] failed to parse invalid index type from error: {msg}')
                    raise
                self._ensure_valid_index(self._index_kwargs)
                LOG.warning(f'[Milvus Store] Unsupported index type: {wrong_index_type}. '
                            f'Fallback to AUTOINDEX and retry (try #{retry + 1}).')
                self._create_collection(client, collection_name, embed_kwargs, retry=retry + 1)
            else:
                raise e

    def _ensure_valid_index(self, index_params: Union[list, dict]):
        embed_index_map = {
            DataType.FLOAT_VECTOR: ['FLAT', 'HNSW', 'IVF_FLAT', 'IVF_SQ8', 'IVF_PQ', 'AUTOINDEX', 'DISKANN'],
            DataType.SPARSE_FLOAT_VECTOR: ['SPARSE_INVERTED_INDEX', 'SPARSE_WAND'],
            DataType.VARCHAR: ['INVERTED_INDEX'],
            DataType.STRING: ['INVERTED_INDEX'],
            DataType.ARRAY: ['INVERTED_INDEX'],
            DataType.INT32: ['INVERTED_INDEX'],
            DataType.INT64: ['INVERTED_INDEX'],
            DataType.FLOAT: ['INVERTED_INDEX'],
            DataType.BOOLEAN: ['INVERTED_INDEX'],
        }

        def _replace_index_type(index_item: dict):
            """
            Raise ValueError if the DataType is not supported by Milvus.
            Raise ValueError if the IndexType is not compatible with the DataType.
            Fallback to the default index type if the IndexType is compatible with the DataType
            but not supported by Milvus.
            """
            embed_key = index_item.get('embed_key')
            dtype = self._embed_datatypes.get(embed_key)
            index_type = index_item.get('index_type').upper()
            if dtype not in embed_index_map:
                raise ValueError(f'[Milvus Store]: Unsupported data type: {DataType(dtype).name}.')
            if index_type not in embed_index_map.get(dtype):
                raise ValueError(f'[Milvus Store] {DataType(dtype).name}: Unsupported index type: {index_type}.')
            else:
                index_type = list(embed_index_map.get(dtype))[0]
                index_item['index_type'] = index_type
                self._ensure_params_defaults(index_item)

        if isinstance(index_params, list):
            for index_item in index_params:
                _replace_index_type(index_item)
        elif isinstance(index_params, dict):
            _replace_index_type(index_params)

    def _ensure_params_defaults(self, index_item: dict):
        """
        Fill in the missing fields (index_type, metric_type, params) of a single index item.
        Do not override the fields explicitly provided by the user (only setdefault)
        params will be filled in with common defaults based on index_type
        (if params already exist, only fill in missing keys)
        """
        if not isinstance(index_item, dict):
            return

        # Normalize index_type
        itype = index_item.get('index_type')
        if itype:
            itype_up = str(itype).upper()
            index_item['index_type'] = itype_up
        else:
            raise ValueError(f'cannot find `index_type` in `index_kwargs` of `{index_item}`')

        defaults = MILVUS_INDEX_TYPE_DEFAULTS.get(index_item['index_type'], None)
        if defaults is None:
            raise ValueError(f'[Milvus Store] Unsupported index type: {index_item["index_type"]}')

        # metric_type default fill (do not override user)
        if 'metric_type' not in index_item and 'metric_type' in defaults:
            index_item['metric_type'] = defaults['metric_type']

        default_params = defaults.get('params', {})
        if 'params' not in index_item or index_item.get('params') is None:
            index_item['params'] = dict(default_params)
        else:
            # fill in the missing keys of params
            if isinstance(index_item['params'], dict):
                for k, v in default_params.items():
                    index_item['params'].setdefault(k, v)
            else:
                # if user passed a non-dict (exception), replace it with the default dict
                index_item['params'] = dict(default_params)

    def _serialize_data(self, d: dict, required_embed_keys: Optional[Set[str]] = None) -> dict:
        # only keep primary_key, embedding and global_meta
        res = {
            self._primary_key: d.get(self._primary_key, '')
        }
        emb = d.get('embedding', {})
        embed_keys = required_embed_keys if required_embed_keys is not None else emb.keys()
        for embed_key in embed_keys:
            if embed_key in emb:
                res[self._gen_embed_key(embed_key)] = emb[embed_key]
        global_meta = d.get('global_meta', {})
        for name, desc in self._global_metadata_desc.items():
            value = global_meta.get(name, desc.default_value)
            if value is not None:
                res[self._gen_global_meta_key(name)] = value
        return res

    def _deserialize_data(self, d: dict) -> dict:
        res = {
            self._primary_key: d.get(self._primary_key, ''),
            'embedding': {}
        }
        for k, v in d.items():
            if k.startswith(EMBED_PREFIX):
                res['embedding'][k[len(EMBED_PREFIX):]] = v
        return res

    def _gen_embed_key(self, k: str) -> str:
        return EMBED_PREFIX + k

    def _gen_global_meta_key(self, k: str) -> str:
        return GLOBAL_META_KEY_PREFIX + k

    def _construct_criteria(self, criteria: dict) -> dict:
        res = {}
        criteria = dict(criteria)
        if self._primary_key in criteria:
            res['ids'] = criteria[self._primary_key]
        else:
            filter_str = ''
            for key, vaule in criteria.items():
                if key not in self._global_metadata_desc:
                    continue
                field_name = self._gen_global_meta_key(key)
                if len(filter_str) > 0:
                    filter_str += ' and '
                if isinstance(vaule, list):
                    filter_str += f'{field_name} in {vaule}'
                elif isinstance(vaule, str):
                    filter_str += f'{field_name} == "{vaule}"'
                else:
                    raise ValueError(f'invalid criteria type: {type(vaule)}')
            res['filter'] = filter_str
        return res

    @override
    def collection_exists(self, collection_name: str) -> bool:
        try:
            with self._client_context() as client:
                return client.has_collection(collection_name)
        except Exception as e:
            LOG.warning(f'[Milvus Store - collection_exists] error checking {collection_name}: {e}')
            return False

    @override
    def search(self, collection_name: str, query_embedding: Union[dict, List[float]], topk: int,
               filters: Optional[Dict[str, Union[List, set]]] = None, embed_key: Optional[str] = None,
               filter_str: Optional[str] = '', **kwargs) -> List[dict]:
        """
执行向量相似度检索,并可按元数据过滤。

Args:
    collection_name (str): 待搜索集合。
    query_embedding (List[float]): 查询向量。
    topk (int): 返回邻近数量。
    filters (Optional[Dict[str, Union[List, Set]]]): 元数据过滤映射。
    embed_key (str): 使用的嵌入字段。
    filter_str (Optional[str], optional): Filter expression string. Defaults to empty string
    kwargs: Other search parameters

**Returns:**

- List[dict]: 每项包含 'uid' 及相似度 'score'。
"""
        with self._client_context() as client:
            if not embed_key or embed_key not in self._embed_datatypes:
                raise ValueError(f'[Milvus Store - search] Not supported or None `embed_key`: {embed_key}')
            if not client.has_collection(collection_name):
                return []
            client.load_collection(collection_name)

            res = []
            filter_expr = self._construct_filter_expr(filters) if filters else ''
            if filter_str:
                filter_expr = f'{filter_expr} and {filter_str}' if filter_expr else filter_str

            results = client.search(collection_name=collection_name, data=[query_embedding], limit=topk,
                                    anns_field=self._gen_embed_key(embed_key),
                                    filter=filter_expr)
            if len(results) != 1:
                raise ValueError(f'number of results [{len(results)}] != expected [1]')
            for result in results[0]:
                score = result.get('distance', 0)
                uid = result.get('id', result.get(self._primary_key, ''))
                if not uid:
                    continue
                res.append({'uid': uid, 'score': score})
        return res

    def _construct_filter_expr(self, filters: Dict[str, Union[str, int, List, Set]]) -> str:
        ret_str = ''
        if not filters:
            return ret_str
        for name, candidates in filters.items():
            desc = self._global_metadata_desc.get(name)
            if not desc:
                raise ValueError(f'cannot find desc of field [{name}]')
            key = self._gen_global_meta_key(name)
            if isinstance(candidates, str):
                candidates = [candidates]
            elif (not isinstance(candidates, list)) and (not isinstance(candidates, set)):
                candidates = list(candidates)
            if desc.data_type == DataType.ARRAY:
                ret_str += f'array_contains_any({key}, {candidates}) and '
            else:
                ret_str += f'{key} in {candidates} and '
        if len(ret_str) > 0:
            return ret_str[:-5]  # truncate the last ' and '
        return ret_str

    def validate_milvus_embed_keys(self, index_kwargs: Optional[Union[List, Dict]]):  # noqa: C901
        """
        Validate and preprocess the index_kwargs of milvus store_conf:
        1. Auto fill the only one missing embed_key into the configuration without embed_key;
        2. The embed_key in self._embed must be a subset of the embed_key in store_conf;
        3. store_conf can contain additional embed_key;
        4. Duplicate embed_key is forbidden;
        5. If multiple embed_key are missing, raise an error.
        """
        if not isinstance(index_kwargs, (list, dict)):
            raise TypeError(f'[Milvus Store] index_kwargs must be a list or dict, but got {type(index_kwargs)}')

        embed_keys = list(self._embed_datatypes.keys())
        if not embed_keys:
            raise ValueError('self._embed is empty, cannot build index configuration')

        normalized_index_kwargs = []
        no_embedkey_entries = []
        seen = set()
        if isinstance(index_kwargs, dict):
            index_kwargs = [index_kwargs]
        for i, idx_conf in enumerate(index_kwargs):
            if not isinstance(idx_conf, dict):
                raise TypeError(f'index_kwargs position {i} must be a dictionary, but got {type(idx_conf)}')

            embed_key = idx_conf.get('embed_key')
            if embed_key:
                if embed_key in seen:
                    raise ValueError(f'duplicate embed_key {embed_key} in index_kwargs position {i}')
                seen.add(embed_key)
            else:
                no_embedkey_entries.append((i, idx_conf))

            normalized_index_kwargs.append(idx_conf)

        store_embed_keys = seen
        missing_keys = set(embed_keys) - store_embed_keys

        if len(missing_keys) > 1:
            raise ValueError(
                f'[Milvus Store] store_conf is missing the following embed_key: {missing_keys} '
                f'(only supports auto filling one missing item)'
            )
        elif len(missing_keys) == 1:
            missing_key = next(iter(missing_keys))

            if len(no_embedkey_entries) == 1:
                idx = no_embedkey_entries[0][1]
                idx['embed_key'] = missing_key
            elif len(no_embedkey_entries) == 0:
                if self._embed_datatypes.get(missing_key) == DataType.FLOAT_VECTOR:
                    normalized_index_kwargs.append({
                        'embed_key': missing_key,
                        'index_type': 'FLAT',
                        'metric_type': 'COSINE'
                    })
                else:
                    normalized_index_kwargs.append({
                        'embed_key': missing_key,
                        'index_type': 'SPARSE_INVERTED_INDEX',
                        'metric_type': 'L2'
                    })
            else:
                raise ValueError(
                    f'[Milvus Store] Found multiple entries without embed_key, cannot determine '
                    f'which one to fill. Missing embed_keys: {missing_keys}'
                )

        return normalized_index_kwargs

dir property

存储目录属性,基于 URI 推断。远程模式返回 None。

Returns:

  • Optional[str]: 本地 milvus.db 文件的目录路径,或 None。

connect(embed_dims=None, embed_datatypes=None, embed=None, global_metadata_desc=None, **kwargs)

初始化 Milvus 客户端,传入向量化模型参数和全局元数据描述。

Parameters:

  • embed_dims (Dict[str, int], default: None ) –

    每个嵌入键对应的向量维度。

  • embed_datatypes (Dict[str, DataType], default: None ) –

    每个嵌入键的数据类型。

  • global_metadata_desc (Dict[str, GlobalMetadataDesc], default: None ) –

    全局元数据字段的描述。

  • kwargs

    其他连接参数

Source code in lazyllm/tools/rag/store/vector/milvus_store.py
    @override
    def connect(self, embed_dims: Optional[Dict[str, int]] = None,
                embed_datatypes: Optional[Dict[str, DataType]] = None,
                embed: Optional[Dict[str, Callable]] = None,
                global_metadata_desc: Optional[Dict[str, GlobalMetadataDesc]] = None, **kwargs):
        """
初始化 Milvus 客户端,传入向量化模型参数和全局元数据描述。

Args:
    embed_dims (Dict[str, int]): 每个嵌入键对应的向量维度。
    embed_datatypes (Dict[str, DataType]): 每个嵌入键的数据类型。
    global_metadata_desc (Dict[str, GlobalMetadataDesc]): 全局元数据字段的描述。
    kwargs: 其他连接参数
"""
        self._embed_dims = embed_dims or {}
        self._embed_datatypes = embed_datatypes or {}
        self._embed = embed or {}
        self._global_metadata_desc = global_metadata_desc or {}
        if self._index_kwargs is not None and self._embed_datatypes:
            self._index_kwargs = self.validate_milvus_embed_keys(self._index_kwargs)
        self._set_constants()

        self._ddl_lock = threading.Lock()
        self._db_ready = False
        self._ensure_database()

        max_pool_size = int(self._client_kwargs.pop('max_pool_size', 8))
        self._client_pool = _ClientPool(self._new_client, max_size=max_pool_size)
        LOG.info('[Milvus Vector Store] init success!')

delete(collection_name, criteria=None, **kwargs)

删除整个集合或按条件删除指定记录。

Parameters:

  • collection_name (str) –

    目标集合名称。

  • criteria (Optional[dict], default: None ) –

    若为 None 则删除整个集合;否则按 uid 列表或元数据条件过滤。

  • kwargs

    其他查询参数

Returns:

  • bool: 如果删除成功返回True,否则返回False。
Source code in lazyllm/tools/rag/store/vector/milvus_store.py
    @override
    def delete(self, collection_name: str, criteria: Optional[dict] = None, **kwargs) -> bool:
        """
删除整个集合或按条件删除指定记录。

Args:
    collection_name (str): 目标集合名称。
    criteria (Optional[dict]): 若为 None 则删除整个集合;否则按 uid 列表或元数据条件过滤。
    kwargs: 其他查询参数

**Returns:**

- bool: 如果删除成功返回True,否则返回False。
"""
        try:
            with self._client_context() as client:
                if not client.has_collection(collection_name):
                    return True
                client.load_collection(collection_name)
                if not criteria:
                    with self._ddl_lock:
                        if client.has_collection(collection_name):
                            client.drop_collection(collection_name=collection_name)
                else:
                    client.delete(collection_name=collection_name, **self._construct_criteria(criteria))
            return True
        except Exception as e:
            LOG.error(f'[Milvus Store - delete] error: {e}')
            LOG.error(traceback.format_exc())
            return False

get(collection_name, criteria=None, **kwargs)

检索匹配主键或元数据过滤条件的记录。

Parameters:

  • collection_name (str) –

    待查询集合。

  • criteria (Optional[dict], default: None ) –

    包含 'uid' 列表或元数据字段过滤条件。

  • kwargs

    其他查询参数

Returns:

  • List[dict]: 每项包含 'uid' 及 'embedding' 映射。
Source code in lazyllm/tools/rag/store/vector/milvus_store.py
    @override
    def get(self, collection_name: str, criteria: Optional[dict] = None, **kwargs) -> List[dict]:  # noqa: C901
        """
检索匹配主键或元数据过滤条件的记录。

Args:
    collection_name (str): 待查询集合。
    criteria (Optional[dict]): 包含 'uid' 列表或元数据字段过滤条件。
    kwargs: 其他查询参数

**Returns:**

- List[dict]: 每项包含 'uid' 及 'embedding' 映射。
"""
        try:
            with self._client_context() as client:
                if not client.has_collection(collection_name):
                    return []
                client.load_collection(collection_name)
                col_desc = client.describe_collection(collection_name=collection_name)
                field_names = [field.get('name') for field in col_desc.get('fields', [])
                               if field.get('name').startswith(EMBED_PREFIX)]
                query_kwargs = self._construct_criteria(criteria) if criteria else {}
                if version.parse(pymilvus.__version__) < version.parse('2.4.11'):
                    # For older versions, batch query manually
                    res = self._batch_query_legacy(client, collection_name, field_names, query_kwargs)
                else:
                    if criteria and self._primary_key in criteria:
                        ids = criteria[self._primary_key]
                        if isinstance(ids, str):
                            ids = [ids]
                        query_kwargs = {'filter': f'{self._primary_key} in {ids}'}
                        # return all fields
                        field_names = None
                    else:
                        query_kwargs.update(**kwargs)

                    iterator = client.query_iterator(collection_name=collection_name,
                                                     batch_size=MILVUS_PAGINATION_OFFSET,
                                                     output_fields=field_names, **query_kwargs)
                    res = []
                    while True:
                        result = iterator.next()
                        if not result:
                            iterator.close()
                            break
                        res += result
            return [self._deserialize_data(r) for r in res]
        except Exception as e:
            LOG.error(f'[Milvus Store - get] error: {e}')
            LOG.error(traceback.format_exc())
            return []

search(collection_name, query_embedding, topk, filters=None, embed_key=None, filter_str='', **kwargs)

执行向量相似度检索,并可按元数据过滤。

Parameters:

  • collection_name (str) –

    待搜索集合。

  • query_embedding (List[float]) –

    查询向量。

  • topk (int) –

    返回邻近数量。

  • filters (Optional[Dict[str, Union[List, Set]]], default: None ) –

    元数据过滤映射。

  • embed_key (str, default: None ) –

    使用的嵌入字段。

  • filter_str (Optional[str], default: '' ) –

    Filter expression string. Defaults to empty string

  • kwargs

    Other search parameters

Returns:

  • List[dict]: 每项包含 'uid' 及相似度 'score'。
Source code in lazyllm/tools/rag/store/vector/milvus_store.py
    @override
    def search(self, collection_name: str, query_embedding: Union[dict, List[float]], topk: int,
               filters: Optional[Dict[str, Union[List, set]]] = None, embed_key: Optional[str] = None,
               filter_str: Optional[str] = '', **kwargs) -> List[dict]:
        """
执行向量相似度检索,并可按元数据过滤。

Args:
    collection_name (str): 待搜索集合。
    query_embedding (List[float]): 查询向量。
    topk (int): 返回邻近数量。
    filters (Optional[Dict[str, Union[List, Set]]]): 元数据过滤映射。
    embed_key (str): 使用的嵌入字段。
    filter_str (Optional[str], optional): Filter expression string. Defaults to empty string
    kwargs: Other search parameters

**Returns:**

- List[dict]: 每项包含 'uid' 及相似度 'score'。
"""
        with self._client_context() as client:
            if not embed_key or embed_key not in self._embed_datatypes:
                raise ValueError(f'[Milvus Store - search] Not supported or None `embed_key`: {embed_key}')
            if not client.has_collection(collection_name):
                return []
            client.load_collection(collection_name)

            res = []
            filter_expr = self._construct_filter_expr(filters) if filters else ''
            if filter_str:
                filter_expr = f'{filter_expr} and {filter_str}' if filter_expr else filter_str

            results = client.search(collection_name=collection_name, data=[query_embedding], limit=topk,
                                    anns_field=self._gen_embed_key(embed_key),
                                    filter=filter_expr)
            if len(results) != 1:
                raise ValueError(f'number of results [{len(results)}] != expected [1]')
            for result in results[0]:
                score = result.get('distance', 0)
                uid = result.get('id', result.get(self._primary_key, ''))
                if not uid:
                    continue
                res.append({'uid': uid, 'score': score})
        return res

upsert(collection_name, data)

批量写入或更新切片数据到 Milvus 集合。

Parameters:

  • collection_name (str) –

    集合名称,通常为 "group_embedKey" 格式。

  • data (List[dict]) –

    切片数据列表。

Returns:

  • bool: 操作成功返回 True,否则 False。
Source code in lazyllm/tools/rag/store/vector/milvus_store.py
    @override
    def upsert(self, collection_name: str, data: List[dict]) -> bool:  # noqa: C901
        """
批量写入或更新切片数据到 Milvus 集合。

Args:
    collection_name (str): 集合名称,通常为 "group_embedKey" 格式。
    data (List[dict]): 切片数据列表。

**Returns:**

- bool: 操作成功返回 True,否则 False。
"""
        try:
            if not data: return True
            with self._client_context() as client:
                collection_exists = client.has_collection(collection_name)
                required_embed_keys = (
                    self._collection_embed_keys(client, collection_name)
                    if collection_exists else self._data_embed_keys(data)
                )
                if not required_embed_keys:
                    return True

                # Only require embeddings that belong to this collection. Different node groups may use
                # different embedding models, e.g. text groups use BGE while image groups use SigLIP.
                valid_data = [d for d in data if self._row_has_valid_embedding(d, required_embed_keys)]
                dropped = len(data) - len(valid_data)
                if dropped:
                    LOG.warning(f'[Milvus Store - upsert] Dropping {dropped} rows with missing/empty embedding for '
                                f'collection {collection_name}, required embeddings: {sorted(required_embed_keys)}.')
                data = valid_data
                if not data:
                    return True

                if not collection_exists:
                    with self._ddl_lock:
                        if not client.has_collection(collection_name):
                            self._resolve_missing_embed_specs(required_embed_keys)
                            if self._index_kwargs is not None:
                                self._index_kwargs = self.validate_milvus_embed_keys(self._index_kwargs)
                            embed_kwargs = {}
                            for embed_key in required_embed_keys:
                                assert self._embed_datatypes.get(embed_key), \
                                    f'cannot find embedding params for embed [{embed_key}]'
                                if embed_key not in embed_kwargs:
                                    embed_kwargs[embed_key] = {
                                        'dtype': self._type2milvus[self._embed_datatypes[embed_key]]
                                    }
                                if self._embed_dims.get(embed_key):
                                    embed_kwargs[embed_key]['dim'] = self._embed_dims[embed_key]
                            self._create_collection(client, collection_name, embed_kwargs)

                for i in range(0, len(data), MILVUS_UPSERT_BATCH_SIZE):
                    client.upsert(collection_name=collection_name,
                                  data=[self._serialize_data(d, required_embed_keys)
                                        for d in data[i:i + MILVUS_UPSERT_BATCH_SIZE]])
            return True
        except Exception as e:
            LOG.error(f'[Milvus Store - upsert] error: {e}')
            LOG.error(traceback.format_exc())
            return False

validate_milvus_embed_keys(index_kwargs)

Validate and preprocess the index_kwargs of milvus store_conf: 1. Auto fill the only one missing embed_key into the configuration without embed_key; 2. The embed_key in self._embed must be a subset of the embed_key in store_conf; 3. store_conf can contain additional embed_key; 4. Duplicate embed_key is forbidden; 5. If multiple embed_key are missing, raise an error.

Source code in lazyllm/tools/rag/store/vector/milvus_store.py
def validate_milvus_embed_keys(self, index_kwargs: Optional[Union[List, Dict]]):  # noqa: C901
    """
    Validate and preprocess the index_kwargs of milvus store_conf:
    1. Auto fill the only one missing embed_key into the configuration without embed_key;
    2. The embed_key in self._embed must be a subset of the embed_key in store_conf;
    3. store_conf can contain additional embed_key;
    4. Duplicate embed_key is forbidden;
    5. If multiple embed_key are missing, raise an error.
    """
    if not isinstance(index_kwargs, (list, dict)):
        raise TypeError(f'[Milvus Store] index_kwargs must be a list or dict, but got {type(index_kwargs)}')

    embed_keys = list(self._embed_datatypes.keys())
    if not embed_keys:
        raise ValueError('self._embed is empty, cannot build index configuration')

    normalized_index_kwargs = []
    no_embedkey_entries = []
    seen = set()
    if isinstance(index_kwargs, dict):
        index_kwargs = [index_kwargs]
    for i, idx_conf in enumerate(index_kwargs):
        if not isinstance(idx_conf, dict):
            raise TypeError(f'index_kwargs position {i} must be a dictionary, but got {type(idx_conf)}')

        embed_key = idx_conf.get('embed_key')
        if embed_key:
            if embed_key in seen:
                raise ValueError(f'duplicate embed_key {embed_key} in index_kwargs position {i}')
            seen.add(embed_key)
        else:
            no_embedkey_entries.append((i, idx_conf))

        normalized_index_kwargs.append(idx_conf)

    store_embed_keys = seen
    missing_keys = set(embed_keys) - store_embed_keys

    if len(missing_keys) > 1:
        raise ValueError(
            f'[Milvus Store] store_conf is missing the following embed_key: {missing_keys} '
            f'(only supports auto filling one missing item)'
        )
    elif len(missing_keys) == 1:
        missing_key = next(iter(missing_keys))

        if len(no_embedkey_entries) == 1:
            idx = no_embedkey_entries[0][1]
            idx['embed_key'] = missing_key
        elif len(no_embedkey_entries) == 0:
            if self._embed_datatypes.get(missing_key) == DataType.FLOAT_VECTOR:
                normalized_index_kwargs.append({
                    'embed_key': missing_key,
                    'index_type': 'FLAT',
                    'metric_type': 'COSINE'
                })
            else:
                normalized_index_kwargs.append({
                    'embed_key': missing_key,
                    'index_type': 'SPARSE_INVERTED_INDEX',
                    'metric_type': 'L2'
                })
        else:
            raise ValueError(
                f'[Milvus Store] Found multiple entries without embed_key, cannot determine '
                f'which one to fill. Missing embed_keys: {missing_keys}'
            )

    return normalized_index_kwargs

lazyllm.tools.rag.store.hybrid.hybrid_store.HybridStore

Bases: LazyLLMStoreBase

混合存储类,结合了分段存储和向量存储的功能。

Parameters:

  • segment_store (LazyLLMStoreBase) –

    分段存储实例,用于存储文档的原始内容。

  • vector_store (LazyLLMStoreBase) –

    向量存储实例,用于存储文档的向量表示。

Source code in lazyllm/tools/rag/store/hybrid/hybrid_store.py
class HybridStore(LazyLLMStoreBase):
    """混合存储类,结合了分段存储和向量存储的功能。

Args:
    segment_store (LazyLLMStoreBase): 分段存储实例,用于存储文档的原始内容。
    vector_store (LazyLLMStoreBase): 向量存储实例,用于存储文档的向量表示。
"""
    capability = StoreCapability.ALL
    need_embedding = True
    supports_index_registration = False

    def __init__(self, segment_store: LazyLLMStoreBase, vector_store: LazyLLMStoreBase):
        self.segment_store: LazyLLMStoreBase = segment_store
        self.vector_store: LazyLLMStoreBase = vector_store

    @property
    def dir(self):
        return self.segment_store.dir

    @override
    def seg_connect(self, *args, **kwargs):
        """连接底层分段存储。

与 ``connect()`` 不同,``seg_connect()`` 仅初始化 ``segment_store``,不会触发向量存储的连接逻辑。
``DocumentStore`` 在 ``_seg_init()`` 阶段调用此方法,通常传入 ``global_metadata_desc`` 以注册全局元数据字段。

Args:
    *args: 传递给 ``segment_store.connect()`` 的位置参数。
    **kwargs: 传递给 ``segment_store.connect()`` 的关键字参数,常见包括
        ``global_metadata_desc``(全局元数据 schema 描述)。

**Returns:**

- None
"""
        self.segment_store.connect(*args, **kwargs)

    @override
    def vec_connect(self, *args, **kwargs):
        """连接底层向量存储。

与 ``connect()`` 不同,``vec_connect()`` 仅初始化 ``vector_store``,不会触发分段存储的连接逻辑。
``DocumentStore`` 在 ``_vec_init()`` 阶段调用此方法,通常传入 ``embed_dims``、``embed_datatypes``、
``collections`` 和 ``global_metadata_desc``,以便向量后端创建或校验 collection schema。

Args:
    *args: 传递给 ``vector_store.connect()`` 的位置参数。
    **kwargs: 传递给 ``vector_store.connect()`` 的关键字参数,常见包括
        ``embed_dims``(各 embed key 的向量维度)、
        ``embed_datatypes``(各 embed key 的数据类型)、
        ``global_metadata_desc``(全局元数据 schema 描述)、
        ``collections``(需要预创建的 collection 名称列表)。

**Returns:**

- None
"""
        self.vector_store.connect(*args, **kwargs)

    @override
    def connect(self, *args, **kwargs):
        """连接到底层的分段存储和向量存储。

Args:
    *args: 传递给存储连接方法的位置参数。
    **kwargs: 传递给存储连接方法的关键字参数。
"""
        self.seg_connect(*args, **kwargs)
        self.vec_connect(*args, **kwargs)

    @override
    def try_read_dims_from_schema(self, collections: List[str]):
        return self.vector_store.try_read_dims_from_schema(collections)

    @override
    def upsert(self, collection_name: str, data: List[dict]) -> bool:
        """向存储中插入或更新数据。

Args:
    collection_name (str): 集合名称。
    data (List[dict]): 要插入或更新的数据列表,每个数据项都是一个字典。

**Returns:**

- bool: 操作成功返回True,否则返回False。
"""
        segments = [{k: v for k, v in segment.items() if k != 'embedding'} for segment in data]
        return self.segment_store.upsert(collection_name=collection_name, data=segments) and \
            self.vector_store.upsert(collection_name=collection_name, data=data)

    @override
    def delete(self, collection_name: str, criteria: Optional[dict] = None, **kwargs) -> bool:
        """从存储中删除数据。

Args:
    collection_name (str): 集合名称。
    criteria (Optional[dict]): 删除条件,默认为None。
    **kwargs: 其他参数。

**Returns:**

- bool: 操作成功返回True,否则返回False。
"""
        return self.segment_store.delete(collection_name=collection_name, criteria=criteria, **kwargs) and \
            self.vector_store.delete(collection_name=collection_name, criteria=criteria, **kwargs)

    def drop_collection(self, collection_name: str) -> bool:
        """删除指定集合,同时从分段存储和向量存储中移除对应数据。

Args:
    collection_name (str): 要删除的集合名称。

**Returns:**

- bool: 若两个底层存储均成功删除则返回 ``True``,任意一个失败则返回 ``False``。
"""
        ok = True
        for store in (self.segment_store, self.vector_store):
            if hasattr(store, 'drop_collection'):
                result = store.drop_collection(collection_name)
                ok = bool(result) and ok
            else:
                lazyllm.LOG.warning(
                    f'[HybridStore] {type(store).__name__} does not implement '
                    f'drop_collection; skipping for collection {collection_name!r}'
                )
        return ok

    @override
    def get(self, collection_name: str, criteria: Optional[dict] = None, **kwargs) -> List[dict]:
        """从存储中获取数据。

Args:
    collection_name (str): 集合名称。
    criteria (Optional[dict]): 查询条件,默认为None。
    **kwargs: 其他参数。

**Returns:**

- List[dict]: 返回符合条件的数据列表。

Raises:
    ValueError: 当向量存储中的uid在分段存储中找不到时抛出。
"""
        res_segments = self.segment_store.get(collection_name=collection_name, criteria=criteria, **kwargs)
        total = None
        if isinstance(res_segments, tuple):
            res_segments, total = res_segments
        if not res_segments:
            return ([], total or 0) if total is not None else []
        uids = [item.get('uid') for item in res_segments]
        res_vectors = self.vector_store.get(collection_name=collection_name, criteria={'uid': uids})

        data = {}
        for item in res_segments:
            data[item.get('uid')] = item
        for item in res_vectors:
            if item.get('uid') in data:
                data[item.get('uid')]['embedding'] = item.get('embedding')
            else:
                raise ValueError(f'[HybridStore - get] uid {item["uid"]} in vector store'
                                 ' but not found in segment store')
        ordered = [data[item.get('uid')] for item in res_segments if item.get('uid') in data]
        return (ordered, total) if total is not None else ordered

    @override
    def collection_exists(self, collection_name: str) -> bool:
        return self.vector_store.collection_exists(collection_name)

    @override
    def search(self, collection_name: str, query: str, query_embedding: Optional[Union[dict, List[float]]] = None,
               topk: int = 10, filters: Optional[Dict[str, Union[str, int, List, Set]]] = None,
               embed_key: Optional[str] = None, **kwargs) -> List[dict]:
        """在存储中搜索数据。

Args:
    collection_name (str): 集合名称。
    query (str): 搜索查询字符串。
    query_embedding (Optional[Union[dict, List[float]]]): 查询的向量表示,默认为None。
    topk (int): 返回的最大结果数量,默认为10。
    filters (Optional[Dict[str, Union[str, int, List, Set]]]): 过滤条件,默认为None。
    embed_key (Optional[str]): 嵌入向量的键名,默认为None。
    **kwargs: 其他参数。

**Returns:**

- List[dict]: 返回搜索结果列表。
"""
        if embed_key:
            # vector store only give uid and score
            res = self.vector_store.search(collection_name=collection_name, query=query, query_embedding=query_embedding,
                                           topk=topk, filters=filters, embed_key=embed_key, **kwargs)
            if not res: return []
            uid2score = {item['uid']: item['score'] for item in res}
            uids = list(uid2score.keys())
            segments = self.segment_store.get(collection_name=collection_name, criteria={'uid': uids})
            uid2segment = {}
            for segment in segments:
                segment['score'] = uid2score.get(segment['uid'], 0)
                uid2segment[segment.get('uid')] = segment
            ordered = [uid2segment[uid] for uid in uids if uid in uid2segment]
            return ordered
        else:
            res = self.segment_store.search(collection_name=collection_name, query=query,
                                            topk=topk, filters=filters, **kwargs)
            return res

connect(*args, **kwargs)

连接到底层的分段存储和向量存储。

Parameters:

  • *args

    传递给存储连接方法的位置参数。

  • **kwargs

    传递给存储连接方法的关键字参数。

Source code in lazyllm/tools/rag/store/hybrid/hybrid_store.py
    @override
    def connect(self, *args, **kwargs):
        """连接到底层的分段存储和向量存储。

Args:
    *args: 传递给存储连接方法的位置参数。
    **kwargs: 传递给存储连接方法的关键字参数。
"""
        self.seg_connect(*args, **kwargs)
        self.vec_connect(*args, **kwargs)

delete(collection_name, criteria=None, **kwargs)

从存储中删除数据。

Parameters:

  • collection_name (str) –

    集合名称。

  • criteria (Optional[dict], default: None ) –

    删除条件,默认为None。

  • **kwargs

    其他参数。

Returns:

  • bool: 操作成功返回True,否则返回False。
Source code in lazyllm/tools/rag/store/hybrid/hybrid_store.py
    @override
    def delete(self, collection_name: str, criteria: Optional[dict] = None, **kwargs) -> bool:
        """从存储中删除数据。

Args:
    collection_name (str): 集合名称。
    criteria (Optional[dict]): 删除条件,默认为None。
    **kwargs: 其他参数。

**Returns:**

- bool: 操作成功返回True,否则返回False。
"""
        return self.segment_store.delete(collection_name=collection_name, criteria=criteria, **kwargs) and \
            self.vector_store.delete(collection_name=collection_name, criteria=criteria, **kwargs)

drop_collection(collection_name)

删除指定集合,同时从分段存储和向量存储中移除对应数据。

Parameters:

  • collection_name (str) –

    要删除的集合名称。

Returns:

  • bool: 若两个底层存储均成功删除则返回 True,任意一个失败则返回 False
Source code in lazyllm/tools/rag/store/hybrid/hybrid_store.py
    def drop_collection(self, collection_name: str) -> bool:
        """删除指定集合,同时从分段存储和向量存储中移除对应数据。

Args:
    collection_name (str): 要删除的集合名称。

**Returns:**

- bool: 若两个底层存储均成功删除则返回 ``True``,任意一个失败则返回 ``False``。
"""
        ok = True
        for store in (self.segment_store, self.vector_store):
            if hasattr(store, 'drop_collection'):
                result = store.drop_collection(collection_name)
                ok = bool(result) and ok
            else:
                lazyllm.LOG.warning(
                    f'[HybridStore] {type(store).__name__} does not implement '
                    f'drop_collection; skipping for collection {collection_name!r}'
                )
        return ok

get(collection_name, criteria=None, **kwargs)

从存储中获取数据。

Parameters:

  • collection_name (str) –

    集合名称。

  • criteria (Optional[dict], default: None ) –

    查询条件,默认为None。

  • **kwargs

    其他参数。

Returns:

  • List[dict]: 返回符合条件的数据列表。

Raises:

  • ValueError

    当向量存储中的uid在分段存储中找不到时抛出。

Source code in lazyllm/tools/rag/store/hybrid/hybrid_store.py
    @override
    def get(self, collection_name: str, criteria: Optional[dict] = None, **kwargs) -> List[dict]:
        """从存储中获取数据。

Args:
    collection_name (str): 集合名称。
    criteria (Optional[dict]): 查询条件,默认为None。
    **kwargs: 其他参数。

**Returns:**

- List[dict]: 返回符合条件的数据列表。

Raises:
    ValueError: 当向量存储中的uid在分段存储中找不到时抛出。
"""
        res_segments = self.segment_store.get(collection_name=collection_name, criteria=criteria, **kwargs)
        total = None
        if isinstance(res_segments, tuple):
            res_segments, total = res_segments
        if not res_segments:
            return ([], total or 0) if total is not None else []
        uids = [item.get('uid') for item in res_segments]
        res_vectors = self.vector_store.get(collection_name=collection_name, criteria={'uid': uids})

        data = {}
        for item in res_segments:
            data[item.get('uid')] = item
        for item in res_vectors:
            if item.get('uid') in data:
                data[item.get('uid')]['embedding'] = item.get('embedding')
            else:
                raise ValueError(f'[HybridStore - get] uid {item["uid"]} in vector store'
                                 ' but not found in segment store')
        ordered = [data[item.get('uid')] for item in res_segments if item.get('uid') in data]
        return (ordered, total) if total is not None else ordered

search(collection_name, query, query_embedding=None, topk=10, filters=None, embed_key=None, **kwargs)

在存储中搜索数据。

Parameters:

  • collection_name (str) –

    集合名称。

  • query (str) –

    搜索查询字符串。

  • query_embedding (Optional[Union[dict, List[float]]], default: None ) –

    查询的向量表示,默认为None。

  • topk (int, default: 10 ) –

    返回的最大结果数量,默认为10。

  • filters (Optional[Dict[str, Union[str, int, List, Set]]], default: None ) –

    过滤条件,默认为None。

  • embed_key (Optional[str], default: None ) –

    嵌入向量的键名,默认为None。

  • **kwargs

    其他参数。

Returns:

  • List[dict]: 返回搜索结果列表。
Source code in lazyllm/tools/rag/store/hybrid/hybrid_store.py
    @override
    def search(self, collection_name: str, query: str, query_embedding: Optional[Union[dict, List[float]]] = None,
               topk: int = 10, filters: Optional[Dict[str, Union[str, int, List, Set]]] = None,
               embed_key: Optional[str] = None, **kwargs) -> List[dict]:
        """在存储中搜索数据。

Args:
    collection_name (str): 集合名称。
    query (str): 搜索查询字符串。
    query_embedding (Optional[Union[dict, List[float]]]): 查询的向量表示,默认为None。
    topk (int): 返回的最大结果数量,默认为10。
    filters (Optional[Dict[str, Union[str, int, List, Set]]]): 过滤条件,默认为None。
    embed_key (Optional[str]): 嵌入向量的键名,默认为None。
    **kwargs: 其他参数。

**Returns:**

- List[dict]: 返回搜索结果列表。
"""
        if embed_key:
            # vector store only give uid and score
            res = self.vector_store.search(collection_name=collection_name, query=query, query_embedding=query_embedding,
                                           topk=topk, filters=filters, embed_key=embed_key, **kwargs)
            if not res: return []
            uid2score = {item['uid']: item['score'] for item in res}
            uids = list(uid2score.keys())
            segments = self.segment_store.get(collection_name=collection_name, criteria={'uid': uids})
            uid2segment = {}
            for segment in segments:
                segment['score'] = uid2score.get(segment['uid'], 0)
                uid2segment[segment.get('uid')] = segment
            ordered = [uid2segment[uid] for uid in uids if uid in uid2segment]
            return ordered
        else:
            res = self.segment_store.search(collection_name=collection_name, query=query,
                                            topk=topk, filters=filters, **kwargs)
            return res

seg_connect(*args, **kwargs)

连接底层分段存储。

connect() 不同,seg_connect() 仅初始化 segment_store,不会触发向量存储的连接逻辑。 DocumentStore_seg_init() 阶段调用此方法,通常传入 global_metadata_desc 以注册全局元数据字段。

Parameters:

  • *args

    传递给 segment_store.connect() 的位置参数。

  • **kwargs

    传递给 segment_store.connect() 的关键字参数,常见包括 global_metadata_desc(全局元数据 schema 描述)。

Returns:

  • None
Source code in lazyllm/tools/rag/store/hybrid/hybrid_store.py
    @override
    def seg_connect(self, *args, **kwargs):
        """连接底层分段存储。

与 ``connect()`` 不同,``seg_connect()`` 仅初始化 ``segment_store``,不会触发向量存储的连接逻辑。
``DocumentStore`` 在 ``_seg_init()`` 阶段调用此方法,通常传入 ``global_metadata_desc`` 以注册全局元数据字段。

Args:
    *args: 传递给 ``segment_store.connect()`` 的位置参数。
    **kwargs: 传递给 ``segment_store.connect()`` 的关键字参数,常见包括
        ``global_metadata_desc``(全局元数据 schema 描述)。

**Returns:**

- None
"""
        self.segment_store.connect(*args, **kwargs)

upsert(collection_name, data)

向存储中插入或更新数据。

Parameters:

  • collection_name (str) –

    集合名称。

  • data (List[dict]) –

    要插入或更新的数据列表,每个数据项都是一个字典。

Returns:

  • bool: 操作成功返回True,否则返回False。
Source code in lazyllm/tools/rag/store/hybrid/hybrid_store.py
    @override
    def upsert(self, collection_name: str, data: List[dict]) -> bool:
        """向存储中插入或更新数据。

Args:
    collection_name (str): 集合名称。
    data (List[dict]): 要插入或更新的数据列表,每个数据项都是一个字典。

**Returns:**

- bool: 操作成功返回True,否则返回False。
"""
        segments = [{k: v for k, v in segment.items() if k != 'embedding'} for segment in data]
        return self.segment_store.upsert(collection_name=collection_name, data=segments) and \
            self.vector_store.upsert(collection_name=collection_name, data=data)

vec_connect(*args, **kwargs)

连接底层向量存储。

connect() 不同,vec_connect() 仅初始化 vector_store,不会触发分段存储的连接逻辑。 DocumentStore_vec_init() 阶段调用此方法,通常传入 embed_dimsembed_datatypescollectionsglobal_metadata_desc,以便向量后端创建或校验 collection schema。

Parameters:

  • *args

    传递给 vector_store.connect() 的位置参数。

  • **kwargs

    传递给 vector_store.connect() 的关键字参数,常见包括 embed_dims(各 embed key 的向量维度)、 embed_datatypes(各 embed key 的数据类型)、 global_metadata_desc(全局元数据 schema 描述)、 collections(需要预创建的 collection 名称列表)。

Returns:

  • None
Source code in lazyllm/tools/rag/store/hybrid/hybrid_store.py
    @override
    def vec_connect(self, *args, **kwargs):
        """连接底层向量存储。

与 ``connect()`` 不同,``vec_connect()`` 仅初始化 ``vector_store``,不会触发分段存储的连接逻辑。
``DocumentStore`` 在 ``_vec_init()`` 阶段调用此方法,通常传入 ``embed_dims``、``embed_datatypes``、
``collections`` 和 ``global_metadata_desc``,以便向量后端创建或校验 collection schema。

Args:
    *args: 传递给 ``vector_store.connect()`` 的位置参数。
    **kwargs: 传递给 ``vector_store.connect()`` 的关键字参数,常见包括
        ``embed_dims``(各 embed key 的向量维度)、
        ``embed_datatypes``(各 embed key 的数据类型)、
        ``global_metadata_desc``(全局元数据 schema 描述)、
        ``collections``(需要预创建的 collection 名称列表)。

**Returns:**

- None
"""
        self.vector_store.connect(*args, **kwargs)

lazyllm.tools.rag.store.hybrid.oceanbase_store.OceanBaseStore

Bases: EmbedResolveMixin, LazyLLMStoreBase

OceanBase 存储类,用于存储和检索文档节点。

Parameters:

  • uri (str, default: '127.0.0.1:2881' ) –

    OceanBase 数据库的 URI。

  • user (str) –

    OceanBase 数据库的用户名。

  • password (str) –

    OceanBase 数据库的密码。

  • db_name (str, default: 'test' ) –

    OceanBase 数据库的名称。

  • drop_old (bool) –

    是否删除旧的表。

  • index_kwargs (List[dict], default: None ) –

    索引配置列表。

  • client_kwargs (Dict, default: None ) –

    客户端配置字典。

  • max_pool_size (int) –

    最大连接池大小。

  • normalize (bool) –

    是否规范化数据。

  • enable_fulltext_index (bool) –

    是否启用全文索引。

Source code in lazyllm/tools/rag/store/hybrid/oceanbase_store.py
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
789
790
791
792
793
794
795
796
797
798
799
800
801
802
803
804
805
806
807
808
809
810
811
812
813
814
815
816
817
818
819
820
821
822
823
824
825
826
827
828
829
830
831
832
833
834
835
836
837
838
839
840
841
842
843
844
845
846
847
848
849
850
851
852
853
854
855
856
857
858
859
860
861
862
863
864
865
866
867
868
869
870
871
872
873
874
875
876
877
878
879
880
881
882
883
884
885
886
887
888
889
890
891
892
893
894
895
896
897
898
899
900
901
902
903
904
905
906
907
908
909
910
911
912
913
914
915
916
917
918
919
920
921
922
923
924
925
926
927
928
929
930
931
932
933
934
935
936
937
938
939
940
941
942
943
944
945
946
947
948
949
950
951
952
953
954
955
956
957
958
959
960
961
962
963
964
965
966
967
968
969
970
971
972
973
974
975
976
977
978
979
980
981
982
983
984
985
986
987
988
989
class OceanBaseStore(EmbedResolveMixin, LazyLLMStoreBase):
    """OceanBase 存储类,用于存储和检索文档节点。

Args:
    uri (str): OceanBase 数据库的 URI。
    user (str): OceanBase 数据库的用户名。
    password (str): OceanBase 数据库的密码。
    db_name (str): OceanBase 数据库的名称。
    drop_old (bool): 是否删除旧的表。
    index_kwargs (List[dict]): 索引配置列表。
    client_kwargs (Dict): 客户端配置字典。
    max_pool_size (int): 最大连接池大小。
    normalize (bool): 是否规范化数据。
    enable_fulltext_index (bool): 是否启用全文索引。
"""
    capability = StoreCapability.ALL
    need_embedding = True
    supports_index_registration = True

    def __init__(self, uri: str = '127.0.0.1:2881', db_name: str = 'test',
                 index_kwargs: Optional[Union[Dict, List]] = None, client_kwargs: Optional[Dict] = None):
        self._uri = uri
        self._db_name = db_name
        self._index_kwargs = index_kwargs or {}
        self._client_kwargs = client_kwargs or {}
        self._ddl_lock = threading.Lock()
        self._embed_datatypes: Union[Dict[str, DataType], Dict[str, Dict]] = {}
        self._global_metadata_desc: Dict[str, GlobalMetadataDesc] = {}
        self._primary_key = 'uid'
        self._hnsw_ef_search = {}

    @contextmanager
    def _client_context(self) -> 'pyobvector.ObVecClient':
        c = self._client_pool.acquire()
        try:
            try:
                c.perform_raw_text_sql('SET SESSION ob_query_timeout = 300000000')
            except Exception as e:
                LOG.warning(f'[OceanBaseStore] Failed to set query timeout in context: {e}')
            yield c
        finally:
            self._client_pool.release(c)

    def _new_client(self):
        kwargs = dict(self._client_kwargs)
        try:
            c = pyobvector.ObVecClient(
                uri=self._uri, db_name=self._db_name,
                user=self._user, password=self._password, **kwargs
            )

            try:
                c.perform_raw_text_sql('SET SESSION ob_query_timeout = 300000000')
                result = c.perform_raw_text_sql("SHOW VARIABLES LIKE 'ob_query_timeout'")
                if result:
                    assert result.fetchone() is not None
            except Exception as e:
                LOG.warning(f'[OceanBaseStore] Failed to set/verify query timeout: {e}')

            LOG.info(f'[OceanBaseStore] Successfully connected to {self._uri}')
            return c
        except Exception as e:
            LOG.error(f'[OceanBaseStore - _new_client] error: {e}')
            raise e

    @override
    def connect(self, embed_dims: Optional[Dict[str, int]] = None,
                embed_datatypes: Optional[Dict[str, DataType]] = None,
                embed: Optional[Dict[str, Callable]] = None,
                global_metadata_desc: Optional[Dict[str, GlobalMetadataDesc]] = None, **kwargs):
        """连接到底层的 OceanBase 数据库。

Args:
    embed_dims (Dict[str, int]): 嵌入维度字典。
    embed_datatypes (Dict[str, DataType]): 嵌入数据类型字典。
    global_metadata_desc (Dict[str, GlobalMetadataDesc]): 全局元数据描述字典。
    **kwargs: 其他参数。
"""
        self._embed_dims = embed_dims or {}
        self._embed_datatypes = embed_datatypes or {}
        self._embed = embed or {}
        self._global_metadata_desc = global_metadata_desc or {}
        self._set_constants()

        # Extract connection parameters from client_kwargs
        self._user = self._parse_user(self._client_kwargs.pop('user', 'root@test'))
        self._password = self._client_kwargs.pop('password', '')
        self._normalize = self._client_kwargs.pop('normalize', False)
        self._enable_fulltext_index = self._client_kwargs.pop('enable_fulltext_index', False)
        max_pool_size = int(self._client_kwargs.pop('max_pool_size', 8))

        self._ensure_database()
        self._client_pool = _ClientPool(self._new_client, max_size=max_pool_size)
        LOG.info('[OceanBaseStore] init success!')

    def upsert(self, collection_name: str, data: List[dict], range_part: Optional['pyobvector.RangeListPartInfo'] = None, **kwargs) -> bool:  # noqa: C901 E501
        """向存储中插入或更新数据。

Args:
    collection_name (str): 集合名称。
    data (List[dict]): 要插入或更新的数据列表,每个数据项都是一个字典。
    range_part (Optional[RangeListPartInfo]): 范围分区信息,暂未实现分区功能。
    **kwargs: 其他参数。

**Returns:**

- bool: 操作成功返回True,否则返回False。
"""
        try:
            if not data:
                return True

            if not collection_name or not isinstance(collection_name, str):
                LOG.error('[OceanBaseStore - upsert] Invalid collection_name')
                return False

            all_embed_keys = set()
            for item in data:
                if 'embedding' in item and isinstance(item['embedding'], dict):
                    all_embed_keys.update(item['embedding'].keys())

            with self._client_context() as client:
                with self._ddl_lock:
                    if not client.check_table_exists(collection_name):
                        self._resolve_missing_embed_specs(all_embed_keys)
                        embed_kwargs = {}
                        if all_embed_keys:
                            for embed_key in all_embed_keys:
                                if not self._embed_datatypes.get(embed_key):
                                    LOG.error(f'[OceanBaseStore - upsert] Cannot find embedding for embed [{embed_key}]')
                                    return False

                                if embed_key not in embed_kwargs:
                                    embed_kwargs[embed_key] = {
                                        'dtype': self._type2oceanbase[self._embed_datatypes[embed_key]]
                                    }
                                if self._embed_dims.get(embed_key):
                                    embed_kwargs[embed_key]['dim'] = self._embed_dims[embed_key]

                        self._create_table_and_index(client, collection_name, embed_kwargs, range_part)

                total_inserted = 0
                failed_batches = []

                serialized_data = [self._serialize_data(d) for d in data]

                for i in range(0, len(serialized_data), OB_UPSERT_BATCH_SIZE):
                    batch_num = i // OB_UPSERT_BATCH_SIZE + 1
                    try:
                        if i == 0 or batch_num % 10 == 0:
                            try:
                                client.perform_raw_text_sql('SET SESSION ob_query_timeout = 300000000')
                            except Exception as timeout_err:
                                LOG.warning(f'[OceanBaseStore - upsert] Failed to set '
                                            f'query timeout for batch {batch_num}: {timeout_err}')

                        batch_data = serialized_data[i:i + OB_UPSERT_BATCH_SIZE]
                        client.upsert(table_name=collection_name, data=batch_data)
                        total_inserted += len(batch_data)

                    except Exception as batch_err:
                        LOG.error(f'[OceanBaseStore - upsert] Failed to insert batch {batch_num}: {batch_err}')
                        LOG.error(f'[OceanBaseStore - upsert] Error details: {traceback.format_exc()}')
                        failed_batches.append(batch_num)
                        continue

                if failed_batches:
                    LOG.warning(f'[OceanBaseStore - upsert] Failed batches: {failed_batches}')
                    if total_inserted == 0:
                        return False
            return True
        except Exception as e:
            LOG.error(f'[OceanBaseStore - upsert] Unexpected error: {e}')
            LOG.error(traceback.format_exc())
            return False

    @override
    def delete(self, collection_name: str, criteria: Optional[dict] = None, **kwargs) -> bool:
        """从存储中删除数据。

Args:
    collection_name (str): 集合名称。
    criteria (Optional[dict]): 删除条件,默认为None。
    **kwargs: 其他参数。

**Returns:**

- bool: 操作成功返回True,否则返回False。
"""
        try:
            with self._client_context() as client:
                if not client.check_table_exists(collection_name):
                    return True

                if not criteria:
                    with self._ddl_lock:
                        client.drop_table_if_exist(table_name=collection_name)
                else:
                    ids, where_clause = self._get_ids_where_clause(criteria)

                    client.delete(
                        table_name=collection_name,
                        ids=ids,
                        where_clause=where_clause,
                        **kwargs
                    )

            return True
        except Exception as e:
            LOG.error(f'[OceanBaseStore - delete] error: {e}')
            LOG.error(traceback.format_exc())
            return False

    @override
    def get(self, collection_name: str, criteria: Optional[dict] = None, **kwargs) -> List[dict]:
        """从存储中获取数据。

Args:
    collection_name (str): 集合名称。
    criteria (Optional[dict]): 查询条件,默认为None。
    **kwargs: 其他参数。

**Returns:**

- List[dict]: 返回符合条件的数据列表。
"""
        try:
            with self._client_context() as client:
                if not client.check_table_exists(collection_name):
                    return []

                ids, where_clause = self._get_ids_where_clause(criteria)

                if ids is None and where_clause is None:
                    return self._get_all_with_pagination(client, collection_name)

                res = client.get(
                    table_name=collection_name,
                    ids=ids,
                    where_clause=where_clause,
                    output_column_name=None,
                    **kwargs
                )

                if not res:
                    return []
                result = [r._mapping for r in res]

                return [self._deserialize_data(r) for r in result]

        except Exception as e:
            LOG.error(f'[OceanBaseStore - get] error: {e}')
            LOG.error(traceback.format_exc())
            return []

    def _get_all_with_pagination(self, client, collection_name: str) -> List[dict]:
        all_results = []
        offset = 0
        batch_size = DEFAULT_OCEANBASE_PAGINATION_OFFSET

        while True:
            try:
                sql = f'SELECT * FROM {collection_name} LIMIT {batch_size} OFFSET {offset}'
                batch_res = client.perform_raw_text_sql(sql)

                if not batch_res:
                    break

                if hasattr(batch_res, 'fetchall'):
                    rows = batch_res.fetchall()
                    if not rows:
                        break
                    columns = batch_res.keys() if hasattr(batch_res, 'keys') else []
                    batch_results = [dict(zip(columns, row)) for row in rows]
                else:
                    batch_results = [r._mapping if hasattr(r, '_mapping') else dict(r) for r in batch_res]

                if not batch_results:
                    break

                all_results.extend([self._deserialize_data(r) for r in batch_results])

                if len(batch_results) < batch_size:
                    break

                offset += batch_size

            except Exception as e:
                LOG.error(f'[OceanBaseStore - _get_all_with_pagination] error at offset {offset}: {e}')
                LOG.error(traceback.format_exc())
                break

        return all_results

    def search(self, collection_name: str, query: str, query_embedding: Union[dict, List[float]], topk: int, filters: Optional[Dict[str, Union[List, set]]] = None, embed_key: Optional[str] = None, filter_str: Optional[str] = '', **kwargs) -> List[dict]:  # noqa: C901 E501
        """在存储中搜索数据。

Args:
    collection_name (str): 集合名称。
    query_embedding (Union[dict, List[float]]): 查询的向量表示。
    topk (int): 返回的最大结果数量。
    filters (Optional[Dict[str, Union[str, int, List, Set]]]): 过滤条件,默认为None。
    embed_key (Optional[str]): 嵌入向量的键名,默认为None。
    filter_str (Optional[str]): 过滤条件字符串,默认为None。
    **kwargs: 其他参数。

**Returns:**

- List[dict]: 返回搜索结果列表。
"""
        if not query_embedding:
            raise NotImplementedError('Query fulltext search is not supported for now')
        try:
            if not collection_name or not isinstance(collection_name, str):
                LOG.error('[OceanBaseStore - search] Invalid collection_name')
                return []

            with self._client_context() as client:
                if not client.check_table_exists(collection_name):
                    LOG.warning(f'[OceanBaseStore - search] Table {collection_name} does not exist')
                    return []
                if embed_key and embed_key not in self._embed_datatypes:
                    LOG.error(f'[OceanBaseStore - search] embed_key: {embed_key} not exists')
                    return []

                if not embed_key:
                    if self._embed_datatypes:
                        embed_key = next(iter(self._embed_datatypes.keys()))
                        LOG.info(f'[OceanBaseStore - search] No embed_key provided, using default: {embed_key}')
                    else:
                        LOG.error('[OceanBaseStore - search] No embedding datatypes available')
                        return []

                where_clause = None
                if filters or filter_str:
                    filter_parts = []
                    if filters:
                        try:
                            filter_expr = self._construct_filter_expr(filters)
                            if filter_expr:
                                filter_parts.append(filter_expr)
                        except Exception as filter_err:
                            LOG.error(f'[OceanBaseStore - search] Failed to construct filter: {filter_err}')
                            LOG.error(traceback.format_exc())
                            raise RuntimeError(f'Failed to construct filter expression: {filter_err}') from filter_err

                    if filter_str:
                        filter_parts.append(filter_str)

                    if filter_parts:
                        combined_filter = ' and '.join(f'({part})' for part in filter_parts)
                        where_clause = [sqlalchemy.text(combined_filter)]

                if (
                    isinstance(query_embedding, dict)
                    and self._embed_datatypes.get(embed_key) != DataType.SPARSE_FLOAT_VECTOR
                ):
                    vec_data = query_embedding.get(embed_key)
                    if vec_data is None:
                        LOG.error(f'[OceanBaseStore - search] embed_key: {embed_key} not found in query_embedding')
                        return []
                else:
                    vec_data = query_embedding

                if vec_data is None:
                    LOG.error('[OceanBaseStore - search] Vector data is None')
                    return []

                if self._embed_datatypes[embed_key] == DataType.SPARSE_FLOAT_VECTOR:
                    if isinstance(vec_data, dict):
                        vec_data = {int(k) if isinstance(k, str) else k: v for k, v in vec_data.items()}
                    distance_func = pyobvector.inner_product
                else:
                    if self._normalize and isinstance(vec_data, list):
                        vec_data = self._normalize_vector(vec_data)

                    metric_type = self._get_metric_type_for_embed_key(embed_key)
                    distance_func = self._get_distance_function(metric_type)

                search_params = kwargs.get('search_params', {})
                index_type = self._get_index_type_for_embed_key(embed_key)

                if index_type in ['HNSW', 'HNSW_SQ']:
                    ef_search = search_params.get('efSearch', 64)  # Default efSearch
                    if self._hnsw_ef_search.get(embed_key) != ef_search:
                        try:
                            client.set_ob_hnsw_ef_search(ef_search)
                            self._hnsw_ef_search[embed_key] = ef_search
                        except Exception as e:
                            LOG.error(f'[OceanBaseStore - search] Failed to set efSearch: {e}')
                            LOG.error(traceback.format_exc())
                            raise RuntimeError(f'Failed to set efSearch parameter for HNSW index: {e}') from e

                results = client.ann_search(
                    table_name=collection_name,
                    vec_data=vec_data,
                    vec_column_name=self._gen_embed_key(embed_key),
                    distance_func=distance_func,
                    topk=topk,
                    with_dist=True,
                    output_column_names=None,
                    where_clause=where_clause,
                    **kwargs
                )

                res = []
                if not results:
                    LOG.info('[OceanBaseStore - search] No results found')
                    return []

                for row in results:
                    try:
                        if hasattr(row, '_mapping'):
                            row_dict = dict(row._mapping)
                            score = row_dict.pop('distance', 0)
                        elif isinstance(row, dict):
                            row_dict = dict(row)
                            score = row_dict.pop('distance', 0)
                        else:
                            LOG.warning(f'[OceanBaseStore - search] Unsupported row type: {type(row)}')
                            continue

                        doc_data = self._deserialize_data(row_dict)

                        if not doc_data.get(self._primary_key):
                            LOG.warning('[OceanBaseStore - search] Row has no valid uid')
                            continue

                        doc_data['score'] = float(score)
                        res.append(doc_data)

                    except Exception as row_err:
                        LOG.warning(f'[OceanBaseStore - search] Failed to process row: {row_err}')
                        LOG.warning(traceback.format_exc())
                        continue

                LOG.info(f'[OceanBaseStore - search] Returning {len(res)} results')
                return res
        except Exception as e:
            LOG.error(f'[OceanBaseStore - search] Unexpected error: {e}')
            LOG.error(traceback.format_exc())
            return []

    def _create_table_and_index(self, client: 'pyobvector.ObVecClient', collection_name: str, embed_kwargs: Dict, partitions: Optional['pyobvector.RangeListPartInfo'] = None) -> bool:  # noqa: C901 E501
        columns = copy.deepcopy(self._constant_columns)
        indexes = [
            sqlalchemy.Index(f'idx_{collection_name}_parent', 'parent'),
            sqlalchemy.Index(f'idx_{collection_name}_number', 'number'),
        ]

        idx_params = client.prepare_index_params()
        original_index_kwargs = copy.deepcopy(self._index_kwargs)

        index_kwargs_lookup = {}
        fts_idxs = []
        has_explicit_fts = False

        if isinstance(original_index_kwargs, dict):
            original_index_kwargs = [original_index_kwargs]

        for item in original_index_kwargs:
            if not isinstance(item, dict):
                LOG.warning(f'[OceanBaseStore - _create_table_and_index] Invalid index_kwargs item: {item}')
                continue

            embed_key = item.get('embed_key', None)
            if not embed_key:
                has_explicit_fts = True
                field_names = item.get('field_names', ['content'])
                if isinstance(field_names, str):
                    field_names = [field_names]
                index_name = item.get('index_name', f'fts_{field_names[0]}')
                fts_idxs.append(
                    pyobvector.client.fts_index_param.FtsIndexParam(
                        index_name=index_name,
                        field_names=field_names,
                        parser_type=item.get('parser_type', pyobvector.client.fts_index_param.FtsParser.IK),
                    )
                )
                continue

            self._ensure_params_defaults(item)
            index_kwargs_lookup[embed_key] = item.copy()
            index_kwargs_lookup[embed_key].pop('embed_key', None)

        for k, kws in embed_kwargs.items():
            embed_field_name = self._gen_embed_key(k)
            dim = kws.get('dim', None)
            dtype = kws.get('dtype')

            if not dtype:
                LOG.error(f'[OceanBaseStore - _create_table_and_index] No dtype specified for embed_key: {k}')
                raise ValueError(f'No dtype specified for embed_key: {k}')

            if dtype == pyobvector.VECTOR and not dim:
                raise ValueError(f'Embedding `{k}` lacks dim parameter (required for VECTOR type)')

            if dim:
                columns.append(sqlalchemy.Column(embed_field_name, dtype(dim)))
            else:
                columns.append(sqlalchemy.Column(embed_field_name, dtype))

            if k in index_kwargs_lookup:
                index_item = index_kwargs_lookup[k]
                index_type_str = index_item.get('index_type', 'HNSW')

                if index_type_str not in self._oceanbase_supported_vector_index_types:
                    LOG.warning(f'[OceanBaseStore - _create_table_and_index] Unsupported index type: {index_type_str}')
                    continue

                if dtype == pyobvector.VECTOR:
                    idx_params.add_index(
                        field_name=embed_field_name,
                        index_type=self._oceanbase_supported_vector_index_types[index_type_str],
                        metric_type=index_item.get('metric_type', 'l2'),
                        index_name=f'vidx_{k}',
                        params=index_item.get('params', {})
                    )
                else:
                    idx_params.add_index(
                        field_name=embed_field_name,
                        index_type='daat',
                        index_name=f'vidx_{k}',
                        metric_type='inner_product',
                    )
                    LOG.info(f'[OceanBaseStore - _create_table_and_index] Added DAAT index for sparse vector {k}')

            if self._enable_fulltext_index and not has_explicit_fts:
                fts_idxs.append(
                    pyobvector.client.fts_index_param.FtsIndexParam(
                        index_name='fts_content',
                        field_names=['content'],
                        parser_type=pyobvector.client.fts_index_param.FtsParser.IK,
                    )
                )

        try:
            client.create_table_with_index_params(
                table_name=collection_name,
                columns=columns,
                indexes=indexes,
                vidxs=idx_params,
                fts_idxs=fts_idxs if fts_idxs else None,
                partitions=partitions,
            )

            LOG.info(f'[OceanBaseStore - _create_table_and_index] Table {collection_name} created successfully')
            return True

        except Exception as e:
            LOG.error(f'[OceanBaseStore - _create_table_and_index] Failed to create table {collection_name}: {e}')
            LOG.error(traceback.format_exc())
            raise e

    def _get_metric_type_for_embed_key(self, embed_key: str) -> str:

        if isinstance(self._index_kwargs, dict):
            index_kwargs_list = [self._index_kwargs]
        else:
            index_kwargs_list = self._index_kwargs or []

        for index_kwarg in index_kwargs_list:
            if index_kwarg.get('embed_key') == embed_key:
                return index_kwarg.get('metric_type', 'l2')
        return 'l2'

    def _get_distance_function(self, metric_type: str):
        metric_type = metric_type.lower()
        if metric_type == 'inner_product':
            return pyobvector.inner_product
        elif metric_type == 'l2':
            return pyobvector.l2_distance
        elif metric_type == 'cosine':
            return pyobvector.cosine_distance
        else:
            raise ValueError(f'Unsupported metric type: {metric_type}')

    def _get_index_type_for_embed_key(self, embed_key: str) -> str:
        if isinstance(self._index_kwargs, dict):
            index_kwargs_list = [self._index_kwargs]
        else:
            index_kwargs_list = self._index_kwargs or []

        for index_kwarg in index_kwargs_list:
            if index_kwarg.get('embed_key') == embed_key:
                return index_kwarg.get('index_type', 'HNSW').upper()
        return 'HNSW'

    def _normalize_vector(self, vector: List[float]) -> List[float]:
        try:
            arr = np.array(vector)
            norm = np.linalg.norm(arr)
            if norm > 0:
                arr = arr / norm
            return arr.tolist()
        except ImportError:
            norm = math.sqrt(sum(x * x for x in vector))
            if norm > 0:
                return [x / norm for x in vector]
            return vector

    def _serialize_data(self, d: dict) -> dict:
        meta = d.get('meta', {})
        meta_str = meta if isinstance(meta, str) else json.dumps(meta, ensure_ascii=False) if meta else '{}'

        image_keys = d.get('image_keys', [])
        image_keys_str = (
            image_keys if isinstance(image_keys, str)
            else json.dumps(image_keys, ensure_ascii=False) if image_keys
            else '[]'
        )

        excluded_embed = d.get('excluded_embed_metadata_keys', [])
        excluded_embed_str = (
            excluded_embed if isinstance(excluded_embed, str)
            else json.dumps(excluded_embed, ensure_ascii=False) if excluded_embed
            else '[]'
        )

        excluded_llm = d.get('excluded_llm_metadata_keys', [])
        excluded_llm_str = (
            excluded_llm if isinstance(excluded_llm, str)
            else json.dumps(excluded_llm, ensure_ascii=False) if excluded_llm
            else '[]'
        )

        res = {
            self._primary_key: d.get(self._primary_key, ''),
            'doc_id': d.get('doc_id', ''),
            'group': d.get('group', ''),
            'content': d.get('content', ''),
            'meta': meta_str,
            'type': d.get('type', SegmentType.TEXT.value),
            'number': d.get('number', 0),
            'kb_id': d.get('kb_id', ''),
            'parent': d.get('parent', ''),
            'answer': d.get('answer', ''),
            'image_keys': image_keys_str,
            'excluded_embed_metadata_keys': excluded_embed_str,
            'excluded_llm_metadata_keys': excluded_llm_str,
        }

        for embed_key, value in d.get('embedding', {}).items():
            if self._embed_datatypes.get(embed_key) == DataType.SPARSE_FLOAT_VECTOR:
                if isinstance(value, dict):
                    value = {int(k) if isinstance(k, str) else k: v for k, v in value.items()}
            else:
                if self._normalize and isinstance(value, list):
                    value = self._normalize_vector(value)

            res[self._gen_embed_key(embed_key)] = value
        global_meta = d.get('global_meta', {})
        for name, desc in self._global_metadata_desc.items():
            value = global_meta.get(name, desc.default_value)
            if value is not None:
                res[self._gen_global_meta_key(name)] = value

        return res

    def _deserialize_data(self, d: dict) -> dict:
        res = {
            self._primary_key: d.get(self._primary_key, ''),
            'doc_id': d.get('doc_id', ''),
            'group': d.get('group', ''),
            'content': d.get('content', ''),
            'meta': json.loads(d.get('meta', '{}')) if isinstance(d.get('meta'), str) else (d.get('meta') or {}),
            'type': d.get('type', SegmentType.TEXT.value),
            'number': d.get('number', 0),
            'kb_id': d.get('kb_id', ''),
            'parent': d.get('parent', ''),
            'answer': d.get('answer', ''),
            'image_keys': (
                json.loads(d.get('image_keys', '[]')) if isinstance(d.get('image_keys'), str)
                else (d.get('image_keys') or [])
            ),
            'excluded_embed_metadata_keys': (
                json.loads(d.get('excluded_embed_metadata_keys', '[]')) if isinstance(d.get('excluded_embed_metadata_keys'), str)  # noqa: E501
                else (d.get('excluded_embed_metadata_keys') or [])
            ),
            'excluded_llm_metadata_keys': (
                json.loads(d.get('excluded_llm_metadata_keys', '[]')) if isinstance(d.get('excluded_llm_metadata_keys'), str)  # noqa: E501
                else (d.get('excluded_llm_metadata_keys') or [])
            ),
            'embedding': {},
            'global_meta': {}
        }

        for k, v in d.items():
            if k.startswith(EMBED_PREFIX):
                res['embedding'][k[len(EMBED_PREFIX):]] = v
            elif k.startswith(GLOBAL_META_KEY_PREFIX):
                meta_key = k[len(GLOBAL_META_KEY_PREFIX):]
                res['global_meta'][meta_key] = v

        return res

    def _gen_global_meta_key(self, k: str) -> str:
        return GLOBAL_META_KEY_PREFIX + k

    def _gen_embed_key(self, k: str) -> str:
        return EMBED_PREFIX + k

    def _ensure_database(self):
        DB_USER = self._user
        DB_PASSWORD = self._password
        uri_parts = self._uri.split(':')
        if len(uri_parts) < 2:
            raise ValueError(f'Invalid URI format: {self._uri}. Expected format: host:port')
        DB_HOST = uri_parts[0]
        DB_PORT = uri_parts[1]
        DB_NAME = self._db_name

        try:
            engine = sqlalchemy.create_engine(f'mysql+pymysql://{DB_USER}:{DB_PASSWORD}@{DB_HOST}:{DB_PORT}/', pool_pre_ping=True)  # noqa: E501

            with engine.connect() as connection:
                LOG.info('Successfully connected to OceanBase database server!')

                result = connection.execute(sqlalchemy.text('SHOW DATABASES'))
                databases = [row[0] for row in result]

                if DB_NAME in databases:
                    LOG.info(f'Database {DB_NAME} already exists.')
                else:
                    connection.execute(sqlalchemy.text(f'CREATE DATABASE {DB_NAME}'))
                    LOG.info(f'Database {DB_NAME} created successfully!')

        except Exception as e:
            LOG.error(f'[OceanBaseStore - _ensure_database] Unexpected error: {e}')
            raise
        finally:
            engine.dispose()

    def _parse_user(self, user: str) -> str:
        if ':' in user:
            _, tenant, username = user.split(':')
            return f'{username}@{tenant}'
        elif '#' in user:
            username, tenant_cluster = user.split('@')
            tenant, cluster = tenant_cluster.split('#')
            return f'{username}@{tenant}'
        else:
            return user

    def _set_constants(self):
        self._oceanbase_supported_vector_index_types = {
            'HNSW': pyobvector.client.index_param.VecIndexType.HNSW,
            'HNSW_SQ': pyobvector.client.index_param.VecIndexType.HNSW_SQ,
            'IVF': pyobvector.client.index_param.VecIndexType.IVFFLAT,
            'IVF_FLAT': pyobvector.client.index_param.VecIndexType.IVFFLAT,
            'IVF_SQ': pyobvector.client.index_param.VecIndexType.IVFSQ,
            'IVF_PQ': pyobvector.client.index_param.VecIndexType.IVFPQ,
            'FLAT': pyobvector.client.index_param.VecIndexType.IVFFLAT,
        }
        self._type2oceanbase = {
            DataType.ARRAY: pyobvector.ARRAY,
            DataType.FLOAT_VECTOR: pyobvector.VECTOR,
            DataType.SPARSE_FLOAT_VECTOR: pyobvector.SPARSE_VECTOR,
            DataType.STRING: sqlalchemy.dialects.mysql.TEXT,
            DataType.VARCHAR: sqlalchemy.String,
            DataType.INT32: sqlalchemy.Integer,
            DataType.INT64: sqlalchemy.Integer,
        }
        self._builtin_keys = {
            'uid': {'dtype': sqlalchemy.String(512), 'primary_key': True, 'autoincrement': False},
            'doc_id': {'dtype': sqlalchemy.String(512)},
            'group': {'dtype': sqlalchemy.String(512)},
            'content': {'dtype': sqlalchemy.dialects.mysql.LONGTEXT},
            'meta': {'dtype': sqlalchemy.dialects.mysql.LONGTEXT},
            'type': {'dtype': sqlalchemy.Integer},
            'number': {'dtype': sqlalchemy.Integer},
            'kb_id': {'dtype': sqlalchemy.String(512)},
            'parent': {'dtype': sqlalchemy.String(512)},
            'answer': {'dtype': sqlalchemy.dialects.mysql.LONGTEXT},
            'image_keys': {'dtype': sqlalchemy.dialects.mysql.LONGTEXT},
            'excluded_embed_metadata_keys': {'dtype': sqlalchemy.dialects.mysql.LONGTEXT},
            'excluded_llm_metadata_keys': {'dtype': sqlalchemy.dialects.mysql.LONGTEXT},
        }
        self._constant_columns = self._get_constant_columns()

    def _get_constant_columns(self) -> list:
        column_list = []
        for k, kws in self._builtin_keys.items():
            kws_copy = dict(kws)
            dtype = kws_copy.pop('dtype')
            column_list.append(sqlalchemy.Column(k, dtype, **kws_copy))
        for k, desc in self._global_metadata_desc.items():
            field_name = self._gen_global_meta_key(k)
            if desc.data_type == DataType.ARRAY:
                if desc.element_type is None:
                    raise ValueError(f'OceanBase field [{field_name}]: '
                                     '`element_type` is required when `data_type` is ARRAY.')
                column_list.append(sqlalchemy.Column(field_name, pyobvector.ARRAY))
            elif desc.data_type == DataType.VARCHAR:
                column_list.append(sqlalchemy.Column(field_name, sqlalchemy.String(desc.max_size)))
            else:
                column_list.append(sqlalchemy.Column(field_name, self._type2oceanbase[desc.data_type]))
        return column_list

    def _ensure_params_defaults(self, index_item: dict):
        itype = index_item.get('index_type')
        if itype:
            itype_up = str(itype).upper()
            index_item['index_type'] = itype_up
        else:
            LOG.error(f'[OceanBaseStore] Cannot find `index_type` in index_kwargs: {index_item}')
            raise ValueError(f'Cannot find `index_type` in `index_kwargs` of `{index_item}`')

        defaults = OCEANBASE_INDEX_TYPE_DEFAULTS.get(index_item['index_type'], None)
        if defaults is None:
            LOG.error(f'[OceanBaseStore] Unsupported index type: {index_item["index_type"]}')
            raise ValueError(f'[OceanBase Store] Unsupported index type: {index_item["index_type"]}')

        if 'metric_type' not in index_item and 'metric_type' in defaults:
            index_item['metric_type'] = defaults['metric_type']
            LOG.info(f'[OceanBaseStore] Using default metric_type: {defaults["metric_type"]}')

        default_params = defaults.get('params', {})
        if 'params' not in index_item or index_item.get('params') is None:
            index_item['params'] = dict(default_params)
        else:
            if isinstance(index_item['params'], dict):
                for k, v in default_params.items():
                    index_item['params'].setdefault(k, v)
            else:
                index_item['params'] = dict(default_params)

    def _construct_where_clause(self, criteria: dict) -> Optional[list]:
        if not criteria:
            return None

        filter_parts = []
        for key, value in criteria.items():
            if key == self._primary_key:
                continue

            if key in self._global_metadata_desc:
                field_name = self._gen_global_meta_key(key)
            elif key in self._builtin_keys:
                field_name = key
            else:
                continue

            if isinstance(value, list):
                if not value:
                    continue
                if isinstance(value[0], str):
                    values_str = ', '.join(f'"{v}"' for v in value)
                else:
                    values_str = ', '.join(str(v) for v in value)
                filter_parts.append(f'{field_name} in ({values_str})')
            elif isinstance(value, str):
                filter_parts.append(f'{field_name} = "{value}"')
            elif isinstance(value, (int, float)):
                filter_parts.append(f'{field_name} = {value}')
            else:
                raise ValueError(f'Unsupported criteria value type: {type(value)} for key: {key}')

        if not filter_parts:
            return None

        combined_filter = ' and '.join(filter_parts)
        return [sqlalchemy.text(combined_filter)]

    def _construct_filter_expr(self, filters: Dict[str, Union[List, set]]) -> str:  # noqa: C901
        if not filters:
            return ''

        filter_parts = []
        for key, value in filters.items():
            try:
                if key not in self._global_metadata_desc.keys():
                    LOG.debug(f'[OceanBaseStore - _construct_filter_expr] Skipping unknown key: {key}')
                    continue

                field_name = self._gen_global_meta_key(key)

                if isinstance(value, (list, set)):
                    value_list = list(value)
                    if not value_list:
                        continue

                    if isinstance(value_list[0], str):
                        escaped_values = [v.replace('"', '\\"') for v in value_list]
                        values_str = ', '.join(f'"{v}"' for v in escaped_values)
                    else:
                        values_str = ', '.join(str(v) for v in value_list)

                    filter_parts.append(f'{field_name} in ({values_str})')

                elif isinstance(value, str):
                    escaped_value = value.replace('"', '\\"')
                    filter_parts.append(f'{field_name} = "{escaped_value}"')

                elif isinstance(value, (int, float)):
                    filter_parts.append(f'{field_name} = {value}')

                elif isinstance(value, bool):
                    filter_parts.append(f'{field_name} = {1 if value else 0}')

                else:
                    continue

            except Exception as e:
                LOG.warning(f'[OceanBaseStore - _construct_filter_expr] Error processing filter {key}={value}: {e}')
                continue

        result = ' and '.join(filter_parts)
        if result:
            LOG.debug(f'[OceanBaseStore - _construct_filter_expr] Filter expression: {result}')
        return result

    def _get_ids_where_clause(self, criteria: dict):
        ids, where_clause = None, None
        if criteria:
            if self._primary_key in criteria:
                ids = (
                    [criteria[self._primary_key]]
                    if isinstance(criteria[self._primary_key], str)
                    else criteria[self._primary_key]
                )
            else:
                where_clause = self._construct_where_clause(criteria)

        return ids, where_clause

connect(embed_dims=None, embed_datatypes=None, embed=None, global_metadata_desc=None, **kwargs)

连接到底层的 OceanBase 数据库。

Parameters:

  • embed_dims (Dict[str, int], default: None ) –

    嵌入维度字典。

  • embed_datatypes (Dict[str, DataType], default: None ) –

    嵌入数据类型字典。

  • global_metadata_desc (Dict[str, GlobalMetadataDesc], default: None ) –

    全局元数据描述字典。

  • **kwargs

    其他参数。

Source code in lazyllm/tools/rag/store/hybrid/oceanbase_store.py
    @override
    def connect(self, embed_dims: Optional[Dict[str, int]] = None,
                embed_datatypes: Optional[Dict[str, DataType]] = None,
                embed: Optional[Dict[str, Callable]] = None,
                global_metadata_desc: Optional[Dict[str, GlobalMetadataDesc]] = None, **kwargs):
        """连接到底层的 OceanBase 数据库。

Args:
    embed_dims (Dict[str, int]): 嵌入维度字典。
    embed_datatypes (Dict[str, DataType]): 嵌入数据类型字典。
    global_metadata_desc (Dict[str, GlobalMetadataDesc]): 全局元数据描述字典。
    **kwargs: 其他参数。
"""
        self._embed_dims = embed_dims or {}
        self._embed_datatypes = embed_datatypes or {}
        self._embed = embed or {}
        self._global_metadata_desc = global_metadata_desc or {}
        self._set_constants()

        # Extract connection parameters from client_kwargs
        self._user = self._parse_user(self._client_kwargs.pop('user', 'root@test'))
        self._password = self._client_kwargs.pop('password', '')
        self._normalize = self._client_kwargs.pop('normalize', False)
        self._enable_fulltext_index = self._client_kwargs.pop('enable_fulltext_index', False)
        max_pool_size = int(self._client_kwargs.pop('max_pool_size', 8))

        self._ensure_database()
        self._client_pool = _ClientPool(self._new_client, max_size=max_pool_size)
        LOG.info('[OceanBaseStore] init success!')

delete(collection_name, criteria=None, **kwargs)

从存储中删除数据。

Parameters:

  • collection_name (str) –

    集合名称。

  • criteria (Optional[dict], default: None ) –

    删除条件,默认为None。

  • **kwargs

    其他参数。

Returns:

  • bool: 操作成功返回True,否则返回False。
Source code in lazyllm/tools/rag/store/hybrid/oceanbase_store.py
    @override
    def delete(self, collection_name: str, criteria: Optional[dict] = None, **kwargs) -> bool:
        """从存储中删除数据。

Args:
    collection_name (str): 集合名称。
    criteria (Optional[dict]): 删除条件,默认为None。
    **kwargs: 其他参数。

**Returns:**

- bool: 操作成功返回True,否则返回False。
"""
        try:
            with self._client_context() as client:
                if not client.check_table_exists(collection_name):
                    return True

                if not criteria:
                    with self._ddl_lock:
                        client.drop_table_if_exist(table_name=collection_name)
                else:
                    ids, where_clause = self._get_ids_where_clause(criteria)

                    client.delete(
                        table_name=collection_name,
                        ids=ids,
                        where_clause=where_clause,
                        **kwargs
                    )

            return True
        except Exception as e:
            LOG.error(f'[OceanBaseStore - delete] error: {e}')
            LOG.error(traceback.format_exc())
            return False

get(collection_name, criteria=None, **kwargs)

从存储中获取数据。

Parameters:

  • collection_name (str) –

    集合名称。

  • criteria (Optional[dict], default: None ) –

    查询条件,默认为None。

  • **kwargs

    其他参数。

Returns:

  • List[dict]: 返回符合条件的数据列表。
Source code in lazyllm/tools/rag/store/hybrid/oceanbase_store.py
    @override
    def get(self, collection_name: str, criteria: Optional[dict] = None, **kwargs) -> List[dict]:
        """从存储中获取数据。

Args:
    collection_name (str): 集合名称。
    criteria (Optional[dict]): 查询条件,默认为None。
    **kwargs: 其他参数。

**Returns:**

- List[dict]: 返回符合条件的数据列表。
"""
        try:
            with self._client_context() as client:
                if not client.check_table_exists(collection_name):
                    return []

                ids, where_clause = self._get_ids_where_clause(criteria)

                if ids is None and where_clause is None:
                    return self._get_all_with_pagination(client, collection_name)

                res = client.get(
                    table_name=collection_name,
                    ids=ids,
                    where_clause=where_clause,
                    output_column_name=None,
                    **kwargs
                )

                if not res:
                    return []
                result = [r._mapping for r in res]

                return [self._deserialize_data(r) for r in result]

        except Exception as e:
            LOG.error(f'[OceanBaseStore - get] error: {e}')
            LOG.error(traceback.format_exc())
            return []

search(collection_name, query, query_embedding, topk, filters=None, embed_key=None, filter_str='', **kwargs)

在存储中搜索数据。

Parameters:

  • collection_name (str) –

    集合名称。

  • query_embedding (Union[dict, List[float]]) –

    查询的向量表示。

  • topk (int) –

    返回的最大结果数量。

  • filters (Optional[Dict[str, Union[str, int, List, Set]]], default: None ) –

    过滤条件,默认为None。

  • embed_key (Optional[str], default: None ) –

    嵌入向量的键名,默认为None。

  • filter_str (Optional[str], default: '' ) –

    过滤条件字符串,默认为None。

  • **kwargs

    其他参数。

Returns:

  • List[dict]: 返回搜索结果列表。
Source code in lazyllm/tools/rag/store/hybrid/oceanbase_store.py
    def search(self, collection_name: str, query: str, query_embedding: Union[dict, List[float]], topk: int, filters: Optional[Dict[str, Union[List, set]]] = None, embed_key: Optional[str] = None, filter_str: Optional[str] = '', **kwargs) -> List[dict]:  # noqa: C901 E501
        """在存储中搜索数据。

Args:
    collection_name (str): 集合名称。
    query_embedding (Union[dict, List[float]]): 查询的向量表示。
    topk (int): 返回的最大结果数量。
    filters (Optional[Dict[str, Union[str, int, List, Set]]]): 过滤条件,默认为None。
    embed_key (Optional[str]): 嵌入向量的键名,默认为None。
    filter_str (Optional[str]): 过滤条件字符串,默认为None。
    **kwargs: 其他参数。

**Returns:**

- List[dict]: 返回搜索结果列表。
"""
        if not query_embedding:
            raise NotImplementedError('Query fulltext search is not supported for now')
        try:
            if not collection_name or not isinstance(collection_name, str):
                LOG.error('[OceanBaseStore - search] Invalid collection_name')
                return []

            with self._client_context() as client:
                if not client.check_table_exists(collection_name):
                    LOG.warning(f'[OceanBaseStore - search] Table {collection_name} does not exist')
                    return []
                if embed_key and embed_key not in self._embed_datatypes:
                    LOG.error(f'[OceanBaseStore - search] embed_key: {embed_key} not exists')
                    return []

                if not embed_key:
                    if self._embed_datatypes:
                        embed_key = next(iter(self._embed_datatypes.keys()))
                        LOG.info(f'[OceanBaseStore - search] No embed_key provided, using default: {embed_key}')
                    else:
                        LOG.error('[OceanBaseStore - search] No embedding datatypes available')
                        return []

                where_clause = None
                if filters or filter_str:
                    filter_parts = []
                    if filters:
                        try:
                            filter_expr = self._construct_filter_expr(filters)
                            if filter_expr:
                                filter_parts.append(filter_expr)
                        except Exception as filter_err:
                            LOG.error(f'[OceanBaseStore - search] Failed to construct filter: {filter_err}')
                            LOG.error(traceback.format_exc())
                            raise RuntimeError(f'Failed to construct filter expression: {filter_err}') from filter_err

                    if filter_str:
                        filter_parts.append(filter_str)

                    if filter_parts:
                        combined_filter = ' and '.join(f'({part})' for part in filter_parts)
                        where_clause = [sqlalchemy.text(combined_filter)]

                if (
                    isinstance(query_embedding, dict)
                    and self._embed_datatypes.get(embed_key) != DataType.SPARSE_FLOAT_VECTOR
                ):
                    vec_data = query_embedding.get(embed_key)
                    if vec_data is None:
                        LOG.error(f'[OceanBaseStore - search] embed_key: {embed_key} not found in query_embedding')
                        return []
                else:
                    vec_data = query_embedding

                if vec_data is None:
                    LOG.error('[OceanBaseStore - search] Vector data is None')
                    return []

                if self._embed_datatypes[embed_key] == DataType.SPARSE_FLOAT_VECTOR:
                    if isinstance(vec_data, dict):
                        vec_data = {int(k) if isinstance(k, str) else k: v for k, v in vec_data.items()}
                    distance_func = pyobvector.inner_product
                else:
                    if self._normalize and isinstance(vec_data, list):
                        vec_data = self._normalize_vector(vec_data)

                    metric_type = self._get_metric_type_for_embed_key(embed_key)
                    distance_func = self._get_distance_function(metric_type)

                search_params = kwargs.get('search_params', {})
                index_type = self._get_index_type_for_embed_key(embed_key)

                if index_type in ['HNSW', 'HNSW_SQ']:
                    ef_search = search_params.get('efSearch', 64)  # Default efSearch
                    if self._hnsw_ef_search.get(embed_key) != ef_search:
                        try:
                            client.set_ob_hnsw_ef_search(ef_search)
                            self._hnsw_ef_search[embed_key] = ef_search
                        except Exception as e:
                            LOG.error(f'[OceanBaseStore - search] Failed to set efSearch: {e}')
                            LOG.error(traceback.format_exc())
                            raise RuntimeError(f'Failed to set efSearch parameter for HNSW index: {e}') from e

                results = client.ann_search(
                    table_name=collection_name,
                    vec_data=vec_data,
                    vec_column_name=self._gen_embed_key(embed_key),
                    distance_func=distance_func,
                    topk=topk,
                    with_dist=True,
                    output_column_names=None,
                    where_clause=where_clause,
                    **kwargs
                )

                res = []
                if not results:
                    LOG.info('[OceanBaseStore - search] No results found')
                    return []

                for row in results:
                    try:
                        if hasattr(row, '_mapping'):
                            row_dict = dict(row._mapping)
                            score = row_dict.pop('distance', 0)
                        elif isinstance(row, dict):
                            row_dict = dict(row)
                            score = row_dict.pop('distance', 0)
                        else:
                            LOG.warning(f'[OceanBaseStore - search] Unsupported row type: {type(row)}')
                            continue

                        doc_data = self._deserialize_data(row_dict)

                        if not doc_data.get(self._primary_key):
                            LOG.warning('[OceanBaseStore - search] Row has no valid uid')
                            continue

                        doc_data['score'] = float(score)
                        res.append(doc_data)

                    except Exception as row_err:
                        LOG.warning(f'[OceanBaseStore - search] Failed to process row: {row_err}')
                        LOG.warning(traceback.format_exc())
                        continue

                LOG.info(f'[OceanBaseStore - search] Returning {len(res)} results')
                return res
        except Exception as e:
            LOG.error(f'[OceanBaseStore - search] Unexpected error: {e}')
            LOG.error(traceback.format_exc())
            return []

upsert(collection_name, data, range_part=None, **kwargs)

向存储中插入或更新数据。

Parameters:

  • collection_name (str) –

    集合名称。

  • data (List[dict]) –

    要插入或更新的数据列表,每个数据项都是一个字典。

  • range_part (Optional[RangeListPartInfo], default: None ) –

    范围分区信息,暂未实现分区功能。

  • **kwargs

    其他参数。

Returns:

  • bool: 操作成功返回True,否则返回False。
Source code in lazyllm/tools/rag/store/hybrid/oceanbase_store.py
    def upsert(self, collection_name: str, data: List[dict], range_part: Optional['pyobvector.RangeListPartInfo'] = None, **kwargs) -> bool:  # noqa: C901 E501
        """向存储中插入或更新数据。

Args:
    collection_name (str): 集合名称。
    data (List[dict]): 要插入或更新的数据列表,每个数据项都是一个字典。
    range_part (Optional[RangeListPartInfo]): 范围分区信息,暂未实现分区功能。
    **kwargs: 其他参数。

**Returns:**

- bool: 操作成功返回True,否则返回False。
"""
        try:
            if not data:
                return True

            if not collection_name or not isinstance(collection_name, str):
                LOG.error('[OceanBaseStore - upsert] Invalid collection_name')
                return False

            all_embed_keys = set()
            for item in data:
                if 'embedding' in item and isinstance(item['embedding'], dict):
                    all_embed_keys.update(item['embedding'].keys())

            with self._client_context() as client:
                with self._ddl_lock:
                    if not client.check_table_exists(collection_name):
                        self._resolve_missing_embed_specs(all_embed_keys)
                        embed_kwargs = {}
                        if all_embed_keys:
                            for embed_key in all_embed_keys:
                                if not self._embed_datatypes.get(embed_key):
                                    LOG.error(f'[OceanBaseStore - upsert] Cannot find embedding for embed [{embed_key}]')
                                    return False

                                if embed_key not in embed_kwargs:
                                    embed_kwargs[embed_key] = {
                                        'dtype': self._type2oceanbase[self._embed_datatypes[embed_key]]
                                    }
                                if self._embed_dims.get(embed_key):
                                    embed_kwargs[embed_key]['dim'] = self._embed_dims[embed_key]

                        self._create_table_and_index(client, collection_name, embed_kwargs, range_part)

                total_inserted = 0
                failed_batches = []

                serialized_data = [self._serialize_data(d) for d in data]

                for i in range(0, len(serialized_data), OB_UPSERT_BATCH_SIZE):
                    batch_num = i // OB_UPSERT_BATCH_SIZE + 1
                    try:
                        if i == 0 or batch_num % 10 == 0:
                            try:
                                client.perform_raw_text_sql('SET SESSION ob_query_timeout = 300000000')
                            except Exception as timeout_err:
                                LOG.warning(f'[OceanBaseStore - upsert] Failed to set '
                                            f'query timeout for batch {batch_num}: {timeout_err}')

                        batch_data = serialized_data[i:i + OB_UPSERT_BATCH_SIZE]
                        client.upsert(table_name=collection_name, data=batch_data)
                        total_inserted += len(batch_data)

                    except Exception as batch_err:
                        LOG.error(f'[OceanBaseStore - upsert] Failed to insert batch {batch_num}: {batch_err}')
                        LOG.error(f'[OceanBaseStore - upsert] Error details: {traceback.format_exc()}')
                        failed_batches.append(batch_num)
                        continue

                if failed_batches:
                    LOG.warning(f'[OceanBaseStore - upsert] Failed batches: {failed_batches}')
                    if total_inserted == 0:
                        return False
            return True
        except Exception as e:
            LOG.error(f'[OceanBaseStore - upsert] Unexpected error: {e}')
            LOG.error(traceback.format_exc())
            return False

lazyllm.tools.rag.store.ElasticSearchStore

Bases: LazyLLMStoreBase

基于 Elasticsearch 的向量存储实现,继承自 StoreBase。支持向量写入、删除、相似度检索,兼容标量过滤。 Args: uris (List[str]): Elasticsearch 连接 URI(如 ["http://localhost:9200"])。 client_kwargs (Optional[Dict]): 传递给 Elasticsearch 客户端的额外参数。 index_kwargs (Optional[Union[Dict, List]]): 索引创建参数(例如 {"index_type": "IVF_FLAT", "metric_type": "CONSINE"} ,支持按向量模型的key配置列表)。 **kwargs: 预留扩展参数。

Examples:

>>> import lazyllm
>>> from lazyllm.tools.rag.store import ElasticSearchStore
>>> store = ElasticSearchStore(uris=["localhost:9200"], client_kwargs={}, index_kwargs={})
>>> store.connect(embed_dims={"vec_dense": 128, "vec_sparse": 128}, embed_datatypes={"vec_dense": DataType.FLOAT32, "vec_sparse": DataType.FLOAT32}, global_metadata_desc={})
>>> store.upsert(collection_name="test", data=[{"uid": "1", "embedding": {"vec_dense": [0.1, 0.2, 0.3], "vec_sparse": {"1": 0.1, "2": 0.2, "3": 0.3}}, "metadata": {"key1": "value1", "key2": "value2"}}])
>>> store.get(collection_name="test", criteria={"uid": "1"})
>>> store.delete(collection_name="test", criteria={"uid": "1"})
Source code in lazyllm/tools/rag/store/segment/elasticsearch_store.py
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
class ElasticSearchStore(LazyLLMStoreBase):
    """
基于 Elasticsearch 的向量存储实现,继承自 StoreBase。支持向量写入、删除、相似度检索,兼容标量过滤。
Args:
    uris (List[str]): Elasticsearch 连接 URI(如 ["http://localhost:9200"])。
    client_kwargs (Optional[Dict]): 传递给 Elasticsearch 客户端的额外参数。
    index_kwargs (Optional[Union[Dict, List]]): 索引创建参数(例如 {"index_type": "IVF_FLAT", "metric_type": "CONSINE"} ,支持按向量模型的key配置列表)。
    **kwargs: 预留扩展参数。


Examples:
    >>> import lazyllm
    >>> from lazyllm.tools.rag.store import ElasticSearchStore
    >>> store = ElasticSearchStore(uris=["localhost:9200"], client_kwargs={}, index_kwargs={})
    >>> store.connect(embed_dims={"vec_dense": 128, "vec_sparse": 128}, embed_datatypes={"vec_dense": DataType.FLOAT32, "vec_sparse": DataType.FLOAT32}, global_metadata_desc={})
    >>> store.upsert(collection_name="test", data=[{"uid": "1", "embedding": {"vec_dense": [0.1, 0.2, 0.3], "vec_sparse": {"1": 0.1, "2": 0.2, "3": 0.3}}, "metadata": {"key1": "value1", "key2": "value2"}}])
    >>> store.get(collection_name="test", criteria={"uid": "1"})
    >>> store.delete(collection_name="test", criteria={"uid": "1"})
    """
    capability = StoreCapability.SEGMENT
    need_embedding = False
    supports_index_registration = False

    def __init__(
        self,
        uris: List[str],
        client_kwargs: Optional[Dict] = None,
        index_kwargs: Optional[Union[Dict, List]] = None,
        **kwargs,
    ):
        if isinstance(uris, str):
            uris = [uris]

        self._uris = uris
        self._client_kwargs = client_kwargs or {}
        self._index_kwargs = index_kwargs or DEFAULT_MAPPING_BODY
        self._primary_key = 'uid'

    @property
    def dir(self):
        """
远程模式返回 None。
**Returns:**

    Optional[str]: None。
"""
        return None

    @override
    def connect(self, global_metadata_desc: Optional[Dict[str, GlobalMetadataDesc]] = None, **kwargs) -> bool:
        """
初始化 Elasticsearch 客户端,传入向量化模型参数和全局元数据描述。
Args:
    embed_dims (Dict[str, int]): 每个嵌入键对应的向量维度。
    embed_datatypes (Dict[str, DataType]): 每个嵌入键的数据类型。
    global_metadata_desc (Dict[str, GlobalMetadataDesc]): 全局元数据字段的描述。
**Returns:**

    bool: 操作成功返回 True,否则 False。
"""
        try:
            self._ddl_lock = threading.Lock()
            # Elastic Cloud
            if self._client_kwargs.get('cloud_id') and self._client_kwargs.get('api_key'):
                cloud_id = self._client_kwargs.get('cloud_id')
                api_key = self._client_kwargs.get('api_key')
                request_timeout = self._client_kwargs.get('request_timeout', 30)

                self._client = elasticsearch.Elasticsearch(
                    cloud_id=cloud_id, api_key=api_key, request_timeout=request_timeout
                )
                if not self._client.ping():
                    raise ConnectionError(f'Failed to ping ES {self._uris}')
            else:
                client_kwargs = dict(self._client_kwargs)
                if 'request_timeout' not in client_kwargs:
                    client_kwargs['request_timeout'] = 30
                self._client = elasticsearch.Elasticsearch(hosts=self._uris, **client_kwargs)

            self._global_metadata_desc = global_metadata_desc

            self._index_kwargs = self._adapt_mapping_for_global_metadata()

            return True
        # connection failed exception handling
        except elasticsearch.NotFoundError as e:
            LOG.error(f'ElasticSearch sever with cloud id {cloud_id} and api key {api_key} does not exist')
            raise e
        except elasticsearch.AuthenticationException as e:
            LOG.error('ElasticSearch needs Authentication')
            raise e
        except elasticsearch.AuthorizationException as e:
            LOG.error('Unauthorized to access')
            raise e
        except Exception as e:
            LOG.error(f'Fail to connect ElasticSearch sever with cloud id {cloud_id} and api key {api_key}')
            raise e

    @override
    def _ensure_index(self, index: str = None) -> bool:
        if not index or self._client.indices.exists(index=index):
            return False
        try:
            self._client.indices.create(index=index, body=self._index_kwargs)
            return True
        except elasticsearch.TransportError as e:
            if getattr(e, 'error', '') != 'resource_already_exists_exception':
                raise e
        except Exception as e:
            LOG.error(f'[ElasticSearch - _ensure_index] Error creating index {index}: {e}')
            raise e

    @override
    def upsert(self, collection_name: str = None, data: List[Dict] = None) -> bool:
        """
批量写入或更新切片数据到 Elasticsearch 集合。
Args:
    collection_name (str): 集合名称,通常为 "group_embedKey" 格式。
    data (List[dict]): 切片数据列表。
**Returns:**

    bool: 操作成功返回 True,否则 False。
"""
        if not data:
            return False
        try:
            self._ensure_index(collection_name)
            for i in range(0, len(data), INSERT_BATCH_SIZE):
                bulk_data = []
                batch_data = data[i: i + INSERT_BATCH_SIZE]
                for segment in batch_data:
                    segment = self._serialize_node(segment)
                    _id = segment.pop(self._primary_key, None)
                    bulk_data.append({'index': {'_index': collection_name, '_id': _id}})
                    bulk_data.append(segment)

                response = self._client.bulk(index=collection_name, body=bulk_data, refresh='wait_for')
                if response.get('errors'):
                    raise ValueError(
                        f'Error upserting data to Elasticsearch: {response}'
                    )
            return True

        except Exception as e:
            LOG.error(f'[ElasticSearchStore - upsert] Error upserting documents to {collection_name}: {e}')
            raise e

    @override
    def delete(self, collection_name: str = None, criteria: Optional[Dict] = None, **kwargs) -> bool:
        """
删除整个集合或按条件删除指定记录。
Args:
    collection_name (str): 目标集合名称。
    criteria (Optional[dict]): 若为 None 则删除整个集合;否则按 uid 列表或元数据条件过滤。
**Returns:**

    bool: 删除成功返回 True,否则 False。
"""
        try:
            if not self._client.indices.exists(index=collection_name):
                LOG.warning(f'[ElasticSearchStore - delete] Index {collection_name} does not exist')
                return True
            if not criteria:
                with self._ddl_lock:
                    if self._client.indices.exists(index=collection_name):
                        self._client.indices.delete(index=collection_name)
                return True
            else:
                resp = self._client.delete_by_query(
                    index=collection_name,
                    body=self._construct_criteria(criteria),
                    refresh=True,
                    conflicts='proceed',
                    request_timeout=30,
                )

                if resp.get('version_conflicts', 0) > 0:
                    LOG.warning(f'[ElasticsearchStore - delete] Version conflicts: {resp.get("version_conflicts")}')
                if resp.get('failures'):
                    raise ValueError(
                        f'Error deleting data from Elasticsearch: {resp["failures"]}'
                    )
                return True

        except Exception as e:
            LOG.error(f'[ElasticSearchStore - delete] Error deleting from {collection_name}: {e}')
            raise e

    @override
    def get(self, collection_name: str, criteria: Optional[dict] = None, **kwargs) -> List[dict]:  # noqa: C901
        """
检索匹配主键或元数据过滤条件的记录。
Args:
    collection_name (str): 待查询集合。
    criteria (Optional[dict]): 包含 'uid' 列表或元数据字段过滤条件。
**Returns:**

    List[dict]: 每项包含 'uid' 及 'embedding' 映射。
"""
        try:
            if not self._client.indices.exists(index=collection_name):
                return []

            results: List[dict] = []
            criteria = dict(criteria) if criteria else {}
            limit = kwargs.get('limit')
            offset = max(kwargs.get('offset', 0) or 0, 0)
            return_total = kwargs.get('return_total', False)
            sort_by_number = kwargs.get('sort_by_number', False)
            # Query by primary key(mget)
            if criteria and self._primary_key in criteria:
                vals = criteria.pop(self._primary_key)
                if not isinstance(vals, list):
                    vals = [vals]

                resp = self._client.mget(index=collection_name, body={'ids': vals})

                for doc in resp['docs']:
                    if doc.get('found', False):
                        seg = self._transform_segment(doc)
                        if seg:
                            results.append(seg)
                if sort_by_number:
                    results = sorted(results, key=lambda item: (item.get('number', 0), item.get('uid', '')))
                total = len(results)
                if offset > 0 or limit is not None:
                    end = None if limit is None else offset + limit
                    results = results[offset:end]
                return (results, total) if return_total else results
            elif sort_by_number and (limit is not None or offset > 0 or return_total):
                body = self._construct_criteria(criteria) or {'query': {'match_all': {}}}
                body['sort'] = [{'number': {'order': 'asc'}}, {'_id': {'order': 'asc'}}]
                if offset > 0:
                    body['from'] = offset
                if limit is not None:
                    body['size'] = limit
                elif offset > 0:
                    body['size'] = 10000
                if return_total:
                    body['track_total_hits'] = True
                resp = self._client.search(index=collection_name, body=body)
                results = [self._transform_segment(hit) for hit in resp['hits']['hits']]
                total = resp['hits']['total']['value'] if return_total else len(results)
                return (results, total) if return_total else results

            else:
                helpers = elasticsearch.helpers
                query = self._construct_criteria(criteria)
                for hit in helpers.scan(
                    client=self._client,
                    index=collection_name,
                    query=query,  # 8.x need to wrap a query
                    scroll='2m',
                    size=500,
                    preserve_order=False,
                ):
                    seg = self._transform_segment(hit)
                    if seg:
                        results.append(seg)
            return results
        except Exception as e:
            LOG.error(f'[ElasticsearchStore - get] Error getting data from Elasticsearch: {e}')
            return []

    @override
    def search(self, collection_name: str, query: str,
               topk: Optional[int] = 10, filters: Optional[dict] = None, **kwargs) -> List[Dict]:  # noqa: C901
        """
执行向量相似度检索,并可按元数据过滤。
Args:
    collection_name (str): 待搜索集合。
    query (Optional[str]): 查询字符串。
    topk (Optional[int]): 返回邻近数量。
    filters (Optional[dict]): 元数据过滤映射。
    kwargs: 其他搜索参数

**Returns:**

- List[dict]: 返回匹配结果列表及相似度 'score'。
"""
        query_fields = ['*']
        try:
            self._ensure_index(collection_name)
            must_clauses = []
            es_query = {}
            text_query = {
                'multi_match': {
                    'query': query,
                    'fields': query_fields,
                }
            }
            must_clauses.append(text_query)

            filter_query = self._construct_criteria(filters) if filters else {}

            if must_clauses and filter_query:
                # combine filter_query and must_clauses
                filter_must = filter_query['query']['bool']['must']
                es_query = {'query': {'bool': {'must': must_clauses + filter_must}}}
            elif must_clauses:
                es_query = {'query': {'bool': {'must': must_clauses}}}
            elif filter_query:
                es_query = filter_query
            else:
                es_query = {'query': {'match_all': {}}}

            es_query['size'] = topk

            resp = self._client.search(index=collection_name, body=es_query)

            res = []
            for hit in resp['hits']['hits']:
                seg = self._transform_segment(hit)
                if seg:
                    seg['score'] = hit.get('_score', 0.0)
                    res.append(seg)
            return res

        except Exception as e:
            LOG.error(f'[ElasticSearchStore - search] Error searching {collection_name}: {e}')
            return []

    def _serialize_node(self, segment: Dict) -> Dict:
        seg = dict(segment)
        seg.pop('embedding', None)

        # Note: Insertion will fail if a dictionary(meta) has an excessively deep nesting level or overly long keys.
        if self._global_metadata_desc and self._global_metadata_desc == BUILDIN_GLOBAL_META_DESC:
            seg['global_meta'] = json.dumps(seg.get('global_meta', {}), ensure_ascii=False)
            seg['meta'] = json.dumps(seg.get('meta', {}), ensure_ascii=False)
            seg['image_keys'] = json.dumps(seg.get('image_keys', []), ensure_ascii=False)
        return seg

    def _deserialize_node(self, segment: Dict) -> Dict:
        seg = dict(segment)
        if self._global_metadata_desc and self._global_metadata_desc == BUILDIN_GLOBAL_META_DESC:
            seg['meta'] = json.loads(seg.get('meta', '{}'))
            seg['global_meta'] = json.loads(seg.get('global_meta', '{}'))
            seg['image_keys'] = json.loads(seg.get('image_keys', '[]'))
        return seg

    def _construct_criteria(self, criteria: Optional[dict] = None) -> dict:  # noqa: C901
        criteria = dict(criteria) if criteria else {}
        if not criteria:
            return {}
        if self._primary_key in criteria:
            vals = criteria.pop(self._primary_key)
            if not isinstance(vals, list):
                vals = [vals]
            return {'query': {'ids': {'values': vals}}}

        exact_match_fields = {'doc_id', 'kb_id', 'group', 'parent'}

        def _add_clause(key, val):
            if key in exact_match_fields:
                clauses = []
                if isinstance(val, list):
                    clauses.append({'terms': {key: val}})
                    clauses.append({'terms': {f'{key}.keyword': val}})
                else:
                    clauses.append({'term': {key: val}})
                    clauses.append({'term': {f'{key}.keyword': val}})
                must_clauses.append({'bool': {'should': clauses, 'minimum_should_match': 1}})
                return
            if isinstance(val, list):
                must_clauses.append({'terms': {key: val}})
            else:
                must_clauses.append({'term': {key: val}})
        must_clauses = []
        if RAG_DOC_ID in criteria:
            val = criteria.pop(RAG_DOC_ID)
            _add_clause('doc_id', val)
        if RAG_KB_ID in criteria:
            _add_clause('kb_id', criteria.pop(RAG_KB_ID))
        if 'parent' in criteria:
            _add_clause('parent', criteria.pop('parent'))
        if 'number' in criteria:
            _add_clause('number', criteria.pop('number'))

        for k, v in criteria.items():
            field_key = k
            # For custom text fields, use .keyword subfield for exact matching
            if (self._global_metadata_desc
                and self._global_metadata_desc != BUILDIN_GLOBAL_META_DESC
                and k in self._global_metadata_desc.keys()
            ):
                field_desc = self._global_metadata_desc[k]

                if field_desc.data_type in (DataType.VARCHAR, DataType.STRING):
                    field_key = f'{k}.keyword'
            _add_clause(field_key, v)

        return {'query': {'bool': {'must': must_clauses}}} if must_clauses else {}

    def _transform_segment(self, record: dict) -> dict:
        src = record['_source']
        src['uid'] = record['_id']
        return self._deserialize_node(src)

    def _check_ik_plugin(self):
        try:
            plugins = self._client.cat.plugins(format='json')
            if any('analysis-ik' in p.get('component', '') for p in plugins):
                return True
            try:
                self._client.indices.analyze(
                    body={
                        'analyzer': 'ik_max_word',
                        'text': 'machine learning'
                    }
                )
                return True
            except Exception as e:
                LOG.warning(f'IK plugin is not installed: {str(e)}')
                return False
        except Exception as e:
            LOG.warning(f'check IK plugin failed: {e}')
            return False

    def _adapt_mapping_for_global_metadata(self) -> dict:
        check_ik = self._check_ik_plugin()
        if check_ik:
            LOG.info('IK plugin is installed')
        else:
            LOG.warning('IK plugin is not installed, ElasticSearch will \
                use ngram analyzer which is English Only Analyzer')

        if not self._global_metadata_desc or self._global_metadata_desc == BUILDIN_GLOBAL_META_DESC:
            mapping = copy.deepcopy(DEFAULT_MAPPING_BODY)
            if not check_ik:
                content_field = mapping['mappings']['properties'].get('content', {})
                if content_field.get('analyzer') == 'ik_max_word':
                    content_field['analyzer'] = 'ngram_analyzer'
                if content_field.get('search_analyzer') == 'ik_smart':
                    content_field['search_analyzer'] = 'ngram_analyzer'
                mapping['mappings']['properties']['content'] = content_field
            return mapping

        mapping = copy.deepcopy(DEFAULT_MAPPING_BODY)
        mapping['mappings']['dynamic'] = 'true'
        props = {'uid': {'type': 'keyword'}}
        self._type2es = {
            DataType.VARCHAR: 'text',
            DataType.ARRAY: 'array',
            DataType.INT32: 'integer',
            DataType.BOOLEAN: 'boolean',
            DataType.FLOAT: 'float',
            DataType.INT64: 'long',
            DataType.STRING: 'text',
        }

        for field_name, desc in self._global_metadata_desc.items():
            field_type = self._type2es[desc.data_type]
            field_def = {'type': field_type, 'store': True, 'index': True}
            if field_type == 'text':
                # Add keyword subfield for exact matching
                field_def['fields'] = {
                    'keyword': {
                        'type': 'keyword',
                        'ignore_above': 256
                    }
                }
                if check_ik:
                    field_def['analyzer'] = 'ik_max_word'
                    field_def['search_analyzer'] = 'ik_smart'
                else:
                    field_def['analyzer'] = 'ngram_analyzer'
                    field_def['search_analyzer'] = 'ngram_analyzer'
            props[field_name] = field_def
        mapping['mappings']['properties'] = props

        return mapping

dir property

远程模式返回 None。 Returns:

Optional[str]: None。

connect(global_metadata_desc=None, **kwargs)

初始化 Elasticsearch 客户端,传入向量化模型参数和全局元数据描述。 Args: embed_dims (Dict[str, int]): 每个嵌入键对应的向量维度。 embed_datatypes (Dict[str, DataType]): 每个嵌入键的数据类型。 global_metadata_desc (Dict[str, GlobalMetadataDesc]): 全局元数据字段的描述。 Returns:

bool: 操作成功返回 True,否则 False。
Source code in lazyllm/tools/rag/store/segment/elasticsearch_store.py
    @override
    def connect(self, global_metadata_desc: Optional[Dict[str, GlobalMetadataDesc]] = None, **kwargs) -> bool:
        """
初始化 Elasticsearch 客户端,传入向量化模型参数和全局元数据描述。
Args:
    embed_dims (Dict[str, int]): 每个嵌入键对应的向量维度。
    embed_datatypes (Dict[str, DataType]): 每个嵌入键的数据类型。
    global_metadata_desc (Dict[str, GlobalMetadataDesc]): 全局元数据字段的描述。
**Returns:**

    bool: 操作成功返回 True,否则 False。
"""
        try:
            self._ddl_lock = threading.Lock()
            # Elastic Cloud
            if self._client_kwargs.get('cloud_id') and self._client_kwargs.get('api_key'):
                cloud_id = self._client_kwargs.get('cloud_id')
                api_key = self._client_kwargs.get('api_key')
                request_timeout = self._client_kwargs.get('request_timeout', 30)

                self._client = elasticsearch.Elasticsearch(
                    cloud_id=cloud_id, api_key=api_key, request_timeout=request_timeout
                )
                if not self._client.ping():
                    raise ConnectionError(f'Failed to ping ES {self._uris}')
            else:
                client_kwargs = dict(self._client_kwargs)
                if 'request_timeout' not in client_kwargs:
                    client_kwargs['request_timeout'] = 30
                self._client = elasticsearch.Elasticsearch(hosts=self._uris, **client_kwargs)

            self._global_metadata_desc = global_metadata_desc

            self._index_kwargs = self._adapt_mapping_for_global_metadata()

            return True
        # connection failed exception handling
        except elasticsearch.NotFoundError as e:
            LOG.error(f'ElasticSearch sever with cloud id {cloud_id} and api key {api_key} does not exist')
            raise e
        except elasticsearch.AuthenticationException as e:
            LOG.error('ElasticSearch needs Authentication')
            raise e
        except elasticsearch.AuthorizationException as e:
            LOG.error('Unauthorized to access')
            raise e
        except Exception as e:
            LOG.error(f'Fail to connect ElasticSearch sever with cloud id {cloud_id} and api key {api_key}')
            raise e

delete(collection_name=None, criteria=None, **kwargs)

删除整个集合或按条件删除指定记录。 Args: collection_name (str): 目标集合名称。 criteria (Optional[dict]): 若为 None 则删除整个集合;否则按 uid 列表或元数据条件过滤。 Returns:

bool: 删除成功返回 True,否则 False。
Source code in lazyllm/tools/rag/store/segment/elasticsearch_store.py
    @override
    def delete(self, collection_name: str = None, criteria: Optional[Dict] = None, **kwargs) -> bool:
        """
删除整个集合或按条件删除指定记录。
Args:
    collection_name (str): 目标集合名称。
    criteria (Optional[dict]): 若为 None 则删除整个集合;否则按 uid 列表或元数据条件过滤。
**Returns:**

    bool: 删除成功返回 True,否则 False。
"""
        try:
            if not self._client.indices.exists(index=collection_name):
                LOG.warning(f'[ElasticSearchStore - delete] Index {collection_name} does not exist')
                return True
            if not criteria:
                with self._ddl_lock:
                    if self._client.indices.exists(index=collection_name):
                        self._client.indices.delete(index=collection_name)
                return True
            else:
                resp = self._client.delete_by_query(
                    index=collection_name,
                    body=self._construct_criteria(criteria),
                    refresh=True,
                    conflicts='proceed',
                    request_timeout=30,
                )

                if resp.get('version_conflicts', 0) > 0:
                    LOG.warning(f'[ElasticsearchStore - delete] Version conflicts: {resp.get("version_conflicts")}')
                if resp.get('failures'):
                    raise ValueError(
                        f'Error deleting data from Elasticsearch: {resp["failures"]}'
                    )
                return True

        except Exception as e:
            LOG.error(f'[ElasticSearchStore - delete] Error deleting from {collection_name}: {e}')
            raise e

get(collection_name, criteria=None, **kwargs)

检索匹配主键或元数据过滤条件的记录。 Args: collection_name (str): 待查询集合。 criteria (Optional[dict]): 包含 'uid' 列表或元数据字段过滤条件。 Returns:

List[dict]: 每项包含 'uid' 及 'embedding' 映射。
Source code in lazyllm/tools/rag/store/segment/elasticsearch_store.py
    @override
    def get(self, collection_name: str, criteria: Optional[dict] = None, **kwargs) -> List[dict]:  # noqa: C901
        """
检索匹配主键或元数据过滤条件的记录。
Args:
    collection_name (str): 待查询集合。
    criteria (Optional[dict]): 包含 'uid' 列表或元数据字段过滤条件。
**Returns:**

    List[dict]: 每项包含 'uid' 及 'embedding' 映射。
"""
        try:
            if not self._client.indices.exists(index=collection_name):
                return []

            results: List[dict] = []
            criteria = dict(criteria) if criteria else {}
            limit = kwargs.get('limit')
            offset = max(kwargs.get('offset', 0) or 0, 0)
            return_total = kwargs.get('return_total', False)
            sort_by_number = kwargs.get('sort_by_number', False)
            # Query by primary key(mget)
            if criteria and self._primary_key in criteria:
                vals = criteria.pop(self._primary_key)
                if not isinstance(vals, list):
                    vals = [vals]

                resp = self._client.mget(index=collection_name, body={'ids': vals})

                for doc in resp['docs']:
                    if doc.get('found', False):
                        seg = self._transform_segment(doc)
                        if seg:
                            results.append(seg)
                if sort_by_number:
                    results = sorted(results, key=lambda item: (item.get('number', 0), item.get('uid', '')))
                total = len(results)
                if offset > 0 or limit is not None:
                    end = None if limit is None else offset + limit
                    results = results[offset:end]
                return (results, total) if return_total else results
            elif sort_by_number and (limit is not None or offset > 0 or return_total):
                body = self._construct_criteria(criteria) or {'query': {'match_all': {}}}
                body['sort'] = [{'number': {'order': 'asc'}}, {'_id': {'order': 'asc'}}]
                if offset > 0:
                    body['from'] = offset
                if limit is not None:
                    body['size'] = limit
                elif offset > 0:
                    body['size'] = 10000
                if return_total:
                    body['track_total_hits'] = True
                resp = self._client.search(index=collection_name, body=body)
                results = [self._transform_segment(hit) for hit in resp['hits']['hits']]
                total = resp['hits']['total']['value'] if return_total else len(results)
                return (results, total) if return_total else results

            else:
                helpers = elasticsearch.helpers
                query = self._construct_criteria(criteria)
                for hit in helpers.scan(
                    client=self._client,
                    index=collection_name,
                    query=query,  # 8.x need to wrap a query
                    scroll='2m',
                    size=500,
                    preserve_order=False,
                ):
                    seg = self._transform_segment(hit)
                    if seg:
                        results.append(seg)
            return results
        except Exception as e:
            LOG.error(f'[ElasticsearchStore - get] Error getting data from Elasticsearch: {e}')
            return []

search(collection_name, query, topk=10, filters=None, **kwargs)

执行向量相似度检索,并可按元数据过滤。 Args: collection_name (str): 待搜索集合。 query (Optional[str]): 查询字符串。 topk (Optional[int]): 返回邻近数量。 filters (Optional[dict]): 元数据过滤映射。 kwargs: 其他搜索参数

Returns:

  • List[dict]: 返回匹配结果列表及相似度 'score'。
Source code in lazyllm/tools/rag/store/segment/elasticsearch_store.py
    @override
    def search(self, collection_name: str, query: str,
               topk: Optional[int] = 10, filters: Optional[dict] = None, **kwargs) -> List[Dict]:  # noqa: C901
        """
执行向量相似度检索,并可按元数据过滤。
Args:
    collection_name (str): 待搜索集合。
    query (Optional[str]): 查询字符串。
    topk (Optional[int]): 返回邻近数量。
    filters (Optional[dict]): 元数据过滤映射。
    kwargs: 其他搜索参数

**Returns:**

- List[dict]: 返回匹配结果列表及相似度 'score'。
"""
        query_fields = ['*']
        try:
            self._ensure_index(collection_name)
            must_clauses = []
            es_query = {}
            text_query = {
                'multi_match': {
                    'query': query,
                    'fields': query_fields,
                }
            }
            must_clauses.append(text_query)

            filter_query = self._construct_criteria(filters) if filters else {}

            if must_clauses and filter_query:
                # combine filter_query and must_clauses
                filter_must = filter_query['query']['bool']['must']
                es_query = {'query': {'bool': {'must': must_clauses + filter_must}}}
            elif must_clauses:
                es_query = {'query': {'bool': {'must': must_clauses}}}
            elif filter_query:
                es_query = filter_query
            else:
                es_query = {'query': {'match_all': {}}}

            es_query['size'] = topk

            resp = self._client.search(index=collection_name, body=es_query)

            res = []
            for hit in resp['hits']['hits']:
                seg = self._transform_segment(hit)
                if seg:
                    seg['score'] = hit.get('_score', 0.0)
                    res.append(seg)
            return res

        except Exception as e:
            LOG.error(f'[ElasticSearchStore - search] Error searching {collection_name}: {e}')
            return []

upsert(collection_name=None, data=None)

批量写入或更新切片数据到 Elasticsearch 集合。 Args: collection_name (str): 集合名称,通常为 "group_embedKey" 格式。 data (List[dict]): 切片数据列表。 Returns:

bool: 操作成功返回 True,否则 False。
Source code in lazyllm/tools/rag/store/segment/elasticsearch_store.py
    @override
    def upsert(self, collection_name: str = None, data: List[Dict] = None) -> bool:
        """
批量写入或更新切片数据到 Elasticsearch 集合。
Args:
    collection_name (str): 集合名称,通常为 "group_embedKey" 格式。
    data (List[dict]): 切片数据列表。
**Returns:**

    bool: 操作成功返回 True,否则 False。
"""
        if not data:
            return False
        try:
            self._ensure_index(collection_name)
            for i in range(0, len(data), INSERT_BATCH_SIZE):
                bulk_data = []
                batch_data = data[i: i + INSERT_BATCH_SIZE]
                for segment in batch_data:
                    segment = self._serialize_node(segment)
                    _id = segment.pop(self._primary_key, None)
                    bulk_data.append({'index': {'_index': collection_name, '_id': _id}})
                    bulk_data.append(segment)

                response = self._client.bulk(index=collection_name, body=bulk_data, refresh='wait_for')
                if response.get('errors'):
                    raise ValueError(
                        f'Error upserting data to Elasticsearch: {response}'
                    )
            return True

        except Exception as e:
            LOG.error(f'[ElasticSearchStore - upsert] Error upserting documents to {collection_name}: {e}')
            raise e

lazyllm.tools.rag.readers.ReaderBase

Bases: ModuleBase

基础文档读取器类,提供文档加载的基本接口。继承自 ModuleBase,使用 LazyLLMRegisterMetaClass 作为元类。

所有 Reader 在 reader(file, ...) 时可选启用算法端内容缓存:将解析完成后的 List[DocNode] 写入 ModuleCache,相同文件内容与 Reader 配置再次调用时直接返回缓存,跳过 _load_data 及下游 OCR 请求。

全局开关为 lazyllm.config['reader_use_cache'](环境变量 LAZYLLM_READER_USE_CACHE,默认 False)。

缓存与 OCR 服务端 use_cache 为两层独立机制:

  • 算法端内容缓存(本类):缓存 DocNode 列表,由 lazyllm.config['reader_use_cache'] 控制。
  • OCR 服务端缓存(如 MineruPDFReader):OCR Reader _load_data(..., use_cache=...),默认 True

缓存键由 Reader 类型、appendix_hash_key(子类配置,如 OCR URL/backend)、文件 mtimest_sizeextra_info 等调用参数共同决定;文件修改后(mtime/size 变化)自动 miss。

存储后端复用 ModuleCache,由全局配置选择:

  • LAZYLLM_CACHE_STRATEGYmemory(默认)/ file / sqlite / redis
  • LAZYLLM_CACHE_MODERW / RO / WO / NONE
  • LAZYLLM_CACHE_DIR:缓存根目录,默认 ~/.lazyllm/cachesqlite 策略下 db 为 {CACHE_DIR}/module/cache.db

Parameters:

  • *args

    位置参数,保留给子类或父类使用。

  • return_trace (bool, default: True ) –

    是否返回处理过程的追踪信息,默认为 True。

  • **kwargs

    关键字参数,保留给子类或父类使用。

Examples:

from lazyllm.tools.rag.readers.readerBase import LazyLLMReaderBase
from lazyllm.tools.rag.doc_node import DocNode
from typing import Iterable

class CustomReader(LazyLLMReaderBase):
    def _lazy_load_data(self, file_paths: list, **kwargs) -> Iterable[DocNode]:
        for file_path in file_paths:
            # Process each file and yield DocNode
            content = self._read_file(file_path)
            yield DocNode(
                text=content,
                metadata={"source": file_path}
            )

# Create reader instance
reader = CustomReader(return_trace=True)

# Load documents
documents = reader.forward(file_paths=["doc1.txt", "doc2.txt"])
Source code in lazyllm/tools/rag/readers/readerBase.py
class LazyLLMReaderBase(ModuleBase, metaclass=LazyLLMRegisterMetaClass):
    """
基础文档读取器类,提供文档加载的基本接口。继承自 ModuleBase,使用 LazyLLMRegisterMetaClass 作为元类。

所有 Reader 在 ``reader(file, ...)`` 时可选启用**算法端内容缓存**:将解析完成后的 ``List[DocNode]`` 写入
``ModuleCache``,相同文件内容与 Reader 配置再次调用时直接返回缓存,跳过 ``_load_data`` 及下游 OCR 请求。

全局开关为 ``lazyllm.config['reader_use_cache']``(环境变量 ``LAZYLLM_READER_USE_CACHE``,默认 False)。

缓存与 OCR 服务端 ``use_cache`` 为两层独立机制:

- **算法端内容缓存**(本类):缓存 ``DocNode`` 列表,由 ``lazyllm.config['reader_use_cache']`` 控制。
- **OCR 服务端缓存**(如 MineruPDFReader):OCR Reader ``_load_data(..., use_cache=...)``,默认 ``True``。

缓存键由 Reader 类型、``appendix_hash_key``(子类配置,如 OCR URL/backend)、文件
``mtime`` 与 ``st_size`` 及 ``extra_info`` 等调用参数共同决定;文件修改后(mtime/size 变化)自动 miss。

存储后端复用 ``ModuleCache``,由全局配置选择:

- ``LAZYLLM_CACHE_STRATEGY``:``memory``(默认)/ ``file`` / ``sqlite`` / ``redis``
- ``LAZYLLM_CACHE_MODE``:``RW`` / ``RO`` / ``WO`` / ``NONE``
- ``LAZYLLM_CACHE_DIR``:缓存根目录,默认 ``~/.lazyllm/cache``;``sqlite`` 策略下 db 为
  ``{CACHE_DIR}/module/cache.db``

Args:
    *args: 位置参数,保留给子类或父类使用。
    return_trace (bool): 是否返回处理过程的追踪信息,默认为 True。
    **kwargs: 关键字参数,保留给子类或父类使用。


Examples:

    from lazyllm.tools.rag.readers.readerBase import LazyLLMReaderBase
    from lazyllm.tools.rag.doc_node import DocNode
    from typing import Iterable

    class CustomReader(LazyLLMReaderBase):
        def _lazy_load_data(self, file_paths: list, **kwargs) -> Iterable[DocNode]:
            for file_path in file_paths:
                # Process each file and yield DocNode
                content = self._read_file(file_path)
                yield DocNode(
                    text=content,
                    metadata={"source": file_path}
                )

    # Create reader instance
    reader = CustomReader(return_trace=True)

    # Load documents
    documents = reader.forward(file_paths=["doc1.txt", "doc2.txt"])
    """
    post_action = None

    _encoding_cache = {}
    _cache_lock = threading.Lock()
    _cache_max_size = 1000

    def __init__(self, *args, return_trace: bool = True, **kwargs):
        super().__init__(return_trace=return_trace)
        self.use_cache(bool(config['reader_use_cache']))

    @property
    def __cache_hash__(self):
        cache_hash = super().__cache_hash__
        if self.post_action is not None:
            cache_hash += f'@post_action:{_callable_cache_signature(self.post_action)}'
        return cache_hash

    def _lazy_load_data(self, *args, **load_kwargs) -> Iterable[DocNode]:
        raise NotImplementedError(f'{self.__class__.__name__} does not implement lazy_load_data method.')

    def _load_data(self, *args, **load_kwargs) -> List[DocNode]:
        return list(self._lazy_load_data(*args, **load_kwargs))

    def forward(self, *args, **kwargs) -> List[DocNode]:
        load_kwargs = {k: v for k, v in kwargs.items() if k not in _READER_CALL_SKIP_KEYS}
        r = self._load_data(*args, **load_kwargs)
        r = [r] if isinstance(r, DocNode) else [] if r is None else r
        if r and self.post_action:
            r = [x for sub in [self.post_action(n) for n in r] for x in (sub if isinstance(sub, list) else [sub])]
        return r

    @classmethod
    def detect_encoding(cls, file_path: Union[str, Path], fs: Optional['fsspec.AbstractFileSystem'] = None,  # noqa: C901
                        sample_size: int = 10000, use_cache: bool = True,
                        enable_chardet: bool = True) -> str:
        """检测文件的编码。

Args:
    file_path (str): 文件路径。
    fs (fsspec.AbstractFileSystem): 文件系统。
    sample_size (int): 样本大小。
    use_cache (bool): 是否使用缓存。
    enable_chardet (bool): 是否启用 chardet。

**Returns:**

- str: 文件的编码。


Examples:
    >>> import lazyllm
    >>> from lazyllm.tools.rag.readers import LazyLLMReaderBase
    >>> reader = LazyLLMReaderBase()
    >>> encoding = reader.detect_encoding("path/to/file.txt")
    >>> print(encoding)
    """
        if not isinstance(file_path, Path):
            file_path = Path(file_path)

        fs = fs or get_default_fs()

        cache_key = str(file_path) if use_cache else None
        if cache_key:
            with cls._cache_lock:
                if cache_key in cls._encoding_cache:
                    cached_encoding = cls._encoding_cache[cache_key]
                    return cached_encoding

        try:
            with fs.open(file_path, 'rb') as f:
                raw_data = f.read(sample_size)
        except Exception as e:
            LOG.warning(f'Failed to read file {file_path}: {e}')
            return 'utf-8'

        if not raw_data:
            return 'utf-8'

        bom_encodings = [
            (b'\xef\xbb\xbf', 'utf-8-sig'),
            (b'\xff\xfe\x00\x00', 'utf-32-le'),
            (b'\x00\x00\xfe\xff', 'utf-32-be'),
            (b'\xff\xfe', 'utf-16-le'),
            (b'\xfe\xff', 'utf-16-be'),
        ]

        for bom, encoding in bom_encodings:
            if raw_data.startswith(bom):
                cls._cache_encoding(cache_key, encoding)
                return encoding

        has_high_bytes = any(b > 127 for b in raw_data[:1000])

        if has_high_bytes:
            chinese_encodings = ['gb18030', 'gbk', 'gb2312', 'big5']
            # Prefer UTF-8 when valid; otherwise fall back to Chinese encodings.
            # Do not require Chinese chars in the first N chars — CSV headers are often long ASCII.
            if cls._try_decode(raw_data, 'utf-8'):
                cls._cache_encoding(cache_key, 'utf-8')
                return 'utf-8'

            for encoding in chinese_encodings:
                if cls._try_decode(raw_data, encoding):
                    cls._cache_encoding(cache_key, encoding)
                    return encoding
        else:
            primary_encodings = ['utf-8', 'gb18030', 'gbk', 'gb2312', 'big5']
            for encoding in primary_encodings:
                if cls._try_decode(raw_data, encoding):
                    cls._cache_encoding(cache_key, encoding)
                    return encoding

        if cls._try_decode(raw_data, 'latin-1'):
            cls._cache_encoding(cache_key, 'latin-1')
            return 'latin-1'

        if enable_chardet:
            try:
                detected = charset_normalizer.from_path(file_path).best().encoding
                if detected:
                    cls._cache_encoding(cache_key, detected)
                    return detected
                else:
                    LOG.warning(f'Charset normalizer detection failed: {detected}')
            except Exception as e:
                LOG.warning(f'Charset normalizer detection failed: {e}')

        try:
            system_encoding = locale.getpreferredencoding(False)
            LOG.warning(f'Using system default encoding {system_encoding} for {file_path}')
            cls._cache_encoding(cache_key, system_encoding)
            return system_encoding
        except Exception:
            pass
        LOG.warning(f'Could not detect encoding for {file_path}, using utf-8 as fallback')
        cls._cache_encoding(cache_key, 'utf-8')
        return 'utf-8'

    @staticmethod
    def _try_decode(data: bytes, encoding: str) -> bool:
        try:
            data.decode(encoding)
            return True
        except (UnicodeDecodeError, LookupError):
            return False

    @classmethod
    def _cache_encoding(cls, cache_key: Optional[str], encoding: str) -> None:
        if cache_key is None:
            return

        with cls._cache_lock:
            if len(cls._encoding_cache) >= cls._cache_max_size:
                old_keys = list(cls._encoding_cache.keys())[:100]
                for key in old_keys:
                    del cls._encoding_cache[key]
                LOG.debug(f'Encoding cache cleaned: removed {len(old_keys)} entries')

            cls._encoding_cache[cache_key] = encoding

    @classmethod
    def clear_encoding_cache(cls) -> None:
        """清空编码缓存。

Args:
    file_path (str): 文件路径。
    fs (fsspec.AbstractFileSystem): 文件系统。
    sample_size (int): 样本大小。
    use_cache (bool): 是否使用缓存。
    enable_chardet (bool): 是否启用 chardet。


Examples:
    >>> import lazyllm
    >>> from lazyllm.tools.rag.readers import LazyLLMReaderBase
    >>> reader = LazyLLMReaderBase()
    >>> reader.clear_encoding_cache()
    """
        with cls._cache_lock:
            cls._encoding_cache.clear()

    @classmethod
    def get_encoding_cache_stats(cls) -> dict:
        """获取编码缓存统计信息。

**Returns:**

- dict: 编码缓存统计信息。


Examples:
    >>> import lazyllm
    >>> from lazyllm.tools.rag.readers import LazyLLMReaderBase
    >>> reader = LazyLLMReaderBase()
    >>> stats = reader.get_encoding_cache_stats()
    >>> print(stats)
    """
        with cls._cache_lock:
            return {
                'cache_size': len(cls._encoding_cache),
                'cache_max_size': cls._cache_max_size,
                'usage_ratio': len(cls._encoding_cache) / cls._cache_max_size if cls._cache_max_size > 0 else 0
            }

clear_encoding_cache() classmethod

清空编码缓存。

Parameters:

  • file_path (str) –

    文件路径。

  • fs (AbstractFileSystem) –

    文件系统。

  • sample_size (int) –

    样本大小。

  • use_cache (bool) –

    是否使用缓存。

  • enable_chardet (bool) –

    是否启用 chardet。

Examples:

>>> import lazyllm
>>> from lazyllm.tools.rag.readers import LazyLLMReaderBase
>>> reader = LazyLLMReaderBase()
>>> reader.clear_encoding_cache()
Source code in lazyllm/tools/rag/readers/readerBase.py
    @classmethod
    def clear_encoding_cache(cls) -> None:
        """清空编码缓存。

Args:
    file_path (str): 文件路径。
    fs (fsspec.AbstractFileSystem): 文件系统。
    sample_size (int): 样本大小。
    use_cache (bool): 是否使用缓存。
    enable_chardet (bool): 是否启用 chardet。


Examples:
    >>> import lazyllm
    >>> from lazyllm.tools.rag.readers import LazyLLMReaderBase
    >>> reader = LazyLLMReaderBase()
    >>> reader.clear_encoding_cache()
    """
        with cls._cache_lock:
            cls._encoding_cache.clear()

detect_encoding(file_path, fs=None, sample_size=10000, use_cache=True, enable_chardet=True) classmethod

检测文件的编码。

Parameters:

  • file_path (str) –

    文件路径。

  • fs (AbstractFileSystem, default: None ) –

    文件系统。

  • sample_size (int, default: 10000 ) –

    样本大小。

  • use_cache (bool, default: True ) –

    是否使用缓存。

  • enable_chardet (bool, default: True ) –

    是否启用 chardet。

Returns:

  • str: 文件的编码。

Examples:

>>> import lazyllm
>>> from lazyllm.tools.rag.readers import LazyLLMReaderBase
>>> reader = LazyLLMReaderBase()
>>> encoding = reader.detect_encoding("path/to/file.txt")
>>> print(encoding)
Source code in lazyllm/tools/rag/readers/readerBase.py
    @classmethod
    def detect_encoding(cls, file_path: Union[str, Path], fs: Optional['fsspec.AbstractFileSystem'] = None,  # noqa: C901
                        sample_size: int = 10000, use_cache: bool = True,
                        enable_chardet: bool = True) -> str:
        """检测文件的编码。

Args:
    file_path (str): 文件路径。
    fs (fsspec.AbstractFileSystem): 文件系统。
    sample_size (int): 样本大小。
    use_cache (bool): 是否使用缓存。
    enable_chardet (bool): 是否启用 chardet。

**Returns:**

- str: 文件的编码。


Examples:
    >>> import lazyllm
    >>> from lazyllm.tools.rag.readers import LazyLLMReaderBase
    >>> reader = LazyLLMReaderBase()
    >>> encoding = reader.detect_encoding("path/to/file.txt")
    >>> print(encoding)
    """
        if not isinstance(file_path, Path):
            file_path = Path(file_path)

        fs = fs or get_default_fs()

        cache_key = str(file_path) if use_cache else None
        if cache_key:
            with cls._cache_lock:
                if cache_key in cls._encoding_cache:
                    cached_encoding = cls._encoding_cache[cache_key]
                    return cached_encoding

        try:
            with fs.open(file_path, 'rb') as f:
                raw_data = f.read(sample_size)
        except Exception as e:
            LOG.warning(f'Failed to read file {file_path}: {e}')
            return 'utf-8'

        if not raw_data:
            return 'utf-8'

        bom_encodings = [
            (b'\xef\xbb\xbf', 'utf-8-sig'),
            (b'\xff\xfe\x00\x00', 'utf-32-le'),
            (b'\x00\x00\xfe\xff', 'utf-32-be'),
            (b'\xff\xfe', 'utf-16-le'),
            (b'\xfe\xff', 'utf-16-be'),
        ]

        for bom, encoding in bom_encodings:
            if raw_data.startswith(bom):
                cls._cache_encoding(cache_key, encoding)
                return encoding

        has_high_bytes = any(b > 127 for b in raw_data[:1000])

        if has_high_bytes:
            chinese_encodings = ['gb18030', 'gbk', 'gb2312', 'big5']
            # Prefer UTF-8 when valid; otherwise fall back to Chinese encodings.
            # Do not require Chinese chars in the first N chars — CSV headers are often long ASCII.
            if cls._try_decode(raw_data, 'utf-8'):
                cls._cache_encoding(cache_key, 'utf-8')
                return 'utf-8'

            for encoding in chinese_encodings:
                if cls._try_decode(raw_data, encoding):
                    cls._cache_encoding(cache_key, encoding)
                    return encoding
        else:
            primary_encodings = ['utf-8', 'gb18030', 'gbk', 'gb2312', 'big5']
            for encoding in primary_encodings:
                if cls._try_decode(raw_data, encoding):
                    cls._cache_encoding(cache_key, encoding)
                    return encoding

        if cls._try_decode(raw_data, 'latin-1'):
            cls._cache_encoding(cache_key, 'latin-1')
            return 'latin-1'

        if enable_chardet:
            try:
                detected = charset_normalizer.from_path(file_path).best().encoding
                if detected:
                    cls._cache_encoding(cache_key, detected)
                    return detected
                else:
                    LOG.warning(f'Charset normalizer detection failed: {detected}')
            except Exception as e:
                LOG.warning(f'Charset normalizer detection failed: {e}')

        try:
            system_encoding = locale.getpreferredencoding(False)
            LOG.warning(f'Using system default encoding {system_encoding} for {file_path}')
            cls._cache_encoding(cache_key, system_encoding)
            return system_encoding
        except Exception:
            pass
        LOG.warning(f'Could not detect encoding for {file_path}, using utf-8 as fallback')
        cls._cache_encoding(cache_key, 'utf-8')
        return 'utf-8'

get_encoding_cache_stats() classmethod

获取编码缓存统计信息。

Returns:

  • dict: 编码缓存统计信息。

Examples:

>>> import lazyllm
>>> from lazyllm.tools.rag.readers import LazyLLMReaderBase
>>> reader = LazyLLMReaderBase()
>>> stats = reader.get_encoding_cache_stats()
>>> print(stats)
Source code in lazyllm/tools/rag/readers/readerBase.py
    @classmethod
    def get_encoding_cache_stats(cls) -> dict:
        """获取编码缓存统计信息。

**Returns:**

- dict: 编码缓存统计信息。


Examples:
    >>> import lazyllm
    >>> from lazyllm.tools.rag.readers import LazyLLMReaderBase
    >>> reader = LazyLLMReaderBase()
    >>> stats = reader.get_encoding_cache_stats()
    >>> print(stats)
    """
        with cls._cache_lock:
            return {
                'cache_size': len(cls._encoding_cache),
                'cache_max_size': cls._cache_max_size,
                'usage_ratio': len(cls._encoding_cache) / cls._cache_max_size if cls._cache_max_size > 0 else 0
            }

lazyllm.tools.rag.readers.readerBase.LazyLLMReaderBase

Bases: ModuleBase

基础文档读取器类,提供文档加载的基本接口。继承自 ModuleBase,使用 LazyLLMRegisterMetaClass 作为元类。

所有 Reader 在 reader(file, ...) 时可选启用算法端内容缓存:将解析完成后的 List[DocNode] 写入 ModuleCache,相同文件内容与 Reader 配置再次调用时直接返回缓存,跳过 _load_data 及下游 OCR 请求。

全局开关为 lazyllm.config['reader_use_cache'](环境变量 LAZYLLM_READER_USE_CACHE,默认 False)。

缓存与 OCR 服务端 use_cache 为两层独立机制:

  • 算法端内容缓存(本类):缓存 DocNode 列表,由 lazyllm.config['reader_use_cache'] 控制。
  • OCR 服务端缓存(如 MineruPDFReader):OCR Reader _load_data(..., use_cache=...),默认 True

缓存键由 Reader 类型、appendix_hash_key(子类配置,如 OCR URL/backend)、文件 mtimest_sizeextra_info 等调用参数共同决定;文件修改后(mtime/size 变化)自动 miss。

存储后端复用 ModuleCache,由全局配置选择:

  • LAZYLLM_CACHE_STRATEGYmemory(默认)/ file / sqlite / redis
  • LAZYLLM_CACHE_MODERW / RO / WO / NONE
  • LAZYLLM_CACHE_DIR:缓存根目录,默认 ~/.lazyllm/cachesqlite 策略下 db 为 {CACHE_DIR}/module/cache.db

Parameters:

  • *args

    位置参数,保留给子类或父类使用。

  • return_trace (bool, default: True ) –

    是否返回处理过程的追踪信息,默认为 True。

  • **kwargs

    关键字参数,保留给子类或父类使用。

Examples:

from lazyllm.tools.rag.readers.readerBase import LazyLLMReaderBase
from lazyllm.tools.rag.doc_node import DocNode
from typing import Iterable

class CustomReader(LazyLLMReaderBase):
    def _lazy_load_data(self, file_paths: list, **kwargs) -> Iterable[DocNode]:
        for file_path in file_paths:
            # Process each file and yield DocNode
            content = self._read_file(file_path)
            yield DocNode(
                text=content,
                metadata={"source": file_path}
            )

# Create reader instance
reader = CustomReader(return_trace=True)

# Load documents
documents = reader.forward(file_paths=["doc1.txt", "doc2.txt"])
Source code in lazyllm/tools/rag/readers/readerBase.py
class LazyLLMReaderBase(ModuleBase, metaclass=LazyLLMRegisterMetaClass):
    """
基础文档读取器类,提供文档加载的基本接口。继承自 ModuleBase,使用 LazyLLMRegisterMetaClass 作为元类。

所有 Reader 在 ``reader(file, ...)`` 时可选启用**算法端内容缓存**:将解析完成后的 ``List[DocNode]`` 写入
``ModuleCache``,相同文件内容与 Reader 配置再次调用时直接返回缓存,跳过 ``_load_data`` 及下游 OCR 请求。

全局开关为 ``lazyllm.config['reader_use_cache']``(环境变量 ``LAZYLLM_READER_USE_CACHE``,默认 False)。

缓存与 OCR 服务端 ``use_cache`` 为两层独立机制:

- **算法端内容缓存**(本类):缓存 ``DocNode`` 列表,由 ``lazyllm.config['reader_use_cache']`` 控制。
- **OCR 服务端缓存**(如 MineruPDFReader):OCR Reader ``_load_data(..., use_cache=...)``,默认 ``True``。

缓存键由 Reader 类型、``appendix_hash_key``(子类配置,如 OCR URL/backend)、文件
``mtime`` 与 ``st_size`` 及 ``extra_info`` 等调用参数共同决定;文件修改后(mtime/size 变化)自动 miss。

存储后端复用 ``ModuleCache``,由全局配置选择:

- ``LAZYLLM_CACHE_STRATEGY``:``memory``(默认)/ ``file`` / ``sqlite`` / ``redis``
- ``LAZYLLM_CACHE_MODE``:``RW`` / ``RO`` / ``WO`` / ``NONE``
- ``LAZYLLM_CACHE_DIR``:缓存根目录,默认 ``~/.lazyllm/cache``;``sqlite`` 策略下 db 为
  ``{CACHE_DIR}/module/cache.db``

Args:
    *args: 位置参数,保留给子类或父类使用。
    return_trace (bool): 是否返回处理过程的追踪信息,默认为 True。
    **kwargs: 关键字参数,保留给子类或父类使用。


Examples:

    from lazyllm.tools.rag.readers.readerBase import LazyLLMReaderBase
    from lazyllm.tools.rag.doc_node import DocNode
    from typing import Iterable

    class CustomReader(LazyLLMReaderBase):
        def _lazy_load_data(self, file_paths: list, **kwargs) -> Iterable[DocNode]:
            for file_path in file_paths:
                # Process each file and yield DocNode
                content = self._read_file(file_path)
                yield DocNode(
                    text=content,
                    metadata={"source": file_path}
                )

    # Create reader instance
    reader = CustomReader(return_trace=True)

    # Load documents
    documents = reader.forward(file_paths=["doc1.txt", "doc2.txt"])
    """
    post_action = None

    _encoding_cache = {}
    _cache_lock = threading.Lock()
    _cache_max_size = 1000

    def __init__(self, *args, return_trace: bool = True, **kwargs):
        super().__init__(return_trace=return_trace)
        self.use_cache(bool(config['reader_use_cache']))

    @property
    def __cache_hash__(self):
        cache_hash = super().__cache_hash__
        if self.post_action is not None:
            cache_hash += f'@post_action:{_callable_cache_signature(self.post_action)}'
        return cache_hash

    def _lazy_load_data(self, *args, **load_kwargs) -> Iterable[DocNode]:
        raise NotImplementedError(f'{self.__class__.__name__} does not implement lazy_load_data method.')

    def _load_data(self, *args, **load_kwargs) -> List[DocNode]:
        return list(self._lazy_load_data(*args, **load_kwargs))

    def forward(self, *args, **kwargs) -> List[DocNode]:
        load_kwargs = {k: v for k, v in kwargs.items() if k not in _READER_CALL_SKIP_KEYS}
        r = self._load_data(*args, **load_kwargs)
        r = [r] if isinstance(r, DocNode) else [] if r is None else r
        if r and self.post_action:
            r = [x for sub in [self.post_action(n) for n in r] for x in (sub if isinstance(sub, list) else [sub])]
        return r

    @classmethod
    def detect_encoding(cls, file_path: Union[str, Path], fs: Optional['fsspec.AbstractFileSystem'] = None,  # noqa: C901
                        sample_size: int = 10000, use_cache: bool = True,
                        enable_chardet: bool = True) -> str:
        """检测文件的编码。

Args:
    file_path (str): 文件路径。
    fs (fsspec.AbstractFileSystem): 文件系统。
    sample_size (int): 样本大小。
    use_cache (bool): 是否使用缓存。
    enable_chardet (bool): 是否启用 chardet。

**Returns:**

- str: 文件的编码。


Examples:
    >>> import lazyllm
    >>> from lazyllm.tools.rag.readers import LazyLLMReaderBase
    >>> reader = LazyLLMReaderBase()
    >>> encoding = reader.detect_encoding("path/to/file.txt")
    >>> print(encoding)
    """
        if not isinstance(file_path, Path):
            file_path = Path(file_path)

        fs = fs or get_default_fs()

        cache_key = str(file_path) if use_cache else None
        if cache_key:
            with cls._cache_lock:
                if cache_key in cls._encoding_cache:
                    cached_encoding = cls._encoding_cache[cache_key]
                    return cached_encoding

        try:
            with fs.open(file_path, 'rb') as f:
                raw_data = f.read(sample_size)
        except Exception as e:
            LOG.warning(f'Failed to read file {file_path}: {e}')
            return 'utf-8'

        if not raw_data:
            return 'utf-8'

        bom_encodings = [
            (b'\xef\xbb\xbf', 'utf-8-sig'),
            (b'\xff\xfe\x00\x00', 'utf-32-le'),
            (b'\x00\x00\xfe\xff', 'utf-32-be'),
            (b'\xff\xfe', 'utf-16-le'),
            (b'\xfe\xff', 'utf-16-be'),
        ]

        for bom, encoding in bom_encodings:
            if raw_data.startswith(bom):
                cls._cache_encoding(cache_key, encoding)
                return encoding

        has_high_bytes = any(b > 127 for b in raw_data[:1000])

        if has_high_bytes:
            chinese_encodings = ['gb18030', 'gbk', 'gb2312', 'big5']
            # Prefer UTF-8 when valid; otherwise fall back to Chinese encodings.
            # Do not require Chinese chars in the first N chars — CSV headers are often long ASCII.
            if cls._try_decode(raw_data, 'utf-8'):
                cls._cache_encoding(cache_key, 'utf-8')
                return 'utf-8'

            for encoding in chinese_encodings:
                if cls._try_decode(raw_data, encoding):
                    cls._cache_encoding(cache_key, encoding)
                    return encoding
        else:
            primary_encodings = ['utf-8', 'gb18030', 'gbk', 'gb2312', 'big5']
            for encoding in primary_encodings:
                if cls._try_decode(raw_data, encoding):
                    cls._cache_encoding(cache_key, encoding)
                    return encoding

        if cls._try_decode(raw_data, 'latin-1'):
            cls._cache_encoding(cache_key, 'latin-1')
            return 'latin-1'

        if enable_chardet:
            try:
                detected = charset_normalizer.from_path(file_path).best().encoding
                if detected:
                    cls._cache_encoding(cache_key, detected)
                    return detected
                else:
                    LOG.warning(f'Charset normalizer detection failed: {detected}')
            except Exception as e:
                LOG.warning(f'Charset normalizer detection failed: {e}')

        try:
            system_encoding = locale.getpreferredencoding(False)
            LOG.warning(f'Using system default encoding {system_encoding} for {file_path}')
            cls._cache_encoding(cache_key, system_encoding)
            return system_encoding
        except Exception:
            pass
        LOG.warning(f'Could not detect encoding for {file_path}, using utf-8 as fallback')
        cls._cache_encoding(cache_key, 'utf-8')
        return 'utf-8'

    @staticmethod
    def _try_decode(data: bytes, encoding: str) -> bool:
        try:
            data.decode(encoding)
            return True
        except (UnicodeDecodeError, LookupError):
            return False

    @classmethod
    def _cache_encoding(cls, cache_key: Optional[str], encoding: str) -> None:
        if cache_key is None:
            return

        with cls._cache_lock:
            if len(cls._encoding_cache) >= cls._cache_max_size:
                old_keys = list(cls._encoding_cache.keys())[:100]
                for key in old_keys:
                    del cls._encoding_cache[key]
                LOG.debug(f'Encoding cache cleaned: removed {len(old_keys)} entries')

            cls._encoding_cache[cache_key] = encoding

    @classmethod
    def clear_encoding_cache(cls) -> None:
        """清空编码缓存。

Args:
    file_path (str): 文件路径。
    fs (fsspec.AbstractFileSystem): 文件系统。
    sample_size (int): 样本大小。
    use_cache (bool): 是否使用缓存。
    enable_chardet (bool): 是否启用 chardet。


Examples:
    >>> import lazyllm
    >>> from lazyllm.tools.rag.readers import LazyLLMReaderBase
    >>> reader = LazyLLMReaderBase()
    >>> reader.clear_encoding_cache()
    """
        with cls._cache_lock:
            cls._encoding_cache.clear()

    @classmethod
    def get_encoding_cache_stats(cls) -> dict:
        """获取编码缓存统计信息。

**Returns:**

- dict: 编码缓存统计信息。


Examples:
    >>> import lazyllm
    >>> from lazyllm.tools.rag.readers import LazyLLMReaderBase
    >>> reader = LazyLLMReaderBase()
    >>> stats = reader.get_encoding_cache_stats()
    >>> print(stats)
    """
        with cls._cache_lock:
            return {
                'cache_size': len(cls._encoding_cache),
                'cache_max_size': cls._cache_max_size,
                'usage_ratio': len(cls._encoding_cache) / cls._cache_max_size if cls._cache_max_size > 0 else 0
            }

clear_encoding_cache() classmethod

清空编码缓存。

Parameters:

  • file_path (str) –

    文件路径。

  • fs (AbstractFileSystem) –

    文件系统。

  • sample_size (int) –

    样本大小。

  • use_cache (bool) –

    是否使用缓存。

  • enable_chardet (bool) –

    是否启用 chardet。

Examples:

>>> import lazyllm
>>> from lazyllm.tools.rag.readers import LazyLLMReaderBase
>>> reader = LazyLLMReaderBase()
>>> reader.clear_encoding_cache()
Source code in lazyllm/tools/rag/readers/readerBase.py
    @classmethod
    def clear_encoding_cache(cls) -> None:
        """清空编码缓存。

Args:
    file_path (str): 文件路径。
    fs (fsspec.AbstractFileSystem): 文件系统。
    sample_size (int): 样本大小。
    use_cache (bool): 是否使用缓存。
    enable_chardet (bool): 是否启用 chardet。


Examples:
    >>> import lazyllm
    >>> from lazyllm.tools.rag.readers import LazyLLMReaderBase
    >>> reader = LazyLLMReaderBase()
    >>> reader.clear_encoding_cache()
    """
        with cls._cache_lock:
            cls._encoding_cache.clear()

detect_encoding(file_path, fs=None, sample_size=10000, use_cache=True, enable_chardet=True) classmethod

检测文件的编码。

Parameters:

  • file_path (str) –

    文件路径。

  • fs (AbstractFileSystem, default: None ) –

    文件系统。

  • sample_size (int, default: 10000 ) –

    样本大小。

  • use_cache (bool, default: True ) –

    是否使用缓存。

  • enable_chardet (bool, default: True ) –

    是否启用 chardet。

Returns:

  • str: 文件的编码。

Examples:

>>> import lazyllm
>>> from lazyllm.tools.rag.readers import LazyLLMReaderBase
>>> reader = LazyLLMReaderBase()
>>> encoding = reader.detect_encoding("path/to/file.txt")
>>> print(encoding)
Source code in lazyllm/tools/rag/readers/readerBase.py
    @classmethod
    def detect_encoding(cls, file_path: Union[str, Path], fs: Optional['fsspec.AbstractFileSystem'] = None,  # noqa: C901
                        sample_size: int = 10000, use_cache: bool = True,
                        enable_chardet: bool = True) -> str:
        """检测文件的编码。

Args:
    file_path (str): 文件路径。
    fs (fsspec.AbstractFileSystem): 文件系统。
    sample_size (int): 样本大小。
    use_cache (bool): 是否使用缓存。
    enable_chardet (bool): 是否启用 chardet。

**Returns:**

- str: 文件的编码。


Examples:
    >>> import lazyllm
    >>> from lazyllm.tools.rag.readers import LazyLLMReaderBase
    >>> reader = LazyLLMReaderBase()
    >>> encoding = reader.detect_encoding("path/to/file.txt")
    >>> print(encoding)
    """
        if not isinstance(file_path, Path):
            file_path = Path(file_path)

        fs = fs or get_default_fs()

        cache_key = str(file_path) if use_cache else None
        if cache_key:
            with cls._cache_lock:
                if cache_key in cls._encoding_cache:
                    cached_encoding = cls._encoding_cache[cache_key]
                    return cached_encoding

        try:
            with fs.open(file_path, 'rb') as f:
                raw_data = f.read(sample_size)
        except Exception as e:
            LOG.warning(f'Failed to read file {file_path}: {e}')
            return 'utf-8'

        if not raw_data:
            return 'utf-8'

        bom_encodings = [
            (b'\xef\xbb\xbf', 'utf-8-sig'),
            (b'\xff\xfe\x00\x00', 'utf-32-le'),
            (b'\x00\x00\xfe\xff', 'utf-32-be'),
            (b'\xff\xfe', 'utf-16-le'),
            (b'\xfe\xff', 'utf-16-be'),
        ]

        for bom, encoding in bom_encodings:
            if raw_data.startswith(bom):
                cls._cache_encoding(cache_key, encoding)
                return encoding

        has_high_bytes = any(b > 127 for b in raw_data[:1000])

        if has_high_bytes:
            chinese_encodings = ['gb18030', 'gbk', 'gb2312', 'big5']
            # Prefer UTF-8 when valid; otherwise fall back to Chinese encodings.
            # Do not require Chinese chars in the first N chars — CSV headers are often long ASCII.
            if cls._try_decode(raw_data, 'utf-8'):
                cls._cache_encoding(cache_key, 'utf-8')
                return 'utf-8'

            for encoding in chinese_encodings:
                if cls._try_decode(raw_data, encoding):
                    cls._cache_encoding(cache_key, encoding)
                    return encoding
        else:
            primary_encodings = ['utf-8', 'gb18030', 'gbk', 'gb2312', 'big5']
            for encoding in primary_encodings:
                if cls._try_decode(raw_data, encoding):
                    cls._cache_encoding(cache_key, encoding)
                    return encoding

        if cls._try_decode(raw_data, 'latin-1'):
            cls._cache_encoding(cache_key, 'latin-1')
            return 'latin-1'

        if enable_chardet:
            try:
                detected = charset_normalizer.from_path(file_path).best().encoding
                if detected:
                    cls._cache_encoding(cache_key, detected)
                    return detected
                else:
                    LOG.warning(f'Charset normalizer detection failed: {detected}')
            except Exception as e:
                LOG.warning(f'Charset normalizer detection failed: {e}')

        try:
            system_encoding = locale.getpreferredencoding(False)
            LOG.warning(f'Using system default encoding {system_encoding} for {file_path}')
            cls._cache_encoding(cache_key, system_encoding)
            return system_encoding
        except Exception:
            pass
        LOG.warning(f'Could not detect encoding for {file_path}, using utf-8 as fallback')
        cls._cache_encoding(cache_key, 'utf-8')
        return 'utf-8'

get_encoding_cache_stats() classmethod

获取编码缓存统计信息。

Returns:

  • dict: 编码缓存统计信息。

Examples:

>>> import lazyllm
>>> from lazyllm.tools.rag.readers import LazyLLMReaderBase
>>> reader = LazyLLMReaderBase()
>>> stats = reader.get_encoding_cache_stats()
>>> print(stats)
Source code in lazyllm/tools/rag/readers/readerBase.py
    @classmethod
    def get_encoding_cache_stats(cls) -> dict:
        """获取编码缓存统计信息。

**Returns:**

- dict: 编码缓存统计信息。


Examples:
    >>> import lazyllm
    >>> from lazyllm.tools.rag.readers import LazyLLMReaderBase
    >>> reader = LazyLLMReaderBase()
    >>> stats = reader.get_encoding_cache_stats()
    >>> print(stats)
    """
        with cls._cache_lock:
            return {
                'cache_size': len(cls._encoding_cache),
                'cache_max_size': cls._cache_max_size,
                'usage_ratio': len(cls._encoding_cache) / cls._cache_max_size if cls._cache_max_size > 0 else 0
            }

lazyllm.tools.rag.readers.readerBase.TxtReader

Bases: LazyLLMReaderBase

TxtReader 类用于从文本文件中加载内容,并将其封装为 DocNode 对象列表。

该类继承自 LazyLLMReaderBase,主要功能包括:

  • 支持指定文本编码读取文件;
  • 可选返回加载过程的跟踪信息;
  • 继承算法端内容缓存(lazyllm.config['reader_use_cache'])。

Parameters:

  • encoding (str, default: None ) –

    文件读取的文本编码,默认值为 'utf-8'。

  • return_trace (bool, default: True ) –

    是否返回加载过程的跟踪信息,默认值为 True。

  • auto_detect_encoding (bool, default: config['auto_detect_encoding'] ) –

    是否自动检测编码,默认读取 LAZYLLM_AUTO_DETECT_ENCODING

  • enable_chardet (bool, default: config['enable_chardet'] ) –

    检测编码时是否启用 chardet,默认读取 LAZYLLM_ENABLE_CHARDET

  • use_encoding_cache (bool, default: config['use_encoding_cache'] ) –

    是否缓存编码检测结果(类级缓存,与内容缓存无关)。

Source code in lazyllm/tools/rag/readers/readerBase.py
class TxtReader(LazyLLMReaderBase):
    """TxtReader 类用于从文本文件中加载内容,并将其封装为 `DocNode` 对象列表。

该类继承自 `LazyLLMReaderBase`,主要功能包括:

- 支持指定文本编码读取文件;
- 可选返回加载过程的跟踪信息;
- 继承算法端内容缓存(``lazyllm.config['reader_use_cache']``)。

Args:
    encoding (str): 文件读取的文本编码,默认值为 'utf-8'。
    return_trace (bool): 是否返回加载过程的跟踪信息,默认值为 True。
    auto_detect_encoding (bool): 是否自动检测编码,默认读取 ``LAZYLLM_AUTO_DETECT_ENCODING``。
    enable_chardet (bool): 检测编码时是否启用 chardet,默认读取 ``LAZYLLM_ENABLE_CHARDET``。
    use_encoding_cache (bool): 是否缓存编码检测结果(类级缓存,与内容缓存无关)。
"""
    def __init__(self, encoding: Optional[str] = None, return_trace: bool = True,
                 auto_detect_encoding: bool = config['auto_detect_encoding'],
                 enable_chardet: bool = config['enable_chardet'],
                 use_encoding_cache: bool = config['use_encoding_cache']) -> None:
        super().__init__(return_trace=return_trace)
        self._encoding = encoding
        self._auto_detect_encoding = auto_detect_encoding
        self._enable_chardet = enable_chardet
        self._use_encoding_cache = use_encoding_cache

    @property
    def appendix_hash_key(self):
        return f'{self._encoding}|{self._auto_detect_encoding}|{self._enable_chardet}'

    def _load_data(self, file: Path, fs: Optional['fsspec.AbstractFileSystem'] = None) -> List[DocNode]:
        if self._encoding:
            encoding = self._encoding
        elif self._auto_detect_encoding:
            encoding = self.detect_encoding(
                file, fs,
                use_cache=self._use_encoding_cache,
                enable_chardet=self._enable_chardet
            )
        else:
            encoding = 'utf-8'

        try:
            with (fs or get_default_fs()).open(file, mode='r', encoding=encoding) as f:
                content = f.read()
            return [DocNode(text=content)]
        except Exception:
            if not self._auto_detect_encoding and self._encoding:
                try:
                    detected_encoding = self.detect_encoding(
                        file, fs,
                        use_cache=self._use_encoding_cache,
                        enable_chardet=self._enable_chardet
                    )
                    with (fs or get_default_fs()).open(file, mode='r', encoding=detected_encoding) as f:
                        content = f.read()
                    return [DocNode(text=content)]
                except Exception as e:
                    LOG.error(f'Auto-detection also failed for {file}: {e}')
            elif self._auto_detect_encoding and self._enable_chardet:
                try:
                    detected = charset_normalizer.from_path(file).best()
                    if detected and detected.encoding and detected.encoding.lower() != encoding.lower():
                        with (fs or get_default_fs()).open(file, mode='r', encoding=detected.encoding) as f:
                            content = f.read()
                        return [DocNode(text=content)]
                except Exception as e2:
                    LOG.error(f'charset_normalizer also failed for {file}: {e2}')
            raise

lazyllm.tools.rag.readers.PandasExcelReader

Bases: LazyLLMReaderBase

用于读取 Excel 文件(.xlsx),并将内容提取为文本。

Parameters:

  • concat_rows (bool, default: True ) –

    是否将所有行拼接为一个文本块。

  • sheet_name (Optional[str], default: None ) –

    要读取的工作表名称。若为 None,则读取所有工作表。

  • pandas_config (Optional[Dict], default: None ) –

    pandas.read_excel 的可选配置项。

  • fill_method (Optional[str], default: 'fillna' ) –

    缺失值填充策略,可选 'fillna'(default) / 'ffill' / 'bfill'。

  • return_trace (bool, default: True ) –

    是否返回处理过程的 trace。

  • col_joiner (str, default: ' ' ) –

    列之间的连接符,默认为空格。

Source code in lazyllm/tools/rag/readers/pandasReader.py
class PandasExcelReader(LazyLLMReaderBase):
    """用于读取 Excel 文件(.xlsx),并将内容提取为文本。

Args:
    concat_rows (bool): 是否将所有行拼接为一个文本块。
    sheet_name (Optional[str]): 要读取的工作表名称。若为 None,则读取所有工作表。
    pandas_config (Optional[Dict]): pandas.read_excel 的可选配置项。
    fill_method (Optional[str]): 缺失值填充策略,可选 'fillna'(default) / 'ffill' / 'bfill'。
    return_trace (bool): 是否返回处理过程的 trace。
    col_joiner (str): 列之间的连接符,默认为空格。
"""
    def __init__(self, concat_rows: bool = True, sheet_name: Optional[str] = None,
                 pandas_config: Optional[Dict] = None, fill_method: Optional[str] = 'fillna',
                 return_trace: bool = True, col_joiner: str = ' ') -> None:
        super().__init__(return_trace=return_trace)
        self._concat_rows = concat_rows
        self._sheet_name = sheet_name
        self._pandas_config = pandas_config or {}
        self._fill_method = fill_method
        self._col_joiner = col_joiner

    def _load_data(self, file: Path, fs: Optional['fsspec.AbstractFileSystem'] = None) -> List[DocNode]:
        openpyxl_spec = importlib.util.find_spec('openpyxl')
        if openpyxl_spec is not None: pass
        else: raise ImportError('Please install openpyxl to read Excel files. '
                                'You can install it with `pip install openpyxl`')

        if not isinstance(file, Path): file = Path(file)
        if fs:
            with fs.open(file) as f:
                dfs = pd.read_excel(f, self._sheet_name, **self._pandas_config)
        else:
            dfs = pd.read_excel(file, self._sheet_name, **self._pandas_config)

        def process_df(df: pd.DataFrame) -> List[DocNode]:
            df = _apply_fill(df, self._fill_method)
            text_list = df.astype(str).apply(lambda row: self._col_joiner.join(row.values), axis=1).tolist()

            if self._concat_rows:
                return [DocNode(text='\n'.join(text_list))]
            return [DocNode(text=text) for text in text_list]

        dfs_list = [dfs] if isinstance(dfs, pd.DataFrame) else dfs.values()
        return [doc for df in dfs_list for doc in process_df(df)]

lazyllm.tools.rag.readers.PDFReader

Bases: _RichReader

用于读取 PDF 文件并提取其中的文本内容。

Parameters:

  • split_doc (bool, default: True ) –

    若为 True(默认),则解析为一个 RichDocNode,可以搭配 RichTransform 解析出带有页信息的节点; 若为 False,则解析为一个纯文本的 DocNode

  • post_func (Optional[Callable[[List[DocNode]], List[DocNode]]], default: None ) –

    结果后处理函数, 需返回 List[DocNode],并会将 extra_info 写入每个节点的 global_metadata

  • return_trace (bool, default: True ) –

    是否返回处理过程的 trace,默认为 True。

  • return_full_document ((bool, 已弃用), default: None ) –

    此参数将在未来版本中删除,请使用 split_doc 替代。

Notes

split_doc=True 时返回 RichDocNode,否则返回 DocNode,两种情况都只返回一个节点。 当 split_doc=True 时,强烈建议搭配 RichTransform 使用,可以解析出带有页信息等 metadata 的节点; 如不使用 RichTransform,则解析出的节点会回退为纯文本节点。

Source code in lazyllm/tools/rag/readers/pdfReader.py
class PDFReader(_RichReader):
    """用于读取 PDF 文件并提取其中的文本内容。

Args:
    split_doc (bool): 若为 True(默认),则解析为一个 `RichDocNode`,可以搭配 `RichTransform` 解析出带有页信息的节点;
        若为 False,则解析为一个纯文本的 `DocNode`。
    post_func (Optional[Callable[[List[DocNode]], List[DocNode]]]): 结果后处理函数,
        需返回 `List[DocNode]`,并会将 `extra_info` 写入每个节点的 `global_metadata`。
    return_trace (bool): 是否返回处理过程的 trace,默认为 True。
    return_full_document (bool, 已弃用): 此参数将在未来版本中删除,请使用 `split_doc` 替代。

Notes:
    当 `split_doc=True` 时返回 `RichDocNode`,否则返回 `DocNode`,两种情况都只返回一个节点。
    当 `split_doc=True` 时,强烈建议搭配 `RichTransform` 使用,可以解析出带有页信息等 metadata 的节点;
    如不使用 `RichTransform`,则解析出的节点会回退为纯文本节点。
"""
    def __init__(self, split_doc: bool = True,
                 post_func: Optional[Callable[[List[DocNode]], List[DocNode]]] = None,
                 return_trace: bool = True, *, return_full_document=None) -> None:
        if return_full_document is not None:
            LOG.warning('return_full_document is deprecated, please use split_doc instead')
            assert split_doc ^ return_full_document, \
                'split_doc and return_full_document cannot be both True or False'
            split_doc = not return_full_document
        super().__init__(post_func=post_func, split_doc=split_doc, return_trace=return_trace)

    @retry(stop_after_attempt=3)
    def _load_data(self, file: Path, fs: Optional['fsspec.AbstractFileSystem'] = None) -> List[DocNode]:
        if not isinstance(file, Path): file = Path(file)

        fs = fs or get_default_fs()
        with fs.open(file, 'rb') as fp:
            stream = fp if is_default_fs(fs) else io.BytesIO(fp.read())
            pdf = pypdf.PdfReader(stream)
            num_pages = len(pdf.pages)
            docs = []
            for page in range(num_pages):
                page_text = pdf.pages[page].extract_text()
                page_label = pdf.page_labels[page]
                metadata = {'page_label': page_label}
                docs.append(DocNode(text=page_text, metadata=metadata))
            return docs

lazyllm.tools.rag.readers.PPTXReader

Bases: LazyLLMReaderBase

用于解析 PPTX(PowerPoint)文件的读取器,能够提取幻灯片中的文本,并对嵌入图像进行视觉描述生成。

Parameters:

  • return_trace (bool, default: True ) –

    是否记录处理过程的 trace,默认为 True。

Source code in lazyllm/tools/rag/readers/pptxReader.py
class PPTXReader(LazyLLMReaderBase):
    """用于解析 PPTX(PowerPoint)文件的读取器,能够提取幻灯片中的文本,并对嵌入图像进行视觉描述生成。

Args:
    return_trace (bool): 是否记录处理过程的 trace,默认为 True。
"""
    def __init__(self, return_trace: bool = True) -> None:
        try:
            thirdparty.check_packages(['python-pptx', 'torch', 'Pillow', 'transformers'])
        except ImportError:
            raise ImportError('Please install extra dependencies that are required for the '
                              'PPTXReader: `pip install torch transformers python-pptx Pillow`')

        super().__init__(return_trace=return_trace)
        model = tf.VisionEncoderDecoderModel.from_pretrained('nlpconnect/vit-gpt2-image-captioning')
        feature_extractor = tf.ViTFeatureExtractor.from_pretrained('nlpconnect/vit-gpt2-image-captioning')
        tokenizer = tf.AutoTokenizer.from_pretrained('nlpconnect/vit-gpt2-image-captioning')

        self._parser_config = {'feature_extractor': feature_extractor, 'model': model, 'tokenizer': tokenizer}

    def _caption_image(self, tmp_image_file: str) -> str:
        from PIL import Image

        model = self._parser_config['model']
        feature_extractor = self._parser_config['feature_extractor']
        tokenizer = self._parser_config['tokenizer']

        device = infer_torch_device()
        model.to(device)

        max_length = 16
        num_beams = 4
        gen_kwargs = {'max_length': max_length, 'num_beams': num_beams}

        i_image = Image.open(tmp_image_file)
        if i_image.mode != 'RGB': i_image = i_image.convert(mode='RGB')

        pixel_values = feature_extractor(images=[i_image], return_tensors='pt').pixel_values
        pixel_values = pixel_values.to(device)

        output_ids = model.generate(pixel_values, **gen_kwargs)

        preds = tokenizer.batch_decode(output_ids, skip_special_tokens=True)
        return preds[0].strip()

    def _load_data(self, file: Path, fs: Optional['fsspec.AbstractFileSystem'] = None) -> List[DocNode]:
        if not isinstance(file, Path): file = Path(file)

        if fs:
            with fs.open(file) as f:
                presentation = pptx.Presentation(f)
        else:
            presentation = pptx.Presentation(file)

        result = ''
        for i, slide in enumerate(presentation.slides):
            result += f'\n\nSlide #{i}: \n'
            for shape in slide.shapes:
                if hasattr(shape, 'image'):
                    image = shape.image
                    image_bytes = image.blob
                    f = tempfile.NamedTemporaryFile('wb', delete=False)
                    try:
                        f.write(image_bytes)
                        f.close()
                        result += f'\n Image: {self._caption_image(f.name)}\n\n'
                    finally:
                        os.unlink(f.name)

                if hasattr(shape, 'text'): result += f'{shape.text}\n'
        return [DocNode(text=result)]

lazyllm.tools.rag.readers.VideoAudioReader

Bases: LazyLLMReaderBase

用于从视频或音频文件中提取语音内容的读取器,依赖 OpenAI 的 Whisper 模型进行语音识别。

Parameters:

  • model_version (str, default: 'base' ) –

    Whisper 模型的版本(如 "base", "small", "medium", "large"),默认为 "base"。

  • return_trace (bool, default: True ) –

    是否返回处理过程的 trace,默认为 True。

Source code in lazyllm/tools/rag/readers/videoAudioReader.py
class VideoAudioReader(LazyLLMReaderBase):
    """用于从视频或音频文件中提取语音内容的读取器,依赖 OpenAI 的 Whisper 模型进行语音识别。

Args:
    model_version (str): Whisper 模型的版本(如 "base", "small", "medium", "large"),默认为 "base"。
    return_trace (bool): 是否返回处理过程的 trace,默认为 True。
"""
    def __init__(self, model_version: str = 'base', return_trace: bool = True) -> None:
        super().__init__(return_trace=return_trace)
        self._model_version = model_version

        try:
            import whisper
        except ImportError:
            raise ImportError('Please install OpenAI whisper model '
                              '`pip install openai-whisper` to use the model')

        model = whisper.load_model(self._model_version)
        self._parser_config = {'model': model}

    def _load_data(self, file: Path, fs: Optional['fsspec.AbstractFileSystem'] = None) -> List[DocNode]:
        import whisper

        if not isinstance(file, Path): file = Path(file)

        if file.name.endswith('mp4'):
            try:
                from pydub import AudioSegment
            except ImportError:
                raise ImportError('Please install pydub `pip install pydub`')

            if fs:
                with fs.open(file, 'rb') as f:
                    video = AudioSegment.from_file(f, format='mp4')
            else:
                video = AudioSegment.from_file(file, format='mp4')

            audio = video.split_to_mono()[0]
            file_str = str(file)[:-4] + '.mp3'
            audio.export(file_str, format='mp3')

        model = cast(whisper.Whisper, self._parser_config['model'])
        result = model.transcribe(str(file))

        transcript = result['text']
        return [DocNode(text=transcript)]

lazyllm.tools.SqlManager

Bases: DBManager

SqlManager是与数据库进行交互的专用工具。它提供了连接数据库,设置、创建、检查数据表,插入数据,执行查询的方法。

Parameters:

  • db_type (str) –

    数据库类型,支持: postgresql, mysql, mssql, sqlite, mysql+pymysql

  • user (str) –

    数据库用户名

  • password (str) –

    数据库密码

  • host (str) –

    数据库主机地址

  • port (int) –

    数据库端口号

  • db_name (str) –

    数据库名称

  • options_str (str, default: None ) –

    连接选项字符串,默认为None

  • tables_info_dict (Dict, default: None ) –

    表结构信息字典,用于初始化表结构,默认为None

Source code in lazyllm/tools/sql/sql_manager.py
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
class SqlManager(DBManager):
    """SqlManager是与数据库进行交互的专用工具。它提供了连接数据库,设置、创建、检查数据表,插入数据,执行查询的方法。

Args:
    db_type (str): 数据库类型,支持: postgresql, mysql, mssql, sqlite, mysql+pymysql
    user (str): 数据库用户名
    password (str): 数据库密码
    host (str): 数据库主机地址
    port (int): 数据库端口号
    db_name (str): 数据库名称
    options_str (str, optional): 连接选项字符串,默认为None
    tables_info_dict (Dict, optional): 表结构信息字典,用于初始化表结构,默认为None
"""
    DB_TYPE_SUPPORTED = set(['postgresql', 'mysql', 'mssql', 'sqlite', 'mysql+pymysql', 'tidb'])
    DB_DRIVER_MAP = {'mysql': 'pymysql', 'tidb': 'pymysql'}
    PYTYPE_TO_SQL_MAP = {
        'integer': sqlalchemy.Integer,
        'string': sqlalchemy.Text,
        'text': sqlalchemy.Text,
        'boolean': sqlalchemy.Boolean,
        'float': sqlalchemy.Float,
        'datetime': sqlalchemy.DateTime,
        'bytes': sqlalchemy.LargeBinary,
        'bool': sqlalchemy.Boolean,
        'date': sqlalchemy.Date,
        'time': sqlalchemy.Time,
        'list': sqlalchemy.ARRAY,
        'dict': sqlalchemy.JSON,
        'uuid': sqlalchemy.Uuid,
    }

    def __init__(self, db_type: str, user: str, password: str, host: str, port: int, db_name: str, *,
                 options_str: str = None, tables_info_dict: Dict = None):
        db_type = db_type.lower()
        if db_type not in self.DB_TYPE_SUPPORTED:
            raise ValueError(f'{db_type} not supported')
        super().__init__(db_type)
        self._user = user
        self._password = password
        self._host = host
        self._port = port
        self._db_name = db_name
        self._tables_desc_dict = {}
        self._visible_tables = None
        self._metadata = sqlalchemy.MetaData()
        self._options_str = options_str
        self._orm_cache = {}
        self._engine = None
        self._Session = None
        if tables_info_dict:
            self._init_tables_by_info(tables_info_dict)

    def _init_tables_by_info(self, tables_info_dict):
        try:
            tables_info = TablesInfo.model_validate(tables_info_dict)
            self._visible_tables = [table_info.name for table_info in tables_info.tables]
            # create table if not exist
            self._create_tables_by_info(tables_info)
            desc_dict = self._gen_desc_by_info(tables_info)
            self.set_desc(desc_dict)
        except pydantic.ValidationError as e:
            raise ValueError(f'Validate tables_info_dict failed: {str(e)}')

    def _sql_type_for(self, py_type: str, *, is_primary_key: bool = False):
        t = py_type.lower()
        if self._db_type in ('mysql', 'tidb', 'mysql+pymysql'):
            # MySQL/TiDB do not allow TEXT/BLOB columns to be used as primary keys
            # without a prefix length. Use VARCHAR for identifier-like key columns.
            if is_primary_key and t in ('string', 'text'):
                return sqlalchemy.String(255)
            if t == 'list':
                return sqlalchemy.JSON
            if t == 'uuid':
                return sqlalchemy.String(36)
        return self.PYTYPE_TO_SQL_MAP.get(t, sqlalchemy.Text)

    def _create_tables_by_info(self, tables_info: TablesInfo):
        for table_info in tables_info.tables:
            attrs = {'__tablename__': table_info.name, '__table_args__': {'extend_existing': True},
                     'metadata': self._metadata}
            for column_info in table_info.columns:
                column_type = column_info.data_type.lower()
                is_nullable = column_info.nullable
                column_name = column_info.name
                is_primary = column_info.is_primary_key
                default_value = column_info.default
                # Keep cross-db compatibility while handling MySQL/TiDB PK restrictions.
                real_type = self._sql_type_for(column_type, is_primary_key=is_primary)
                # Handle default value
                # For non-integer primary keys, disable autoincrement so SQLAlchemy
                # always includes the column in INSERT statements.
                autoincrement = 'auto' if not is_primary else (
                    'auto' if column_type in ('integer', 'int') else False
                )
                if default_value is not None:
                    attrs[column_name] = sqlalchemy.Column(real_type, nullable=is_nullable,
                                                           primary_key=is_primary, default=default_value,
                                                           autoincrement=autoincrement)
                else:
                    attrs[column_name] = sqlalchemy.Column(real_type, nullable=is_nullable,
                                                           primary_key=is_primary,
                                                           autoincrement=autoincrement)
            # When create dynamic class with same name, old version will be replaced
            TableClass = type(table_info.name.capitalize(), (TableBase,), attrs)
            self.create_table(TableClass)

    def _gen_desc_by_info(self, tables_info: TablesInfo) -> dict:
        desc_dict = {}
        for table_info in tables_info.tables:
            table_comment = ''
            if table_info.comment:
                table_comment += f'COMMENT ON TABLE "{table_info.name}": {table_info.comment}\n'
            for column_info in table_info.columns:
                table_comment += f'COMMENT ON COLUMN "{table_info.name}.{column_info.name}": {column_info.comment}\n'
            if table_comment:
                desc_dict[table_info.name] = table_comment
        return desc_dict

    def _gen_conn_url(self, db_name: str = None) -> str:
        db_name = self._db_name if db_name is None else db_name
        if self._db_type == 'sqlite':
            conn_url = f'sqlite:///{db_name}{("?" + self._options_str) if self._options_str else ""}'
        else:
            driver = self.DB_DRIVER_MAP.get(self._db_type if self._db_type != 'tidb' else 'mysql', '')
            password = quote_plus(self._password)
            prefix = 'mysql' if self._db_type == 'tidb' else self._db_type
            db_path = f'/{db_name}' if db_name else '/'
            conn_url = (f'{prefix}{("+" + driver) if driver else ""}://{self._user}:{password}@{self._host}'
                        f':{self._port}{db_path}{("?" + self._options_str) if self._options_str else ""}')
        return conn_url

    def _mysql_engine_kwargs(self) -> dict:
        kwargs = {
            'pool_size': 10,
            'max_overflow': 20,
            'pool_pre_ping': True,
        }
        if self._db_type == 'tidb':
            kwargs.update({'pool_recycle': 300, 'connect_args': {}, 'echo': False})
        else:
            kwargs.update({'pool_recycle': 3600})
        return kwargs

    @staticmethod
    def _default_engine_kwargs() -> dict:
        return {
            'pool_size': 10,
            'max_overflow': 20,
            'pool_pre_ping': True,
            'pool_recycle': 3600,
        }

    @staticmethod
    def _get_operational_error_code(error: OperationalError):
        args = getattr(getattr(error, 'orig', None), 'args', ())
        return args[0] if args else None

    @staticmethod
    def _get_operational_error_pgcode(error: OperationalError):
        orig = getattr(error, 'orig', None)
        return getattr(orig, 'pgcode', None) or getattr(orig, 'sqlstate', None)

    def _is_database_not_found_error(self, error: OperationalError) -> bool:
        if self._db_type in ('mysql', 'mysql+pymysql', 'tidb'):
            return self._get_operational_error_code(error) == 1049
        if self._db_type == 'postgresql':
            if self._get_operational_error_pgcode(error) == '3D000':
                return True
            error_msg = str(getattr(error, 'orig', error)).lower()
            return 'does not exist' in error_msg and 'database' in error_msg
        return False

    def _ensure_database_exists(self, conn_url: str):
        if self._db_type not in ('mysql', 'mysql+pymysql', 'tidb', 'postgresql'):
            return
        engine_kwargs = self._mysql_engine_kwargs() if self._db_type in ('mysql', 'mysql+pymysql', 'tidb') \
            else self._default_engine_kwargs()
        probe_engine = sqlalchemy.create_engine(conn_url, **engine_kwargs)
        try:
            with probe_engine.connect():
                return
        except OperationalError as e:
            if not self._is_database_not_found_error(e):
                raise
        finally:
            probe_engine.dispose()

        if self._db_type == 'postgresql':
            admin_engine = sqlalchemy.create_engine(
                self._gen_conn_url('postgres'),
                isolation_level='AUTOCOMMIT',
                **self._default_engine_kwargs()
            )
        else:
            admin_engine = sqlalchemy.create_engine(self._gen_conn_url(''), **self._mysql_engine_kwargs())
        try:
            with admin_engine.connect() as conn:
                if self._db_type == 'postgresql':
                    exists = conn.execute(
                        sqlalchemy.text('SELECT 1 FROM pg_database WHERE datname = :db_name'),
                        {'db_name': self._db_name}
                    ).scalar()
                    if not exists:
                        escaped_db_name = self._db_name.replace('"', '""')
                        conn.execute(sqlalchemy.text(f'CREATE DATABASE "{escaped_db_name}"'))
                else:
                    escaped_db_name = self._db_name.replace('`', '``')
                    conn.execute(sqlalchemy.text(f'CREATE DATABASE IF NOT EXISTS `{escaped_db_name}`'))
                    conn.commit()
        finally:
            admin_engine.dispose()

    @property
    def engine(self):
        if self._engine is None:
            conn_url = self._gen_conn_url()
            if self._db_type == 'sqlite':
                self._engine = sqlalchemy.create_engine(
                    conn_url,
                    connect_args={'check_same_thread': False, 'timeout': 30},
                    poolclass=sqlalchemy.pool.QueuePool,
                    echo=False
                )
                with self._engine.connect() as conn:
                    conn.execute(sqlalchemy.text('PRAGMA journal_mode=WAL'))
                    conn.execute(sqlalchemy.text('PRAGMA synchronous=NORMAL'))
                    conn.execute(sqlalchemy.text('PRAGMA busy_timeout=30000'))
                    conn.commit()
            elif self._db_type in ('mysql', 'mysql+pymysql', 'tidb'):
                self._ensure_database_exists(conn_url)
                self._engine = sqlalchemy.create_engine(conn_url, **self._mysql_engine_kwargs())
            elif self._db_type == 'postgresql':
                self._ensure_database_exists(conn_url)
                self._engine = sqlalchemy.create_engine(conn_url, **self._default_engine_kwargs())
            else:
                self._engine = sqlalchemy.create_engine(conn_url, **self._default_engine_kwargs())
        return self._engine

    @property
    def Session(self):
        if self._Session is None:
            self._Session = sessionmaker(bind=self.engine, expire_on_commit=False)
        return self._Session

    def dispose(self):
        """Release the underlying engine's connection pool.

        Needed so callers (e.g. DocServer._Impl.stop) can close sqlite file handles
        before removing the containing directory; on Windows this is required, because
        open handles block ``TemporaryDirectory`` cleanup.
        """
        if self._engine is not None:
            try:
                self._engine.dispose()
            except Exception:
                pass
            self._engine = None
        self._Session = None

    @contextmanager
    def get_session(self, session=None):
        """一个数据库会话上下文管理器。

默认(``session=None``)会创建一个新的 SQLAlchemy 会话并在上下文退出时自动提交;若上下文内部抛出异常则会自动回滚;无论是否成功,会话最终都会被关闭。

当传入一个外部 ``session`` 时,本上下文是透明的:会话被原样 yield,提交/回滚/关闭由最初打开它的 ``get_session`` 负责。上下文内部抛出的异常仍会向上传播,由外层 ``get_session`` 触发回滚,因此辅助方法可以通过可选的 ``session=None`` 参数参与调用方驱动的多步事务,而无需关心自己是否拥有该会话。

Args:
    session (Optional[Session]): 外部已打开的 SQLAlchemy 会话。为 ``None`` 时创建并自管理一个新会话,否则透传使用。
"""
        if session is not None:
            yield session
            return
        session = self.Session()
        try:
            yield session
            session.commit()
        except Exception:
            session.rollback()
            raise
        finally:
            session.close()

    @staticmethod
    def paginate(query, *, page: int = 1, page_size: int = 20) -> Dict[str, Any]:
        """对一个 SQLAlchemy ``Query`` 应用基于 page 的分页。

会先将 ``page`` 与 ``page_size`` 截断到不小于 1,再对传入的查询做一次 ``COUNT`` 拿到不分页的总数,然后用 ``OFFSET``/``LIMIT`` 取当前页。返回形如 ``{'items', 'total', 'page', 'page_size'}`` 的字典。

返回的 ``items`` 是原始 SQLAlchemy 行对象——不做任何业务转换。调用方自行负责把行转换成所需格式(例如通过 ``_orm_to_dict``),因此本工具可以被任意列表接口复用。

Args:
    query (Query): 已经应用好所需 ``filter``/``order_by`` 的 SQLAlchemy ``Query``。
    page (int): 页码,从 1 开始;小于 1 时会被截断为 1。
    page_size (int): 每页条数;小于 1 时会被截断为 1。
"""
        page = max(page, 1)
        page_size = max(page_size, 1)
        total = query.count()
        rows = query.offset((page - 1) * page_size).limit(page_size).all()
        return {'items': rows, 'total': total, 'page': page, 'page_size': page_size}

    def check_connection(self) -> DBResult:
        """检查数据库连接状态。

测试与数据库的连接是否正常建立。

**Returns:**

- DBResult: DBResult.status 连接成功(True), 连接失败(False)。DBResult.detail 包含失败信息
"""
        try:
            with self.engine.connect() as _:
                return DBResult()
        except SQLAlchemyError as e:
            return DBResult(status=DBStatus.FAIL, detail=str(e))

    @property
    def desc(self) -> str:
        if self._desc is None:
            self.set_desc(tables_desc_dict={})
        return self._desc

    def set_desc(self, tables_desc_dict: dict = {}):  # noqa B006
        """对于SqlManager搭配LLM使用自然语言查询的表项设置其描述,尤其当其表名、列名及取值不具有自解释能力时。
例如:
数据表Document的status列取值包括: "waiting", "working", "success", "failed",tables_desc_dict参数应为 {"Document": "status列取值包括: waiting, working, success, failed"}

Args:
    tables_desc_dict (dict): 表项的补充说明
"""
        self._desc = ''
        if not isinstance(tables_desc_dict, dict):
            raise ValueError(f'desc type {type(tables_desc_dict)} not supported')
        self._tables_desc_dict = tables_desc_dict
        if len(self.visible_tables) == 0:
            return
        # Generate desc according to table schema and comment
        self._desc = 'The tables description is as follows\n```\n'
        for table_name in self.visible_tables:
            self._desc += f'Table {table_name}\n(\n'
            TableCls = self.get_table_orm_class(table_name)
            if TableCls is None:
                # The table could be dropped in other session
                continue
            table_columns = TableCls.__table__.columns
            for i, column in enumerate(table_columns):
                self._desc += f' {column.name} {column.type}'
                if i != len(table_columns) - 1:
                    self._desc += ','
                self._desc += '\n'
            self._desc += ');\n'
            if table_name in tables_desc_dict:
                self._desc += tables_desc_dict[table_name] + '\n\n'
        self._desc += '```\n'

    @property
    def visible_tables(self):
        if self._visible_tables is None:
            self._visible_tables = self.get_all_tables()
        return self._visible_tables

    @visible_tables.setter
    def visible_tables(self, visible_tables: list):
        all_tables = set(self.get_all_tables())
        for ele in visible_tables:
            if ele not in all_tables:
                raise ValueError(f'Table {ele} not found in database')
        self._visible_tables = visible_tables
        self.set_desc(self._tables_desc_dict)

    def _refresh_metadata(self, only=None):
        # refresh metadata in case of deleting/creating table in other session
        try:
            if only:
                self._metadata.reflect(bind=self.engine, only=only, extend_existing=True)
            elif not self._metadata.tables:
                self._metadata.reflect(bind=self.engine)
        except Exception as e:
            raise ValueError(f'Refresh metadata failed: {e}')

    def get_all_tables(self) -> list:
        """获取数据库中所有表的列表。

刷新元数据后返回当前数据库中的所有表名。

**Returns:**

- List[str]: 数据库中所有表名的列表
"""
        self._refresh_metadata()
        return list(self._metadata.tables.keys())

    def get_table_orm_class(self, table_name):
        """根据表名获取对应的ORM类。

通过表名反射获取SQLAlchemy自动映射的ORM类。

Args:
    table_name (str): 要获取的表名

**Returns:**

- sqlalchemy.ext.automap.Class: 对应的ORM类,如果表不存在返回None
"""
        if table_name in self._orm_cache:
            return self._orm_cache[table_name]
        self._refresh_metadata(only=[table_name])
        # SQLAlchemy automap sets autoincrement='auto' for all primary keys, including
        # non-integer ones (TEXT/VARCHAR). This causes SQLAlchemy to omit the PK column
        # from INSERT statements when the value is provided, leading to NOT NULL errors.
        # Explicitly disable autoincrement for non-integer primary key columns before
        # calling Base.prepare() so the compiled INSERT includes the PK column.
        table = self._metadata.tables.get(table_name)
        if table is not None:
            for col in table.primary_key.columns:
                if not isinstance(col.type, sqlalchemy.Integer):
                    col.autoincrement = False
        Base = automap_base(metadata=self._metadata)
        Base.prepare()
        class_obj = getattr(Base.classes, table_name, None)
        self._orm_cache[table_name] = class_obj
        return class_obj

    def execute_commit(self, statement: str):
        """执行SQL提交语句。

执行DDL或DML语句并自动提交事务,适用于CREATE、ALTER、INSERT、UPDATE、DELETE等操作。

Args:
    statement (str): 要执行的SQL语句
"""
        with self.get_session() as session:
            session.execute(sqlalchemy.text(statement))

    def execute_query(self, statement: str) -> str:
        """执行sql查询脚本并以JSON字符串返回结果。
"""
        statement = re.sub(r'/\*.*?\*/', '', statement, flags=re.DOTALL).strip()
        create_table_pattern = r'.*\s*create\s+table\s+.*'
        drop_table_pattern = r'.*\s*drop\s+table\s+.*'
        statement_lower = statement.lower()
        if re.match(create_table_pattern, statement_lower):
            return f'Create table not supported. Original statement: {statement}'
        elif re.match(drop_table_pattern, statement_lower):
            return f'Drop table not supported. Original statement: {statement}'
        try:
            result = []
            session = self.Session()
            # Use original session without post commit
            with session as session:
                cursor_result = session.execute(sqlalchemy.text(statement))
                columns = list(cursor_result.keys())
                result = [dict(zip(columns, row)) for row in cursor_result]
            str_result = json.dumps(result, ensure_ascii=False, default=self._serialize_uncommon_type)
        except Exception as e:
            str_result = f'Execute SQL ERROR: {str(e)}'
        return str_result

    def _create_by_script(self, table: str) -> DBResult:
        status = DBStatus.SUCCESS
        detail = 'Success'
        try:
            with self.engine.connect() as conn:
                conn.execute(sqlalchemy.text(table))
                conn.commit()
        except OperationalError as e:
            status = DBStatus.FAIL
            detail = f'ERROR: {str(e)}'
        return DBResult(status=status, detail=detail)

    def _create_by_api(self, table: Union[DeclarativeBase, DeclarativeMeta]) -> DBResult:
        table.metadata.create_all(bind=self.engine, checkfirst=True)
        return DBResult()

    def create_table(self, table: Union[str, Type[DeclarativeBase], DeclarativeMeta]) -> DBResult:
        """创建数据表

Args:
    table (str/Type[DeclarativeBase]/DeclarativeMeta): 数据表schema。支持三种参数类型:类型为str的sql语句,继承自DeclarativeBase或继承自declarative_base()的ORM类
"""
        status = DBStatus.SUCCESS
        detail = 'Success'
        if isinstance(table, str):
            return self._create_by_script(table)
        # Support DeclarativeMeta created by declarative_base() which is deprecated since: 2.0
        elif issubclass(table, DeclarativeBase) or isinstance(table, DeclarativeMeta):
            return self._create_by_api(table)
        else:
            status = DBStatus.FAIL
            detail += f'Failed: Unsupported Type: {table}'
        return DBResult(status=status, detail=detail)

    def drop_table(self, table: Union[str, Type[DeclarativeBase], DeclarativeMeta]) -> DBResult:
        """删除数据表

Args:
    table (str/Type[DeclarativeBase]/DeclarativeMeta): 数据表schema。支持三种参数类型:类型为str的数据表名,继承自DeclarativeBase或继承自declarative_base()的ORM类
"""
        metadata = self._metadata
        if isinstance(table, str):
            tablename = table
        elif issubclass(table, DeclarativeBase) or isinstance(table, DeclarativeMeta):
            tablename = table.__tablename__
        else:
            return DBResult(status=DBStatus.FAIL, detail=f'{table} type unsupported')
        Table = sqlalchemy.Table(tablename, metadata, autoload_with=self.engine)
        Table.drop(self.engine, checkfirst=True)
        return DBResult()

    def insert_values(self, table_name: str, vals: List[dict]) -> DBResult:
        """批量数据插入

Args:
    table_name (str): 数据表名
    vals (List[dict]): 待插入数据,格式为[{"col_name1": v01, "col_name2": v02, ...}, {"col_name1": v11, "col_name2": v12, ...}, ...]
"""
        TableCls = self.get_table_orm_class(table_name)
        if TableCls is None:
            return DBResult(status=DBStatus.FAIL, detail=f'{table_name} not found in database')
        try:
            with self.get_session() as session:
                objects = [TableCls(**v) for v in vals]
                session.add_all(objects)
            return DBResult()
        except Exception as e:
            return DBResult(status=DBStatus.FAIL, detail=f'Insert failed: {e}')

check_connection()

检查数据库连接状态。

测试与数据库的连接是否正常建立。

Returns:

  • DBResult: DBResult.status 连接成功(True), 连接失败(False)。DBResult.detail 包含失败信息
Source code in lazyllm/tools/sql/sql_manager.py
    def check_connection(self) -> DBResult:
        """检查数据库连接状态。

测试与数据库的连接是否正常建立。

**Returns:**

- DBResult: DBResult.status 连接成功(True), 连接失败(False)。DBResult.detail 包含失败信息
"""
        try:
            with self.engine.connect() as _:
                return DBResult()
        except SQLAlchemyError as e:
            return DBResult(status=DBStatus.FAIL, detail=str(e))

create_table(table)

创建数据表

Parameters:

  • table (str / Type[DeclarativeBase] / DeclarativeMeta) –

    数据表schema。支持三种参数类型:类型为str的sql语句,继承自DeclarativeBase或继承自declarative_base()的ORM类

Source code in lazyllm/tools/sql/sql_manager.py
    def create_table(self, table: Union[str, Type[DeclarativeBase], DeclarativeMeta]) -> DBResult:
        """创建数据表

Args:
    table (str/Type[DeclarativeBase]/DeclarativeMeta): 数据表schema。支持三种参数类型:类型为str的sql语句,继承自DeclarativeBase或继承自declarative_base()的ORM类
"""
        status = DBStatus.SUCCESS
        detail = 'Success'
        if isinstance(table, str):
            return self._create_by_script(table)
        # Support DeclarativeMeta created by declarative_base() which is deprecated since: 2.0
        elif issubclass(table, DeclarativeBase) or isinstance(table, DeclarativeMeta):
            return self._create_by_api(table)
        else:
            status = DBStatus.FAIL
            detail += f'Failed: Unsupported Type: {table}'
        return DBResult(status=status, detail=detail)

dispose()

Release the underlying engine's connection pool.

Needed so callers (e.g. DocServer._Impl.stop) can close sqlite file handles before removing the containing directory; on Windows this is required, because open handles block TemporaryDirectory cleanup.

Source code in lazyllm/tools/sql/sql_manager.py
def dispose(self):
    """Release the underlying engine's connection pool.

    Needed so callers (e.g. DocServer._Impl.stop) can close sqlite file handles
    before removing the containing directory; on Windows this is required, because
    open handles block ``TemporaryDirectory`` cleanup.
    """
    if self._engine is not None:
        try:
            self._engine.dispose()
        except Exception:
            pass
        self._engine = None
    self._Session = None

drop_table(table)

删除数据表

Parameters:

  • table (str / Type[DeclarativeBase] / DeclarativeMeta) –

    数据表schema。支持三种参数类型:类型为str的数据表名,继承自DeclarativeBase或继承自declarative_base()的ORM类

Source code in lazyllm/tools/sql/sql_manager.py
    def drop_table(self, table: Union[str, Type[DeclarativeBase], DeclarativeMeta]) -> DBResult:
        """删除数据表

Args:
    table (str/Type[DeclarativeBase]/DeclarativeMeta): 数据表schema。支持三种参数类型:类型为str的数据表名,继承自DeclarativeBase或继承自declarative_base()的ORM类
"""
        metadata = self._metadata
        if isinstance(table, str):
            tablename = table
        elif issubclass(table, DeclarativeBase) or isinstance(table, DeclarativeMeta):
            tablename = table.__tablename__
        else:
            return DBResult(status=DBStatus.FAIL, detail=f'{table} type unsupported')
        Table = sqlalchemy.Table(tablename, metadata, autoload_with=self.engine)
        Table.drop(self.engine, checkfirst=True)
        return DBResult()

execute_commit(statement)

执行SQL提交语句。

执行DDL或DML语句并自动提交事务,适用于CREATE、ALTER、INSERT、UPDATE、DELETE等操作。

Parameters:

  • statement (str) –

    要执行的SQL语句

Source code in lazyllm/tools/sql/sql_manager.py
    def execute_commit(self, statement: str):
        """执行SQL提交语句。

执行DDL或DML语句并自动提交事务,适用于CREATE、ALTER、INSERT、UPDATE、DELETE等操作。

Args:
    statement (str): 要执行的SQL语句
"""
        with self.get_session() as session:
            session.execute(sqlalchemy.text(statement))

execute_query(statement)

执行sql查询脚本并以JSON字符串返回结果。

Source code in lazyllm/tools/sql/sql_manager.py
    def execute_query(self, statement: str) -> str:
        """执行sql查询脚本并以JSON字符串返回结果。
"""
        statement = re.sub(r'/\*.*?\*/', '', statement, flags=re.DOTALL).strip()
        create_table_pattern = r'.*\s*create\s+table\s+.*'
        drop_table_pattern = r'.*\s*drop\s+table\s+.*'
        statement_lower = statement.lower()
        if re.match(create_table_pattern, statement_lower):
            return f'Create table not supported. Original statement: {statement}'
        elif re.match(drop_table_pattern, statement_lower):
            return f'Drop table not supported. Original statement: {statement}'
        try:
            result = []
            session = self.Session()
            # Use original session without post commit
            with session as session:
                cursor_result = session.execute(sqlalchemy.text(statement))
                columns = list(cursor_result.keys())
                result = [dict(zip(columns, row)) for row in cursor_result]
            str_result = json.dumps(result, ensure_ascii=False, default=self._serialize_uncommon_type)
        except Exception as e:
            str_result = f'Execute SQL ERROR: {str(e)}'
        return str_result

get_all_tables()

获取数据库中所有表的列表。

刷新元数据后返回当前数据库中的所有表名。

Returns:

  • List[str]: 数据库中所有表名的列表
Source code in lazyllm/tools/sql/sql_manager.py
    def get_all_tables(self) -> list:
        """获取数据库中所有表的列表。

刷新元数据后返回当前数据库中的所有表名。

**Returns:**

- List[str]: 数据库中所有表名的列表
"""
        self._refresh_metadata()
        return list(self._metadata.tables.keys())

get_session(session=None)

一个数据库会话上下文管理器。

默认(session=None)会创建一个新的 SQLAlchemy 会话并在上下文退出时自动提交;若上下文内部抛出异常则会自动回滚;无论是否成功,会话最终都会被关闭。

当传入一个外部 session 时,本上下文是透明的:会话被原样 yield,提交/回滚/关闭由最初打开它的 get_session 负责。上下文内部抛出的异常仍会向上传播,由外层 get_session 触发回滚,因此辅助方法可以通过可选的 session=None 参数参与调用方驱动的多步事务,而无需关心自己是否拥有该会话。

Parameters:

  • session (Optional[Session], default: None ) –

    外部已打开的 SQLAlchemy 会话。为 None 时创建并自管理一个新会话,否则透传使用。

Source code in lazyllm/tools/sql/sql_manager.py
    @contextmanager
    def get_session(self, session=None):
        """一个数据库会话上下文管理器。

默认(``session=None``)会创建一个新的 SQLAlchemy 会话并在上下文退出时自动提交;若上下文内部抛出异常则会自动回滚;无论是否成功,会话最终都会被关闭。

当传入一个外部 ``session`` 时,本上下文是透明的:会话被原样 yield,提交/回滚/关闭由最初打开它的 ``get_session`` 负责。上下文内部抛出的异常仍会向上传播,由外层 ``get_session`` 触发回滚,因此辅助方法可以通过可选的 ``session=None`` 参数参与调用方驱动的多步事务,而无需关心自己是否拥有该会话。

Args:
    session (Optional[Session]): 外部已打开的 SQLAlchemy 会话。为 ``None`` 时创建并自管理一个新会话,否则透传使用。
"""
        if session is not None:
            yield session
            return
        session = self.Session()
        try:
            yield session
            session.commit()
        except Exception:
            session.rollback()
            raise
        finally:
            session.close()

get_table_orm_class(table_name)

根据表名获取对应的ORM类。

通过表名反射获取SQLAlchemy自动映射的ORM类。

Parameters:

  • table_name (str) –

    要获取的表名

Returns:

  • sqlalchemy.ext.automap.Class: 对应的ORM类,如果表不存在返回None
Source code in lazyllm/tools/sql/sql_manager.py
    def get_table_orm_class(self, table_name):
        """根据表名获取对应的ORM类。

通过表名反射获取SQLAlchemy自动映射的ORM类。

Args:
    table_name (str): 要获取的表名

**Returns:**

- sqlalchemy.ext.automap.Class: 对应的ORM类,如果表不存在返回None
"""
        if table_name in self._orm_cache:
            return self._orm_cache[table_name]
        self._refresh_metadata(only=[table_name])
        # SQLAlchemy automap sets autoincrement='auto' for all primary keys, including
        # non-integer ones (TEXT/VARCHAR). This causes SQLAlchemy to omit the PK column
        # from INSERT statements when the value is provided, leading to NOT NULL errors.
        # Explicitly disable autoincrement for non-integer primary key columns before
        # calling Base.prepare() so the compiled INSERT includes the PK column.
        table = self._metadata.tables.get(table_name)
        if table is not None:
            for col in table.primary_key.columns:
                if not isinstance(col.type, sqlalchemy.Integer):
                    col.autoincrement = False
        Base = automap_base(metadata=self._metadata)
        Base.prepare()
        class_obj = getattr(Base.classes, table_name, None)
        self._orm_cache[table_name] = class_obj
        return class_obj

insert_values(table_name, vals)

批量数据插入

Parameters:

  • table_name (str) –

    数据表名

  • vals (List[dict]) –

    待插入数据,格式为[{"col_name1": v01, "col_name2": v02, ...}, {"col_name1": v11, "col_name2": v12, ...}, ...]

Source code in lazyllm/tools/sql/sql_manager.py
    def insert_values(self, table_name: str, vals: List[dict]) -> DBResult:
        """批量数据插入

Args:
    table_name (str): 数据表名
    vals (List[dict]): 待插入数据,格式为[{"col_name1": v01, "col_name2": v02, ...}, {"col_name1": v11, "col_name2": v12, ...}, ...]
"""
        TableCls = self.get_table_orm_class(table_name)
        if TableCls is None:
            return DBResult(status=DBStatus.FAIL, detail=f'{table_name} not found in database')
        try:
            with self.get_session() as session:
                objects = [TableCls(**v) for v in vals]
                session.add_all(objects)
            return DBResult()
        except Exception as e:
            return DBResult(status=DBStatus.FAIL, detail=f'Insert failed: {e}')

paginate(query, *, page=1, page_size=20) staticmethod

对一个 SQLAlchemy Query 应用基于 page 的分页。

会先将 pagepage_size 截断到不小于 1,再对传入的查询做一次 COUNT 拿到不分页的总数,然后用 OFFSET/LIMIT 取当前页。返回形如 {'items', 'total', 'page', 'page_size'} 的字典。

返回的 items 是原始 SQLAlchemy 行对象——不做任何业务转换。调用方自行负责把行转换成所需格式(例如通过 _orm_to_dict),因此本工具可以被任意列表接口复用。

Parameters:

  • query (Query) –

    已经应用好所需 filter/order_by 的 SQLAlchemy Query

  • page (int, default: 1 ) –

    页码,从 1 开始;小于 1 时会被截断为 1。

  • page_size (int, default: 20 ) –

    每页条数;小于 1 时会被截断为 1。

Source code in lazyllm/tools/sql/sql_manager.py
    @staticmethod
    def paginate(query, *, page: int = 1, page_size: int = 20) -> Dict[str, Any]:
        """对一个 SQLAlchemy ``Query`` 应用基于 page 的分页。

会先将 ``page`` 与 ``page_size`` 截断到不小于 1,再对传入的查询做一次 ``COUNT`` 拿到不分页的总数,然后用 ``OFFSET``/``LIMIT`` 取当前页。返回形如 ``{'items', 'total', 'page', 'page_size'}`` 的字典。

返回的 ``items`` 是原始 SQLAlchemy 行对象——不做任何业务转换。调用方自行负责把行转换成所需格式(例如通过 ``_orm_to_dict``),因此本工具可以被任意列表接口复用。

Args:
    query (Query): 已经应用好所需 ``filter``/``order_by`` 的 SQLAlchemy ``Query``。
    page (int): 页码,从 1 开始;小于 1 时会被截断为 1。
    page_size (int): 每页条数;小于 1 时会被截断为 1。
"""
        page = max(page, 1)
        page_size = max(page_size, 1)
        total = query.count()
        rows = query.offset((page - 1) * page_size).limit(page_size).all()
        return {'items': rows, 'total': total, 'page': page, 'page_size': page_size}

set_desc(tables_desc_dict={})

对于SqlManager搭配LLM使用自然语言查询的表项设置其描述,尤其当其表名、列名及取值不具有自解释能力时。 例如: 数据表Document的status列取值包括: "waiting", "working", "success", "failed",tables_desc_dict参数应为 {"Document": "status列取值包括: waiting, working, success, failed"}

Parameters:

  • tables_desc_dict (dict, default: {} ) –

    表项的补充说明

Source code in lazyllm/tools/sql/sql_manager.py
    def set_desc(self, tables_desc_dict: dict = {}):  # noqa B006
        """对于SqlManager搭配LLM使用自然语言查询的表项设置其描述,尤其当其表名、列名及取值不具有自解释能力时。
例如:
数据表Document的status列取值包括: "waiting", "working", "success", "failed",tables_desc_dict参数应为 {"Document": "status列取值包括: waiting, working, success, failed"}

Args:
    tables_desc_dict (dict): 表项的补充说明
"""
        self._desc = ''
        if not isinstance(tables_desc_dict, dict):
            raise ValueError(f'desc type {type(tables_desc_dict)} not supported')
        self._tables_desc_dict = tables_desc_dict
        if len(self.visible_tables) == 0:
            return
        # Generate desc according to table schema and comment
        self._desc = 'The tables description is as follows\n```\n'
        for table_name in self.visible_tables:
            self._desc += f'Table {table_name}\n(\n'
            TableCls = self.get_table_orm_class(table_name)
            if TableCls is None:
                # The table could be dropped in other session
                continue
            table_columns = TableCls.__table__.columns
            for i, column in enumerate(table_columns):
                self._desc += f' {column.name} {column.type}'
                if i != len(table_columns) - 1:
                    self._desc += ','
                self._desc += '\n'
            self._desc += ');\n'
            if table_name in tables_desc_dict:
                self._desc += tables_desc_dict[table_name] + '\n\n'
        self._desc += '```\n'

lazyllm.tools.SqlCall

Bases: ModuleBase

SqlCall 是一个扩展自 ModuleBase 的类,提供了使用语言模型(LLM)生成和执行 SQL 查询的接口。 它设计用于与 SQL 数据库交互,从语言模型的响应中提取 SQL 查询,执行这些查询,并返回结果或解释。

Parameters:

  • llm

    用于生成和解释 SQL 查询及解释的大语言模型。

  • sql_manager (DBManager) –

    数据库管理器实例,包含数据库连接和描述信息

  • sql_examples (str, default: '' ) –

    SQL示例字符串,用于提示工程。默认为空字符串

  • sql_post_func (Callable, default: None ) –

    对生成的SQL语句进行后处理的函数。默认为 None

  • use_llm_for_sql_result (bool, default: True ) –

    是否使用LLM来解释SQL执行结果。默认为 True

  • return_trace (bool, default: False ) –

    是否返回执行跟踪信息。默认为 False

Examples:

>>> # First, run SqlManager example
>>> import lazyllm
>>> from lazyllm.tools import SQLManger, SqlCall
>>> sql_tool = SQLManger("personal.db")
>>> sql_llm = lazyllm.OnlineChatModule(model="gpt-4o", source="openai", base_url="***")
>>> sql_call = SqlCall(sql_llm, sql_tool, use_llm_for_sql_result=True)
>>> print(sql_call("去年一整年销售额最多的员工是谁?"))
Source code in lazyllm/tools/sql_call/sql_call.py
class SqlCall(ModuleBase):
    """SqlCall 是一个扩展自 ModuleBase 的类,提供了使用语言模型(LLM)生成和执行 SQL 查询的接口。
它设计用于与 SQL 数据库交互,从语言模型的响应中提取 SQL 查询,执行这些查询,并返回结果或解释。

Args:
    llm: 用于生成和解释 SQL 查询及解释的大语言模型。
    sql_manager (DBManager): 数据库管理器实例,包含数据库连接和描述信息
    sql_examples (str, optional): SQL示例字符串,用于提示工程。默认为空字符串
    sql_post_func (Callable, optional): 对生成的SQL语句进行后处理的函数。默认为 ``None``
    use_llm_for_sql_result (bool, optional): 是否使用LLM来解释SQL执行结果。默认为 ``True``
    return_trace (bool, optional): 是否返回执行跟踪信息。默认为 ``False``


Examples:
        >>> # First, run SqlManager example
        >>> import lazyllm
        >>> from lazyllm.tools import SQLManger, SqlCall
        >>> sql_tool = SQLManger("personal.db")
        >>> sql_llm = lazyllm.OnlineChatModule(model="gpt-4o", source="openai", base_url="***")
        >>> sql_call = SqlCall(sql_llm, sql_tool, use_llm_for_sql_result=True)
        >>> print(sql_call("去年一整年销售额最多的员工是谁?"))
    """
    EXAMPLE_TITLE = 'Here are some example: '

    def __init__(self, llm, sql_manager: DBManager, sql_examples: str = '', sql_post_func: Callable = None,
                 use_llm_for_sql_result=True, return_trace: bool = False) -> None:
        super().__init__(return_trace=return_trace)
        if not sql_manager.desc:
            raise ValueError('Error: sql_manager found empty description.')
        self._sql_tool = sql_manager
        self.sql_post_func = sql_post_func

        if sql_manager.db_type == 'mongodb':
            self._query_prompter = ChatPrompter(instruction=mongodb_query_instruct_template).pre_hook(
                self.sql_query_promt_hook
            )
            statement_type = 'mongodb json pipeline'
            self._pattern = re.compile(r'```json(.+?)```', re.DOTALL)
        else:
            self._query_prompter = ChatPrompter(instruction=sql_query_instruct_template).pre_hook(
                self.sql_query_promt_hook
            )
            statement_type = 'sql query'
            self._pattern = re.compile(r'```sql(.+?)```', re.DOTALL)

        self._llm_query = llm.share(prompt=self._query_prompter).used_by(self._module_id)
        self._answer_prompter = ChatPrompter(
            instruction=db_explain_instruct_template.format(statement_type=statement_type, db_type=sql_manager.db_type)
        ).pre_hook(self.sql_explain_prompt_hook)
        self._llm_answer = llm.share(prompt=self._answer_prompter).used_by(self._module_id)
        self.example = sql_examples
        with pipeline() as sql_execute_ppl:
            sql_execute_ppl.exec = self._sql_tool.execute_query
            if use_llm_for_sql_result:
                sql_execute_ppl.concate = (lambda q, r: [q, r]) | bind(sql_execute_ppl.input, _0)
                sql_execute_ppl.llm_answer = self._llm_answer
        with pipeline() as ppl:
            ppl.llm_query = self._llm_query
            ppl.sql_extractor = self.extract_sql_from_response
            with switch(judge_on_full_input=False) as ppl.sw:
                ppl.sw.case[False, lambda x: x]
                ppl.sw.case[True, sql_execute_ppl]
        self._impl = ppl

    def sql_query_promt_hook(
        self,
        input: Union[str, List, Dict[str, str], None] = None,
        history: Optional[List[Union[List[str], Dict[str, Any]]]] = None,
        tools: Union[List[Dict[str, Any]], None] = None,
        label: Union[str, None] = None,
    ):
        """为从用户输入生成数据库查询准备 prompt 的 hook。

Args:
    input (Union[str, List, Dict[str, str], None]): 用户的自然语言查询。
    history (List[Union[List[str], Dict[str, Any]]]): 会话历史。
    tools (Union[List[Dict[str, Any]], None]): 可用工具描述。
    label (Union[str, None]): 可选标签。

**Returns:**

- Tuple: 包含格式化后的 prompt 字典(包括 current_date、db_type、desc、user_query)、history、tools 和 label。
"""
        current_date = datetime.datetime.now().strftime('%Y-%m-%d')
        schema_desc = self._sql_tool.desc
        if self.example:
            schema_desc += f'\n{self.EXAMPLE_TITLE}\n{self.example}\n'
        if not isinstance(input, str):
            raise ValueError(f'Unexpected type for input: {type(input)}')
        return (
            dict(current_date=current_date, db_type=self._sql_tool.db_type, desc=schema_desc, user_query=input),
            history or [],
            tools,
            label,
        )

    def sql_explain_prompt_hook(
        self,
        input: Union[str, List, Dict[str, str], None] = None,
        history: List[Union[List[str], Dict[str, Any]]] = [],  # noqa B006
        tools: Union[List[Dict[str, Any]], None] = None,
        label: Union[str, None] = None,
    ):
        """为解释数据库查询执行结果准备 prompt 的 hook。

Args:
    input (Union[str, List, Dict[str, str], None]): 包含查询和结果的列表。
    history (List[Union[List[str], Dict[str, Any]]]): 会话历史。
    tools (Union[List[Dict[str, Any]], None]): 可用工具描述。
    label (Union[str, None]): 可选标签。

**Returns:**

- Tuple: 包含格式化后的 prompt 字典(history_info、desc、query、result、explain_query)、history、tools 和 label。
"""
        explain_query = 'Tell the user based on the execution results, making sure to keep the language consistent \
            with the user\'s input and don\'t translate original result.'
        if not isinstance(input, list) and len(input) != 2:
            raise ValueError(f'Unexpected type for input: {type(input)}')
        assert 'root_input' in globals and self._llm_answer._module_id in globals['root_input']
        user_query = globals['root_input'][self._llm_answer._module_id]
        globals.pop('root_input')
        history_info = chat_history_to_str(history, user_query)
        return (
            dict(
                history_info=history_info,
                desc=self._sql_tool.desc,
                query=input[0],
                result=input[1],
                explain_query=explain_query,
            ),
            history,
            tools,
            label,
        )

    def extract_sql_from_response(self, str_response: str) -> tuple[bool, str]:
        """从原始 LLM 响应中提取 SQL(或 MongoDB pipeline)语句。

Args:
    str_response (str): LLM 返回的原始文本,可能包含代码块。

**Returns:**

- tuple[bool, str]: 第一个元素表示是否成功提取,第二个是清洗后的或原始内容。如果提供了 sql_post_func,则会应用于提取结果。
"""
        # Remove the triple backticks if present
        matches = self._pattern.findall(str_response)
        if matches:
            # Return the first match
            extracted_content = matches[0].strip()
            if self._sql_tool.db_type != 'mongodb':
                stmts = [s.strip() for s in extracted_content.split(';') if s.strip()]
                if stmts:
                    extracted_content = stmts[0]
            return True, extracted_content if not self.sql_post_func else self.sql_post_func(extracted_content)
        else:
            return False, str_response

    def forward(self, input: str, llm_chat_history: List[Dict[str, Any]] = None):
        globals['root_input'] = {self._llm_answer._module_id: input}
        if self._module_id in globals['chat_history']:
            globals['chat_history'][self._llm_query._module_id] = globals['chat_history'][self._module_id]
        return self._impl(input)

extract_sql_from_response(str_response)

从原始 LLM 响应中提取 SQL(或 MongoDB pipeline)语句。

Parameters:

  • str_response (str) –

    LLM 返回的原始文本,可能包含代码块。

Returns:

  • tuple[bool, str]: 第一个元素表示是否成功提取,第二个是清洗后的或原始内容。如果提供了 sql_post_func,则会应用于提取结果。
Source code in lazyllm/tools/sql_call/sql_call.py
    def extract_sql_from_response(self, str_response: str) -> tuple[bool, str]:
        """从原始 LLM 响应中提取 SQL(或 MongoDB pipeline)语句。

Args:
    str_response (str): LLM 返回的原始文本,可能包含代码块。

**Returns:**

- tuple[bool, str]: 第一个元素表示是否成功提取,第二个是清洗后的或原始内容。如果提供了 sql_post_func,则会应用于提取结果。
"""
        # Remove the triple backticks if present
        matches = self._pattern.findall(str_response)
        if matches:
            # Return the first match
            extracted_content = matches[0].strip()
            if self._sql_tool.db_type != 'mongodb':
                stmts = [s.strip() for s in extracted_content.split(';') if s.strip()]
                if stmts:
                    extracted_content = stmts[0]
            return True, extracted_content if not self.sql_post_func else self.sql_post_func(extracted_content)
        else:
            return False, str_response

sql_explain_prompt_hook(input=None, history=[], tools=None, label=None)

为解释数据库查询执行结果准备 prompt 的 hook。

Parameters:

  • input (Union[str, List, Dict[str, str], None], default: None ) –

    包含查询和结果的列表。

  • history (List[Union[List[str], Dict[str, Any]]], default: [] ) –

    会话历史。

  • tools (Union[List[Dict[str, Any]], None], default: None ) –

    可用工具描述。

  • label (Union[str, None], default: None ) –

    可选标签。

Returns:

  • Tuple: 包含格式化后的 prompt 字典(history_info、desc、query、result、explain_query)、history、tools 和 label。
Source code in lazyllm/tools/sql_call/sql_call.py
    def sql_explain_prompt_hook(
        self,
        input: Union[str, List, Dict[str, str], None] = None,
        history: List[Union[List[str], Dict[str, Any]]] = [],  # noqa B006
        tools: Union[List[Dict[str, Any]], None] = None,
        label: Union[str, None] = None,
    ):
        """为解释数据库查询执行结果准备 prompt 的 hook。

Args:
    input (Union[str, List, Dict[str, str], None]): 包含查询和结果的列表。
    history (List[Union[List[str], Dict[str, Any]]]): 会话历史。
    tools (Union[List[Dict[str, Any]], None]): 可用工具描述。
    label (Union[str, None]): 可选标签。

**Returns:**

- Tuple: 包含格式化后的 prompt 字典(history_info、desc、query、result、explain_query)、history、tools 和 label。
"""
        explain_query = 'Tell the user based on the execution results, making sure to keep the language consistent \
            with the user\'s input and don\'t translate original result.'
        if not isinstance(input, list) and len(input) != 2:
            raise ValueError(f'Unexpected type for input: {type(input)}')
        assert 'root_input' in globals and self._llm_answer._module_id in globals['root_input']
        user_query = globals['root_input'][self._llm_answer._module_id]
        globals.pop('root_input')
        history_info = chat_history_to_str(history, user_query)
        return (
            dict(
                history_info=history_info,
                desc=self._sql_tool.desc,
                query=input[0],
                result=input[1],
                explain_query=explain_query,
            ),
            history,
            tools,
            label,
        )

sql_query_promt_hook(input=None, history=None, tools=None, label=None)

为从用户输入生成数据库查询准备 prompt 的 hook。

Parameters:

  • input (Union[str, List, Dict[str, str], None], default: None ) –

    用户的自然语言查询。

  • history (List[Union[List[str], Dict[str, Any]]], default: None ) –

    会话历史。

  • tools (Union[List[Dict[str, Any]], None], default: None ) –

    可用工具描述。

  • label (Union[str, None], default: None ) –

    可选标签。

Returns:

  • Tuple: 包含格式化后的 prompt 字典(包括 current_date、db_type、desc、user_query)、history、tools 和 label。
Source code in lazyllm/tools/sql_call/sql_call.py
    def sql_query_promt_hook(
        self,
        input: Union[str, List, Dict[str, str], None] = None,
        history: Optional[List[Union[List[str], Dict[str, Any]]]] = None,
        tools: Union[List[Dict[str, Any]], None] = None,
        label: Union[str, None] = None,
    ):
        """为从用户输入生成数据库查询准备 prompt 的 hook。

Args:
    input (Union[str, List, Dict[str, str], None]): 用户的自然语言查询。
    history (List[Union[List[str], Dict[str, Any]]]): 会话历史。
    tools (Union[List[Dict[str, Any]], None]): 可用工具描述。
    label (Union[str, None]): 可选标签。

**Returns:**

- Tuple: 包含格式化后的 prompt 字典(包括 current_date、db_type、desc、user_query)、history、tools 和 label。
"""
        current_date = datetime.datetime.now().strftime('%Y-%m-%d')
        schema_desc = self._sql_tool.desc
        if self.example:
            schema_desc += f'\n{self.EXAMPLE_TITLE}\n{self.example}\n'
        if not isinstance(input, str):
            raise ValueError(f'Unexpected type for input: {type(input)}')
        return (
            dict(current_date=current_date, db_type=self._sql_tool.db_type, desc=schema_desc, user_query=input),
            history or [],
            tools,
            label,
        )

lazyllm.tools.rag.component.bm25.BM25

A BM25 retriever that uses the BM25 algorithm to retrieve nodes.

Source code in lazyllm/tools/rag/component/bm25.py
class BM25:
    """A BM25 retriever that uses the BM25 algorithm to retrieve nodes."""

    def __init__(
        self,
        nodes: List[DocNode],
        language: str = 'en',
        topk: int = 2,
        **kwargs,
    ) -> None:
        if language == 'en':
            self._stemmer = Stemmer.Stemmer('english')
            self._stopwords = language
            self._tokenizer = lambda t: t
        elif language == 'zh':
            self._stemmer = None
            # TODO(ywt): after bm25s supports cn stopwards, update this
            self._stopwords = STOPWORDS_CHINESE
            self._tokenizer = lambda t: ' '.join(jieba.lcut(t))
        self.topk = min(topk, len(nodes))
        self.nodes = nodes

        corpus_tokens = bm25s.tokenize(
            [self._tokenizer(node.get_text()) for node in nodes],
            stopwords=self._stopwords,
            stemmer=self._stemmer,
        )
        self.bm25 = bm25s.BM25()
        self.bm25.index(corpus_tokens)

    def retrieve(self, query: str, topk: Optional[int] = None) -> List[Tuple[DocNode, float]]:
        """使用BM25算法检索与查询最相关的文档节点。

Args:
    query (str): 查询文本。

**Returns:**

- List[Tuple[DocNode, float]]: 返回一个列表,每个元素为(文档节点, 相关度分数)的元组。
"""
        if topk is None:
            topk = self.topk
        else:
            topk = min(topk, len(self.nodes))
        tokenized_query = bm25s.tokenize(
            self._tokenizer(query), stopwords=self._stopwords, stemmer=self._stemmer
        )
        indexs, scores = self.bm25.retrieve(tokenized_query, k=topk)
        results = []
        for idx, score in zip(indexs[0], scores[0]):
            results.append((self.nodes[idx], score))
        return results

retrieve(query, topk=None)

使用BM25算法检索与查询最相关的文档节点。

Parameters:

  • query (str) –

    查询文本。

Returns:

  • List[Tuple[DocNode, float]]: 返回一个列表,每个元素为(文档节点, 相关度分数)的元组。
Source code in lazyllm/tools/rag/component/bm25.py
    def retrieve(self, query: str, topk: Optional[int] = None) -> List[Tuple[DocNode, float]]:
        """使用BM25算法检索与查询最相关的文档节点。

Args:
    query (str): 查询文本。

**Returns:**

- List[Tuple[DocNode, float]]: 返回一个列表,每个元素为(文档节点, 相关度分数)的元组。
"""
        if topk is None:
            topk = self.topk
        else:
            topk = min(topk, len(self.nodes))
        tokenized_query = bm25s.tokenize(
            self._tokenizer(query), stopwords=self._stopwords, stemmer=self._stemmer
        )
        indexs, scores = self.bm25.retrieve(tokenized_query, k=topk)
        results = []
        for idx, score in zip(indexs[0], scores[0]):
            results.append((self.nodes[idx], score))
        return results

lazyllm.tools.rag.doc_to_db.SchemaExtractor

Bases: ModuleBase

Schema aware extractor that materializes BaseModel schemas into database tables.

Source code in lazyllm/tools/rag/doc_to_db/extractor.py
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
class SchemaExtractor(ModuleBase):
    """Schema aware extractor that materializes BaseModel schemas into database tables."""

    TABLE_PREFIX = 'lazyllm_schema'
    SYS_KB_ID = 'kb_id'
    SYS_DOC_ID = 'doc_id'

    TYPE_MAP = {
        str: sqlalchemy.Text,
        int: sqlalchemy.Integer,
        float: sqlalchemy.Float,
        bool: sqlalchemy.Boolean,
        list: sqlalchemy.JSON,
        dict: sqlalchemy.JSON,
    }
    TYPE_NAME_MAP = {
        'string': str,
        'text': str,
        'int': int,
        'integer': int,
        'float': float,
        'number': float,
        'boolean': bool,
        'bool': bool,
        'list': list,
        'array': list,
        'dict': dict,
        'object': dict,
        'map': dict,
    }

    def __init__(self, db_config: Dict[str, Any], llm: LLMBase, *, name: Optional[str] = None,
                 table_prefix: Optional[str] = None,
                 force_refresh: bool = False, extraction_mode: ExtractionMode = ExtractionMode.TEXT,
                 max_len: int = ONE_DOC_LENGTH_LIMIT, num_workers: int = 4):
        super().__init__()
        if not isinstance(llm, LLMBase):
            raise TypeError('llm must be an instance of LLMBase')
        self._name = name
        self._llm = llm
        self._table_prefix = table_prefix or self.TABLE_PREFIX
        self._sql_manager = None
        self._db_config = db_config
        self._table_cache: Dict[str, Type[_TableBase]] = {}
        self._schema_registry: Dict[str, Type[BaseModel]] = {}
        self._active_schema_set_id: Optional[str] = None
        self._force_refresh = force_refresh
        self._extraction_mode = extraction_mode
        self._max_len = max_len
        self._num_workers = num_workers

    @property
    def name(self) -> Optional[str]:
        return self._name

    @property
    def sql_manager(self) -> SqlManager:
        self._lazy_init()
        return self._sql_manager

    def sql_manager_for_nl2sql(self,  # noqa: C901
                               kb_ids: Union[str, List[str]] = None) -> SqlManager:
        """
基于已注册的 schema,生成一个仅暴露相关表的 SqlManager,用于 SqlCall 模块中 NL2SQL 查询;会附带表结构描述和可见表列表。

Args:
    kb_ids (Union[str, List[str]], optional): 过滤的知识库 ID,可单个或列表。

**Returns:**

- SqlManager: 仅包含可见表、列信息及说明的 SqlManager 实例,用于 NL2SQL。
"""
        self._lazy_init()
        if not self._sql_manager:
            raise ValueError('SqlManager is not initialized')
        if not self._db_config:
            raise ValueError('db_config is required to build SqlManager')
        if not self._active_schema_set_id:
            raise ValueError('No active schema set registered')

        schema_info_table = TABLE_SCHEMA_SET_INFO['name']
        desc_map: Dict[str, str] = {}

        def _schema_table_desc(model: Type[BaseModel]) -> str:
            schema_desc = self._get_schema_set_str(model)
            return '\n'.join([s for s in [
                (model.__doc__ or '').strip(),
                schema_desc,
                f'System columns: {self.SYS_KB_ID}, {self.SYS_DOC_ID}, extract_meta',
            ] if s])

        kb_id_list = None
        if kb_ids is not None:
            if isinstance(kb_ids, (list, tuple, set)):
                kb_id_list = [str(k) for k in kb_ids if k is not None]
            else:
                kb_id_list = [str(kb_ids)]
            if not kb_id_list:
                kb_id_list = None

        schema_set_id = self._active_schema_set_id
        if not self.has_schema_set(schema_set_id):
            raise ValueError(f'Schema set {schema_set_id} not found')
        schema_model = self._schema_registry[schema_set_id]
        table_name = self._ensure_table(schema_set_id, schema_model)
        target_tables = {table_name}
        desc_map[table_name] = _schema_table_desc(schema_model)

        target_tables.discard(schema_info_table)
        tables_info_dict = {'tables': []}
        for table_name in target_tables:
            table_cls = self._sql_manager.get_table_orm_class(table_name)
            if table_cls is None:
                continue
            columns = []
            for col in table_cls.__table__.columns:
                columns.append({
                    'name': col.name,
                    'data_type': _col_type_name(col),
                    'nullable': bool(col.nullable),
                    'is_primary_key': bool(col.primary_key),
                    'comment': getattr(col, 'comment', '') or '',
                })
            tables_info_dict['tables'].append({'name': table_name, 'columns': columns, 'comment': ''})
        new_manager = self._init_sql_manager({**self._db_config, 'tables_info_dict': tables_info_dict})
        new_manager.visible_tables = list(target_tables)
        if desc_map:
            new_manager.set_desc(desc_map)
        return new_manager

    @once_wrapper
    def _lazy_init(self):
        self._sql_manager = self._init_sql_manager(self._db_config) if self._db_config else None
        if self._sql_manager:
            self._ensure_management_tables()

    def register_schema_set(self, schema_set: Type[BaseModel], schema_set_id: str = None,   # noqa: C901
                            force_refresh: bool = False) -> str:
        """schema set registration, idempotent"""
        try:
            self._lazy_init()
            self._validate_schema_model(schema_set)

            fields = getattr(schema_set, 'model_fields', None) or getattr(schema_set, '__fields__', {})

            def _safe_default(val: Any):
                if val is None:
                    return None
                if val.__class__.__name__ in ('PydanticUndefinedType', 'UndefinedType'):
                    return None
                if isinstance(val, (str, int, float, bool)):
                    return val
                return str(val)

            signature = [
                (name, str(getattr(f, 'annotation', None) or getattr(f, 'outer_type_', None)),
                 _safe_default(getattr(f, 'default', None)), getattr(f, 'default_factory', None) is not None,
                 getattr(f, 'is_required', lambda: False)())
                for name, f in fields.items()
            ]
            signature.sort(key=lambda x: x[0])
            idem_key = json.dumps(signature, ensure_ascii=False)

            if self._sql_manager:
                table_cls = self._sql_manager.get_table_orm_class(TABLE_SCHEMA_SET_INFO['name'])
                if table_cls is None:
                    raise ValueError('Schema set table not initialized')
                with self._sql_manager.get_session() as session:
                    existing = session.query(table_cls).filter_by(idem_key=idem_key).first()
                    if existing:
                        existing_id = str(existing.schema_set_id if hasattr(existing, 'schema_set_id') else existing.id)
                        if schema_set_id and str(schema_set_id) != existing_id:
                            raise ValueError(
                                f'schema_set_id mismatch for idem_key, expect {existing_id}, got {schema_set_id}'
                            )
                        schema_set_id = schema_set_id or existing_id
                    else:
                        schema_json = (schema_set.model_json_schema() if hasattr(schema_set, 'model_json_schema')
                                       else schema_set.schema())
                        desc = (schema_set.__doc__ or '').strip() or 'Schema set'
                        obj_kwargs = dict(schema_set_json=json.dumps(schema_json, ensure_ascii=False),
                                          desc=desc, idem_key=idem_key, created_at=datetime.now(),
                                          updated_at=datetime.now())
                        if schema_set_id is None:
                            schema_set_id = str(uuid4().hex)
                        obj_kwargs['schema_set_id'] = str(schema_set_id)
                        new_obj = table_cls(**obj_kwargs)
                        session.add(new_obj)
                        session.flush()
                        schema_set_id = str(new_obj.schema_set_id if hasattr(new_obj, 'schema_set_id') else new_obj.id)

            if schema_set_id is None:
                raise ValueError('schema_set_id is required and could not be derived')

            schema_set_id = str(schema_set_id)
            self._schema_registry[schema_set_id] = schema_set
            if self._sql_manager:
                self._ensure_table(schema_set_id, schema_set)
            self._active_schema_set_id = schema_set_id
            return schema_set_id
        except Exception as e:
            LOG.error(f'Failed to register schema set: {e}')
            raise e

    def _model_from_schema_json(self, schema_json: str, model_name: str = 'RecoveredSchema') -> Type[BaseModel]:
        """Reconstruct a minimal BaseModel subclass from stored JSON schema."""
        try:
            schema_dict = json.loads(schema_json)
        except Exception as exc:
            raise ValueError(f'Invalid schema json: {exc}') from exc
        properties = schema_dict.get('properties', {})
        required = set(schema_dict.get('required', []) or [])
        type_map = {
            'string': str,
            'integer': int,
            'number': float,
            'boolean': bool,
            'array': list,
            'object': dict,
        }
        fields_def: Dict[str, Tuple[Any, Any]] = {}
        for name, prop in properties.items():
            t_name = prop.get('type')
            py_type = type_map.get(t_name, str)
            desc = prop.get('description', '')
            default = ... if name in required else None
            fields_def[name] = (py_type, Field(default=default, description=desc))
        return create_model(model_name, **fields_def)  # type: ignore[arg-type]

    def has_schema_set(self, schema_set_id: str) -> bool:
        """
检查指定 schema_set_id 是否已注册,缺失时会尝试从数据库恢复模型并建表。

Args:
    schema_set_id (str): 目标 schema 集合 ID。

**Returns:**

- bool: 是否已存在。
"""
        self._lazy_init()
        if self._sql_manager:
            table_cls = self._sql_manager.get_table_orm_class(TABLE_SCHEMA_SET_INFO['name'])
            if table_cls is None:
                raise ValueError('Schema set table not initialized')
            with self._sql_manager.get_session() as session:
                existing = session.query(table_cls).filter_by(schema_set_id=schema_set_id).first()
                if not existing:
                    return False
                if schema_set_id not in self._schema_registry:
                    recovered_schema = self._model_from_schema_json(existing.schema_set_json,
                                                                    model_name=f'Schema_{schema_set_id}')
                    self._schema_registry[schema_set_id] = recovered_schema
                    self._ensure_table(schema_set_id, recovered_schema)
                return True
        return schema_set_id in self._schema_registry

    def _get_schema_set_str(self, schema_set) -> str:
        """Return a human readable schema description: name, description, data type."""
        model = None
        if isinstance(schema_set, str):
            model = self._schema_registry.get(schema_set)
        else:
            model = schema_set
        if not model:
            raise ValueError(f'Unknown schema_set: {schema_set}')
        fields = getattr(model, 'model_fields', None) or getattr(model, '__fields__', {})

        def _field_type_str(field_obj: Any) -> str:
            anno = getattr(field_obj, 'annotation', None) or getattr(field_obj, 'outer_type_', None)
            origin = get_origin(anno)
            args = get_args(anno)
            if origin is Union and args:
                non_none = [arg for arg in args if arg is not type(None)]  # noqa: E721
                anno = non_none[0] if non_none else anno
            return getattr(anno, '__name__', str(anno))

        lines: List[str] = []
        for name, field in fields.items():
            desc = getattr(field, 'description', None)
            if desc is None:
                field_info = getattr(field, 'field_info', None)
                desc = getattr(field_info, 'description', None) if field_info else None
            type_str = _field_type_str(field)
            lines.append(f"name: {name}, description: {desc or ''}, type: {type_str}")
        return '\n'.join(lines)

    def analyze_schema_and_register(self, data: Union[str, List[DocNode]],
                                    schema_set_id: Optional[str] = None) -> SchemaSetInfo:
        """Infer a schema from sample data, register it, and return the registration info."""
        self._lazy_init()
        if not self._llm:
            raise ValueError('LLM not initialized')
        if not data:
            raise ValueError('data is empty')

        if isinstance(data, str):
            sample_text = data[:self._max_len]
        else:
            chunks = self._gen_text_list_from_nodes(data)
            sample_text = '\n\n'.join(chunks)[:self._max_len] if chunks else ''
        if not sample_text:
            raise ValueError('No content available for schema analysis')

        llm = self._llm.share(prompt=SCHEMA_ANALYZE_PROMPT, format=JsonFormatter())
        payload = Template(SCHEMA_ANALYZE_INPUT_FORMAT).substitute(text=sample_text)
        res = llm(payload)
        fields_def: Dict[str, Tuple[Any, Any]] = {}
        for item in res:
            if not isinstance(item, dict):
                continue
            name = item.get('name')
            if not name:
                continue
            desc = item.get('description') or ''
            py_type = self._normalize_py_type(item.get('type'))
            fields_def[name] = (py_type, Field(default=None, description=desc))
        if not fields_def:
            # Fallback: single generic field capturing text content
            fields_def['content'] = (str, Field(default=None, description='Raw content snippet'))

        model_name = f'AutoSchema{uuid4().hex}'
        schema_model = create_model(model_name, **fields_def)  # type: ignore[arg-type]
        reg_id = self.register_schema_set(schema_model, schema_set_id)
        return SchemaSetInfo(schema_set_id=reg_id, schema_model=schema_model)

    def _gen_text_list_from_nodes(self, nodes: List[DocNode]) -> list[str]:
        """Generate full text blocks with metadata, each capped by `self._max_len`."""
        if not nodes:
            return []
        template = 'File Info:\n{file_metas}\nFile Content:\n{file_content}\n\n'
        metas = '\n'.join([f'{k}: {v}' for k, v in nodes[0].global_metadata.items()])

        # Reserve space for metadata and static prompt text.
        base_len = len(template.format(file_metas=metas, file_content=''))
        content_limit = max(self._max_len - base_len, 0)
        if content_limit == 0:
            return [template.format(file_metas=metas, file_content='')]

        chunks: List[str] = []
        current = ''
        for node in nodes:
            node_text = node.text
            sep_len = 1 if current else 0
            if len(current) + sep_len + len(node_text) <= content_limit:
                current = f'{current}\n{node_text}' if current else node_text
                continue

            if current:
                chunks.append(current)
            start = 0
            while start < len(node_text):
                end = start + content_limit
                chunks.append(node_text[start:end])
                start = end
            current = ''

        if current:
            chunks.append(current)

        return [template.format(file_metas=metas, file_content=chunk) for chunk in chunks]

    def _text_extract_impl(self, data: Union[str, List[DocNode]], schema_set_id: str) -> ExtractResult:  # noqa: C901
        if not self._llm:
            raise ValueError('LLM not initialized')
        schema_set = self._schema_registry.get(schema_set_id)
        llm = self._llm.share(prompt=SCHEMA_EXTRACT_PROMPT, format=JsonFormatter())
        content_list = self._gen_text_list_from_nodes(data) if isinstance(data, list) else [data]
        schema_str = self._get_schema_set_str(schema_set)
        input_list = [
            Template(SCHEMA_EXTRACT_INPUT_FORMAT).substitute(schema=schema_str, text=content)
            for content in content_list
        ]
        if self._num_workers > 1:
            pool = ThreadPoolExecutor(max_workers=self._num_workers)
            fs = [pool.submit(llm, text) for text in input_list]
            res = [f.result() for f in fs]
        else:
            res = [llm(text) for text in input_list]
        # process res by vote
        schema_val_clues: Dict[str, Dict[str, List[str]]] = {}
        for res_item in res:
            if not isinstance(res_item, list):
                LOG.error(f'[Schema Extractor - _text_extract_impl] invalid format {res_item}')
                continue
            for info in res_item:
                if not isinstance(info, dict):
                    continue
                schema_name = info.get('schema_name') or info.get('field_name')
                if not schema_name:
                    continue
                val_js = json.dumps(info.get('value'), ensure_ascii=False)
                if val_js is None:
                    continue
                clues = info.get('clues') or []
                schema_val_clues.setdefault(schema_name, {}).setdefault(val_js, []).extend(clues)

        data: Dict[str, Any] = {}
        clue_meta: Dict[str, ExtractClue] = {}
        for name, val_map in schema_val_clues.items():
            best_val_js = None
            best_clues: List[str] = []
            for v_js, clues in val_map.items():
                if len(clues) > len(best_clues):
                    best_val_js = v_js
                    best_clues = clues
            if best_val_js is None:
                continue
            try:
                best_val = json.loads(best_val_js)
            except Exception:
                best_val = best_val_js
            data[name] = best_val
            clue_meta[name] = ExtractClue(reason='selected_by_max_clues', citation=best_clues)

        meta = ExtractMeta(
            schema_set_id=schema_set_id,
            mode=self._extraction_mode,
            kb_id='',
            doc_id='',
            clues=clue_meta,
        )
        return [ExtractResult(data=data, metadata=meta)]

    def _multimodal_extract_impl(self, doc_nodes: List[DocNode], schemet_set_id: str) -> ExtractResult:
        # TODO: currently only support text extract
        raise NotImplementedError('Multimodal extract not implemented')

    def _schema_extract_impl(self, doc_nodes: List[DocNode]):
        raise NotImplementedError('Schema extract not implemented')

    def _validate_extract_params(self, data: Union[str, List[DocNode]]) -> Tuple[str, str, str]:
        self._lazy_init()
        if not data:
            raise ValueError('data is empty')
        if not self._active_schema_set_id:
            raise ValueError('No active schema set registered')
        kb_id = doc_id = None
        if isinstance(data, str):
            kb_id = DEFAULT_KB_ID
            doc_id = hashlib.sha256(data.encode('utf-8')).hexdigest()
        else:
            for node in data:
                meta = getattr(node, 'global_metadata', {}) or {}
                cur_kb_id = meta.get(RAG_KB_ID)
                cur_doc_id = meta.get(RAG_DOC_ID)
                if cur_kb_id is None or cur_doc_id is None:
                    raise ValueError('node.global_metadata must contain kb_id and doc_id')
                cur_kb_id = str(cur_kb_id)
                cur_doc_id = str(cur_doc_id)
                if kb_id is None:
                    kb_id = cur_kb_id
                elif kb_id != cur_kb_id:
                    raise ValueError('kb_id in data must be unique')
                if doc_id is None:
                    doc_id = cur_doc_id
                elif doc_id != cur_doc_id:
                    raise ValueError('doc_id in data must be unique')
        return kb_id, doc_id, self._active_schema_set_id

    def extract_and_store(self, data: Union[str, List[DocNode]],  # noqa: C901
                          schema_set_id: str = None, schema_set: Type[BaseModel] = None) -> ExtractResult:
        """
按已注册的 schema 抽取文本/DocNode 内容并写入对应表,若传入 schema_set 会先注册;同文档重复调用会返回缓存结果。

Args:
    data (Union[str, List[DocNode]]): 文本或 DocNode 列表(需同一文档)。
    schema_set_id (str, optional): 指定使用的 schema 集合 ID。
    schema_set (Type[BaseModel], optional): 动态注册并使用的 schema。

**Returns:**

- ExtractResult: 抽取结果,`data` 为字段名到值的字典,`metadata` 包含 schema_set_id、kb_id、doc_id 及按字段的线索信息;可能为 None 表示无可写入。
"""
        self._lazy_init()
        if schema_set is not None:
            schema_set_id = self.register_schema_set(schema_set, schema_set_id)
        if schema_set_id and not self.has_schema_set(schema_set_id):
            raise ValueError(f'schema_set_id {schema_set_id} not found')
        if not isinstance(data, (str, list)):
            raise TypeError(f'data must be a string or a list of DocNode, got {type(data)}')
        if isinstance(data, list) and any(not isinstance(n, DocNode) for n in data):
            raise TypeError('data list must contain DocNode instances')
        kb_id, doc_id, active_set_id = self._validate_extract_params(data)
        schema_set_id = schema_set_id or active_set_id
        search_res = self._get_extract_data(kb_id=kb_id, doc_ids=[doc_id])
        if search_res: return search_res[0]
        if schema_set_id not in self._schema_registry:
            raise ValueError(f'Unknown schema_set_id: {schema_set_id}')
        if self._extraction_mode == ExtractionMode.TEXT:
            res = self._text_extract_impl(data, schema_set_id)
        elif self._extraction_mode == ExtractionMode.MULTIMODAL:
            res = self._multimodal_extract_impl(data, schema_set_id)
        else:
            raise ValueError(f'Unknown extraction mode: {self._extraction_mode}')
        if not res:
            return None
        res_item = res[0] if isinstance(res, list) else res
        res_item.metadata.kb_id = kb_id
        res_item.metadata.doc_id = doc_id

        schema_model = self._schema_registry[schema_set_id]
        table_name = self._ensure_table(schema_set_id, schema_model)
        table_cls = self._sql_manager.get_table_orm_class(table_name)
        if table_cls is None:
            raise ValueError(f'Target table {table_name} not initialized')
        payload = {
            self.SYS_KB_ID: kb_id,
            self.SYS_DOC_ID: doc_id,
        }
        payload.update(self._to_model_dict(res_item.data, schema_model))
        meta_obj = getattr(res_item, 'metadata', None) or {}
        if isinstance(meta_obj, BaseModel):
            try:
                meta_payload = meta_obj.model_dump(mode='json')
            except AttributeError:
                meta_payload = meta_obj.dict(use_enum_values=True)
        elif isinstance(meta_obj, dict):
            meta_payload = self._json_safe(meta_obj)
        else:
            meta_payload = {}
        payload['extract_meta'] = self._json_safe(meta_payload)

        with self._sql_manager.get_session() as session:
            session.query(table_cls).filter_by(
                **{self.SYS_KB_ID: kb_id, self.SYS_DOC_ID: doc_id}
            ).delete()
            session.add(table_cls(**payload))
        return res_item

    def _delete_extract_data(self, doc_ids: List[str], kb_id: str = None) -> bool:
        try:
            self._lazy_init()
            if not self._sql_manager:
                raise ValueError('SqlManager is not initialized')
            if not doc_ids:
                return True
            if not self._active_schema_set_id:
                return True

            kb_id = kb_id or DEFAULT_KB_ID
            doc_ids = [str(d) for d in doc_ids]

            schema_set_id = self._active_schema_set_id
            table_name = self._table_name(schema_set_id)
            table_cls = self._sql_manager.get_table_orm_class(table_name)
            if table_cls is None:
                return True

            with self._sql_manager.get_session() as session:
                session.query(table_cls).filter_by(
                    **{self.SYS_KB_ID: kb_id}
                ).filter(
                    table_cls.doc_id.in_(doc_ids)
                ).delete(synchronize_session=False)
            return True
        except Exception as e:
            LOG.error(f'Failed to delete doc_ids={doc_ids} from kb_id={kb_id}', e)
            return False

    def _get_extract_data(self, doc_ids: List[str],  # noqa: C901
                          kb_id: str = None) -> List[ExtractResult]:
        self._lazy_init()
        if not self._sql_manager:
            raise ValueError('SqlManager is not initialized')
        if not doc_ids:
            return []
        if not self._active_schema_set_id:
            return []

        schema_set_id = self._active_schema_set_id
        self.has_schema_set(schema_set_id)
        table_name = self._table_name(schema_set_id)
        table_cls = self._sql_manager.get_table_orm_class(table_name)
        if table_cls is None:
            return []

        schema_model = self._schema_registry.get(schema_set_id)
        with self._sql_manager.get_session() as session:
            rows = session.query(table_cls).filter_by(
                **{self.SYS_KB_ID: kb_id}
            ).filter(
                table_cls.doc_id.in_(doc_ids)
            ).all()

        results: List[ExtractResult] = []
        sys_fields = {self.SYS_KB_ID, self.SYS_DOC_ID, 'extract_meta'}
        for row in rows:
            row_data = {}
            for col in table_cls.__table__.columns:
                name = col.name
                if name in sys_fields:
                    continue
                row_data[name] = getattr(row, name)
            if schema_model:
                try:
                    row_data = self._to_model_dict(row_data, schema_model)
                except Exception:
                    pass

            meta_payload = getattr(row, 'extract_meta', {}) or {}
            if not isinstance(meta_payload, dict):
                try:
                    meta_payload = json.loads(meta_payload)
                except Exception:
                    meta_payload = {}
            meta_payload = meta_payload if isinstance(meta_payload, dict) else {}
            meta_payload.setdefault('schema_set_id', schema_set_id)
            meta_payload.setdefault('kb_id', kb_id)
            meta_payload.setdefault('doc_id', str(getattr(row, self.SYS_DOC_ID, '')))
            try:
                meta = ExtractMeta(**meta_payload)
            except Exception:
                meta = ExtractMeta(schema_set_id=schema_set_id, kb_id=kb_id,
                                   doc_id=str(getattr(row, self.SYS_DOC_ID, '')))
            results.append(ExtractResult(data=row_data, metadata=meta))
        return results

    def forward(self, data: Union[str, List[DocNode]]) -> ExtractResult:
        self._lazy_init()
        res = self.extract_and_store(data=data)
        LOG.info(f'[Schema Extractor] extract res: {res}')
        return res

    def _init_sql_manager(self, db_config: Dict[str, Any]) -> SqlManager:
        return SqlManager(**db_config)

    def _table_name(self, schema_set_id: str) -> str:
        return f'{self._table_prefix}_{schema_set_id}'

    def _ensure_management_tables(self) -> None:
        tables_info_dict = {'tables': [TABLE_SCHEMA_SET_INFO]}
        try:
            self._sql_manager._init_tables_by_info(tables_info_dict)
        except Exception as e:
            LOG.warning(f'Ensure management tables failed: {e}')

    def _ensure_table(self, schema_set_id: str, schema_model: Optional[Type[BaseModel]] = None) -> str:
        if not self._sql_manager:
            raise ValueError('SqlManager is not initialized')
        table_name = self._table_name(schema_set_id)
        if table_name in self._table_cache:
            return table_name
        if schema_model is None:
            schema_model = self._schema_registry.get(schema_set_id)
        if schema_model is None:
            raise ValueError(f'No schema model registered for {schema_set_id}')

        attrs: Dict[str, Any] = {
            '__tablename__': table_name,
            '__table_args__': (
                sqlalchemy.Index(f'idx_{table_name}_kb', self.SYS_KB_ID),
                {'extend_existing': True},
            ),
        }
        attrs[self.SYS_KB_ID] = sqlalchemy.Column(sqlalchemy.String(128), primary_key=True, nullable=False)
        attrs[self.SYS_DOC_ID] = sqlalchemy.Column(sqlalchemy.String(128), primary_key=True, nullable=False)
        attrs['extract_meta'] = sqlalchemy.Column(sqlalchemy.JSON, nullable=True)

        for field_name, field_type in self._iter_schema_fields(schema_model):
            if field_name in attrs:
                continue
            attrs[field_name] = sqlalchemy.Column(field_type, nullable=True)

        table_cls = type(table_name.capitalize(), (_TableBase,), attrs)
        db_result = self._sql_manager.create_table(table_cls)
        if db_result.status != DBStatus.SUCCESS:
            LOG.warning(f'Create table failed: {db_result.detail}')
        else:
            self._table_cache[table_name] = table_cls
        return table_name

    def _iter_schema_fields(self, model: Type[BaseModel]) -> List[tuple[str, Any]]:
        try:
            fields = model.model_fields  # pydantic v2
        except AttributeError:
            fields = model.__fields__  # type: ignore[attr-defined]  # pydantic v1
        result = []
        for name, field in fields.items():
            annotation = getattr(field, 'annotation', None) or getattr(field, 'outer_type_', None)
            result.append((name, self._column_type(annotation)))
        return result

    def _normalize_py_type(self, type_hint: Any):
        if isinstance(type_hint, str):
            return self.TYPE_NAME_MAP.get(type_hint.lower(), str)
        return type_hint or str

    def _column_type(self, annotation: Any):
        origin = get_origin(annotation)
        args = get_args(annotation)
        if origin is Union and args:
            non_none = [arg for arg in args if arg is not type(None)]  # noqa: E721
            annotation = non_none[0] if non_none else str
            origin = get_origin(annotation)
        if origin in (list, set, tuple):
            return sqlalchemy.JSON
        resolved = self._normalize_py_type(annotation)
        if resolved in self.TYPE_MAP:
            return self.TYPE_MAP[resolved]
        if resolved in (list, set, tuple):
            return sqlalchemy.JSON
        return sqlalchemy.Text

    def _to_model_dict(self, payload: Union[BaseModel, Dict[str, Any]], model_cls: Type[BaseModel]) -> Dict[str, Any]:
        if isinstance(payload, BaseModel):
            try:
                return payload.model_dump()
            except AttributeError:
                return payload.dict()
        validated = model_cls(**payload)
        try:
            return validated.model_dump()
        except AttributeError:
            return validated.dict()

    def _validate_schema_model(self, model: Type[BaseModel]) -> None:
        if not model or not issubclass(model, BaseModel):
            raise TypeError('schema_set must be a pydantic BaseModel subclass')

    def _json_safe(self, obj: Any) -> Any:
        """Convert common objects (Enum/BaseModel) to JSON-serializable primitives."""
        if isinstance(obj, Enum):
            return obj.value
        if isinstance(obj, BaseModel):
            try:
                return obj.model_dump(mode='json')
            except AttributeError:
                return obj.dict(use_enum_values=True)
        if isinstance(obj, dict):
            return {k: self._json_safe(v) for k, v in obj.items()}
        if isinstance(obj, (list, tuple, set)):
            return [self._json_safe(v) for v in obj]
        return obj

analyze_schema_and_register(data, schema_set_id=None)

Infer a schema from sample data, register it, and return the registration info.

Source code in lazyllm/tools/rag/doc_to_db/extractor.py
def analyze_schema_and_register(self, data: Union[str, List[DocNode]],
                                schema_set_id: Optional[str] = None) -> SchemaSetInfo:
    """Infer a schema from sample data, register it, and return the registration info."""
    self._lazy_init()
    if not self._llm:
        raise ValueError('LLM not initialized')
    if not data:
        raise ValueError('data is empty')

    if isinstance(data, str):
        sample_text = data[:self._max_len]
    else:
        chunks = self._gen_text_list_from_nodes(data)
        sample_text = '\n\n'.join(chunks)[:self._max_len] if chunks else ''
    if not sample_text:
        raise ValueError('No content available for schema analysis')

    llm = self._llm.share(prompt=SCHEMA_ANALYZE_PROMPT, format=JsonFormatter())
    payload = Template(SCHEMA_ANALYZE_INPUT_FORMAT).substitute(text=sample_text)
    res = llm(payload)
    fields_def: Dict[str, Tuple[Any, Any]] = {}
    for item in res:
        if not isinstance(item, dict):
            continue
        name = item.get('name')
        if not name:
            continue
        desc = item.get('description') or ''
        py_type = self._normalize_py_type(item.get('type'))
        fields_def[name] = (py_type, Field(default=None, description=desc))
    if not fields_def:
        # Fallback: single generic field capturing text content
        fields_def['content'] = (str, Field(default=None, description='Raw content snippet'))

    model_name = f'AutoSchema{uuid4().hex}'
    schema_model = create_model(model_name, **fields_def)  # type: ignore[arg-type]
    reg_id = self.register_schema_set(schema_model, schema_set_id)
    return SchemaSetInfo(schema_set_id=reg_id, schema_model=schema_model)

extract_and_store(data, schema_set_id=None, schema_set=None)

按已注册的 schema 抽取文本/DocNode 内容并写入对应表,若传入 schema_set 会先注册;同文档重复调用会返回缓存结果。

Parameters:

  • data (Union[str, List[DocNode]]) –

    文本或 DocNode 列表(需同一文档)。

  • schema_set_id (str, default: None ) –

    指定使用的 schema 集合 ID。

  • schema_set (Type[BaseModel], default: None ) –

    动态注册并使用的 schema。

Returns:

  • ExtractResult: 抽取结果,data 为字段名到值的字典,metadata 包含 schema_set_id、kb_id、doc_id 及按字段的线索信息;可能为 None 表示无可写入。
Source code in lazyllm/tools/rag/doc_to_db/extractor.py
    def extract_and_store(self, data: Union[str, List[DocNode]],  # noqa: C901
                          schema_set_id: str = None, schema_set: Type[BaseModel] = None) -> ExtractResult:
        """
按已注册的 schema 抽取文本/DocNode 内容并写入对应表,若传入 schema_set 会先注册;同文档重复调用会返回缓存结果。

Args:
    data (Union[str, List[DocNode]]): 文本或 DocNode 列表(需同一文档)。
    schema_set_id (str, optional): 指定使用的 schema 集合 ID。
    schema_set (Type[BaseModel], optional): 动态注册并使用的 schema。

**Returns:**

- ExtractResult: 抽取结果,`data` 为字段名到值的字典,`metadata` 包含 schema_set_id、kb_id、doc_id 及按字段的线索信息;可能为 None 表示无可写入。
"""
        self._lazy_init()
        if schema_set is not None:
            schema_set_id = self.register_schema_set(schema_set, schema_set_id)
        if schema_set_id and not self.has_schema_set(schema_set_id):
            raise ValueError(f'schema_set_id {schema_set_id} not found')
        if not isinstance(data, (str, list)):
            raise TypeError(f'data must be a string or a list of DocNode, got {type(data)}')
        if isinstance(data, list) and any(not isinstance(n, DocNode) for n in data):
            raise TypeError('data list must contain DocNode instances')
        kb_id, doc_id, active_set_id = self._validate_extract_params(data)
        schema_set_id = schema_set_id or active_set_id
        search_res = self._get_extract_data(kb_id=kb_id, doc_ids=[doc_id])
        if search_res: return search_res[0]
        if schema_set_id not in self._schema_registry:
            raise ValueError(f'Unknown schema_set_id: {schema_set_id}')
        if self._extraction_mode == ExtractionMode.TEXT:
            res = self._text_extract_impl(data, schema_set_id)
        elif self._extraction_mode == ExtractionMode.MULTIMODAL:
            res = self._multimodal_extract_impl(data, schema_set_id)
        else:
            raise ValueError(f'Unknown extraction mode: {self._extraction_mode}')
        if not res:
            return None
        res_item = res[0] if isinstance(res, list) else res
        res_item.metadata.kb_id = kb_id
        res_item.metadata.doc_id = doc_id

        schema_model = self._schema_registry[schema_set_id]
        table_name = self._ensure_table(schema_set_id, schema_model)
        table_cls = self._sql_manager.get_table_orm_class(table_name)
        if table_cls is None:
            raise ValueError(f'Target table {table_name} not initialized')
        payload = {
            self.SYS_KB_ID: kb_id,
            self.SYS_DOC_ID: doc_id,
        }
        payload.update(self._to_model_dict(res_item.data, schema_model))
        meta_obj = getattr(res_item, 'metadata', None) or {}
        if isinstance(meta_obj, BaseModel):
            try:
                meta_payload = meta_obj.model_dump(mode='json')
            except AttributeError:
                meta_payload = meta_obj.dict(use_enum_values=True)
        elif isinstance(meta_obj, dict):
            meta_payload = self._json_safe(meta_obj)
        else:
            meta_payload = {}
        payload['extract_meta'] = self._json_safe(meta_payload)

        with self._sql_manager.get_session() as session:
            session.query(table_cls).filter_by(
                **{self.SYS_KB_ID: kb_id, self.SYS_DOC_ID: doc_id}
            ).delete()
            session.add(table_cls(**payload))
        return res_item

has_schema_set(schema_set_id)

检查指定 schema_set_id 是否已注册,缺失时会尝试从数据库恢复模型并建表。

Parameters:

  • schema_set_id (str) –

    目标 schema 集合 ID。

Returns:

  • bool: 是否已存在。
Source code in lazyllm/tools/rag/doc_to_db/extractor.py
    def has_schema_set(self, schema_set_id: str) -> bool:
        """
检查指定 schema_set_id 是否已注册,缺失时会尝试从数据库恢复模型并建表。

Args:
    schema_set_id (str): 目标 schema 集合 ID。

**Returns:**

- bool: 是否已存在。
"""
        self._lazy_init()
        if self._sql_manager:
            table_cls = self._sql_manager.get_table_orm_class(TABLE_SCHEMA_SET_INFO['name'])
            if table_cls is None:
                raise ValueError('Schema set table not initialized')
            with self._sql_manager.get_session() as session:
                existing = session.query(table_cls).filter_by(schema_set_id=schema_set_id).first()
                if not existing:
                    return False
                if schema_set_id not in self._schema_registry:
                    recovered_schema = self._model_from_schema_json(existing.schema_set_json,
                                                                    model_name=f'Schema_{schema_set_id}')
                    self._schema_registry[schema_set_id] = recovered_schema
                    self._ensure_table(schema_set_id, recovered_schema)
                return True
        return schema_set_id in self._schema_registry

register_schema_set(schema_set, schema_set_id=None, force_refresh=False)

schema set registration, idempotent

Source code in lazyllm/tools/rag/doc_to_db/extractor.py
def register_schema_set(self, schema_set: Type[BaseModel], schema_set_id: str = None,   # noqa: C901
                        force_refresh: bool = False) -> str:
    """schema set registration, idempotent"""
    try:
        self._lazy_init()
        self._validate_schema_model(schema_set)

        fields = getattr(schema_set, 'model_fields', None) or getattr(schema_set, '__fields__', {})

        def _safe_default(val: Any):
            if val is None:
                return None
            if val.__class__.__name__ in ('PydanticUndefinedType', 'UndefinedType'):
                return None
            if isinstance(val, (str, int, float, bool)):
                return val
            return str(val)

        signature = [
            (name, str(getattr(f, 'annotation', None) or getattr(f, 'outer_type_', None)),
             _safe_default(getattr(f, 'default', None)), getattr(f, 'default_factory', None) is not None,
             getattr(f, 'is_required', lambda: False)())
            for name, f in fields.items()
        ]
        signature.sort(key=lambda x: x[0])
        idem_key = json.dumps(signature, ensure_ascii=False)

        if self._sql_manager:
            table_cls = self._sql_manager.get_table_orm_class(TABLE_SCHEMA_SET_INFO['name'])
            if table_cls is None:
                raise ValueError('Schema set table not initialized')
            with self._sql_manager.get_session() as session:
                existing = session.query(table_cls).filter_by(idem_key=idem_key).first()
                if existing:
                    existing_id = str(existing.schema_set_id if hasattr(existing, 'schema_set_id') else existing.id)
                    if schema_set_id and str(schema_set_id) != existing_id:
                        raise ValueError(
                            f'schema_set_id mismatch for idem_key, expect {existing_id}, got {schema_set_id}'
                        )
                    schema_set_id = schema_set_id or existing_id
                else:
                    schema_json = (schema_set.model_json_schema() if hasattr(schema_set, 'model_json_schema')
                                   else schema_set.schema())
                    desc = (schema_set.__doc__ or '').strip() or 'Schema set'
                    obj_kwargs = dict(schema_set_json=json.dumps(schema_json, ensure_ascii=False),
                                      desc=desc, idem_key=idem_key, created_at=datetime.now(),
                                      updated_at=datetime.now())
                    if schema_set_id is None:
                        schema_set_id = str(uuid4().hex)
                    obj_kwargs['schema_set_id'] = str(schema_set_id)
                    new_obj = table_cls(**obj_kwargs)
                    session.add(new_obj)
                    session.flush()
                    schema_set_id = str(new_obj.schema_set_id if hasattr(new_obj, 'schema_set_id') else new_obj.id)

        if schema_set_id is None:
            raise ValueError('schema_set_id is required and could not be derived')

        schema_set_id = str(schema_set_id)
        self._schema_registry[schema_set_id] = schema_set
        if self._sql_manager:
            self._ensure_table(schema_set_id, schema_set)
        self._active_schema_set_id = schema_set_id
        return schema_set_id
    except Exception as e:
        LOG.error(f'Failed to register schema set: {e}')
        raise e

sql_manager_for_nl2sql(kb_ids=None)

基于已注册的 schema,生成一个仅暴露相关表的 SqlManager,用于 SqlCall 模块中 NL2SQL 查询;会附带表结构描述和可见表列表。

Parameters:

  • kb_ids (Union[str, List[str]], default: None ) –

    过滤的知识库 ID,可单个或列表。

Returns:

  • SqlManager: 仅包含可见表、列信息及说明的 SqlManager 实例,用于 NL2SQL。
Source code in lazyllm/tools/rag/doc_to_db/extractor.py
    def sql_manager_for_nl2sql(self,  # noqa: C901
                               kb_ids: Union[str, List[str]] = None) -> SqlManager:
        """
基于已注册的 schema,生成一个仅暴露相关表的 SqlManager,用于 SqlCall 模块中 NL2SQL 查询;会附带表结构描述和可见表列表。

Args:
    kb_ids (Union[str, List[str]], optional): 过滤的知识库 ID,可单个或列表。

**Returns:**

- SqlManager: 仅包含可见表、列信息及说明的 SqlManager 实例,用于 NL2SQL。
"""
        self._lazy_init()
        if not self._sql_manager:
            raise ValueError('SqlManager is not initialized')
        if not self._db_config:
            raise ValueError('db_config is required to build SqlManager')
        if not self._active_schema_set_id:
            raise ValueError('No active schema set registered')

        schema_info_table = TABLE_SCHEMA_SET_INFO['name']
        desc_map: Dict[str, str] = {}

        def _schema_table_desc(model: Type[BaseModel]) -> str:
            schema_desc = self._get_schema_set_str(model)
            return '\n'.join([s for s in [
                (model.__doc__ or '').strip(),
                schema_desc,
                f'System columns: {self.SYS_KB_ID}, {self.SYS_DOC_ID}, extract_meta',
            ] if s])

        kb_id_list = None
        if kb_ids is not None:
            if isinstance(kb_ids, (list, tuple, set)):
                kb_id_list = [str(k) for k in kb_ids if k is not None]
            else:
                kb_id_list = [str(kb_ids)]
            if not kb_id_list:
                kb_id_list = None

        schema_set_id = self._active_schema_set_id
        if not self.has_schema_set(schema_set_id):
            raise ValueError(f'Schema set {schema_set_id} not found')
        schema_model = self._schema_registry[schema_set_id]
        table_name = self._ensure_table(schema_set_id, schema_model)
        target_tables = {table_name}
        desc_map[table_name] = _schema_table_desc(schema_model)

        target_tables.discard(schema_info_table)
        tables_info_dict = {'tables': []}
        for table_name in target_tables:
            table_cls = self._sql_manager.get_table_orm_class(table_name)
            if table_cls is None:
                continue
            columns = []
            for col in table_cls.__table__.columns:
                columns.append({
                    'name': col.name,
                    'data_type': _col_type_name(col),
                    'nullable': bool(col.nullable),
                    'is_primary_key': bool(col.primary_key),
                    'comment': getattr(col, 'comment', '') or '',
                })
            tables_info_dict['tables'].append({'name': table_name, 'columns': columns, 'comment': ''})
        new_manager = self._init_sql_manager({**self._db_config, 'tables_info_dict': tables_info_dict})
        new_manager.visible_tables = list(target_tables)
        if desc_map:
            new_manager.set_desc(desc_map)
        return new_manager

lazyllm.tools.rag.readers.DocxReader

Bases: _RichReader

docx格式文件解析器,从 .docx 文件中读取文本内容并封装为文档节点(DocNode)列表。

Parameters:

  • file (Path) –

    .docx 文件路径。

  • fs (Optional[AbstractFileSystem]) –

    可选的文件系统对象,支持自定义读取方式。

Returns:

  • List[DocNode]: 包含文档中所有文本内容的节点列表。
Source code in lazyllm/tools/rag/readers/docxReader.py
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
class DocxReader(_RichReader):
    """docx格式文件解析器,从 `.docx` 文件中读取文本内容并封装为文档节点(DocNode)列表。

Args:
    file (Path): `.docx` 文件路径。
    fs (Optional[AbstractFileSystem]): 可选的文件系统对象,支持自定义读取方式。

**Returns:**

- List[DocNode]: 包含文档中所有文本内容的节点列表。
"""
    def __init__(self, split_doc: Optional[bool] = False, extra_info: Optional[Dict] = None,
                 extract_process: Optional[Callable] = None, post_func: Optional[Callable] = None,
                 extract_global_info: bool = True, image_save_path: Optional[str] = None,
                 save_image: bool = True, return_trace: bool = True):
        super().__init__(split_doc=split_doc, return_trace=return_trace, post_func=None)
        self._post_func = post_func or self._default_post
        self.extract_process = extract_process or self._default_extract
        self.extract_global_info = extract_global_info
        self._extra_info = extra_info or {}
        self._image_save_path = image_save_path
        self._save_image = save_image

    def _extract_global_info(self, doc: 'docx.Document', file_path: Path) -> Dict[str, Any]:
        global_info = dict(self._extra_info)

        if self.extract_global_info and self._split_doc:
            try:
                props = doc.core_properties

                str_props = ['author', 'title', 'subject', 'keywords', 'comments']

                special_props = {
                    'created': lambda x: x.isoformat(),
                    'modified': lambda x: x.isoformat(),
                    'revision': lambda x: x,
                }

                for prop_name in str_props:
                    prop_value = getattr(props, prop_name, None)
                    if prop_value:
                        global_info[prop_name] = str(prop_value)

                for prop_name, converter in special_props.items():
                    prop_value = getattr(props, prop_name, None)
                    if prop_value is not None:
                        try:
                            global_info[prop_name] = converter(prop_value)
                        except (AttributeError, ValueError) as e:
                            LOG.debug(f'Failed to convert {prop_name}: {e}')

                global_info['file_path'] = str(file_path)
                global_info['file_name'] = file_path.name
                global_info['file_size'] = file_path.stat().st_size if file_path.exists() else 0

            except Exception as e:
                LOG.warning(f'Failed to extract global info from {file_path}: {e}')

        return global_info

    def _load_data(self, file: Path, fs: Optional['fsspec.AbstractFileSystem'] = None,
                   **kwargs) -> List[DocNode]:
        if not isinstance(file, Path):
            file = Path(file)

        if file.name.endswith('.doc'):
            raise ValueError(f'Only expected docx file, but got {file.name}')

        if self._split_doc:
            try:
                return self._enhanced_load(file, fs, **kwargs)

            except Exception:
                try:
                    return self._load(file, fs, **kwargs)
                except Exception as e:
                    raise e
        return self._load(file, fs, **kwargs)

    def _load(self, file: Path, fs: Optional['fsspec.AbstractFileSystem'] = None, **kwargs) -> List[DocNode]:
        try:
            if fs:
                with fs.open(file) as f:
                    text = docx2txt.process(f)
            else:
                text = docx2txt.process(file)
            if not text:
                raise ValueError(f"Fail loading file {file.name}, maybe it's empty")
            return [DocNode(text=text)]
        except Exception as docx2txt_error:
            LOG.error(f'Failed for {file}: {str(docx2txt_error)}')
            raise

    def _enhanced_load(self, file: Path, fs: Optional['fsspec.AbstractFileSystem'] = None, **kwargs) -> List[DocNode]:
        with pipeline() as p:
            p.f1 = self._read_file
            p.f2 = bind(self.extract_process, file, _0)
            p.f3 = bind(self._post_func, _0, **kwargs)

        nodes = p(file, fs)
        return nodes

    def _read_file(self, file: Path, fs: Optional['fsspec.AbstractFileSystem'] = None) -> 'docx.Document':
        fs = fs or get_default_fs()

        try:
            file_size = fs.size(file)
            if file_size == 0:
                raise ValueError(f'Input file {file.name} is empty')

        except Exception as e:
            LOG.error(f'Fail to load file for {file}: {e}')
            raise e

        temp_files_to_cleanup = []
        temp_path = None

        try:
            if is_default_fs(fs):
                doc = docx.Document(docx=str(file))
            else:
                with tempfile.NamedTemporaryFile(suffix='.docx', delete=False) as tmp_file:
                    temp_path = tmp_file.name

                with fs.open(file, 'rb') as remote_file:
                    tmp_file.write(remote_file.read())

                doc = docx.Document(docx=temp_path)
            return doc

        except Exception as e:
            LOG.error(f'[ERROR] file--{file.name}--Wrong file, failed to read file: {e}')
            raise
        finally:
            for temp_file in temp_files_to_cleanup:
                try:
                    if os.path.exists(temp_file):
                        os.unlink(temp_file)
                except Exception as e:
                    LOG.warning(f'Failed to clean up temporary file {temp_file}: {e}')

    def _default_extract(self, file, doc) -> List[DocNode]:  # noqa: C901
        global_info = self._extract_global_info(doc, file)

        base_metadata = {'file_name': file.name}

        doc_list = []

        paragraphs = list(doc.paragraphs)
        tables = list(doc.tables)
        paragraph_idx = 0
        table_idx = 0

        elements = list(doc.element.body)
        for element in elements:
            if element.tag.endswith('tbl'):
                if table_idx < len(tables):
                    table = tables[table_idx]
                    table_idx += 1

                    table_node = self._process_table(table, base_metadata, global_info)
                    doc_list.append(table_node)

            elif element.tag.endswith('p'):
                if paragraph_idx >= len(paragraphs):
                    continue

                para = paragraphs[paragraph_idx]
                paragraph_idx += 1

                has_image = False
                for run in para.runs:
                    has_drawing = run.element.xpath('.//*[local-name()="drawing"]')
                    has_imagedata = run.element.xpath('.//*[local-name()="imagedata"]')
                    if has_drawing or has_imagedata:
                        has_image = True

                    if has_image:
                        image_nodes = self._extract_images_from_paragraph(
                            para, doc, base_metadata, global_info
                        )
                        if image_nodes:
                            doc_list.extend(image_nodes)

                        if self._content_clean(element):
                            continue

                    math_text = self._extract_math_from_element(element)
                    if math_text:
                        math_node = self._process_math(math_text, base_metadata, global_info)
                        doc_list.append(math_node)
                        if self._content_clean(element):
                            continue

                    if self._content_clean(element):
                        continue

                    content = element.text.replace('\u3000', ' ').strip('\n') if element.text else ''

                    para_node = self._process_paragraph(para, content, base_metadata, global_info)
                    doc_list.append(para_node)
        if not doc_list:
            raise ValueError('file Extraction failed')
        return doc_list

    def _default_post(self, doc_list, **kwargs) -> List[DocNode]:
        for index, node in enumerate(doc_list):
            node.metadata['index'] = index

        for node in doc_list:
            node.excluded_embed_metadata_keys = ['style_dict', 'type', 'index', 'text_level', 'lines']
            node.excluded_llm_metadata_keys = ['style_dict', 'type', 'index', 'text_level', 'lines']

        return doc_list

    def _content_clean(self, element) -> bool:
        content = element.text.replace('\u3000', ' ').strip('\n') if element.text else ''
        content_clean = content.replace(' ', '').replace('\n', '')

        return not content_clean

    def _get_aligned_type(self, para: 'docx.text.paragraph') -> str:
        aligned_type = {
            docx.enum.text.WD_ALIGN_PARAGRAPH.LEFT: 'left',
            docx.enum.text.WD_ALIGN_PARAGRAPH.CENTER: 'center',
            docx.enum.text.WD_ALIGN_PARAGRAPH.RIGHT: 'right',
            docx.enum.text.WD_ALIGN_PARAGRAPH.JUSTIFY: 'both_ends',
            docx.enum.text.WD_ALIGN_PARAGRAPH.DISTRIBUTE: 'distribute',
            None: '',
        }

        try:
            alignment = para.alignment
        except (AttributeError, TypeError):
            alignment = None

        return aligned_type.get(alignment, '')

    def _get_style_info(self, style: 'docx.styles.style.ParagraphStyle') -> dict:
        try:
            font = style.font
            style_dict = {
                'style_name': style.name,
                'style_type': style.type,
                'font_name': font.name if font.name else None,
                'font_bold': bool(font.bold) if font.bold is not None else False,
                'font_size': font.size.pt if font.size else None,
            }
        except Exception:
            style_dict = {
                'style_name': style.name if style else '',
                'style_type': style.type if style else None,
                'font_name': None,
                'font_bold': False,
                'font_size': None,
            }
        return style_dict

    def _extract_images_from_paragraph(self, para, doc: 'docx.Document', base_metadata: dict,
                                       extra_info: Optional[Dict]) -> List[DocNode]:
        image_nodes = []

        for run in para.runs:
            has_drawing = run.element.xpath('.//*[local-name()="drawing"]')
            has_imagedata = run.element.xpath('.//*[local-name()="imagedata"]')

            if not (has_drawing or has_imagedata):
                continue

            for shape in run.element.findall('.//{http://schemas.openxmlformats.org/drawingml/2006/main}blip'):
                embed_id = shape.get('{http://schemas.openxmlformats.org/officeDocument/2006/relationships}embed')
                if not embed_id:
                    continue

                try:
                    image_part = doc.part.related_parts[embed_id]

                    if self._save_image:
                        image_data = image_part.blob
                        original_filename = os.path.basename(image_part.partname)
                        file_extension = os.path.splitext(original_filename)[1] or '.png'
                        image_filename = f'{uuid.uuid4()}{file_extension}'

                        try:
                            os.makedirs(self._image_save_path, exist_ok=True)
                        except Exception:
                            LOG.warning('use default image save path ~/.lazyllm/image')
                            image_path = os.path.join(os.path.expanduser('~'), '.lazyllm')
                            self._image_save_path = Path(image_path) / 'image'
                            continue

                        image_save_path = os.path.join(self._image_save_path, image_filename)
                        with open(image_save_path, 'wb') as img_file:
                            img_file.write(image_data)
                    else:
                        original_filename = os.path.basename(image_part.partname)
                        file_extension = os.path.splitext(original_filename)[1] or '.png'
                        image_filename = f'image_{uuid.uuid4()}{file_extension}'

                    if self._save_image:
                        image_text = f'![]({image_filename})'
                    else:
                        image_text = ''

                    metadata = copy.deepcopy(base_metadata)
                    metadata['type'] = 'image'

                    image_nodes.append(DocNode(
                        text=image_text,
                        metadata=metadata,
                        global_metadata=extra_info
                    ))
                except Exception as e:
                    LOG.error(f'[Docx Reader] Failed to extract image: {e}')

        return image_nodes

    def _extract_math_from_element(self, element) -> Optional[str]:
        math_nodes = element.xpath('.//*[local-name()="oMath"] | .//*[local-name()="oMathPara"]')
        if not math_nodes:
            return None

        try:
            math = math_nodes[0]
            return self._math_to_text(math)
        except Exception as e:
            LOG.error(f'[Docx Reader] Failed to extract math: {e}')
            return None

    def _math_to_text(self, math_node) -> str:
        try:
            def text_generator(node):
                if node.text:
                    yield node.text

                for child in node:
                    yield from text_generator(child)

                if node.tail:
                    yield node.tail

            return ''.join(text_generator(math_node))
        except Exception as e:
            LOG.error(f'[Docx Reader] Failed to convert math to text: {e}')
            return ''

    def _table_to_markdown(self, table: 'docx.table.Table') -> str:
        if not table.rows:
            return '\n[empty table]\n'

        try:
            col_size = len(table.rows[0].cells)
            md_lines = []

            for row_idx, row in enumerate(table.rows):
                cells = []
                for i in range(min(col_size, len(row.cells))):
                    cell = row.cells[i]
                    text = getattr(cell, 'text', '') or ''.join(
                        p.text for p in getattr(cell, 'paragraphs', [])
                    )
                    cells.append(text.replace('\n', ' ').replace('\r', ' ').strip())

                cells.extend([''] * (col_size - len(cells)))

                md_lines.append('| ' + ' | '.join(cells) + ' |')

                if row_idx == 0:
                    md_lines.append('|' + '|'.join([' --- '] * col_size) + '|')

            return '\n' + '\n'.join(md_lines) + '\n'

        except Exception:
            return '\n[Table parse failed]\n'

    def _process_table(self, table: 'docx.table.Table', base_metadata: dict,
                       extra_info: Optional[Dict] = None) -> DocNode:
        metadata = copy.deepcopy(base_metadata)
        metadata['type'] = 'table'

        table_md = self._table_to_markdown(table)
        table_text = table_md

        return DocNode(text=table_text, metadata=metadata, global_metadata=extra_info)

    def _check_run_bold(self, para) -> bool:
        try:
            runs = [run for run in para.runs if run.text.strip()]
            return bool(runs and all(run.font.bold for run in runs))
        except Exception:
            return False

    def _process_paragraph(self, para, content: str, base_metadata: dict,
                           extra_info: Optional[Dict]) -> DocNode:
        metadata = copy.deepcopy(base_metadata)
        style_dict = self._get_style_info(para.style)

        if not style_dict.get('font_bold') and self._check_run_bold(para):
            style_dict['font_bold'] = True

        aligned_type = self._get_aligned_type(para)
        style_dict.update({'aligned_type': aligned_type})

        number_title_pattern = r'^(\d{1,2}(?:\.\d{1,2})+)\s*([^\d].*)$'
        if style_dict.get('font_bold') and re.match(number_title_pattern, content):
            content = f'**{content}**'

        metadata['style_dict'] = style_dict
        metadata['type'] = 'text'

        return DocNode(text=content, metadata=metadata, global_metadata=extra_info)

    def _process_math(self, math_text: str, base_metadata: dict,
                      extra_info: Optional[Dict]) -> DocNode:
        metadata = copy.deepcopy(base_metadata)
        metadata['type'] = 'equation'
        return DocNode(text=math_text, metadata=metadata, global_metadata=extra_info)

lazyllm.tools.rag.readers.EpubReader

Bases: LazyLLMReaderBase

用于读取 .epub 格式电子书的文件读取器。

继承自 LazyLLMReaderBase,只需实现 _load_data 方法,即可通过 Document 组件自动加载 .epub 文件中的内容。

注意:当前版本不支持通过 fsspec 文件系统(如远程路径)加载 epub 文件,若提供 fs 参数,将回退到本地文件读取。

Returns:

  • List[DocNode]: 所有章节内容合并后的文本节点列表。
Source code in lazyllm/tools/rag/readers/epubReader.py
class EpubReader(LazyLLMReaderBase):
    """用于读取 `.epub` 格式电子书的文件读取器。

继承自 `LazyLLMReaderBase`,只需实现 `_load_data` 方法,即可通过 `Document` 组件自动加载 `.epub` 文件中的内容。

注意:当前版本不支持通过 fsspec 文件系统(如远程路径)加载 epub 文件,若提供 `fs` 参数,将回退到本地文件读取。

**Returns:**

- List[DocNode]: 所有章节内容合并后的文本节点列表。
"""
    def _load_data(self, file: Path, fs: Optional['fsspec.AbstractFileSystem'] = None) -> List[DocNode]:
        if not isinstance(file, Path): file = Path(file)

        if fs:
            LOG.warning('fs was specified but EpubReader doesn\'t support loading from '
                        'fsspec filesystems. Will load from local filesystem instead.')

        text_list = []

        spec = importlib.util.find_spec('ebooklib.epub')
        if spec is None:
            raise ImportError(
                'Please install ebooklib to use ebooklib module. '
                'You can install it with `pip install ebooklib`'
            )
        epub_module = importlib.util.module_from_spec(spec)
        spec.loader.exec_module(epub_module)

        book = epub_module.read_epub(file, options={'ignore_ncs': True})

        for item in book.get_items():
            if item.get_type() == ebooklib.ITEM_DOCUMENT:
                text_list.append(html2text.html2text(item.get_content().decode('utf-8')))
        text = '\n'.join(text_list)
        return [DocNode(text=text)]

lazyllm.tools.rag.readers.HWPReader

Bases: LazyLLMReaderBase

HWP文件解析器,支持从本地文件系统读取 HWP 文件。它会从文档中提取正文部分的文本内容,返回 DocNode 列表。

HWP 是一种专有的二进制格式,主要在韩国使用。由于格式封闭,因此只能解析部分内容(如文本段落),但对常规文本提取已经足够使用。

Parameters:

  • return_trace (bool, default: True ) –

    是否启用 trace 日志记录,默认为 True

Source code in lazyllm/tools/rag/readers/hwpReader.py
class HWPReader(LazyLLMReaderBase):
    """HWP文件解析器,支持从本地文件系统读取 HWP 文件。它会从文档中提取正文部分的文本内容,返回 DocNode 列表。

HWP 是一种专有的二进制格式,主要在韩国使用。由于格式封闭,因此只能解析部分内容(如文本段落),但对常规文本提取已经足够使用。

Args:
    return_trace (bool): 是否启用 trace 日志记录,默认为 ``True``。
"""
    def __init__(self, return_trace: bool = True) -> None:
        super().__init__(return_trace=return_trace)
        self._FILE_HEADER_SECTION = 'FileHeader'
        self._HWP_SUMMARY_SECTION = '\x05HwpSummaryInformation'
        self._SECTION_NAME_LENGTH = len('Section')
        self._BODYTEXT_SECTION = 'BodyText'
        self._HWP_TEXT_TAGS = [67]
        self._text = ''

    def _load_data(self, file: Path, fs: Optional['fsspec.AbstractFileSystem'] = None) -> List[DocNode]:
        if fs:
            LOG.warning('fs was specified but HWPReader doesn\'t support loading from '
                        'fsspec filesystems. Will load from local filesystem instead.')

        if not isinstance(file, Path): file = Path(file)

        load_file = olefile.OleFileIO(file)
        file_dir = load_file.listdir()
        if self._is_valid(file_dir) is False: raise Exception('Not Valid HwpFile')

        result_text = self._get_text(load_file, file_dir)
        return [DocNode(text=result_text)]

    def _is_valid(self, dirs: List[str]) -> bool:
        if [self._FILE_HEADER_SECTION] not in dirs: return False
        return [self._HWP_SUMMARY_SECTION] in dirs

    def _get_text(self, load_file: Any, file_dirs: List[str]) -> str:
        sections = self._get_body_sections(file_dirs)
        text = ''
        for section in sections:
            text += self._get_text_from_section(load_file, section)
            text += '\n'

        self._text = text
        return self._text

    def _get_body_sections(self, dirs: List[str]) -> List[str]:
        m = []
        for d in dirs:
            if d[0] == self._BODYTEXT_SECTION:
                m.append(int(d[1][self._SECTION_NAME_LENGTH:]))

        return ['BodyText/Section' + str(x) for x in sorted(m)]

    def _is_compressed(self, load_file: Any) -> bool:
        header = load_file.openstream('FileHeader')
        header_data = header.read()
        return (header_data[36] & 1) == 1

    def _get_text_from_section(self, load_file: Any, section: str) -> str:
        bodytext = load_file.openstream(section)
        data = bodytext.read()

        unpacked_data = (zlib.decompress(data, -15) if self._is_compressed(load_file) else data)
        size = len(unpacked_data)

        i = 0
        text = ''
        while i < size:
            header = struct.unpack_from('<I', unpacked_data, i)[0]
            rec_type = header & 0x3FF
            (header >> 10) & 0x3FF
            rec_len = (header >> 20) & 0xFFF

            if rec_type in self._HWP_TEXT_TAGS:
                rec_data = unpacked_data[i + 4: i + 4 + rec_len]
                text += rec_data.decode('utf-16')
                text += '\n'

            i += 4 + rec_len
        return text

lazyllm.tools.rag.readers.ImageReader

Bases: LazyLLMReaderBase

用于从图片文件中读取内容的模块。支持保留图片、解析图片中的文本(基于OCR或预训练视觉模型),并返回文本和图片路径的节点列表。

Parameters:

  • parser_config (Optional[Dict], default: None ) –

    解析器配置,包含模型和处理器,默认为 None。当设置 parse_text=True 且 parser_config=None 时,会自动根据 text_type 加载相应模型。

  • keep_image (bool, default: False ) –

    是否保留图片的 base64 编码,默认为 False。

  • parse_text (bool, default: False ) –

    是否解析图片中的文本,默认为 False。

  • text_type (str, default: 'text' ) –

    解析文本的类型,支持 text(默认)和 plain_text。当为 plain_text 时,使用 pytesseract 进行OCR;否则使用预训练视觉编码解码模型。

  • pytesseract_model_kwargs (Optional[Dict], default: None ) –

    传递给 pytesseract OCR 的可选参数,默认为空字典。

  • return_trace (bool, default: True ) –

    是否记录处理过程的 trace,默认为 True。

Source code in lazyllm/tools/rag/readers/imageReader.py
class ImageReader(LazyLLMReaderBase):
    """用于从图片文件中读取内容的模块。支持保留图片、解析图片中的文本(基于OCR或预训练视觉模型),并返回文本和图片路径的节点列表。

Args:
    parser_config (Optional[Dict]): 解析器配置,包含模型和处理器,默认为 None。当设置 parse_text=True 且 parser_config=None 时,会自动根据 text_type 加载相应模型。
    keep_image (bool): 是否保留图片的 base64 编码,默认为 False。
    parse_text (bool): 是否解析图片中的文本,默认为 False。
    text_type (str): 解析文本的类型,支持 ``text``(默认)和 ``plain_text``。当为 ``plain_text`` 时,使用 pytesseract 进行OCR;否则使用预训练视觉编码解码模型。
    pytesseract_model_kwargs (Optional[Dict]): 传递给 pytesseract OCR 的可选参数,默认为空字典。
    return_trace (bool): 是否记录处理过程的 trace,默认为 True。
"""
    def __init__(self, parser_config: Optional[Dict] = None, keep_image: bool = False, parse_text: bool = False,
                 text_type: str = 'text', pytesseract_model_kwargs: Optional[Dict] = None,
                 return_trace: bool = True) -> None:
        super().__init__(return_trace=return_trace)
        self._text_type = text_type
        if parser_config is None and parse_text:
            if text_type == 'plain_text':
                try:
                    import pytesseract
                except ImportError:
                    raise ImportError('Please install extra dependencies that are required for the ImageReader '
                                      'when text_type is "plain_text": `pip install pytesseract`')

                processor = None
                model = pytesseract
            else:
                thirdparty.check_packages(['sentencepiece', 'torch', 'transformers'])

                processor = tf.DonutProcessor.from_pretrained('naver-clova-ix/donut-base-finetuned-cord-v2')
                model = tf.VisionEncoderDecoderModel.from_pretrained('naver-clova-ix/donut-base-finetuned-cord-v2')
            parser_config = {'processor': processor, 'model': model}

        self._parser_config = parser_config
        self._keep_image = keep_image
        self._parse_text = parse_text
        self._pytesseract_model_kwargs = pytesseract_model_kwargs or {}

    def _load_data(self, file: Path, fs: Optional['fsspec.AbstractFileSystem'] = None) -> List[ImageDocNode]:
        if not isinstance(file, Path): file = Path(file)

        if fs:
            with fs.open(path=file) as f:
                image = PIL.Image.open(f.read())
        else:
            image = PIL.Image.open(file)

        if image.mode != 'RGB': image = image.convert('RGB')

        image_str: Optional[str] = None  # noqa
        if self._keep_image: image_str = img_2_b64(image)  # noqa

        text_str: str = ''
        if self._parse_text:
            assert self._parser_config is not None
            model = self._parser_config['model']
            processor = self._parser_config['processor']

            if processor:
                device = infer_torch_device()
                model.to(device)

                task_prompt = '<s_cord-v2>'
                decoder_input_ids = processor.tokenizer(task_prompt, add_special_tokens=False,
                                                        return_tensors='pt').input_ids
                pixel_values = processor(image, return_tensors='pt').pixel_values

                output = model.generate(pixel_values.to(device), decoder_input_ids=decoder_input_ids.to(device),
                                        max_length=model.decoder.config.max_position_embeddings, early_stopping=True,
                                        pad_token_id=processor.tokenizer.pad_token_id,
                                        eos_token_id=processor.tokenizer.eos_token_id, use_cache=True, num_beams=3,
                                        bad_words_ids=[[processor.tokenizer.unk_token_id]],
                                        return_dict_in_generate=True)

                sequence = processor.batch_decode(output.sequences)[0]
                sequence = sequence.replace(processor.tokenizer.eos_token, '').replace(processor.tokenizer.pad_token, '')
                text_str = re.sub(r'<.*?>', '', sequence, count=1).strip()
            else:
                import pytesseract

                model = cast(pytesseract, self._parser_config['model'])
                text_str = model.image_to_string(image, **self._pytesseract_model_kwargs)

        return [ImageDocNode(text=text_str, image_path=str(file))]

lazyllm.tools.rag.readers.IPYNBReader

Bases: LazyLLMReaderBase

用于读取和解析 Jupyter Notebook (.ipynb) 文件的模块。将 notebook 转换成脚本文本后,按代码单元划分为多个文档节点,或合并为单一文本节点。

Parameters:

  • parser_config (Optional[Dict], default: None ) –

    预留的解析器配置参数,当前未使用,默认为 None。

  • concatenate (bool, default: False ) –

    是否将所有代码单元合并成一个整体文本节点,默认为 False,即分割为多个节点。

  • return_trace (bool, default: True ) –

    是否记录处理过程的 trace,默认为 True。

Source code in lazyllm/tools/rag/readers/ipynbReader.py
class IPYNBReader(LazyLLMReaderBase):
    """用于读取和解析 Jupyter Notebook (.ipynb) 文件的模块。将 notebook 转换成脚本文本后,按代码单元划分为多个文档节点,或合并为单一文本节点。

Args:
    parser_config (Optional[Dict]): 预留的解析器配置参数,当前未使用,默认为 None。
    concatenate (bool): 是否将所有代码单元合并成一个整体文本节点,默认为 False,即分割为多个节点。
    return_trace (bool): 是否记录处理过程的 trace,默认为 True。
"""
    def __init__(self, parser_config: Optional[Dict] = None, concatenate: bool = False, return_trace: bool = True):
        super().__init__(return_trace=return_trace)
        self._parser_config = parser_config
        self._concatenate = concatenate

    def _load_data(self, file: Path, fs: Optional['fsspec.AbstractFileSystem'] = None) -> List[DocNode]:
        if not isinstance(file, Path): file = Path(file)

        try:
            import nbformat
        except ImportError:
            raise ImportError('Please install nbformat: pip install nbformat')

        if fs:
            with fs.open(file, 'r', encoding='utf-8') as f:
                notebook = nbformat.read(f, as_version=4)
        else:
            with open(file, 'r', encoding='utf-8') as f:
                notebook = nbformat.read(f, as_version=4)

        cell_texts = []
        for cell in notebook.cells:
            source = getattr(cell, 'source', '')
            if isinstance(source, list):
                source = ''.join(source)
            if not source or not source.strip():
                continue

            cell_texts.append(source)

        if not cell_texts:
            return []

        if self._concatenate:
            return [DocNode(text=('\n\n' + '-' * 80 + '\n\n').join(cell_texts))]
        return [DocNode(text=text) for text in cell_texts]

lazyllm.tools.rag.readers.MineruPDFReader

Bases: _OcrReaderBase

基于Mineru服务的PDF解析器,通过调用Mineru服务的API来解析PDF文件,支持丰富的文档结构识别。

Parameters:

  • url (str, default: None ) –

    Mineru服务的完整API端点URL。

  • backend (str, default: None ) –

    解析引擎类型。可选值: - 'pipeline': 标准处理流水线 - 'vlm-transformers': 基于Transformers的视觉语言模型 - 'vlm-vllm-async-engine': 基于异步VLLM的视觉语言模型 默认为 'pipeline'。

  • extract_table (bool, default: True ) –

    是否提取表格内容并转换为Markdown格式。默认为 True。

  • extract_formula (bool, default: True ) –

    是否提取公式文本。 - True: 提取为LaTeX等文本格式 - False: 将公式保留为图片 默认为 True。

  • split_doc (bool, default: True ) –

    若为 True(默认),则解析为一个 RichDocNode,可以搭配 RichTransform 解析出带有结构信息的节点; 若为 False,则解析为一个纯文本的 DocNode

  • clean_content (bool, default: True ) –

    是否清理冗余内容(页眉、页脚、页码等)。默认为 True。

  • post_func (Optional[Callable[[List[DocNode]], Any]], default: None ) –

    后处理函数, 接收DocNode列表作为参数,用于自定义结果处理。默认为 None。

  • api_key (str, default: None ) –

    初始化时使用的静态鉴权 token。

  • dynamic_auth (bool, default: False ) –

    是否启用动态鉴权。启用后 token 从 globals.config['dynamic_ocr_auth']['mineru'] 读取。

  • auth_strategy (AuthStrategy, default: None ) –

    自定义鉴权注入策略。默认使用 Bearer token。

Notes

split_doc=True 时返回 RichDocNode,否则返回 DocNode,两种情况都只返回一个节点。 当 split_doc=True 时,强烈建议搭配 RichTransform 使用,可以解析出带有结构信息等 metadata 的节点; 如不使用 RichTransform,则解析出的节点会回退为纯文本节点。 请求级 token:通过 inject_reader_config(ocr_config={'ocr_auth': {'mineru': '...'}}) 写入 globals.config['dynamic_ocr_auth'],由 CredentialMixin 在每次 HTTP 请求时读取。 静态默认 token:globals['config']['mineru_api_key'](仅 dynamic_auth=False 时)。 OCR 服务端缓存由 _load_data(..., use_cache=...) 单独控制(默认 True); 算法端 DocNode 内容缓存由全局 lazyllm.config['reader_use_cache'] 控制。

Examples:

from lazyllm.tools.rag.readers import MineruPDFReader reader = MineruPDFReader("http://0.0.0.0:8888") # Mineru server address nodes = reader("path/to/pdf")

Source code in lazyllm/tools/rag/readers/ocrReader/mineru_pdf_reader.py
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
class MineruPDFReader(_OcrReaderBase):
    """基于Mineru服务的PDF解析器,通过调用Mineru服务的API来解析PDF文件,支持丰富的文档结构识别。

Args:
    url (str): Mineru服务的完整API端点URL。
    backend (str, optional): 解析引擎类型。可选值:
        - 'pipeline': 标准处理流水线
        - 'vlm-transformers': 基于Transformers的视觉语言模型
        - 'vlm-vllm-async-engine': 基于异步VLLM的视觉语言模型
        默认为 'pipeline'。
    extract_table (bool, optional): 是否提取表格内容并转换为Markdown格式。默认为 True。
    extract_formula (bool, optional): 是否提取公式文本。
        - True: 提取为LaTeX等文本格式
        - False: 将公式保留为图片
        默认为 True。
    split_doc (bool, optional): 若为 True(默认),则解析为一个 `RichDocNode`,可以搭配 `RichTransform` 解析出带有结构信息的节点;
        若为 False,则解析为一个纯文本的 `DocNode`。
    clean_content (bool, optional): 是否清理冗余内容(页眉、页脚、页码等)。默认为 True。
    post_func (Optional[Callable[[List[DocNode]], Any]], optional): 后处理函数,
        接收DocNode列表作为参数,用于自定义结果处理。默认为 None。
    api_key (str, optional): 初始化时使用的静态鉴权 token。
    dynamic_auth (bool, optional): 是否启用动态鉴权。启用后 token 从
        globals.config['dynamic_ocr_auth']['mineru'] 读取。
    auth_strategy (AuthStrategy, optional): 自定义鉴权注入策略。默认使用 Bearer token。

Notes:
    当 `split_doc=True` 时返回 `RichDocNode`,否则返回 `DocNode`,两种情况都只返回一个节点。
    当 `split_doc=True` 时,强烈建议搭配 `RichTransform` 使用,可以解析出带有结构信息等 metadata 的节点;
    如不使用 `RichTransform`,则解析出的节点会回退为纯文本节点。
    请求级 token:通过 ``inject_reader_config(ocr_config={'ocr_auth': {'mineru': '...'}})`` 写入
    ``globals.config['dynamic_ocr_auth']``,由 CredentialMixin 在每次 HTTP 请求时读取。
    静态默认 token:``globals['config']['mineru_api_key']``(仅 ``dynamic_auth=False`` 时)。
    OCR **服务端**缓存由 ``_load_data(..., use_cache=...)`` 单独控制(默认 ``True``);
    算法端 ``DocNode`` 内容缓存由全局 ``lazyllm.config['reader_use_cache']`` 控制。


Examples:
    from lazyllm.tools.rag.readers import MineruPDFReader
    reader = MineruPDFReader("http://0.0.0.0:8888")  # Mineru server address
    nodes = reader("path/to/pdf")
    """
    def __init__(self,
                 url: Optional[str] = None,
                 backend: Optional[str] = None,
                 upload_mode: Optional[bool] = None,
                 extract_table: bool = True,
                 extract_formula: bool = True,
                 split_doc: bool = True,
                 clean_content: bool = True,
                 timeout: Optional[int] = None,
                 post_func: Optional[Callable] = None,
                 return_trace: bool = True,
                 dropped_types: Optional[Set[str]] = None,
                 api_key: Optional[str] = None,
                 dynamic_auth: bool = False,
                 auth_strategy: Optional[AuthStrategy] = None,
                 **kwargs):
        if dynamic_auth:
            token = None
        else:
            token = api_key if api_key is not None else lazyllm.config['mineru_api_key']
        super().__init__(url=url or default_online_url('mineru'),
                         dropped_types=dropped_types or {
                             'header', 'footer', 'page_number', 'aside_text', 'page_footnote'},
                         return_trace=return_trace,
                         post_func=post_func,
                         image_cache_dir=kwargs.pop('image_cache_dir', os.path.join(
                             lazyllm.config['home'], 'mineru_cache')),
                         token=token,
                         dynamic_auth=dynamic_auth,
                         auth_strategy=auth_strategy,
                         **kwargs)
        self._backend = backend or lazyllm.config['mineru_backend']
        self._timeout = timeout if (timeout is not None and timeout > 0) else None
        self._variant = resolve_ocr_variant('mineru', self._url)
        self._offline_mode = self._variant == OcrServiceVariant.OFFLINE
        self._upload_mode = upload_mode if upload_mode is not None else self._offline_mode
        if self._variant == OcrServiceVariant.ONLINE and not lazyllm.config['mineru_ssl_verify']:
            LOG.warning(
                '[MineruPDFReader] SSL verification disabled for official MinerU API; '
                'set MINERU_SSL_VERIFY=true after mineru.net certificate is fixed.'
            )

    def _online_request_kwargs(self) -> Dict:
        return {'verify': lazyllm.config['mineru_ssl_verify']}

    def _http_execute(self, method: str, url: str, **kwargs):
        kwargs.setdefault('verify', lazyllm.config['mineru_ssl_verify'])
        return super()._http_execute(method, url, **kwargs)

    @property
    def appendix_hash_key(self):
        dropped = ','.join(sorted(self._dropped_types))
        return f'{self._url}|{self._backend}|{self._upload_mode}|{dropped}|{self._split_doc}'

    @override
    def _load_data(self, file, extra_info: Optional[Dict] = None, use_cache: bool = True,
                   **kwargs) -> List['DocNode']:
        file_path = Path(file)
        merged_info = dict(extra_info) if extra_info else {}
        _t0 = time.time()
        if self._offline_mode:
            response_text = self._fetch_sync(file_path, use_cache)
            task_dir = self._image_cache_dir / str(uuid.uuid4())
            task_dir.mkdir(parents=True, exist_ok=True)
            self._download_offline_images(response_text, cache_dir=task_dir)
        else:
            response_text, task_dir = self._fetch_async(file_path, use_cache)
        _t_fetch = time.time() - _t0
        if task_dir is not None:
            merged_info['image_cache_dir'] = str(task_dir)
        _t1 = time.time()
        nodes = self._build_nodes_from_response(response_text, file_path, merged_info)
        _t_build = time.time() - _t1
        LOG.info(f'[BENCHMARK] file={file_path.name} phase=fetch elapsed={_t_fetch:.3f}s')
        LOG.info(f'[BENCHMARK] file={file_path.name} phase=parse elapsed={_t_build:.3f}s')
        return nodes

    def _fetch_sync(self, file: Path, use_cache: bool) -> str:
        payload = {
            'return_content_list': 'true',
            'use_cache': 'false' if not use_cache else 'true',
            'backend': self._backend,
            'table_enable': 'true',
            'formula_enable': 'true',
        }
        if not self._upload_mode:
            payload['files'] = str(file)
            response = post_sync(self._url, payload=payload, timeout=self._timeout)
        else:
            with open(file, 'rb') as f:
                files = {'upload_files': (os.path.basename(file), f)}
                response = post_sync(self._url, payload=payload, files=files, timeout=self._timeout)
        return response.text

    @staticmethod
    def _parse_service_endpoint(url: str) -> tuple:
        raw = (url or '').strip().rstrip('/')
        if not raw:
            raise ValueError('[MineruPDFReader] url is required for offline image download')
        # Scheme-less host:port (e.g. localhost:8000) defaults to http; use https:// explicitly for TLS.
        if '://' not in raw:
            raw = f'http://{raw}'
        parsed = urlparse(raw)
        scheme = parsed.scheme or 'http'
        host = parsed.hostname
        port = parsed.port
        if not host and parsed.path:
            reparsed = urlparse(f'{scheme}://{parsed.path.lstrip("/")}')
            host = reparsed.hostname
            port = reparsed.port
        if not host:
            raise ValueError(f'[MineruPDFReader] cannot parse host from url: {url!r}')
        return scheme, host, port

    def _image_base_url(self) -> str:
        scheme, host, port = self._parse_service_endpoint(self._url)
        if port:
            return f'{scheme}://{host}:{port}'
        return f'{scheme}://{host}'

    @staticmethod
    def _normalize_image_rel_path(path: str) -> str:
        if not path:
            return ''
        path = str(path).replace('\\', '/').strip()
        match = _IMAGE_REF_PATTERN.search(path)
        if match:
            return match.group(0)
        if path.startswith('images/'):
            return path.split('?')[0]
        name = Path(path).name
        if name and '.' in name:
            return f'images/{name}'
        return ''

    def _offline_content_list(self, raw) -> List[dict]:
        if isinstance(raw, dict):
            return raw.get('result', [{}])[0].get('content_list', []) or []
        if isinstance(raw, list):
            return raw
        return []

    @staticmethod
    def _add_image_paths_from_mapping(mapping: dict, rel_paths: Set[str]) -> None:
        for key in ('img_path', 'image_path', 'image_url'):
            normalized = MineruPDFReader._normalize_image_rel_path(mapping.get(key, ''))
            if normalized:
                rel_paths.add(normalized)

    def _add_image_paths_from_item(self, item: dict, rel_paths: Set[str]) -> None:
        self._add_image_paths_from_mapping(item, rel_paths)
        for line in item.get('lines') or []:
            if isinstance(line, dict):
                self._add_image_paths_from_mapping(line, rel_paths)

    def _collect_offline_image_paths(self, response_text: str) -> Set[str]:
        rel_paths = set(_IMAGE_REF_PATTERN.findall(response_text))
        try:
            raw = json.loads(response_text)
        except json.JSONDecodeError:
            return rel_paths
        for item in self._offline_content_list(raw):
            if isinstance(item, dict):
                self._add_image_paths_from_item(item, rel_paths)
        return rel_paths

    @staticmethod
    def _resolve_path_under_base(base_dir: Path, rel_path: str) -> Optional[Path]:
        try:
            base = os.path.realpath(base_dir)
            joined = os.path.realpath(base_dir / rel_path)
        except OSError:
            return None
        if joined == base or joined.startswith(base + os.sep):
            return Path(joined)
        return None

    def _copy_cached_offline_image(self, rel_path: str, save_path: Path) -> bool:
        image_path = self._resolve_path_under_base(self._image_cache_dir, rel_path)
        if image_path is None:
            return False
        try:
            if not image_path.is_file() or image_path.stat().st_size <= 0:
                return False
            save_path.parent.mkdir(parents=True, exist_ok=True)
            shutil.copy2(image_path, save_path)
            return True
        except OSError as exc:
            LOG.warning(
                f'[MineruPDFReader] failed to copy cached image {image_path} -> {save_path}: {exc}'
            )
            return False

    def _download_offline_images(self, response_text: str, cache_dir: Path) -> None:
        if not cache_dir or not response_text:
            return
        rel_paths = self._collect_offline_image_paths(response_text)
        if not rel_paths:
            LOG.warning('[MineruPDFReader] no image paths found in offline OCR response')
            return
        base_url = self._image_base_url().rstrip('/')
        image_tasks = []
        for rel_path in sorted(rel_paths):
            save_path = self._resolve_path_under_base(cache_dir, rel_path)
            if save_path is None:
                LOG.warning(
                    f'[MineruPDFReader] path traversal detected in image path: {rel_path}'
                )
                continue
            try:
                if save_path.is_file() and save_path.stat().st_size > 0:
                    continue
            except OSError:
                pass
            if self._copy_cached_offline_image(rel_path, save_path):
                continue
            save_path.parent.mkdir(parents=True, exist_ok=True)
            image_tasks.append((f'{base_url}/{rel_path}', save_path))
        if not image_tasks:
            return
        LOG.info(
            f'[MineruPDFReader] downloading {len(image_tasks)} offline images '
            f'to {cache_dir} from {base_url}'
        )
        self._download_images(image_tasks)

    @staticmethod
    def _download_images(image_tasks: List[tuple]) -> None:
        def _download_one(task: tuple) -> None:
            img_url, save_path = task
            try:
                resp = requests.get(img_url, timeout=120)
                resp.raise_for_status()
                save_path.write_bytes(resp.content)
            except Exception as exc:
                LOG.warning(
                    f'[MineruPDFReader] failed to download image {img_url} -> {save_path}: {exc}'
                )

        with ThreadPoolExecutor(max_workers=8) as executor:
            list(executor.map(_download_one, image_tasks))

    def _fetch_async(self, file, use_cache: bool = True):
        file_str = str(file)
        splits = self._split_large_pdf(file_str)
        task_dir = self._image_cache_dir / str(uuid.uuid4())

        if len(splits) == 1:
            return retry_transient(
                self._fetch_async_by_upload,
                log_prefix=f'[MineruPDFReader] {os.path.basename(file_str)} ')(
                    splits[0][0], task_dir=task_dir)

        results = {}
        with ThreadPoolExecutor(max_workers=min(len(splits), 5)) as executor:
            futures = {
                executor.submit(
                    retry_transient(
                        self._fetch_async_by_upload,
                        log_prefix=f'[MineruPDFReader] {os.path.basename(sub_path)} '),
                    sub_path, task_dir=task_dir,
                ): start_page
                for sub_path, start_page in splits
            }
            for future in as_completed(futures):
                start_page = futures[future]
                results[start_page] = future.result()

        return self._merge_split_results(results)

    def _merge_split_results(self, results: dict):
        sorted_pages = sorted(results.keys())
        all_content = []
        first_task_dir = None

        for start_page in sorted_pages:
            json_str, task_dir = results[start_page]
            if first_task_dir is None:
                first_task_dir = task_dir
            content = json.loads(json_str)
            items = content

            for item in items:
                if 'page_idx' in item:
                    item['page_idx'] += start_page
                all_content.append(item)

        merged_json = json.dumps(all_content)

        return merged_json, first_task_dir

    def _fetch_async_by_upload(self, file_path: str, task_dir: Optional['Path'] = None):
        """Upload a local file via batch presigned URL and fetch result."""
        fname = os.path.basename(file_path)

        # Step 1: Request presigned upload URL
        payload = {
            'files': [{'name': fname}],
            'model_version': 'vlm',
        }
        resp = self._request(
            'POST',
            'https://mineru.net/api/v4/file-urls/batch',
            json=payload,
            headers={'Content-Type': 'application/json'},
            timeout=self._timeout,
        )
        data = resp.json()
        batch_id = data['data']['batch_id']
        file_url = data['data']['file_urls'][0]
        auth_key = self._auth_key_after_successful_request()

        # Step 2: Upload file to OSS
        _t2 = time.time()
        with open(file_path, 'rb') as f:
            upload_resp = requests.put(
                file_url, data=f, timeout=self._timeout or 300, **self._online_request_kwargs())
            upload_resp.raise_for_status()
        _t_upload = time.time() - _t2
        LOG.info(f'[BENCHMARK] file={fname} phase=upload elapsed={_t_upload:.3f}s')

        # Step 3: Poll batch results
        _t3 = time.time()
        status_url = f'https://mineru.net/api/v4/extract-results/batch/{batch_id}'
        for _ in range(120):
            status_resp = self._request_with_pinned_auth(
                'GET',
                status_url,
                auth_key,
                timeout=self._timeout or 30,
            )
            status_data = status_resp.json()
            extract_result = status_data.get('data', {}).get('extract_result', [])
            if extract_result:
                state = extract_result[0].get('state')
                if state == 'done':
                    full_zip_url = extract_result[0].get('full_zip_url')
                    zip_resp = requests.get(
                        full_zip_url, timeout=self._timeout or 120, **self._online_request_kwargs())
                    zip_resp.raise_for_status()
                    _t_wait = time.time() - _t3
                    LOG.info(f'[BENCHMARK] file={fname} phase=wait elapsed={_t_wait:.3f}s')
                    return self._extract_content_from_zip(zip_resp.content, task_dir=task_dir)
                elif state == 'failed':
                    raise RuntimeError(
                        f'[MineruPDFReader] Batch task failed: '
                        f'{extract_result[0].get("err_msg", "Unknown error")}')
            time.sleep(3)

        raise TimeoutError('[MineruPDFReader] Batch polling timed out')

    def _extract_content_from_zip(self, zip_bytes: bytes, task_dir: Optional['Path'] = None):
        if task_dir is None:
            task_dir = self._image_cache_dir / str(uuid.uuid4())
        task_dir.mkdir(parents=True, exist_ok=True)
        with zipfile.ZipFile(io.BytesIO(zip_bytes)) as zf:
            for member in zf.infolist():
                member_path = Path(member.filename)
                if member_path.is_absolute() or '..' in member_path.parts:
                    raise ValueError(f'Path traversal detected in zip: {member.filename}')
            json_members = [m for m in zf.infolist() if m.filename.endswith('_content_list.json')]
            if not json_members:
                raise ValueError('No *_content_list.json found in zip')
            content = json.loads(zf.read(json_members[0]))
            layout = None
            model = None
            for member in zf.infolist():
                name = Path(member.filename).name
                if name == 'layout.json':
                    layout = json.loads(zf.read(member))
                elif name.endswith('_model.json'):
                    model = json.loads(zf.read(member))
            if layout is not None and model is not None:
                content = self._normalize_online_content_bboxes(content, layout, model)
            for member in zf.infolist():
                if not member.filename.endswith('_content_list.json'):
                    zf.extract(member, task_dir)
        return json.dumps(content), task_dir

    @staticmethod
    def _normalize_online_content_bboxes(content_list: List[dict], layout: dict,
                                         model_pages: List) -> List[dict]:
        """Rewrite official content_list bboxes from OCR raster space into PDF points.

        layout.json provides PDF page_size; model.json provides 0-1 normalized bboxes.
        content_list absolute coords ≈ normalized * OCR canvas, so
        pdf_bbox = content_bbox * page_size / canvas.
        """
        page_sizes = {
            int(p['page_idx']): p['page_size']
            for p in (layout.get('pdf_info') or [])
            if 'page_idx' in p and p.get('page_size')
        }
        if not page_sizes or not isinstance(model_pages, list):
            return content_list

        by_page: Dict[int, List[dict]] = {}
        for item in content_list:
            if not isinstance(item, dict) or item.get('bbox') is None:
                continue
            page_idx = item.get('page_idx')
            if page_idx is None:
                continue
            by_page.setdefault(int(page_idx), []).append(item)

        for page_idx, items in by_page.items():
            page_size = page_sizes.get(page_idx)
            if not page_size or page_idx >= len(model_pages):
                continue
            dets = model_pages[page_idx]
            if not isinstance(dets, list) or not dets:
                continue
            try:
                max_nx = max(float(d['bbox'][2]) for d in dets if isinstance(d, dict) and d.get('bbox'))
                max_ny = max(float(d['bbox'][3]) for d in dets if isinstance(d, dict) and d.get('bbox'))
                max_cx = max(float(it['bbox'][2]) for it in items)
                max_cy = max(float(it['bbox'][3]) for it in items)
            except (KeyError, TypeError, ValueError):
                continue
            if max_nx <= 0 or max_ny <= 0:
                continue
            src_size = (max_cx / max_nx, max_cy / max_ny)
            dst_size = (float(page_size[0]), float(page_size[1]))
            for it in items:
                it['bbox'] = normalize_bbox(it['bbox'], src_size, dst_size)
        return content_list

    @override
    def _adapt_json_to_IR(self, raw, file=None) -> List[Block]:
        # Online API (zip extraction) returns a list directly.
        # Local server returns {'result': [{'content_list': [...]}]}.
        if isinstance(raw, dict):
            content_list = raw['result'][0]['content_list']
        else:
            content_list = raw

        blocks: List[Block] = []
        for item in content_list:
            block = self._adapt_one(item)
            if block is not None:
                if self._offline_mode and 'lines' in item:
                    block.lines = self._normalize_content(item['lines'])
                blocks.append(block)
        return blocks

    def _normalize_content(self, content) -> List:
        if isinstance(content, str):
            return [content.encode('utf-8', 'replace').decode('utf-8')]
        elif isinstance(content, list):
            result = []
            for item in content:
                if isinstance(item, str):
                    result.append(item.encode('utf-8', 'replace').decode('utf-8'))
                elif isinstance(item, dict):
                    normalized = dict(item)
                    if 'content' in normalized and isinstance(normalized['content'], str):
                        normalized['content'] = normalized['content'].encode(
                            'utf-8', 'replace').decode('utf-8')
                    result.append(normalized)
                else:
                    result.append(item)
            return result
        raise TypeError(f'Not supported type: {type(content)}.')

    def _adapt_one(self, item: dict) -> Optional[Block]:  # noqa: C901
        ty = item.get('type')
        if ty is None:
            LOG.warning(f'[MineruPDFReader] content item missing type field, skipped: {item}')
            return None
        if ty in self._dropped_types:
            return None

        text_level = item.get('text_level', -1)
        text = item.get('text', '')
        page_idx = item.get('page_idx')
        if page_idx is None:
            LOG.warning(f'[MineruPDFReader] content item missing page_idx field, skipped: {item}')
            return None
        bbox = item.get('bbox')
        if bbox is None:
            LOG.warning(f'[MineruPDFReader] content item missing bbox field, skipped: {item}')
            return None
        page = PageRef(index=page_idx, bbox=BBox.from_list(bbox))

        if ty == 'title':
            return HeadingBlock(page=page, level=text_level, text=text)
        elif ty in ('text', 'ref_text', 'phonetic'):
            return ParagraphBlock(page=page, text=text)
        elif ty == 'image':
            return self._adapt_image(item, page, page_idx)
        elif ty == 'table':
            return self._adapt_table(item, page, page_idx)
        elif ty == 'equation':
            return FormulaBlock(page=page, latex=text, inline=False)
        elif ty == 'code':
            return self._adapt_code(item, page, page_idx)
        elif ty == 'list':
            return self._adapt_list(item, page, page_idx)
        return None

    def _adapt_image(self, item: dict, page: PageRef, page_idx: int) -> Optional[Block]:
        img_path = item.get('img_path')
        if img_path is None:
            LOG.warning(f'[MineruPDFReader] image block on page {page_idx} missing img_path, skipped')
            return None
        return FigureBlock(
            page=page,
            image_path=Path(img_path),
            caption=self._first(item.get('image_caption')),
            footnote=self._first(item.get('image_footnote')),
        )

    def _adapt_table(self, item: dict, page: PageRef, page_idx: int) -> TableBlock:
        table_body = item.get('table_body')
        if table_body is None:
            LOG.warning(f'[MineruPDFReader] table block on page {page_idx} missing table_body, '
                        f'caption={self._first(item.get("table_caption"))}')
        return TableBlock(
            page=page,
            caption=self._first(item.get('table_caption')),
            footnote=self._first(item.get('table_footnote')),
            cells=self._parse_table_html(table_body or ''),
            page_range=(page_idx, page_idx),
        )

    def _adapt_code(self, item: dict, page: PageRef, page_idx: int) -> Optional[Block]:
        code_body = item.get('code_body')
        if code_body is None:
            LOG.warning(f'[MineruPDFReader] code block on page {page_idx} missing code_body, skipped')
            return None
        return CodeBlock(
            page=page, text=code_body,
            language=item.get('guess_lang'),
            caption=self._first(item.get('code_caption')),
        )

    def _adapt_list(self, item: dict, page: PageRef, page_idx: int) -> Optional[Block]:
        list_items = item.get('list_items')
        if list_items is None:
            LOG.warning(f'[MineruPDFReader] list block on page {page_idx} missing list_items, skipped')
            return None
        return ListBlock(page=page, items=list_items, ordered=False)

    @override
    def _build_nodes_from_blocks(self, blocks: List[Block], file,
                                 extra_info: Optional[Dict] = None) -> List[DocNode]:
        docs = []

        global_metadata = dict(extra_info) if extra_info else {}
        # image_cache_dir is injected into extra_info by _load_data for async requests
        if 'image_cache_dir' not in global_metadata:
            global_metadata['image_cache_dir'] = str(self._image_cache_dir)

        file_name = Path(file).name if not isinstance(file, str) else Path(file).name
        file_path = str(file)

        for b in blocks:
            text = b.text_content()
            metadata = {
                'file_name': file_name,
                'file_path': file_path,
                'type': b.ty,
                'page': b.page.index,
                'bbox': b.page.bbox.to_list(),
                'section_path': b.section.anchors,
            }
            b.update_metadata(metadata)
            if b.lines:
                metadata['lines'] = b.lines
            node = DocNode(text=text, metadata=metadata, global_metadata=global_metadata)
            node.excluded_embed_metadata_keys = [k for k in metadata if k not in ('file_name', 'text')]
            node.excluded_llm_metadata_keys = [k for k in metadata if k not in ('file_name', 'text')]
            docs.append(node)

        return docs

lazyllm.tools.rag.readers.MarkdownReader

Bases: LazyLLMReaderBase

用于读取和解析 Markdown 文件的模块。支持去除超链接和图片,按标题和内容将 Markdown 划分成若干文本段落节点。

Parameters:

  • remove_hyperlinks (bool, default: True ) –

    是否移除超链接,默认 True。

  • remove_images (bool, default: True ) –

    是否移除图片标记,默认 True。

  • return_trace (bool, default: True ) –

    是否记录处理过程的 trace,默认为 True。

Source code in lazyllm/tools/rag/readers/markdownReader.py
class MarkdownReader(LazyLLMReaderBase):
    """用于读取和解析 Markdown 文件的模块。支持去除超链接和图片,按标题和内容将 Markdown 划分成若干文本段落节点。

Args:
    remove_hyperlinks (bool): 是否移除超链接,默认 True。
    remove_images (bool): 是否移除图片标记,默认 True。
    return_trace (bool): 是否记录处理过程的 trace,默认为 True。
"""
    def __init__(self, remove_hyperlinks: bool = True, remove_images: bool = True, return_trace: bool = True) -> None:
        super().__init__(return_trace=return_trace)
        self._remove_hyperlinks = remove_hyperlinks
        self._remove_images = remove_images

    def _markdown_to_tups(self, markdown_text: str) -> List[Tuple[Optional[str], str]]:
        markdown_tups: List[Tuple[Optional[str], str]] = []
        lines = markdown_text.split('\n')

        current_header = None
        current_lines = []
        in_code_block = False

        for line in lines:
            if line.startswith('```'): in_code_block = not in_code_block

            header_match = re.match(r'^#+\s', line)
            if not in_code_block and header_match:
                if current_header is not None or len(current_lines) > 0:
                    markdown_tups.append((current_header, '\n'.join(current_lines)))
                current_header = line
                current_lines.clear()
            else:
                current_lines.append(line)

        markdown_tups.append((current_header, '\n'.join(current_lines)))
        return [(key.strip() if key is not None else None, re.sub(r'<.*?>', '', value))
                for key, value in markdown_tups]

    def remove_images(self, content: str) -> str:
        """移除内容中形如 ![[...]] 的自定义图片标签。

Args:
    content (str): 输入的 markdown 内容。

**Returns:**

- str: 移除图片标签后的内容。
"""
        pattern = r'!{1}\[\[(.*)\]\]'
        return re.sub(pattern, '', content)

    def remove_hyperlinks(self, content: str) -> str:
        """移除 Markdown 超链接,将 [文本](链接) 转换为纯文本。

Args:
    content (str): 输入的 markdown 内容。

**Returns:**

- str: 移除超链接后的内容,仅保留链接文本。
"""
        pattern = r'\[(.*)\]\((.*)\)'
        return re.sub(pattern, r'\1', content)

    def _parse_tups(self, filepath: Path, errors: str = 'ignore',
                    fs: Optional['fsspec.AbstractFileSystem'] = None) -> List[Tuple[Optional[str], str]]:
        fs = fs or fsspec.implementations.local.LocalFileSystem()

        with fs.open(filepath, encoding='utf-8') as f:
            content = f.read().decode(encoding='utf-8')

        if self._remove_hyperlinks: content = self.remove_hyperlinks(content)
        if self._remove_images: content = self.remove_images(content)
        return self._markdown_to_tups(content)

    def _load_data(self, file: Path, fs: Optional['fsspec.AbstractFileSystem'] = None) -> List[DocNode]:
        if not isinstance(file, Path): file = Path(file)

        tups = self._parse_tups(file, fs=fs)
        results = [DocNode(
            content=[value if header is None else f'\n\n{header}\n{value}' for header, value in tups])]
        return results

移除 Markdown 超链接,将 文本 转换为纯文本。

Parameters:

  • content (str) –

    输入的 markdown 内容。

Returns:

  • str: 移除超链接后的内容,仅保留链接文本。
Source code in lazyllm/tools/rag/readers/markdownReader.py
    def remove_hyperlinks(self, content: str) -> str:
        """移除 Markdown 超链接,将 [文本](链接) 转换为纯文本。

Args:
    content (str): 输入的 markdown 内容。

**Returns:**

- str: 移除超链接后的内容,仅保留链接文本。
"""
        pattern = r'\[(.*)\]\((.*)\)'
        return re.sub(pattern, r'\1', content)

remove_images(content)

移除内容中形如 ![[...]] 的自定义图片标签。

Parameters:

  • content (str) –

    输入的 markdown 内容。

Returns:

  • str: 移除图片标签后的内容。
Source code in lazyllm/tools/rag/readers/markdownReader.py
    def remove_images(self, content: str) -> str:
        """移除内容中形如 ![[...]] 的自定义图片标签。

Args:
    content (str): 输入的 markdown 内容。

**Returns:**

- str: 移除图片标签后的内容。
"""
        pattern = r'!{1}\[\[(.*)\]\]'
        return re.sub(pattern, '', content)

lazyllm.tools.rag.readers.MboxReader

Bases: LazyLLMReaderBase

用于解析 Mbox 邮件存档文件的模块。读取邮件内容并格式化为文本,支持限制最大邮件数和自定义消息格式。

Parameters:

  • max_count (int, default: 0 ) –

    最大读取的邮件数量,默认 0 表示读取全部邮件。

  • message_format (str, default: DEFAULT_MESSAGE_FORMAT ) –

    邮件文本格式模板,支持使用 {_date}{_from}{_to}{_subject}{_content} 占位符。

  • return_trace (bool, default: True ) –

    是否记录处理过程的 trace,默认为 True。

Source code in lazyllm/tools/rag/readers/mboxreader.py
class MboxReader(LazyLLMReaderBase):
    """用于解析 Mbox 邮件存档文件的模块。读取邮件内容并格式化为文本,支持限制最大邮件数和自定义消息格式。

Args:
    max_count (int): 最大读取的邮件数量,默认 0 表示读取全部邮件。
    message_format (str): 邮件文本格式模板,支持使用 ``{_date}``、``{_from}``、``{_to}``、``{_subject}`` 和 ``{_content}`` 占位符。
    return_trace (bool): 是否记录处理过程的 trace,默认为 True。
"""
    DEFAULT_MESSAGE_FORMAT: str = (
        'Date: {_date}\n'
        'From: {_from}\n'
        'To: {_to}\n'
        'Subject: {_subject}\n'
        'Content: {_content}'
    )

    def __init__(self, max_count: int = 0, message_format: str = DEFAULT_MESSAGE_FORMAT,
                 return_trace: bool = True) -> None:
        try:
            from bs4 import BeautifulSoup  # noqa
        except ImportError:
            raise ImportError('`BeautifulSoup` package not found: `pip install beautifulsoup4`')

        super().__init__(return_trace=return_trace)
        self._max_count = max_count
        self._message_format = message_format

    def _load_data(self, file: Path, fs: Optional['fsspec.AbstractFileSystem'] = None) -> List[DocNode]:
        import mailbox
        from email.parser import BytesParser
        from email.policy import default
        from bs4 import BeautifulSoup

        if fs:
            LOG.warning('fs was specified but MboxReader doesn\'t support loading from '
                        'fsspec filesystems. Will load from local filesystem instead.')

        i = 0
        results: List[str] = []
        bytes_parser = BytesParser(policy=default).parse
        mbox = mailbox.mbox(file, factory=bytes_parser)

        for _, _msg in enumerate(mbox):
            try:
                msg: mailbox.mboxMessage = _msg
                if msg.is_multipart():
                    for part in msg.walk():
                        ctype = part.get_content_type()
                        cdispo = str(part.get('Content-Disposition'))
                        if ctype == 'text/plain' and 'attachment' not in cdispo:
                            content = part.get_payload(decode=True)
                            break
                else:
                    content = msg.get_payload(decode=True)

                soup = BeautifulSoup(content)
                stripped_content = ' '.join(soup.get_text().split())
                msg_string = self._message_format.format(_date=msg['date'], _from=msg['from'], _to=msg['to'],
                                                         _subject=msg['subject'], _content=stripped_content)
                results.append(msg_string)
            except Exception as e:
                LOG.warning(f'Failed to parse message:\n{_msg}\n with exception {e}')

            i += 1
            if self._max_count > 0 and i >= self._max_count: break
        return [DocNode(text=result) for result in results]

lazyllm.tools.rag.default_index.DefaultIndex

Bases: IndexBase

默认的索引实现,负责通过 embedding 和文本相似度在底层存储中查询、更新和删除文档节点。支持多种相似度度量方式,并在必要时对查询和节点进行 embedding 计算与更新。

Parameters:

  • embed (Dict[str, Callable]) –

    用于生成查询和节点 embedding 的字典,key 是 embedding 名称,value 是接收字符串返回向量的函数。

  • store (StoreBase) –

    底层存储,用于持久化和检索 DocNode 节点。

  • **kwargs

    预留扩展参数。

Returns:

  • DefaultIndex: 默认索引实例。
Source code in lazyllm/tools/rag/default_index.py
class DefaultIndex(IndexBase):
    """默认的索引实现,负责通过 embedding 和文本相似度在底层存储中查询、更新和删除文档节点。支持多种相似度度量方式,并在必要时对查询和节点进行 embedding 计算与更新。

Args:
    embed (Dict[str, Callable]): 用于生成查询和节点 embedding 的字典,key 是 embedding 名称,value 是接收字符串返回向量的函数。
    store (StoreBase): 底层存储,用于持久化和检索 DocNode 节点。
    **kwargs: 预留扩展参数。

**Returns:**

- DefaultIndex: 默认索引实例。
"""
    def __init__(self, embed: Dict[str, Callable], store, **kwargs):
        self.embed = embed
        self.store = store

    @override
    def update(self, nodes: List[DocNode]) -> None:
        """根据提供的节点列表更新索引中的内容。具体行为由子类或外部实现填充(此处为空实现,需在实际使用中覆盖/扩展)。

Args:
    nodes (List[DocNode]): 需要更新(新增或替换)的文档节点列表。
"""
        pass

    @override
    def remove(self, uids: List[str], group_name: Optional[str] = None) -> None:
        """从索引中删除指定 UID 的节点,可选指定分组名称以限定作用域。当前为空实现,使用时需要补全逻辑。

Args:
    uids (List[str]): 要删除的节点唯一标识列表。
    group_name (Optional[str]): 可选的分组名称,用于限定删除范围。
"""
        pass

    @override
    def query(
        self,
        query: str,
        group_name: str,
        similarity_name: str,
        similarity_cut_off: Union[float, Dict[str, float]],
        topk: int,
        embed_keys: Optional[List[str]] = None,
        filters: Optional[Dict[str, List]] = None,
        **kwargs,
    ) -> List[DocNode]:
        """执行一次查询,支持 embedding 和文本两种模式,依据相似度函数过滤并返回符合条件的 DocNode 结果。

Args:
    query (str): 原始查询文本。
    group_name (str): 要检索的节点组名称。
    similarity_name (str): 使用的相似度度量名称,必须在 registered_similarities 中注册。
    similarity_cut_off (Union[float, Dict[str, float]]): 相似度阈值或每个 embedding 对应的阈值字典,用于过滤结果。
    topk (int): 每个相似度渠道最多保留的候选数量。
    embed_keys (Optional[List[str]]): 指定用于 embedding 的 key 列表,若为空则使用所有可用 embedding。
    filters (Optional[Dict[str, List]]): 额外的节点过滤器,应用在计算相似度前。
    **kwargs: 传递给相似度函数的额外参数。

**Returns:**

- list: List[DocNode]: 经过相似度计算与阈值过滤后去重的文档节点列表。
"""
        if similarity_name not in registered_similarities:
            raise ValueError(
                f'{similarity_name} not registered, please check your input. '
                f'Available options now: {registered_similarities.keys()}'
            )
        similarity_func, mode, descend = registered_similarities[similarity_name]

        nodes = self.store.get_nodes(group=group_name)
        if filters:
            nodes = generic_process_filters(nodes, filters)

        if mode == 'embedding':
            assert self.embed, 'Chosen similarity needs embed model.'
            assert len(query) > 0, 'Query should not be empty.'
            if not embed_keys:
                embed_keys = list(self.embed.keys())
            query_embedding = {k: self.embed[k](query) for k in embed_keys}
            self._check_supported(similarity_name, query_embedding)
            modified_nodes = parallel_do_embedding(self.embed, embed_keys, nodes)
            self.store.update_nodes(modified_nodes)
            similarities = similarity_func(query_embedding, nodes, topk=topk, **kwargs)
        elif mode == 'text':
            similarities = similarity_func(query, nodes, topk=topk, **kwargs)
        else:
            raise NotImplementedError(f'Mode {mode} is not supported.')

        if not isinstance(similarities, dict):
            results = self._filter_nodes_by_score(similarities, topk, similarity_cut_off, descend)
        else:
            results = []
            for key in (embed_keys or similarities.keys()):
                sims = similarities[key]
                sim_cut_off = similarity_cut_off if isinstance(similarity_cut_off, float) else similarity_cut_off[key]
                results.extend(self._filter_nodes_by_score(sims, topk, sim_cut_off, descend))
        results = list(set(results))
        LOG.debug(f'Retrieving query `{query}` and get results: {results}')
        return results

    def _filter_nodes_by_score(self, similarities: List[Tuple[DocNode, float]], topk: int,
                               similarity_cut_off: float, descend) -> List[DocNode]:
        similarities.sort(key=lambda x: x[1], reverse=descend)
        if topk is not None:
            similarities = similarities[:topk]

        return [node.with_sim_score(score) for node, score in similarities if score > similarity_cut_off]

    def _check_supported(self, similarity_name: str, query_embedding: Dict[str, Any]) -> None:
        if similarity_name.lower() == 'cosine':
            for k, e in query_embedding.items():
                if is_sparse(e):
                    raise NotImplementedError(f'embed `{k}`, which is sparse, is not supported.')

query(query, group_name, similarity_name, similarity_cut_off, topk, embed_keys=None, filters=None, **kwargs)

执行一次查询,支持 embedding 和文本两种模式,依据相似度函数过滤并返回符合条件的 DocNode 结果。

Parameters:

  • query (str) –

    原始查询文本。

  • group_name (str) –

    要检索的节点组名称。

  • similarity_name (str) –

    使用的相似度度量名称,必须在 registered_similarities 中注册。

  • similarity_cut_off (Union[float, Dict[str, float]]) –

    相似度阈值或每个 embedding 对应的阈值字典,用于过滤结果。

  • topk (int) –

    每个相似度渠道最多保留的候选数量。

  • embed_keys (Optional[List[str]], default: None ) –

    指定用于 embedding 的 key 列表,若为空则使用所有可用 embedding。

  • filters (Optional[Dict[str, List]], default: None ) –

    额外的节点过滤器,应用在计算相似度前。

  • **kwargs

    传递给相似度函数的额外参数。

Returns:

  • list: List[DocNode]: 经过相似度计算与阈值过滤后去重的文档节点列表。
Source code in lazyllm/tools/rag/default_index.py
    @override
    def query(
        self,
        query: str,
        group_name: str,
        similarity_name: str,
        similarity_cut_off: Union[float, Dict[str, float]],
        topk: int,
        embed_keys: Optional[List[str]] = None,
        filters: Optional[Dict[str, List]] = None,
        **kwargs,
    ) -> List[DocNode]:
        """执行一次查询,支持 embedding 和文本两种模式,依据相似度函数过滤并返回符合条件的 DocNode 结果。

Args:
    query (str): 原始查询文本。
    group_name (str): 要检索的节点组名称。
    similarity_name (str): 使用的相似度度量名称,必须在 registered_similarities 中注册。
    similarity_cut_off (Union[float, Dict[str, float]]): 相似度阈值或每个 embedding 对应的阈值字典,用于过滤结果。
    topk (int): 每个相似度渠道最多保留的候选数量。
    embed_keys (Optional[List[str]]): 指定用于 embedding 的 key 列表,若为空则使用所有可用 embedding。
    filters (Optional[Dict[str, List]]): 额外的节点过滤器,应用在计算相似度前。
    **kwargs: 传递给相似度函数的额外参数。

**Returns:**

- list: List[DocNode]: 经过相似度计算与阈值过滤后去重的文档节点列表。
"""
        if similarity_name not in registered_similarities:
            raise ValueError(
                f'{similarity_name} not registered, please check your input. '
                f'Available options now: {registered_similarities.keys()}'
            )
        similarity_func, mode, descend = registered_similarities[similarity_name]

        nodes = self.store.get_nodes(group=group_name)
        if filters:
            nodes = generic_process_filters(nodes, filters)

        if mode == 'embedding':
            assert self.embed, 'Chosen similarity needs embed model.'
            assert len(query) > 0, 'Query should not be empty.'
            if not embed_keys:
                embed_keys = list(self.embed.keys())
            query_embedding = {k: self.embed[k](query) for k in embed_keys}
            self._check_supported(similarity_name, query_embedding)
            modified_nodes = parallel_do_embedding(self.embed, embed_keys, nodes)
            self.store.update_nodes(modified_nodes)
            similarities = similarity_func(query_embedding, nodes, topk=topk, **kwargs)
        elif mode == 'text':
            similarities = similarity_func(query, nodes, topk=topk, **kwargs)
        else:
            raise NotImplementedError(f'Mode {mode} is not supported.')

        if not isinstance(similarities, dict):
            results = self._filter_nodes_by_score(similarities, topk, similarity_cut_off, descend)
        else:
            results = []
            for key in (embed_keys or similarities.keys()):
                sims = similarities[key]
                sim_cut_off = similarity_cut_off if isinstance(similarity_cut_off, float) else similarity_cut_off[key]
                results.extend(self._filter_nodes_by_score(sims, topk, sim_cut_off, descend))
        results = list(set(results))
        LOG.debug(f'Retrieving query `{query}` and get results: {results}')
        return results

remove(uids, group_name=None)

从索引中删除指定 UID 的节点,可选指定分组名称以限定作用域。当前为空实现,使用时需要补全逻辑。

Parameters:

  • uids (List[str]) –

    要删除的节点唯一标识列表。

  • group_name (Optional[str], default: None ) –

    可选的分组名称,用于限定删除范围。

Source code in lazyllm/tools/rag/default_index.py
    @override
    def remove(self, uids: List[str], group_name: Optional[str] = None) -> None:
        """从索引中删除指定 UID 的节点,可选指定分组名称以限定作用域。当前为空实现,使用时需要补全逻辑。

Args:
    uids (List[str]): 要删除的节点唯一标识列表。
    group_name (Optional[str]): 可选的分组名称,用于限定删除范围。
"""
        pass

update(nodes)

根据提供的节点列表更新索引中的内容。具体行为由子类或外部实现填充(此处为空实现,需在实际使用中覆盖/扩展)。

Parameters:

  • nodes (List[DocNode]) –

    需要更新(新增或替换)的文档节点列表。

Source code in lazyllm/tools/rag/default_index.py
    @override
    def update(self, nodes: List[DocNode]) -> None:
        """根据提供的节点列表更新索引中的内容。具体行为由子类或外部实现填充(此处为空实现,需在实际使用中覆盖/扩展)。

Args:
    nodes (List[DocNode]): 需要更新(新增或替换)的文档节点列表。
"""
        pass

lazyllm.tools.Reranker

Bases: ModuleBase, _PostProcess

用于创建节点(文档)后处理和重排序的模块。

Parameters:

  • name (str, default: 'ModuleReranker' ) –

    用于后处理和重排序过程的排序器类型。默认为 'ModuleReranker'。

  • target (str, default: None ) –

    已废弃参数,仅用于提示用户。

  • output_format (Optional[str], default: None ) –

    代表输出格式,默认为None,可选值有 'content' 和 'dict',其中 content 对应输出格式为字符串,dict 对应字典。

  • join (Union[bool, str], default: False ) –

    是否联合输出的 k 个节点,当输出格式为 content 时,如果设置该值为 True,则输出一个长字符串,如果设置为 False 则输出一个字符串列表,其中每个字符串对应每个节点的文本内容。当输出格式是 dict 时,不能联合输出,此时join默认为False,,将输出一个字典,包括'content、'embedding'、'metadata'三个key。

  • kwargs

    传递给重新排序器实例化的其他关键字参数。

详细解释排序器类型

  • Reranker: 实例化一个具有待排序的文档节点node列表和 query的 SentenceTransformerRerank 重排序器。
  • KeywordFilter: 实例化一个具有指定必需和排除关键字的 KeywordNodePostprocessor。它根据这些关键字的存在或缺失来过滤节点。

Examples:

>>> import lazyllm
>>> from lazyllm.tools import Document, Reranker, Retriever, DocNode
>>> m = lazyllm.OnlineEmbeddingModule()
>>> documents = Document(dataset_path='/path/to/user/data', embed=m, manager=False)
>>> retriever = Retriever(documents, group_name='CoarseChunk', similarity='bm25', similarity_cut_off=0.01, topk=6)
>>> reranker = Reranker(DocNode(text=user_data),query="user query")
>>> ppl = lazyllm.ActionModule(retriever, reranker)
>>> ppl.start()
>>> print(ppl("user query"))
Source code in lazyllm/tools/rag/rerank.py
class Reranker(ModuleBase, _PostProcess):
    """用于创建节点(文档)后处理和重排序的模块。

Args:
    name: 用于后处理和重排序过程的排序器类型。默认为 'ModuleReranker'。
    target(str):已废弃参数,仅用于提示用户。
    output_format: 代表输出格式,默认为None,可选值有 'content' 和 'dict',其中 content 对应输出格式为字符串,dict 对应字典。
    join: 是否联合输出的 k 个节点,当输出格式为 content 时,如果设置该值为 True,则输出一个长字符串,如果设置为 False 则输出一个字符串列表,其中每个字符串对应每个节点的文本内容。当输出格式是 dict 时,不能联合输出,此时join默认为False,,将输出一个字典,包括'content、'embedding'、'metadata'三个key。
    kwargs: 传递给重新排序器实例化的其他关键字参数。

详细解释排序器类型

  - Reranker: 实例化一个具有待排序的文档节点node列表和 query的 SentenceTransformerRerank 重排序器。
  - KeywordFilter: 实例化一个具有指定必需和排除关键字的 KeywordNodePostprocessor。它根据这些关键字的存在或缺失来过滤节点。


Examples:

    >>> import lazyllm
    >>> from lazyllm.tools import Document, Reranker, Retriever, DocNode
    >>> m = lazyllm.OnlineEmbeddingModule()
    >>> documents = Document(dataset_path='/path/to/user/data', embed=m, manager=False)
    >>> retriever = Retriever(documents, group_name='CoarseChunk', similarity='bm25', similarity_cut_off=0.01, topk=6)
    >>> reranker = Reranker(DocNode(text=user_data),query="user query")
    >>> ppl = lazyllm.ActionModule(retriever, reranker)
    >>> ppl.start()
    >>> print(ppl("user query"))
    """
    registered_reranker = dict()

    def __new__(cls, name: str = 'ModuleReranker', *args, **kwargs):
        assert name in cls.registered_reranker, f'Reranker: {name} is not registered, please register first.'
        item = cls.registered_reranker[name]
        if isinstance(item, type) and issubclass(item, Reranker):
            return super(Reranker, cls).__new__(item)
        else:
            return super(Reranker, cls).__new__(cls)

    def __init__(self, name: str = 'ModuleReranker', target: Optional[str] = None,
                 output_format: Optional[str] = None, join: Union[bool, str] = False, **kwargs) -> None:
        super().__init__()
        self._name = name
        self._kwargs = kwargs
        lazyllm.deprecated(bool(target), '`target` parameter of reranker')
        _PostProcess.__init__(self, output_format, join)

    def forward(self, nodes: List[DocNode], query: str = '', topk: Optional[int] = None) -> List[DocNode]:
        kwargs = dict(self._kwargs)
        if topk is not None:
            kwargs['topk'] = topk
        results = self.registered_reranker[self._name](nodes, query=query, **kwargs)
        LOG.debug(f'Rerank use `{self._name}` and get nodes: {results}')
        return self._post_process(results)

    @classmethod
    def register_reranker(
        cls: 'Reranker', func: Optional[Callable] = None, batch: bool = False
    ):
        """是一个类装饰器工厂方法,它的核心作用是为 Reranker 类提供灵活的排序算法注册机制

Args:
    func (Optional[Callable]):  要注册的排序函数或排序器类。当使用装饰器语法(@)时可省略。
    batch (bool):是否批量处理节点。默认为False,表示逐节点处理。


Examples:

    @Reranker.register_reranker
    def my_reranker(node: DocNode, **kwargs):
        return node.score * 0.8  # 自定义分数计算
    """
        def decorator(f):
            if isinstance(f, type):
                cls.registered_reranker[f.__name__] = f
                return f
            else:
                def wrapper(nodes, **kwargs):
                    if batch:
                        return f(nodes, **kwargs)
                    else:
                        results = [f(node, **kwargs) for node in nodes]
                        return [result for result in results if result]

                cls.registered_reranker[f.__name__] = wrapper
                return wrapper

        return decorator(func) if func else decorator

register_reranker(func=None, batch=False) classmethod

是一个类装饰器工厂方法,它的核心作用是为 Reranker 类提供灵活的排序算法注册机制

Parameters:

  • func (Optional[Callable], default: None ) –

    要注册的排序函数或排序器类。当使用装饰器语法(@)时可省略。

  • batch (bool, default: False ) –

    是否批量处理节点。默认为False,表示逐节点处理。

Examples:

@Reranker.register_reranker
def my_reranker(node: DocNode, **kwargs):
    return node.score * 0.8  # 自定义分数计算
Source code in lazyllm/tools/rag/rerank.py
    @classmethod
    def register_reranker(
        cls: 'Reranker', func: Optional[Callable] = None, batch: bool = False
    ):
        """是一个类装饰器工厂方法,它的核心作用是为 Reranker 类提供灵活的排序算法注册机制

Args:
    func (Optional[Callable]):  要注册的排序函数或排序器类。当使用装饰器语法(@)时可省略。
    batch (bool):是否批量处理节点。默认为False,表示逐节点处理。


Examples:

    @Reranker.register_reranker
    def my_reranker(node: DocNode, **kwargs):
        return node.score * 0.8  # 自定义分数计算
    """
        def decorator(f):
            if isinstance(f, type):
                cls.registered_reranker[f.__name__] = f
                return f
            else:
                def wrapper(nodes, **kwargs):
                    if batch:
                        return f(nodes, **kwargs)
                    else:
                        results = [f(node, **kwargs) for node in nodes]
                        return [result for result in results if result]

                cls.registered_reranker[f.__name__] = wrapper
                return wrapper

        return decorator(func) if func else decorator

lazyllm.tools.Retriever

Bases: _RetrieverBase, _PostProcess

创建一个用于文档查询和检索的检索模块。此构造函数初始化一个检索模块,该模块根据指定的相似度度量配置文档检索过程。

Parameters:

  • doc (object) –

    文档模块实例。该文档模块可以是单个实例,也可以是一个实例的列表。如果是单个实例,表示对单个Document进行检索,如果是实例的列表,则表示对多个Document进行检索。

  • group_name (str) –

    在哪个 node group 上进行检索。

  • similarity (Optional[str], default: None ) –

    用于设置文档检索的相似度函数。默认为 'dummy'。候选集包括 ["bm25", "bm25_chinese", "cosine"]。

  • similarity_cut_off (Union[float, Dict[str, float]], default: float('-inf') ) –

    当相似度低于指定值时丢弃该文档。在多 embedding 场景下,如果需要对不同的 embedding 指定不同的值,则需要使用字典的方式指定,key 表示指定的是哪个 embedding,value 表示相应的阈值。如果所有的 embedding 使用同一个阈值,则只指定一个数值即可。

  • index (str, default: 'default' ) –

    用于文档检索的索引类型。目前仅支持 'default'。

  • topk (int, default: 6 ) –

    表示取相似度最高的多少篇文档。

  • embed_keys (Optional[List[str]], default: None ) –

    表示通过哪些 embedding 做检索,不指定表示用全部 embedding 进行检索。

  • output_format (Optional[str], default: None ) –

    代表输出格式,默认为None,可选值有 'content' 和 'dict',其中 content 对应输出格式为字符串,dict 对应字典。

  • join (Union[bool, str], default: False ) –

    是否联合输出的 k 个节点,当输出格式为 content 时,如果设置该值为 True,则输出一个长字符串,如果设置为 False 则输出一个字符串列表,其中每个字符串对应每个节点的文本内容。当输出格式是 dict 时,不能联合输出,此时join默认为False,,将输出一个字典,包括'content、'embedding'、'metadata'三个key。

其中 group_name 有三个内置的切分策略,都是使用 SentenceSplitter 做切分,区别在于块大小不同:

  • CoarseChunk: 块大小为 1024,重合长度为 100
  • MediumChunk: 块大小为 256,重合长度为 25
  • FineChunk: 块大小为 128,重合长度为 12

此外,LazyLLM提供了内置的Image节点组存储了所有图像节点,支持图像嵌入和检索。

Examples:

>>> import lazyllm
>>> from lazyllm.tools import Retriever, Document, SentenceSplitter
>>> m = lazyllm.OnlineEmbeddingModule()
>>> documents = Document(dataset_path='/path/to/user/data', embed=m, manager=False)
>>> rm = Retriever(documents, group_name='CoarseChunk', similarity='bm25', similarity_cut_off=0.01, topk=6)
>>> rm.start()
>>> print(rm("user query"))
>>> m1 = lazyllm.TrainableModule('bge-large-zh-v1.5').start()
>>> document1 = Document(dataset_path='/path/to/user/data', embed={'online':m , 'local': m1}, manager=False)
>>> document1.create_node_group(name='sentences', transform=SentenceSplitter, chunk_size=1024, chunk_overlap=100)
>>> retriever = Retriever(document1, group_name='sentences', similarity='cosine', similarity_cut_off=0.4, embed_keys=['local'], topk=3)
>>> print(retriever("user query"))
>>> document2 = Document(dataset_path='/path/to/user/data', embed={'online':m , 'local': m1}, manager=False)
>>> document2.create_node_group(name='sentences', transform=SentenceSplitter, chunk_size=512, chunk_overlap=50)
>>> retriever2 = Retriever([document1, document2], group_name='sentences', similarity='cosine', similarity_cut_off=0.4, embed_keys=['local'], topk=3)
>>> print(retriever2("user query"))
>>>
>>> filters = {
>>>     "author": ["A", "B", "C"],
>>>     "public_year": [2002, 2003, 2004],
>>> }
>>> document3 = Document(dataset_path='/path/to/user/data', embed={'online':m , 'local': m1}, manager=False)
>>> document3.create_node_group(name='sentences', transform=SentenceSplitter, chunk_size=512, chunk_overlap=50)
>>> retriever3 = Retriever([document1, document3], group_name='sentences', similarity='cosine', similarity_cut_off=0.4, embed_keys=['local'], topk=3)
>>> print(retriever3(query="user query", filters=filters))
>>> document4 = Document(dataset_path='/path/to/user/data', embed=lazyllm.TrainableModule('siglip'))
>>> retriever4 = Retriever(document4, group_name='Image', similarity='cosine')
>>> nodes = retriever4("user query")
>>> print([node.get_content() for node in nodes])
>>> document5 = Document(dataset_path='/path/to/user/data', embed=m, manager=False)
>>> rm = Retriever(document5, group_name='CoarseChunk', similarity='bm25_chinese', similarity_cut_off=0.01, topk=3, output_format='content')
>>> rm.start()
>>> print(rm("user query"))
>>> document6 = Document(dataset_path='/path/to/user/data', embed=m, manager=False)
>>> rm = Retriever(document6, group_name='CoarseChunk', similarity='bm25_chinese', similarity_cut_off=0.01, topk=3, output_format='content', join=True)
>>> rm.start()
>>> print(rm("user query"))
>>> document7 = Document(dataset_path='/path/to/user/data', embed=m, manager=False)
>>> rm = Retriever(document7, group_name='CoarseChunk', similarity='bm25_chinese', similarity_cut_off=0.01, topk=3, output_format='dict')
>>> rm.start()
>>> print(rm("user query"))
Source code in lazyllm/tools/rag/retriever.py
class Retriever(_RetrieverBase, _PostProcess):
    """
创建一个用于文档查询和检索的检索模块。此构造函数初始化一个检索模块,该模块根据指定的相似度度量配置文档检索过程。

Args:
    doc: 文档模块实例。该文档模块可以是单个实例,也可以是一个实例的列表。如果是单个实例,表示对单个Document进行检索,如果是实例的列表,则表示对多个Document进行检索。
    group_name: 在哪个 node group 上进行检索。
    similarity: 用于设置文档检索的相似度函数。默认为 'dummy'。候选集包括 ["bm25", "bm25_chinese", "cosine"]。
    similarity_cut_off: 当相似度低于指定值时丢弃该文档。在多 embedding 场景下,如果需要对不同的 embedding 指定不同的值,则需要使用字典的方式指定,key 表示指定的是哪个 embedding,value 表示相应的阈值。如果所有的 embedding 使用同一个阈值,则只指定一个数值即可。
    index: 用于文档检索的索引类型。目前仅支持 'default'。
    topk: 表示取相似度最高的多少篇文档。
    embed_keys: 表示通过哪些 embedding 做检索,不指定表示用全部 embedding 进行检索。
    target:目标组名,将结果转换到目标组。
    output_format: 代表输出格式,默认为None,可选值有 'content' 和 'dict',其中 content 对应输出格式为字符串,dict 对应字典。
    join: 是否联合输出的 k 个节点,当输出格式为 content 时,如果设置该值为 True,则输出一个长字符串,如果设置为 False 则输出一个字符串列表,其中每个字符串对应每个节点的文本内容。当输出格式是 dict 时,不能联合输出,此时join默认为False,,将输出一个字典,包括'content、'embedding'、'metadata'三个key。

其中 `group_name` 有三个内置的切分策略,都是使用 `SentenceSplitter` 做切分,区别在于块大小不同:

- CoarseChunk: 块大小为 1024,重合长度为 100
- MediumChunk: 块大小为 256,重合长度为 25
- FineChunk: 块大小为 128,重合长度为 12

此外,LazyLLM提供了内置的`Image`节点组存储了所有图像节点,支持图像嵌入和检索。


Examples:

    >>> import lazyllm
    >>> from lazyllm.tools import Retriever, Document, SentenceSplitter
    >>> m = lazyllm.OnlineEmbeddingModule()
    >>> documents = Document(dataset_path='/path/to/user/data', embed=m, manager=False)
    >>> rm = Retriever(documents, group_name='CoarseChunk', similarity='bm25', similarity_cut_off=0.01, topk=6)
    >>> rm.start()
    >>> print(rm("user query"))
    >>> m1 = lazyllm.TrainableModule('bge-large-zh-v1.5').start()
    >>> document1 = Document(dataset_path='/path/to/user/data', embed={'online':m , 'local': m1}, manager=False)
    >>> document1.create_node_group(name='sentences', transform=SentenceSplitter, chunk_size=1024, chunk_overlap=100)
    >>> retriever = Retriever(document1, group_name='sentences', similarity='cosine', similarity_cut_off=0.4, embed_keys=['local'], topk=3)
    >>> print(retriever("user query"))
    >>> document2 = Document(dataset_path='/path/to/user/data', embed={'online':m , 'local': m1}, manager=False)
    >>> document2.create_node_group(name='sentences', transform=SentenceSplitter, chunk_size=512, chunk_overlap=50)
    >>> retriever2 = Retriever([document1, document2], group_name='sentences', similarity='cosine', similarity_cut_off=0.4, embed_keys=['local'], topk=3)
    >>> print(retriever2("user query"))
    >>>
    >>> filters = {
    >>>     "author": ["A", "B", "C"],
    >>>     "public_year": [2002, 2003, 2004],
    >>> }
    >>> document3 = Document(dataset_path='/path/to/user/data', embed={'online':m , 'local': m1}, manager=False)
    >>> document3.create_node_group(name='sentences', transform=SentenceSplitter, chunk_size=512, chunk_overlap=50)
    >>> retriever3 = Retriever([document1, document3], group_name='sentences', similarity='cosine', similarity_cut_off=0.4, embed_keys=['local'], topk=3)
    >>> print(retriever3(query="user query", filters=filters))
    >>> document4 = Document(dataset_path='/path/to/user/data', embed=lazyllm.TrainableModule('siglip'))
    >>> retriever4 = Retriever(document4, group_name='Image', similarity='cosine')
    >>> nodes = retriever4("user query")
    >>> print([node.get_content() for node in nodes])
    >>> document5 = Document(dataset_path='/path/to/user/data', embed=m, manager=False)
    >>> rm = Retriever(document5, group_name='CoarseChunk', similarity='bm25_chinese', similarity_cut_off=0.01, topk=3, output_format='content')
    >>> rm.start()
    >>> print(rm("user query"))
    >>> document6 = Document(dataset_path='/path/to/user/data', embed=m, manager=False)
    >>> rm = Retriever(document6, group_name='CoarseChunk', similarity='bm25_chinese', similarity_cut_off=0.01, topk=3, output_format='content', join=True)
    >>> rm.start()
    >>> print(rm("user query"))
    >>> document7 = Document(dataset_path='/path/to/user/data', embed=m, manager=False)
    >>> rm = Retriever(document7, group_name='CoarseChunk', similarity='bm25_chinese', similarity_cut_off=0.01, topk=3, output_format='dict')
    >>> rm.start()
    >>> print(rm("user query"))
    """
    def __init__(self, doc: object, group_name: str, similarity: Optional[str] = None,
                 similarity_cut_off: Union[float, Dict[str, float]] = float('-inf'), index: str = 'default',
                 topk: int = 6, embed_keys: Optional[List[str]] = None, target: Optional[str] = None,
                 output_format: Optional[str] = None, join: Union[bool, str] = False,
                 weight: Optional[float] = None, priority: Optional[_RetrieverBase.Priority] = None, **kwargs):
        super().__init__()
        if similarity:
            if similarity not in registered_similarities:
                raise ValueError(
                    f"Unregistered similarity: '{similarity}'. "
                    f'Available options are: {list(registered_similarities.keys())}'
                )
            _, mode, _ = registered_similarities[similarity]
        else:
            similarity = 'cosine'
            mode = 'embedding'  # TODO FIXME XXX should be removed after similarity args refactor
        group_name, target = str(group_name), (str(target) if target else None)

        self._docs: List[Document] = [doc] if isinstance(doc, Document) else doc
        # NOTE: multi docs is deprecated and will be removed in the future
        if len(self._docs) > 1:
            LOG.warning('[Retriever] Multi docs is deprecated and will be removed in the future,'
                        ' please use multiple Retrievers instead.')
        self._group_name = group_name
        if index == 'smart_embedding_index':
            index = 'default'
            LOG.warning('[Retriever] `smart_embedding_index` is deprecated, converted to `default`')
        self._mode = mode
        self._index = index
        self._topk = topk
        self._similarity = similarity  # similarity function str
        self._similarity_kw = kwargs  # kw parameters
        self._similarity_cut_off = similarity_cut_off
        self._embed_keys = embed_keys
        self._per_doc_embed_keys = False
        self._target = target
        self._weight, self._priority = weight, priority
        if weight or priority:
            assert not (weight and priority), f'Cannot provide weight({weight}) and priority({priority}) together!'
            assert not output_format or not join, 'shouldn\'t provide output_format/join when weight or priority is set'

        self._init_submodules_and_embed_keys()
        _PostProcess.__init__(self, output_format, join)

    weight = property(lambda self: self._weight)
    priority = property(lambda self: self._priority)

    @once_wrapper
    def _lazy_init(self):
        docs = []
        per_doc_embed_keys = [] if self._per_doc_embed_keys else None
        for idx, doc in enumerate(self._docs):
            if isinstance(doc, UrlDocument) or self._group_name in doc._impl.node_groups \
                    or self._group_name in DocImpl._builtin_node_groups \
                    or self._group_name in DocImpl._global_node_groups:
                docs.append(doc)
                if self._per_doc_embed_keys:
                    per_doc_embed_keys.append(self._embed_keys[idx])
        if not docs: raise RuntimeError(f'Group {self._group_name} not found in document {self._docs}')
        self._docs = docs
        if self._per_doc_embed_keys:
            self._embed_keys = per_doc_embed_keys

    def _init_submodules_and_embed_keys(self):
        group_name = self._group_name
        embed_keys = self._embed_keys
        self._per_doc_embed_keys = (not embed_keys and self._mode == 'embedding')
        if self._per_doc_embed_keys:
            # NOTE: store per-doc embed keys aligned with self._docs order
            self._embed_keys = []
        for doc in self._docs:
            assert isinstance(doc, (Document, UrlDocument)), 'Only Document or List[Document] are supported'
            if isinstance(doc, UrlDocument):
                if embed_keys:
                    self._validate_remote_vec_retr_params(doc, group_name, embed_keys)
                else:
                    group_name, doc_embed_keys = self._validate_remote_vec_retr_params(doc, group_name, None)
                    if self._per_doc_embed_keys:
                        self._embed_keys.append(doc_embed_keys)
                continue
            self._submodules.append(doc)
            if self._per_doc_embed_keys:
                doc_embed_keys = list(doc._impl.embed.keys())
                self._embed_keys.append(doc_embed_keys)
            else:
                doc_embed_keys = embed_keys
            doc.activate_group(group_name, doc_embed_keys)
            if self._target: doc.activate_group(self._target)

    def __getstate__(self):
        state = {'group_name': self._group_name, 'similarity': self._similarity,
                 'similarity_cut_off': self._similarity_cut_off, 'index': self._index, 'topk': self._topk,
                 'similarity_kw': self._similarity_kw, 'embed_keys': self._embed_keys, 'target': self._target,
                 'output_format': self._output_format, 'join': self._join,
                 'per_doc_embed_keys': self._per_doc_embed_keys}
        docs = []
        for doc in self._docs:
            if isinstance(doc, UrlDocument):
                docs.append({'url': doc._manager._url, 'name': doc._curr_group})
            else:
                assert isinstance(doc._manager._kbs, lazyllm.ServerModule), \
                    'Only UrlDocument and Document with ServerModule are supported'
                docs.append({'url': doc._manager._kbs._url, 'name': doc._curr_group})
        state['docs'] = docs
        return state

    def __setstate__(self, state):
        ModuleBase.__init__(self)
        self._group_name = state['group_name']
        self._similarity = state['similarity']
        self._similarity_cut_off = state['similarity_cut_off']
        self._index = state['index']
        self._topk = state['topk']
        self._similarity_kw = state['similarity_kw']
        self._embed_keys = state['embed_keys']
        self._per_doc_embed_keys = state.get('per_doc_embed_keys', False)
        self._target = state['target']
        self._output_format = state['output_format']
        self._join = state['join']
        self._docs = [Document(url=doc['url'], name=doc['name']) for doc in state['docs']]
        _PostProcess.__init__(self, self._output_format, self._join)

    def _validate_remote_vec_retr_params(self, doc: UrlDocument, group_name, embed_keys: Optional[List[str]] = None):
        active_groups = doc.active_node_groups
        if not active_groups:
            raise RuntimeError(f'No active groups found in document {doc._manager._url}')
        if group_name not in active_groups:
            raise RuntimeError(f'Group {group_name} not found or not activated in document {doc._manager._url}')
        if not embed_keys:
            resolved_embed_keys = list(active_groups[group_name])
            return group_name, resolved_embed_keys
        else:
            for k in embed_keys:
                if k not in active_groups[group_name]:
                    raise RuntimeError(f'Embedding key {k} not found in group {group_name} '
                                       f'from document {doc._manager._url},'
                                       f'available keys: {list(active_groups[group_name])}')
            return group_name, embed_keys

    def forward(
            self, query: str, filters: Optional[Dict[str, Union[str, int, List, Set]]] = None,
            topk: Optional[int] = None, **kwargs
    ) -> Union[List[DocNode], str]:
        self._lazy_init()
        resolved_topk = self._topk if topk is None else topk
        all_nodes: List[DocNode] = []
        if self._per_doc_embed_keys:
            if len(self._embed_keys) != len(self._docs):
                raise RuntimeError('Per-doc embed_keys misaligned with docs after lazy init')
        for idx, doc in enumerate(self._docs):
            embed_keys = self._embed_keys[idx] if self._per_doc_embed_keys else self._embed_keys
            nodes = doc.forward(query=query, group_name=self._group_name, similarity=self._similarity,
                                similarity_cut_off=self._similarity_cut_off, index=self._index,
                                topk=resolved_topk, similarity_kws=self._similarity_kw, embed_keys=embed_keys,
                                filters=filters, **kwargs)
            if nodes and self._target and self._target != nodes[0]._group:
                nodes = doc.find(self._target)(nodes)
            all_nodes.extend(nodes)
        return self._post_process(all_nodes)

lazyllm.tools.rag.retriever.TempDocRetriever

Bases: TempRetriever

临时文档检索器,继承自TempRetriever,用于快速处理临时文件并执行检索任务。

Parameters:

  • embed (Callable, default: None ) –

    嵌入函数。

  • output_format (Optional[str], default: None ) –

    结果输出格式(如json),可选默认为None

  • join (Union[bool, str], default: False ) –

    是否合并多段结果(True或用分隔符如"

")

Examples:

>>> import lazyllm
>>> from lazyllm.tools import TempDocRetriever, Document, SentenceSplitter
>>> retriever = TempDocRetriever(output_format="text", join="
---------------
")
    retriever.create_node_group(transform=lambda text: [s.strip() for s in text.split("。") if s] )
    retriever.add_subretriever(group=Document.MediumChunk, topk=3)
    files = ["/path/to/file.txt"]
    results = retriever.forward(files, "什么是机器学习?")
    print(results)
Source code in lazyllm/tools/rag/retriever.py
class TempDocRetriever(TempRetriever):
    """
临时文档检索器,继承自TempRetriever,用于快速处理临时文件并执行检索任务。

Args:
    embed:嵌入函数。
    output_format:结果输出格式(如json),可选默认为None
    join:是否合并多段结果(True或用分隔符如"
")


Examples:

    >>> import lazyllm
    >>> from lazyllm.tools import TempDocRetriever, Document, SentenceSplitter
    >>> retriever = TempDocRetriever(output_format="text", join="
    ---------------
    ")
        retriever.create_node_group(transform=lambda text: [s.strip() for s in text.split("。") if s] )
        retriever.add_subretriever(group=Document.MediumChunk, topk=3)
        files = ["/path/to/file.txt"]
        results = retriever.forward(files, "什么是机器学习?")
        print(results)
    """
    @functools.lru_cache(maxsize=128)  # noqa B019
    def _get_retrievers(self, doc_files: List[str]):
        return self._get_retrievers_impl(doc_files)

    def __del__(self):
        self._get_retrievers.cache_clear()

lazyllm.tools.rag.retriever.UrlDocument

Bases: ModuleBase

UrlDocument类继承自ModuleBase,用于通过指定的URL和名称管理远程文档资源。 内部通过lazyllm的UrlModule代理实际调用,支持文档查找、检索和活跃节点分组查询。

Parameters:

  • url (str) –

    远程文档资源的访问URL。

  • name (str, default: None ) –

    当前文档分组名称,用于标识文档分组。

Source code in lazyllm/tools/rag/document.py
class UrlDocument(ModuleBase):
    """UrlDocument类继承自ModuleBase,用于通过指定的URL和名称管理远程文档资源。
内部通过lazyllm的UrlModule代理实际调用,支持文档查找、检索和活跃节点分组查询。

Args:
    url (str): 远程文档资源的访问URL。
    name (str): 当前文档分组名称,用于标识文档分组。
"""
    def __init__(self, url: str, name: str = None):
        super().__init__()
        self._missing_keys = set(dir(Document)) - set(dir(UrlDocument))
        self._manager = lazyllm.UrlModule(url=ensure_call_endpoint(url))
        self._curr_group = name or RAG_DEFAULT_GROUP_NAME

    def _forward(self, func_name: str, *args, **kwargs):
        args = (self._curr_group, func_name, *args)
        return self._manager._call('__call__', *args, **kwargs)

    def find(self, target) -> Callable:
        """生成一个部分应用函数,用于在当前文档组中查找指定目标。

Args:
    target (str): 需要查找的目标标识。

**Returns:**

- Callable: 调用时会执行查找操作的部分应用函数。
"""
        return functools.partial(self._forward, 'find', group=target)

    def forward(self, *args, **kw):
        return self._forward('retrieve', *args, **kw)

    def get_nodes(self, uids: Optional[List[str]] = None, doc_ids: Optional[Set] = None,
                  group: Optional[str] = None, kb_id: Optional[str] = None, numbers: Optional[Set] = None,
                  limit: Optional[int] = None, offset: int = 0, return_total: bool = False,
                  sort_by_number: bool = False) -> Union[List[DocNode], Tuple[List[DocNode], int]]:
        """按条件获取远程文档节点列表。

Args:
    uids (Optional[List[str]]): 指定节点 uid 列表。
    doc_ids (Optional[Set]): 指定文档 id 集合。
    group (Optional[str]): 节点组名。
    kb_id (Optional[str]): 知识库 id。
    numbers (Optional[Set]): 节点编号集合。

**Returns:**

- List[DocNode]: 命中的节点列表。
"""
        return self._forward(
            '_get_nodes', uids, doc_ids, group, kb_id, numbers, limit, offset, return_total, sort_by_number,
        )

    def get_window_nodes(self, node: DocNode, span: tuple[int, int] = (-5, 5),
                         merge: bool = False) -> Union[List[DocNode], DocNode]:
        """获取远程文档中指定节点的窗口节点。

Args:
    node (DocNode): 目标节点。
    span (tuple[int, int]): 窗口范围,基于 node.number 的相对偏移。
    merge (bool): 是否将窗口节点合并为一个节点返回。

**Returns:**

- Union[List[DocNode], DocNode]: 窗口节点列表,或合并后的单节点。
"""
        return self._forward('_get_window_nodes', node, span, merge)

    def keyword_search(self, group, keyword, doc_id='', kb_id=None,
                       phrase=True, sort_by='score', size=10, file_name=None):
        """在远程文档中执行关键词精准匹配。

与 :meth:`Document.keyword_search` 接口一致,通过 RPC 代理到远端 Document 服务。

Args:
    group (str): 节点组名。
    keyword (str): 待匹配的关键词。
    doc_id (str): 目标文档 ID,默认为空字符串。与 ``file_name`` 二选一,若同时提供则 ``file_name`` 优先。
    kb_id (Optional[str]): 知识库过滤条件。
    phrase (bool): True 为精确子串匹配,False 为单词级匹配。
    sort_by (str): ``"score"`` 按相关性排序,``"number"`` 按文档顺序排序。
    size (int): 最大返回条数。
    file_name (Optional[str]): 按文件名过滤,与 ``doc_id`` 二选一。提供此参数时 ``doc_id`` 被忽略。

Returns:
    List[dict]: 命中的切片列表。
"""
        return self._forward('_keyword_search', group, keyword, doc_id, kb_id, phrase, sort_by, size, file_name)

    @cached_property
    def active_node_groups(self):
        return self._forward('active_node_groups')

    def __getattr__(self, name):
        if name in self.__dict__.get('_missing_keys', []):
            raise AttributeError(f'Document generated with url and name has no attribute `{name}`')

find(target)

生成一个部分应用函数,用于在当前文档组中查找指定目标。

Parameters:

  • target (str) –

    需要查找的目标标识。

Returns:

  • Callable: 调用时会执行查找操作的部分应用函数。
Source code in lazyllm/tools/rag/document.py
    def find(self, target) -> Callable:
        """生成一个部分应用函数,用于在当前文档组中查找指定目标。

Args:
    target (str): 需要查找的目标标识。

**Returns:**

- Callable: 调用时会执行查找操作的部分应用函数。
"""
        return functools.partial(self._forward, 'find', group=target)

get_nodes(uids=None, doc_ids=None, group=None, kb_id=None, numbers=None, limit=None, offset=0, return_total=False, sort_by_number=False)

按条件获取远程文档节点列表。

Parameters:

  • uids (Optional[List[str]], default: None ) –

    指定节点 uid 列表。

  • doc_ids (Optional[Set], default: None ) –

    指定文档 id 集合。

  • group (Optional[str], default: None ) –

    节点组名。

  • kb_id (Optional[str], default: None ) –

    知识库 id。

  • numbers (Optional[Set], default: None ) –

    节点编号集合。

Returns:

  • List[DocNode]: 命中的节点列表。
Source code in lazyllm/tools/rag/document.py
    def get_nodes(self, uids: Optional[List[str]] = None, doc_ids: Optional[Set] = None,
                  group: Optional[str] = None, kb_id: Optional[str] = None, numbers: Optional[Set] = None,
                  limit: Optional[int] = None, offset: int = 0, return_total: bool = False,
                  sort_by_number: bool = False) -> Union[List[DocNode], Tuple[List[DocNode], int]]:
        """按条件获取远程文档节点列表。

Args:
    uids (Optional[List[str]]): 指定节点 uid 列表。
    doc_ids (Optional[Set]): 指定文档 id 集合。
    group (Optional[str]): 节点组名。
    kb_id (Optional[str]): 知识库 id。
    numbers (Optional[Set]): 节点编号集合。

**Returns:**

- List[DocNode]: 命中的节点列表。
"""
        return self._forward(
            '_get_nodes', uids, doc_ids, group, kb_id, numbers, limit, offset, return_total, sort_by_number,
        )

get_window_nodes(node, span=(-5, 5), merge=False)

获取远程文档中指定节点的窗口节点。

Parameters:

  • node (DocNode) –

    目标节点。

  • span (tuple[int, int], default: (-5, 5) ) –

    窗口范围,基于 node.number 的相对偏移。

  • merge (bool, default: False ) –

    是否将窗口节点合并为一个节点返回。

Returns:

  • Union[List[DocNode], DocNode]: 窗口节点列表,或合并后的单节点。
Source code in lazyllm/tools/rag/document.py
    def get_window_nodes(self, node: DocNode, span: tuple[int, int] = (-5, 5),
                         merge: bool = False) -> Union[List[DocNode], DocNode]:
        """获取远程文档中指定节点的窗口节点。

Args:
    node (DocNode): 目标节点。
    span (tuple[int, int]): 窗口范围,基于 node.number 的相对偏移。
    merge (bool): 是否将窗口节点合并为一个节点返回。

**Returns:**

- Union[List[DocNode], DocNode]: 窗口节点列表,或合并后的单节点。
"""
        return self._forward('_get_window_nodes', node, span, merge)

在远程文档中执行关键词精准匹配。

与 :meth:Document.keyword_search 接口一致,通过 RPC 代理到远端 Document 服务。

Parameters:

  • group (str) –

    节点组名。

  • keyword (str) –

    待匹配的关键词。

  • doc_id (str, default: '' ) –

    目标文档 ID,默认为空字符串。与 file_name 二选一,若同时提供则 file_name 优先。

  • kb_id (Optional[str], default: None ) –

    知识库过滤条件。

  • phrase (bool, default: True ) –

    True 为精确子串匹配,False 为单词级匹配。

  • sort_by (str, default: 'score' ) –

    "score" 按相关性排序,"number" 按文档顺序排序。

  • size (int, default: 10 ) –

    最大返回条数。

  • file_name (Optional[str], default: None ) –

    按文件名过滤,与 doc_id 二选一。提供此参数时 doc_id 被忽略。

Returns:

  • List[dict]: 命中的切片列表。

Source code in lazyllm/tools/rag/document.py
    def keyword_search(self, group, keyword, doc_id='', kb_id=None,
                       phrase=True, sort_by='score', size=10, file_name=None):
        """在远程文档中执行关键词精准匹配。

与 :meth:`Document.keyword_search` 接口一致,通过 RPC 代理到远端 Document 服务。

Args:
    group (str): 节点组名。
    keyword (str): 待匹配的关键词。
    doc_id (str): 目标文档 ID,默认为空字符串。与 ``file_name`` 二选一,若同时提供则 ``file_name`` 优先。
    kb_id (Optional[str]): 知识库过滤条件。
    phrase (bool): True 为精确子串匹配,False 为单词级匹配。
    sort_by (str): ``"score"`` 按相关性排序,``"number"`` 按文档顺序排序。
    size (int): 最大返回条数。
    file_name (Optional[str]): 按文件名过滤,与 ``doc_id`` 二选一。提供此参数时 ``doc_id`` 被忽略。

Returns:
    List[dict]: 命中的切片列表。
"""
        return self._forward('_keyword_search', group, keyword, doc_id, kb_id, phrase, sort_by, size, file_name)

lazyllm.tools.rag.doc_service.DocServer

Bases: ModuleBase

文档服务的主入口。

DocServer 负责文档上传/添加/重解析/删除、任务跟踪、知识库管理、chunk 查看,以及跨知识库文档转移。 它是 legacy DocManager / DocListManager API 的推荐替代方案。

Parameters:

  • port (Optional[int], default: None ) –

    本地启动服务时使用的端口。

  • url (Optional[str], default: None ) –

    已存在的 doc_service 地址;提供后当前实例作为远程客户端使用。

  • parser_url (Optional[str], default: None ) –

    本地 doc_service 使用的 parsing service 地址。

  • db_config (Optional[Dict[str, Any]], default: None ) –

    doc_service 元数据数据库配置。

  • parser_db_config (Optional[Dict[str, Any]], default: None ) –

    parsing service 任务数据库配置。

  • parser_poll_interval (float, default: 0.05 ) –

    本地解析协调使用的轮询间隔。

  • storage_dir (Optional[str], default: None ) –

    上传文件保存目录。

  • callback_url (Optional[str], default: None ) –

    接收解析任务回调的地址。

  • launcher

    本地服务启动器。

Source code in lazyllm/tools/rag/doc_service/doc_server.py
  48
  49
  50
  51
  52
  53
  54
  55
  56
  57
  58
  59
  60
  61
  62
  63
  64
  65
  66
  67
  68
  69
  70
  71
  72
  73
  74
  75
  76
  77
  78
  79
  80
  81
  82
  83
  84
  85
  86
  87
  88
  89
  90
  91
  92
  93
  94
  95
  96
  97
  98
  99
 100
 101
 102
 103
 104
 105
 106
 107
 108
 109
 110
 111
 112
 113
 114
 115
 116
 117
 118
 119
 120
 121
 122
 123
 124
 125
 126
 127
 128
 129
 130
 131
 132
 133
 134
 135
 136
 137
 138
 139
 140
 141
 142
 143
 144
 145
 146
 147
 148
 149
 150
 151
 152
 153
 154
 155
 156
 157
 158
 159
 160
 161
 162
 163
 164
 165
 166
 167
 168
 169
 170
 171
 172
 173
 174
 175
 176
 177
 178
 179
 180
 181
 182
 183
 184
 185
 186
 187
 188
 189
 190
 191
 192
 193
 194
 195
 196
 197
 198
 199
 200
 201
 202
 203
 204
 205
 206
 207
 208
 209
 210
 211
 212
 213
 214
 215
 216
 217
 218
 219
 220
 221
 222
 223
 224
 225
 226
 227
 228
 229
 230
 231
 232
 233
 234
 235
 236
 237
 238
 239
 240
 241
 242
 243
 244
 245
 246
 247
 248
 249
 250
 251
 252
 253
 254
 255
 256
 257
 258
 259
 260
 261
 262
 263
 264
 265
 266
 267
 268
 269
 270
 271
 272
 273
 274
 275
 276
 277
 278
 279
 280
 281
 282
 283
 284
 285
 286
 287
 288
 289
 290
 291
 292
 293
 294
 295
 296
 297
 298
 299
 300
 301
 302
 303
 304
 305
 306
 307
 308
 309
 310
 311
 312
 313
 314
 315
 316
 317
 318
 319
 320
 321
 322
 323
 324
 325
 326
 327
 328
 329
 330
 331
 332
 333
 334
 335
 336
 337
 338
 339
 340
 341
 342
 343
 344
 345
 346
 347
 348
 349
 350
 351
 352
 353
 354
 355
 356
 357
 358
 359
 360
 361
 362
 363
 364
 365
 366
 367
 368
 369
 370
 371
 372
 373
 374
 375
 376
 377
 378
 379
 380
 381
 382
 383
 384
 385
 386
 387
 388
 389
 390
 391
 392
 393
 394
 395
 396
 397
 398
 399
 400
 401
 402
 403
 404
 405
 406
 407
 408
 409
 410
 411
 412
 413
 414
 415
 416
 417
 418
 419
 420
 421
 422
 423
 424
 425
 426
 427
 428
 429
 430
 431
 432
 433
 434
 435
 436
 437
 438
 439
 440
 441
 442
 443
 444
 445
 446
 447
 448
 449
 450
 451
 452
 453
 454
 455
 456
 457
 458
 459
 460
 461
 462
 463
 464
 465
 466
 467
 468
 469
 470
 471
 472
 473
 474
 475
 476
 477
 478
 479
 480
 481
 482
 483
 484
 485
 486
 487
 488
 489
 490
 491
 492
 493
 494
 495
 496
 497
 498
 499
 500
 501
 502
 503
 504
 505
 506
 507
 508
 509
 510
 511
 512
 513
 514
 515
 516
 517
 518
 519
 520
 521
 522
 523
 524
 525
 526
 527
 528
 529
 530
 531
 532
 533
 534
 535
 536
 537
 538
 539
 540
 541
 542
 543
 544
 545
 546
 547
 548
 549
 550
 551
 552
 553
 554
 555
 556
 557
 558
 559
 560
 561
 562
 563
 564
 565
 566
 567
 568
 569
 570
 571
 572
 573
 574
 575
 576
 577
 578
 579
 580
 581
 582
 583
 584
 585
 586
 587
 588
 589
 590
 591
 592
 593
 594
 595
 596
 597
 598
 599
 600
 601
 602
 603
 604
 605
 606
 607
 608
 609
 610
 611
 612
 613
 614
 615
 616
 617
 618
 619
 620
 621
 622
 623
 624
 625
 626
 627
 628
 629
 630
 631
 632
 633
 634
 635
 636
 637
 638
 639
 640
 641
 642
 643
 644
 645
 646
 647
 648
 649
 650
 651
 652
 653
 654
 655
 656
 657
 658
 659
 660
 661
 662
 663
 664
 665
 666
 667
 668
 669
 670
 671
 672
 673
 674
 675
 676
 677
 678
 679
 680
 681
 682
 683
 684
 685
 686
 687
 688
 689
 690
 691
 692
 693
 694
 695
 696
 697
 698
 699
 700
 701
 702
 703
 704
 705
 706
 707
 708
 709
 710
 711
 712
 713
 714
 715
 716
 717
 718
 719
 720
 721
 722
 723
 724
 725
 726
 727
 728
 729
 730
 731
 732
 733
 734
 735
 736
 737
 738
 739
 740
 741
 742
 743
 744
 745
 746
 747
 748
 749
 750
 751
 752
 753
 754
 755
 756
 757
 758
 759
 760
 761
 762
 763
 764
 765
 766
 767
 768
 769
 770
 771
 772
 773
 774
 775
 776
 777
 778
 779
 780
 781
 782
 783
 784
 785
 786
 787
 788
 789
 790
 791
 792
 793
 794
 795
 796
 797
 798
 799
 800
 801
 802
 803
 804
 805
 806
 807
 808
 809
 810
 811
 812
 813
 814
 815
 816
 817
 818
 819
 820
 821
 822
 823
 824
 825
 826
 827
 828
 829
 830
 831
 832
 833
 834
 835
 836
 837
 838
 839
 840
 841
 842
 843
 844
 845
 846
 847
 848
 849
 850
 851
 852
 853
 854
 855
 856
 857
 858
 859
 860
 861
 862
 863
 864
 865
 866
 867
 868
 869
 870
 871
 872
 873
 874
 875
 876
 877
 878
 879
 880
 881
 882
 883
 884
 885
 886
 887
 888
 889
 890
 891
 892
 893
 894
 895
 896
 897
 898
 899
 900
 901
 902
 903
 904
 905
 906
 907
 908
 909
 910
 911
 912
 913
 914
 915
 916
 917
 918
 919
 920
 921
 922
 923
 924
 925
 926
 927
 928
 929
 930
 931
 932
 933
 934
 935
 936
 937
 938
 939
 940
 941
 942
 943
 944
 945
 946
 947
 948
 949
 950
 951
 952
 953
 954
 955
 956
 957
 958
 959
 960
 961
 962
 963
 964
 965
 966
 967
 968
 969
 970
 971
 972
 973
 974
 975
 976
 977
 978
 979
 980
 981
 982
 983
 984
 985
 986
 987
 988
 989
 990
 991
 992
 993
 994
 995
 996
 997
 998
 999
1000
1001
1002
1003
1004
1005
1006
1007
1008
1009
1010
1011
1012
1013
1014
1015
1016
1017
1018
1019
1020
1021
1022
1023
1024
1025
1026
1027
1028
1029
1030
1031
1032
1033
1034
1035
1036
1037
1038
1039
1040
1041
1042
1043
1044
1045
1046
1047
1048
1049
1050
1051
1052
1053
1054
1055
1056
1057
1058
1059
1060
1061
1062
1063
1064
1065
1066
1067
1068
1069
1070
1071
1072
1073
1074
1075
1076
1077
1078
1079
1080
1081
1082
1083
1084
1085
1086
1087
1088
1089
1090
1091
1092
1093
1094
1095
1096
1097
1098
1099
1100
1101
1102
1103
1104
1105
1106
1107
1108
1109
1110
1111
1112
1113
1114
1115
1116
1117
1118
1119
1120
1121
1122
1123
1124
1125
1126
1127
1128
1129
1130
1131
1132
1133
1134
1135
1136
1137
1138
1139
1140
1141
1142
1143
1144
1145
1146
1147
1148
1149
1150
1151
1152
1153
1154
1155
1156
1157
1158
1159
1160
1161
1162
1163
1164
1165
1166
1167
1168
1169
1170
1171
1172
1173
1174
1175
1176
1177
1178
1179
1180
1181
1182
1183
1184
1185
1186
1187
1188
1189
1190
1191
1192
1193
1194
1195
1196
1197
1198
1199
1200
1201
1202
1203
1204
1205
1206
1207
1208
1209
1210
1211
1212
1213
1214
1215
1216
1217
1218
1219
1220
1221
1222
1223
1224
1225
1226
1227
1228
1229
1230
1231
1232
1233
1234
1235
1236
1237
1238
1239
1240
1241
1242
1243
1244
1245
1246
1247
1248
1249
1250
1251
1252
1253
1254
1255
1256
1257
1258
1259
1260
1261
1262
1263
1264
1265
1266
1267
1268
1269
1270
1271
1272
1273
1274
1275
1276
1277
1278
1279
1280
1281
1282
1283
1284
1285
1286
1287
1288
1289
1290
1291
1292
1293
1294
1295
1296
1297
1298
1299
1300
1301
1302
1303
1304
1305
1306
1307
1308
1309
1310
1311
1312
1313
1314
1315
1316
1317
1318
1319
1320
1321
1322
1323
1324
1325
1326
1327
1328
1329
1330
1331
1332
1333
1334
1335
1336
1337
1338
1339
1340
1341
1342
1343
1344
1345
1346
1347
1348
1349
1350
1351
1352
1353
1354
1355
1356
1357
1358
1359
1360
1361
1362
1363
1364
1365
1366
1367
1368
1369
1370
1371
1372
1373
1374
1375
1376
1377
1378
1379
1380
1381
1382
1383
1384
1385
1386
1387
1388
1389
1390
1391
1392
1393
1394
1395
1396
1397
1398
1399
1400
1401
1402
1403
1404
1405
1406
1407
1408
1409
1410
1411
1412
1413
1414
1415
1416
1417
1418
1419
class DocServer(ModuleBase):
    """文档服务的主入口。

``DocServer`` 负责文档上传/添加/重解析/删除、任务跟踪、知识库管理、chunk 查看,以及跨知识库文档转移。
它是 legacy ``DocManager`` / ``DocListManager`` API 的推荐替代方案。

Args:
    port (Optional[int]): 本地启动服务时使用的端口。
    url (Optional[str]): 已存在的 doc_service 地址;提供后当前实例作为远程客户端使用。
    parser_url (Optional[str]): 本地 doc_service 使用的 parsing service 地址。
    db_config (Optional[Dict[str, Any]]): doc_service 元数据数据库配置。
    parser_db_config (Optional[Dict[str, Any]]): parsing service 任务数据库配置。
    parser_poll_interval (float): 本地解析协调使用的轮询间隔。
    storage_dir (Optional[str]): 上传文件保存目录。
    callback_url (Optional[str]): 接收解析任务回调的地址。
    launcher: 本地服务启动器。
"""
    class _Impl:
        def __init__(
            self,
            storage_dir: str,
            db_config: Optional[Dict[str, Any]] = None,
            parser_db_config: Optional[Dict[str, Any]] = None,
            parser_poll_interval: float = 0.05,
            parser_url: Optional[str] = None,
            callback_url: Optional[str] = None,
            enable_scan: bool = False,
            scan_interval: int = 10,
        ):
            if not parser_url:
                raise ValueError('parser_url is required; doc_service no longer starts a mock parsing server')
            self._storage_dir = storage_dir
            self._db_config = db_config
            self._parser_db_config = parser_db_config
            self._parser_poll_interval = parser_poll_interval
            self._parser_url = parser_url
            self._callback_url = callback_url
            self._parser = None
            self._manager = None
            self._enable_scan = enable_scan
            self._scan_interval = scan_interval
            self._scan_thread = None
            self._scan_continue = False
            self._owned_kbs: Set[str] = set()

        @once_wrapper(reset_on_pickle=True)
        def _lazy_init(self):
            if self._storage_dir and not os.path.exists(self._storage_dir):
                os.makedirs(self._storage_dir, exist_ok=True)
            self._manager = DocManager(
                db_config=self._db_config,
                parser_url=self._parser_url,
                callback_url=self._callback_url,
            )
            # NOTE: scanning is NOT started here.  Use ``enable_scanning()`` after
            # all KB registrations and parser algorithm registrations are complete
            # so the first scan sees a consistent _owned_kbs set and can route
            # requests to algorithms that actually exist on the parser.
            # For standalone / direct DocServer usage with enable_scan=True and no
            # explicit ``enable_scanning()`` call, the scan thread is started lazily
            # on the first ``_sync_dataset()`` invocation if still not running.

        def stop(self):
            self._scan_continue = False
            if self._scan_thread and self._scan_thread.is_alive():
                self._scan_thread.join(timeout=2)
            # Release the DB engine so callers can clean up the backing directory
            # (sqlite on Windows keeps an exclusive handle until dispose()).
            if self._manager is not None:
                try:
                    self._manager.close()
                except Exception:
                    pass
            return None

        def _sync_dataset_for_kb(self, kb_id: str, algo_id: str, disk_files: list, disk_set: set):
            """Sync one KB: diff disk vs documents table -> upload new / delete stale via unified pipeline."""
            # For retry: exclude FAILED/CANCELED so they get re-uploaded
            synced_docs = self._manager._list_kb_docs_by_path(kb_id, exclude_failed=True)
            # For stale cleanup: include FAILED/CANCELED so removed files get cleaned up
            all_known_docs = self._manager._list_kb_docs_by_path(kb_id, exclude_failed=False)

            # New files (or previously failed) → upload(source_type=SCAN)
            new_paths = [p for p in disk_files if p not in synced_docs]
            if new_paths:
                try:
                    request = UploadRequest(
                        items=[AddFileItem(file_path=p) for p in new_paths],
                        kb_id=kb_id, algo_id=algo_id, source_type=SourceType.SCAN,
                    )
                    self._manager.upload(request)
                except Exception as exc:
                    LOG.error(f'[Scan] upload failed for kb={kb_id}: {len(new_paths)} files: {exc}')

            # Stale files (including failed ones whose source file was removed) → delete
            stale_ids = [did for path, did in all_known_docs.items() if path not in disk_set]
            if stale_ids:
                try:
                    request = DeleteRequest(doc_ids=stale_ids, kb_id=kb_id)
                    self._manager.delete(request)
                except Exception as exc:
                    LOG.error(f'[Scan] delete failed for kb={kb_id}: {len(stale_ids)} docs: {exc}')

            if new_paths or stale_ids:
                LOG.info(f'[Scan] kb={kb_id} sync done: added={len(new_paths)}, deleted={len(stale_ids)}')

        def _sync_dataset(self):
            """One-shot scan: list dir -> sync all active KB+algo pairs."""
            from .utils import list_dataset_files
            disk_files = list_dataset_files(self._storage_dir)
            disk_set = set(disk_files)

            kb_algo_pairs = self._manager._list_active_kb_algo_pairs()
            if not kb_algo_pairs:
                kb_algo_pairs = [('__default__', '__default__')]

            # When this instance has explicitly registered KBs, only scan those
            # to avoid processing KBs that belong to other Document instances
            # sharing the same global DB.
            owned = self._owned_kbs.copy()
            if owned:
                kb_algo_pairs = [(kb, algo) for kb, algo in kb_algo_pairs if kb in owned]

            for kb_id, algo_id in kb_algo_pairs:
                try:
                    self._sync_dataset_for_kb(kb_id, algo_id, disk_files, disk_set)
                except Exception as exc:
                    LOG.error(f'[Scan] sync failed for kb={kb_id}, algo={algo_id}: {exc}')

        def _scan_worker(self):
            """Daemon thread: periodically scan dataset directory."""
            while self._scan_continue:
                try:
                    self._sync_dataset()
                except Exception as exc:
                    LOG.error(f'[Scan] sync failed: {exc}')
                time.sleep(self._scan_interval)

        def _start_scan_monitoring(self):
            if self._scan_thread and self._scan_thread.is_alive():
                return
            self._scan_continue = True
            self._scan_thread = threading.Thread(target=self._scan_worker, daemon=True)
            self._scan_thread.start()

        def enable_scanning(self):
            """Start scanning after all KB registrations and parser algo registrations
            are complete.  Safe to call multiple times (idempotent).

            This is the intended way for ``Document._Manager`` to trigger the first
            scan: it ensures ``_owned_kbs`` is fully populated and all algorithms
            have been registered with the parser before any file-level sync happens.
            """
            self._lazy_init()
            if not self._enable_scan:
                return
            if not (self._storage_dir and os.path.isdir(self._storage_dir)):
                return
            self._sync_dataset()
            self._start_scan_monitoring()

        def ensure_kb_registered(self, kb_id: str, algo_id: Optional[str] = None):
            """Lightweight KB registration: ensure KB + algo binding rows exist in DB.

            Unlike ``create_kb`` this does NOT validate algorithm existence
            against the parser, so it can be called before the algorithm is registered
            (e.g. during ``add_kb_group`` which creates a DocImpl that will register
            its algorithm later during ``_lazy_init``).
            """
            self._lazy_init()
            algo_id = algo_id or kb_id
            self._manager._ensure_kb(kb_id, display_name=kb_id)
            self._manager._ensure_kb_algorithm(kb_id, algo_id)
            self._owned_kbs.add(kb_id)

        def set_runtime_callback_url(self, callback_url: str):
            # The callback URL is configuration, not a reason to initialize the
            # service.  DocServer.start() invokes this as soon as its HTTP socket
            # opens, which can precede parser readiness during parallel startup.
            self._callback_url = callback_url
            if self._manager is not None:
                self._manager.set_callback_url(callback_url)

        @staticmethod
        def _response(data=None, code=200, msg='success', status_code=200):
            payload = BaseResponse(code=code, msg=msg, data=data).model_dump(mode='json')
            return fastapi.responses.JSONResponse(status_code=status_code, content=payload)

        def _run(self, func, *args, success_msg='success', **kwargs):
            try:
                data = func(*args, **kwargs)
                return self._response(data=data, msg=success_msg)
            except DocServiceError as exc:
                data = dict(exc.data or {})
                data.setdefault('biz_code', exc.biz_code)
                return self._response(data=data, code=exc.http_status, msg=exc.msg, status_code=exc.http_status)
            except fastapi.HTTPException as exc:
                detail = exc.detail if isinstance(exc.detail, dict) else {}
                data = detail.get('data')
                if isinstance(data, dict) and 'biz_code' not in data and detail.get('code'):
                    data['biz_code'] = detail['code']
                code = exc.status_code
                msg = detail.get('msg', str(exc.detail))
                return self._response(data=data, code=code, msg=msg, status_code=exc.status_code)

        @staticmethod
        def _build_upload_payload(request: UploadRequest, file_identities: Optional[List[Dict[str, Any]]] = None):
            source_type = request.source_type or SourceType.API
            items = file_identities
            if items is None:
                items = []
                for idx, item in enumerate(request.items):
                    content_hash = None
                    size_bytes = None
                    if os.path.exists(item.file_path):
                        content_hash = sha256_file(item.file_path)
                        size_bytes = os.path.getsize(item.file_path)
                    items.append({
                        'filename': os.path.basename(item.file_path),
                        'content_hash': content_hash,
                        'size_bytes': size_bytes,
                        'doc_id': item.doc_id if idx == 0 else None,
                    })
            return {
                'kb_id': request.kb_id,
                'algo_id': None,  # deprecated; kept for idempotency-key stability
                'source_type': source_type.value,
                'idempotency_key': request.idempotency_key,
                'items': items,
            }

        @staticmethod
        def _build_update_kb_payload(kb_id: str, request: KbUpdateRequest):
            payload = request.model_dump(mode='json', exclude_unset=True)
            payload['kb_id'] = kb_id
            payload['explicit_fields'] = sorted(field for field in request.model_fields_set if field != 'kb_id')
            return payload

        def _gen_unique_upload_path(
            self, filename: str, reserved_paths: Optional[set] = None,
            *, base_dir: Optional[str] = None,
        ):
            safe_name = os.path.basename(filename) or 'upload.bin'
            target_dir = base_dir or self._storage_dir
            file_path = os.path.join(target_dir, safe_name)
            reserved_paths = reserved_paths or set()
            if file_path not in reserved_paths and not os.path.exists(file_path):
                return file_path

            suffix = os.path.splitext(safe_name)[1]
            prefix = safe_name[:-len(suffix)] if suffix else safe_name
            for idx in range(1, 10000):
                candidate = os.path.join(target_dir, f'{prefix}-{idx}{suffix}')
                if candidate not in reserved_paths and not os.path.exists(candidate):
                    return candidate
            digest = hashlib.sha256(safe_name.encode()).hexdigest()[:8]
            return os.path.join(target_dir, f'{prefix}-{digest}{suffix}')

        @staticmethod
        async def _save_upload_file(upload_file: 'fastapi.UploadFile', file_path: str):
            with open(file_path, 'wb') as fh:
                while True:
                    chunk = await upload_file.read(1024 * 1024)
                    if not chunk:
                        break
                    fh.write(chunk)
            await upload_file.close()

        async def _persist_uploads(
            self, files: List[fastapi.UploadFile], *,
            override: bool = False, sub_dir: str = '',
        ):
            saved_paths = []
            file_identities = []
            reserved_paths: Set[str] = set()
            target_dir = os.path.join(self._storage_dir, sub_dir) if sub_dir else self._storage_dir
            if sub_dir:
                os.makedirs(target_dir, exist_ok=True)
            for upload_file in files:
                filename = getattr(upload_file, 'filename', None) or 'upload.bin'
                if override:
                    # Legacy /upload_files behavior: write to
                    # ``storage_dir[/sub_dir]/<name>``, overwriting any
                    # existing file. ``DocManager.upload()`` derives
                    # ``doc_id`` from ``file_path``, so this lets a
                    # re-upload replace the existing document instead of
                    # creating a new one.
                    safe_name = os.path.basename(filename) or 'upload.bin'
                    file_path = os.path.join(target_dir, safe_name)
                    if file_path in reserved_paths:
                        # Two uploads in the same request collided on the same name;
                        # keep the unique-path fallback so we don't lose one of them.
                        file_path = self._gen_unique_upload_path(
                            filename, reserved_paths, base_dir=target_dir,
                        )
                else:
                    file_path = self._gen_unique_upload_path(
                        filename, reserved_paths, base_dir=target_dir,
                    )
                await self._save_upload_file(upload_file, file_path)
                reserved_paths.add(file_path)
                saved_paths.append(file_path)
                file_identities.append({
                    'filename': os.path.basename(file_path),
                    'content_hash': sha256_file(file_path),
                    'size_bytes': os.path.getsize(file_path),
                    'doc_id': None,
                })
            return saved_paths, file_identities

        def _run_upload(self, request: UploadRequest, payload: Optional[Dict[str, Any]] = None):
            idem_payload = payload or self._build_upload_payload(request)
            return self._run(lambda: self._manager.run_idempotent(
                '/v1/docs/upload', request.idempotency_key, idem_payload,
                lambda: {'items': self._manager.upload(request)}
            ))

        @staticmethod
        def _normalize_task_callback(callback: Any) -> TaskCallbackRequest:
            if isinstance(callback, TaskCallbackRequest):
                return callback
            if not isinstance(callback, dict):
                raise DocServiceError('E_INVALID_PARAM', 'invalid callback payload')

            payload = dict(callback.get('payload') or {})
            for field in ('task_type', 'doc_id', 'kb_id', 'algo_id'):
                if callback.get(field) is not None and field not in payload:
                    payload[field] = callback[field]

            event_type = callback.get('event_type')
            status = callback.get('status')
            task_status = callback.get('task_status')

            try:
                if status is not None:
                    normalized_status = DocStatus(status)
                    normalized_event_type = CallbackEventType(event_type) if event_type else (
                        CallbackEventType.START
                        if normalized_status in (DocStatus.WAITING, DocStatus.WORKING) else CallbackEventType.FINISH
                    )
                elif task_status is not None:
                    normalized_status = DocStatus(task_status)
                    normalized_event_type = CallbackEventType(event_type) if event_type else (
                        CallbackEventType.START
                        if normalized_status in (DocStatus.WAITING, DocStatus.WORKING)
                        else CallbackEventType.FINISH
                    )
                else:
                    raise DocServiceError('E_INVALID_PARAM', 'status or task_status is required')
            except ValueError as exc:
                raise DocServiceError('E_INVALID_PARAM', str(exc)) from exc

            callback_data = {
                'callback_id': callback.get('callback_id'),
                'task_id': callback.get('task_id'),
                'event_type': normalized_event_type,
                'status': normalized_status,
                'error_code': callback.get('error_code'),
                'error_msg': callback.get('error_msg'),
                'payload': payload,
            }
            return TaskCallbackRequest.model_validate({k: v for k, v in callback_data.items() if v is not None})

        @staticmethod
        def _format_task_view(task: Optional[Dict[str, Any]]):
            return task

        def _format_task_response_data(self, data: Any):
            if isinstance(data, dict) and isinstance(data.get('items'), list):
                payload = dict(data)
                payload['items'] = [self._format_task_view(item) for item in data['items']]
                return payload
            return self._format_task_view(data)

        def upload_request(self, request: UploadRequest):
            self._lazy_init()
            return self._run_upload(request)

        # Reserved metadata keys that the legacy /upload_files rejected; if a
        # client smuggles them into ``metadatas`` they would silently overwrite
        # internal docid/path tracking via ``parsing_service/impl.py``'s
        # ``setdefault`` calls and break later filter/delete/reparse flows.
        _LEGACY_RESERVED_META_KEYS = frozenset({'docid', 'doc_id', 'lazyllm_doc_path'})

        @classmethod
        def _parse_legacy_metadatas(cls, metadatas: Optional[str], expected_len: int) -> List[Dict[str, Any]]:
            """Validate and parse the legacy ``metadatas`` query param.

            Returns the parsed list (empty when ``metadatas`` is falsy). Raises
            HTTPException(400) on JSON / shape / length / reserved-key issues so
            the caller doesn't have to repeat each guard inline.
            """
            if not metadatas:
                return []
            try:
                parsed = json.loads(metadatas) or []
            except (ValueError, TypeError) as exc:
                raise fastapi.HTTPException(
                    status_code=400, detail=f'metadatas must be valid JSON: {exc}',
                )
            if not isinstance(parsed, list):
                raise fastapi.HTTPException(
                    status_code=400, detail='metadatas must be a JSON array',
                )
            # Legacy contract: rejects arrays with the wrong length or non-dict
            # entries with 400, instead of silently dropping/padding (which
            # would attach the wrong metadata to uploads) or letting AddFileItem
            # raise a 500.
            if len(parsed) != expected_len:
                raise fastapi.HTTPException(
                    status_code=400,
                    detail=f'metadatas length {len(parsed)} does not match files length {expected_len}',
                )
            for entry in parsed:
                if not isinstance(entry, dict):
                    raise fastapi.HTTPException(
                        status_code=400, detail='each metadatas entry must be a JSON object',
                    )
                bad = cls._LEGACY_RESERVED_META_KEYS.intersection(entry.keys())
                if bad:
                    raise fastapi.HTTPException(
                        status_code=400,
                        detail=f'metadatas contains reserved keys: {sorted(bad)}',
                    )
            return parsed

        @staticmethod
        def _normalize_legacy_user_path(user_path: Optional[str]) -> str:
            """Validate and normalize the legacy ``user_path`` query param.

            Rejects absolute paths or any value that climbs out of
            ``storage_dir`` (``..``, ``../x``, ``..\\\\x`` on Windows). Returns
            the relative subdirectory string ('' when no user_path).
            """
            if not user_path:
                return ''
            if os.path.isabs(user_path):
                raise fastapi.HTTPException(
                    status_code=400, detail=f'invalid user_path: {user_path!r}',
                )
            normalized = os.path.normpath(user_path)
            climbed_out = (
                normalized in ('.', '..')
                or normalized.startswith('../')
                or normalized.startswith('..' + os.sep)
                or os.path.isabs(normalized)
            )
            if climbed_out:
                raise fastapi.HTTPException(
                    status_code=400, detail=f'invalid user_path: {user_path!r}',
                )
            sub_dir = normalized.strip('/').strip(os.sep)
            if not sub_dir:
                raise fastapi.HTTPException(
                    status_code=400, detail=f'invalid user_path: {user_path!r}',
                )
            return sub_dir

        async def _legacy_upload(
            self,
            files: List['fastapi.UploadFile'],
            override: bool,
            metadatas: Optional[str],
            group_name: Optional[str],
            user_path: Optional[str],
            *,
            response_shape: str = 'ids_and_results',  # 'ids_and_results' or 'ids_only'
        ):
            """Shared implementation for legacy DocManager-style upload endpoints.

            Kept so DocWebModule and external callers from before the doc_service
            refactor (which used /upload_files and /add_files_to_group on the old
            ``ServerModule(DocManager(...))``) keep working against the new
            DocServer. New callers should target the /v1/docs/* surface instead.

            Compatibility behaviors preserved here:
            - ``override=True`` writes files at deterministic paths
              (``storage_dir/<user_path>/<filename>``) so DocManager.upload
              derives the same ``doc_id`` and reparses the existing document
              rather than creating a duplicate.
            - ``user_path`` namespaces uploads under a subdirectory, so two
              callers can post the same filename without colliding.
            - ``algo_id`` mirrors the kb-binding convention used by
              ``DocServer._Impl.ensure_kb_registered`` (``algo_id == kb_id``)
              so non-default groups don't get rejected by the algorithm
              validator.
            - ``metadatas`` rejects reserved internal keys (``docid``,
              ``doc_id``, ``lazyllm_doc_path``) instead of silently shadowing
              them downstream, and 400s on length / element-type mismatch
              instead of mis-attaching metadata or 500-ing inside AddFileItem.
            - Response body matches the legacy shape: ``data=[ids, results]``
              for ``/upload_files``; flat ``data=ids`` for ``/add_files_to_group``.
              ``results`` propagates per-item ``error_code`` (or ``'ok'``) from
              ``DocManager.upload``, so synchronous failures (e.g. parser
              outage -> ``PARSER_SUBMIT_FAILED``) don't masquerade as success.

            Override→reparse routing: when ``override=True`` and a doc already
            exists at the destination path in the kb (excluding FAILED /
            CANCELED, which still take the upload-retry path so caller
            metadata is applied), the shim sends those items through
            ``DocManager.reparse`` instead of ``upload``, so the legacy
            "replace + reparse" workflow doesn't get rejected by the new
            ``_assert_action_allowed(..., 'upload')`` 409 on SUCCESS docs.
            New paths in the same request still flow through ``upload``.
            Caller-supplied metadata for existing docs is merged directly
            into the documents row before reparse via
            ``_legacy_apply_metadata_to_existing_docs`` -- this avoids the
            ``patch_metadata``→reparse race that would orphan a
            DOC_UPDATE_META task and intermittently 409 the reparse.

            Known compat gaps tracked in #1090 (not exercised by the migrated
            tests so the PR ships as-is; future PRs should harden these):
            - The new file bytes are written before the doc-service validation
              runs; a 4xx from ``upload()``/``reparse()`` after override leaves
              the new bytes on disk while DB attrs still describe the old
              file. New code should prefer the staged ``/v1/docs/upload``.
            - The override+metadata write happens before reparse validation,
              so a reparse rejection (e.g. WORKING/DELETING state) commits
              the metadata change without rolling back. Mitigated for
              FAILED/CANCELED docs by routing them through upload instead.
            - The override+reparse path doesn't refresh ``content_hash`` /
              ``size_bytes`` / ``file_type`` for the replaced file —
              ``list_docs`` / ``get_doc_detail`` will continue showing the
              old file's attrs until a separate metadata patch lands.
            - Validation errors raise ``HTTPException`` and surface FastAPI's
              ``{"detail": ...}`` body, not the standard ``{code,msg,data}``
              envelope.
            """
            self._lazy_init()
            if not files:
                raise fastapi.HTTPException(status_code=400, detail='files is required')
            kb_id = group_name or '__default__'
            # Match Document._Manager / ensure_kb_registered: each kb is bound
            # to an algorithm of the same name. Hard-coding '__default__' here
            # would 400 every non-default-group upload at validation time.
            algo_id = kb_id
            parsed_metadatas = self._parse_legacy_metadatas(metadatas, len(files))
            sub_dir = self._normalize_legacy_user_path(user_path)
            saved_paths, file_identities = await self._persist_uploads(
                files, override=override, sub_dir=sub_dir,
            )
            return self._run(lambda: self._legacy_dispatch_uploads(
                saved_paths=saved_paths,
                file_identities=file_identities,
                metadatas=parsed_metadatas,
                kb_id=kb_id,
                algo_id=algo_id,
                override=override,
                response_shape=response_shape,
            ))

        def _legacy_apply_metadata_to_existing_docs(self, doc_id_meta_pairs):
            """Synchronously merge new metadata into the documents table for the
            given (doc_id, metadata) pairs. Used by the legacy override path
            instead of ``DocManager.patch_metadata`` so we don't enqueue a
            DOC_UPDATE_META task that would race the subsequent reparse.

            We MERGE rather than replace, matching the legacy "metadatas
            updates the named keys" semantic — callers that send a partial
            metadata dict shouldn't lose existing keys.
            """
            if not doc_id_meta_pairs:
                return
            db = self._manager._db_manager
            Doc = db.get_table_orm_class(DOCUMENTS_TABLE_INFO['name'])
            with db.get_session() as session:
                for doc_id, patch in doc_id_meta_pairs:
                    row = session.query(Doc).filter(Doc.doc_id == doc_id).first()
                    if row is None:
                        continue
                    existing = from_json(row.meta) if row.meta else {}
                    existing.update(patch)
                    row.meta = to_json(existing)
                    row.updated_at = datetime.now()
                    session.add(row)

        def _legacy_dispatch_uploads(
            self, *, saved_paths, file_identities, metadatas, kb_id, algo_id,
            override, response_shape,
        ):
            """Split saved files into "new doc -> upload" and (override only)
            "existing doc -> reparse + metadata patch", then merge the per-item
            results in input order so the legacy response shape stays correct.

            ``DocManager.upload`` rejects an already-SUCCESS doc with 409 (via
            ``_assert_action_allowed``), so a re-upload of the same file path
            with ``override=True`` would otherwise break the legacy
            "replace + reparse" workflow. We look up existing doc_ids by path
            and route those through ``DocManager.reparse``, then ``patch_metadata``
            for any updated metadata payload.
            """
            # Pair each saved path with its metadata + the file_identity used
            # to build the idempotency payload.
            inputs = list(zip(saved_paths, metadatas + [{}] * (len(saved_paths) - len(metadatas)), file_identities))
            # ``exclude_failed=True`` keeps FAILED/CANCELED docs OUT of the
            # reparse pool: ``_assert_action_allowed(..., 'upload')`` already
            # accepts those states, and the upload path applies the caller's
            # fresh ``metadatas``. Reparse instead reloads metadata from the
            # existing doc row, so retrying a failed upload with new tags
            # via reparse would silently lose those tags.
            existing_by_path = (
                self._manager._list_kb_docs_by_path(kb_id, exclude_failed=True)
                if override else {}
            )
            new_inputs = [(p, m, fi) for p, m, fi in inputs if p not in existing_by_path]
            reparse_inputs = [(p, m, fi, existing_by_path[p]) for p, m, fi in inputs if p in existing_by_path]

            result_by_path = {}
            # Run reparse FIRST so its validation (``_prepare_reparse_items``
            # → ``_assert_action_allowed``) raises before any new-file upload
            # has been enqueued. If we ran upload first and then reparse
            # raised on a WORKING/DELETING existing doc, the new-file uploads
            # would already be committed -- a partial-commit failure mode.
            if reparse_inputs:
                # Apply the caller's metadata directly to the documents row
                # so the reparse worker (which reloads ``doc.meta`` in
                # ``_prepare_reparse_items``) picks up the new values.
                # We deliberately do NOT call ``DocManager.patch_metadata``
                # here: it enqueues a separate DOC_UPDATE_META task that
                # races ``reparse`` -- on a fast parser the metadata task's
                # START callback flips the doc to WORKING and the subsequent
                # ``reparse()`` 409s via ``_assert_action_allowed``, and
                # even when the race doesn't fire the metadata task is
                # orphaned in WAITING because ``reparse`` overwrites the
                # snapshot's ``current_task_id``. A direct row update
                # avoids both pitfalls while preserving the legacy
                # "overwrite refreshes tags" semantic.
                self._legacy_apply_metadata_to_existing_docs(
                    [(eid, m) for _, m, _, eid in reparse_inputs if m]
                )
                reparse_request = ReparseRequest(
                    doc_ids=[eid for _, _, _, eid in reparse_inputs],
                    kb_id=kb_id, algo_id=algo_id,
                )
                task_ids = self._manager.reparse(reparse_request)
                for (p, _, _, eid), task_id in zip(reparse_inputs, task_ids):
                    result_by_path[p] = {
                        'doc_id': eid, 'task_id': task_id,
                        'accepted': True, 'error_code': None,
                    }

            if new_inputs:
                upload_request = UploadRequest(
                    items=[AddFileItem(file_path=p, metadata=m) for p, m, _ in new_inputs],
                    kb_id=kb_id, algo_id=algo_id, source_type=SourceType.API,
                )
                upload_payload = self._build_upload_payload(
                    upload_request, [fi for _, _, fi in new_inputs],
                )
                upload_result = self._manager.run_idempotent(
                    '/v1/docs/upload', upload_request.idempotency_key, upload_payload,
                    lambda: self._manager.upload(upload_request),
                )
                if isinstance(upload_result, dict) and 'items' in upload_result:
                    upload_items = upload_result['items']
                else:
                    upload_items = upload_result or []
                for (p, _, _), item in zip(new_inputs, upload_items):
                    result_by_path[p] = item

            ordered = [result_by_path.get(p, {'doc_id': None, 'accepted': False,
                                              'error_code': 'MISSING'})
                       for p in saved_paths]
            doc_ids = [it.get('doc_id') for it in ordered]
            if response_shape == 'ids_and_results':
                # Propagate per-item status from DocManager.upload so callers
                # of /upload_files see synchronous failures (e.g. parser outage
                # -> accepted=False / error_code=PARSER_SUBMIT_FAILED) instead
                # of a misleading 'ok' for every file.
                results = [
                    'ok' if it.get('accepted', True) else (
                        it.get('error_code') or it.get('error_msg') or 'failed'
                    )
                    for it in ordered
                ]
                return [doc_ids, results]
            return doc_ids

        @app.post('/upload_files')
        async def upload_files_legacy(
            self,
            files: List['fastapi.UploadFile'] = fastapi.File(...),  # noqa: B008
            # Match legacy DocManager default: ``False`` keeps the unique-path
            # fallback so a re-upload without ``?override=true`` does not
            # silently replace an existing document.
            override: bool = fastapi.Query(False),  # noqa: B008
            metadatas: Optional[str] = fastapi.Query(None),  # noqa: B008
            group_name: Optional[str] = fastapi.Query(None),  # noqa: B008
            user_path: Optional[str] = fastapi.Query(None),  # noqa: B008
        ):
            return await self._legacy_upload(
                files, override, metadatas, group_name, user_path,
                response_shape='ids_and_results',
            )

        @app.post('/add_files_to_group')
        async def add_files_to_group_legacy(
            self,
            files: List['fastapi.UploadFile'] = fastapi.File(...),  # noqa: B008
            group_name: str = fastapi.Query(...),  # noqa: B008
            # Same legacy default as /upload_files; explicit opt-in required
            # to overwrite.
            override: bool = fastapi.Query(False),  # noqa: B008
            metadatas: Optional[str] = fastapi.Query(None),  # noqa: B008
            user_path: Optional[str] = fastapi.Query(None),  # noqa: B008
        ):
            return await self._legacy_upload(
                files, override, metadatas, group_name, user_path,
                response_shape='ids_only',
            )

        @app.post('/v1/docs/upload')
        async def upload(
            self,
            files: List['fastapi.UploadFile'] = fastapi.File(...),  # noqa: B008
            kb_id: Optional[str] = fastapi.Form(None),  # noqa: B008
            source_type: Optional[SourceType] = fastapi.Form(None),  # noqa: B008
            doc_id: Optional[str] = fastapi.Form(None),  # noqa: B008
            idempotency_key: Optional[str] = fastapi.Form(None),  # noqa: B008
        ):
            self._lazy_init()
            if not files:
                raise fastapi.HTTPException(status_code=400, detail='files is required')
            kb_id = kb_id or '__default__'
            source_type = source_type or SourceType.API
            saved_paths, file_identities = await self._persist_uploads(files)
            upload_request = UploadRequest(
                items=[
                    AddFileItem(file_path=path, doc_id=(doc_id if idx == 0 else None))
                    for idx, path in enumerate(saved_paths)
                ],
                kb_id=kb_id,
                source_type=source_type,
                idempotency_key=idempotency_key,
            )
            if file_identities:
                file_identities[0]['doc_id'] = doc_id
            return self._run_upload(upload_request, self._build_upload_payload(upload_request, file_identities))

        @app.post('/v1/docs/add')
        def add(self, request: AddRequest):
            self._lazy_init()
            payload = request.model_dump(mode='json')
            return self._run(lambda: self._manager.run_idempotent(
                '/v1/docs/add', request.idempotency_key, payload, lambda: {'items': self._manager.add_files(request)}
            ))

        @app.post('/v1/docs/reparse')
        def reparse(self, request: ReparseRequest):
            self._lazy_init()
            payload = request.model_dump(mode='json')
            return self._run(lambda: self._manager.run_idempotent(
                '/v1/docs/reparse', request.idempotency_key, payload,
                lambda: {'task_ids': self._manager.reparse(request)}
            ))

        @app.post('/v1/docs/delete')
        def delete(self, request: DeleteRequest):
            self._lazy_init()
            payload = request.model_dump(mode='json')
            return self._run(lambda: self._manager.run_idempotent(
                '/v1/docs/delete', request.idempotency_key, payload, lambda: {'items': self._manager.delete(request)}
            ))

        @app.post('/v1/docs/transfer')
        def transfer(self, request: TransferRequest):
            self._lazy_init()
            payload = request.model_dump(mode='json')
            return self._run(lambda: self._manager.run_idempotent(
                '/v1/docs/transfer', request.idempotency_key, payload,
                lambda: {'items': self._manager.transfer(request)}
            ))

        @app.get('/v1/docs')
        def list_docs(
            self,
            status: Optional[List[str]] = None,
            kb_id: Optional[str] = None,
            keyword: Optional[str] = None,
            include_deleted_or_canceled: bool = True,
            page: int = 1,
            page_size: int = 20,
        ):
            self._lazy_init()
            return self._run(lambda: self._manager.list_docs(
                status=status,
                kb_id=kb_id,
                keyword=keyword,
                include_deleted_or_canceled=include_deleted_or_canceled,
                page=page,
                page_size=page_size,
            ))

        @app.get('/v1/docs/node_groups')
        def list_doc_node_groups(self, kb_id: str, doc_id: str):
            self._lazy_init()
            return self._run(lambda: self._manager.list_doc_node_groups(kb_id=kb_id, doc_id=doc_id))

        @app.get('/v1/docs/{doc_id}/ng-status')
        def get_doc_ng_status(self, doc_id: str, kb_id: str):
            self._lazy_init()
            return self._run(lambda: self._manager.get_doc_ng_status(kb_id=kb_id, doc_id=doc_id))

        @app.get('/v1/docs/{doc_id}')
        def get_doc(self, doc_id: str):
            self._lazy_init()
            return self._run(lambda: self._manager.get_doc_detail(doc_id))

        @app.post('/v1/docs/metadata/patch')
        def patch_metadata(self, request: MetadataPatchRequest):
            self._lazy_init()
            payload = request.model_dump(mode='json')
            return self._run(lambda: self._manager.run_idempotent(
                '/v1/docs/metadata/patch', request.idempotency_key, payload,
                lambda: self._manager.patch_metadata(request)
            ))

        @app.get('/v1/tasks')
        def list_tasks(self, status: Optional[List[str]] = None, page: int = 1, page_size: int = 20):
            self._lazy_init()
            resp = self._manager.list_tasks(status, page, page_size)
            return self._response(
                data=self._format_task_response_data(resp.data),
                code=resp.code,
                msg=resp.msg,
                status_code=resp.code,
            )

        @app.get('/v1/tasks/{task_id}')
        def get_task(self, task_id: str):
            self._lazy_init()
            resp = self._manager.get_task(task_id)
            return self._response(
                data=self._format_task_response_data(resp.data),
                code=resp.code,
                msg=resp.msg,
                status_code=resp.code,
            )

        @app.post('/v1/tasks/cancel')
        def cancel_task(self, request: TaskCancelRequest):
            self._lazy_init()
            payload = request.model_dump(mode='json')

            def _cancel():
                resp = self._manager.cancel_task(request.task_id)
                if resp.code == 404:
                    raise DocServiceError('E_NOT_FOUND', resp.msg, resp.data)
                if resp.code == 409:
                    raise DocServiceError('E_STATE_CONFLICT', resp.msg, resp.data)
                if resp.code != 200:
                    raise DocServiceError('E_INVALID_PARAM', resp.msg, resp.data)
                return resp.data
            return self._run(lambda: self._manager.run_idempotent(
                '/v1/tasks/cancel', request.idempotency_key, payload, _cancel
            ))

        def task_callback(self, callback: Any):
            self._lazy_init()
            return self._run(lambda: self._manager.on_task_callback(self._normalize_task_callback(callback)))

        @app.post('/v1/internal/callbacks/tasks')
        def task_callback_http(self, request: TaskCallbackPayload):
            return self.task_callback(request.model_dump(mode='json', exclude_none=True))

        @app.get('/v1/algo/{algo_id}/groups')
        def get_algo_groups(self, algo_id: str):
            self._lazy_init()
            return self._run(lambda: self._manager.get_algo_groups(algo_id))

        @app.get('/v1/algo/list')
        def list_algorithms(self):
            self._lazy_init()
            return self._run(lambda: {'items': self._manager.list_algorithms()})

        @app.post('/v1/algorithms/info')
        def get_algorithm_info(self, request: AlgorithmInfoRequest):
            self._lazy_init()
            return self._run(lambda: self._manager.get_algorithm_info(request.algo_id))

        @app.get('/v1/chunks')
        def list_chunks(
            self,
            kb_id: str,
            doc_id: str,
            group: str,
            algo_id: Optional[str] = None,
            page: int = 1,
            page_size: int = 20,
            offset: Optional[int] = None,
        ):
            self._lazy_init()
            return self._run(lambda: self._manager.list_chunks(
                kb_id=kb_id, doc_id=doc_id, group=group, algo_id=algo_id,
                page=page, page_size=page_size, offset=offset,
            ))

        @app.post('/v1/tasks/batch')
        def get_tasks_batch(self, request: TaskBatchRequest):
            self._lazy_init()
            return self._run(lambda: self._manager.get_tasks_batch(request.task_ids))

        @app.get('/v1/kbs')
        def list_kbs(
            self,
            page: int = 1,
            page_size: int = 20,
            keyword: Optional[str] = None,
            status: Optional[List[str]] = None,
            owner_id: Optional[str] = None,
        ):
            self._lazy_init()
            return self._run(lambda: self._manager.list_kbs(
                page=page,
                page_size=page_size,
                keyword=keyword,
                status=status,
                owner_id=owner_id,
            ))

        @app.get('/v1/kbs/{kb_id}')
        def get_kb(self, kb_id: str):
            self._lazy_init()
            return self._run(lambda: self._manager.get_kb(kb_id))

        @app.post('/v1/kbs')
        def create_kb(self, request: KbCreateRequest):
            self._lazy_init()
            if not request.kb_id:
                raise DocServiceError('E_INVALID_PARAM', 'kb_id is required')
            payload = request.model_dump(mode='json')
            return self._run(lambda: self._manager.run_idempotent(
                '/v1/kbs', request.idempotency_key, payload,
                lambda: self._manager.create_kb(
                    request.kb_id,
                    display_name=request.display_name,
                    description=request.description,
                    owner_id=request.owner_id,
                    meta=request.meta,
                    algo_id=request.algo_id,
                )
            ))

        @app.post('/v1/kbs/{kb_id}')
        def update_kb(self, kb_id: str, request: KbUpdateRequest):
            self._lazy_init()
            if request.kb_id and request.kb_id != kb_id:
                raise DocServiceError(
                    'E_INVALID_PARAM',
                    f'kb_id mismatch: path={kb_id}, body={request.kb_id}',
                    {'kb_id': kb_id, 'request_kb_id': request.kb_id},
                )
            payload = self._build_update_kb_payload(kb_id, request)
            return self._run(lambda: self._manager.run_idempotent(
                f'/v1/kbs/{kb_id}:patch', request.idempotency_key, payload,
                lambda: self._manager.update_kb(
                    kb_id,
                    display_name=request.display_name,
                    description=request.description,
                    owner_id=request.owner_id,
                    meta=request.meta,
                    algo_id=request.algo_id,
                    explicit_fields=set(request.model_fields_set),
                )
            ))

        @app.delete('/v1/kbs/{kb_id}/algos/{algo_id}')
        def unbind_algo(self, kb_id: str, algo_id: str, dry_run: bool = False):
            self._lazy_init()
            return self._run(lambda: self._manager.unbind_algo(kb_id, algo_id, dry_run=dry_run))

        @app.delete('/v1/kbs/{kb_id}')
        def delete_kb(self, kb_id: str, idempotency_key: Optional[str] = None):
            self._lazy_init()
            payload = {'kb_id': kb_id}
            return self._run(lambda: self._manager.run_idempotent(
                '/v1/kbs/{kb_id}:delete', idempotency_key, payload, lambda: self._manager.delete_kb(kb_id)
            ))

        @app.delete('/v1/kbs')
        def delete_kbs(self, request: KbDeleteBatchRequest):
            self._lazy_init()
            payload = request.model_dump(mode='json')
            return self._run(lambda: self._manager.run_idempotent(
                '/v1/kbs:delete', request.idempotency_key, payload, lambda: self._manager.delete_kbs(request.kb_ids)
            ))

        @app.post('/v1/ng/{group_name}/lazy_mode')
        def set_node_group_lazy_mode(self, group_name: str,
                                     lazy_mode: Optional[str] = None):
            self._lazy_init()
            return self._run(lambda: self._manager.set_node_group_lazy_mode(group_name, lazy_mode))

        @app.get('/v1/live')
        def live(self):
            return BaseResponse(code=200, msg='success', data={'status': 'ok', 'version': 'v1'})

        @app.get('/v1/ready')
        def ready(self):
            # Do not enter the once-only initializer until the parser is ready.
            # A failed once_wrapper call caches its exception, so probing first is
            # required for dependency-tolerant parallel startup.
            try:
                ParserClient(parser_url=self._parser_url).health()
            except Exception:
                return fastapi.responses.JSONResponse(
                    status_code=503,
                    content=BaseResponse(
                        code=503,
                        msg='parser is not ready',
                        data={'status': 'starting', 'version': 'v1', 'deps': {'sql': False, 'parser': False}},
                    ).model_dump(mode='json'),
                )
            try:
                self._lazy_init()
            except Exception as exc:
                # The dependency may disappear between the preflight and the
                # initializer.  Clear once_wrapper's cached exception so a later
                # readiness probe can retry safely.
                self._lazy_init.flag.reset()
                LOG.warning(f'[DocServer] readiness initialization failed: {exc}')
                return fastapi.responses.JSONResponse(
                    status_code=503,
                    content=BaseResponse(
                        code=503,
                        msg='service initialization is not ready',
                        data={'status': 'starting', 'version': 'v1'},
                    ).model_dump(mode='json'),
                )
            health = self._manager.health()
            if not health.get('deps', {}).get('parser'):
                return fastapi.responses.JSONResponse(
                    status_code=503,
                    content=BaseResponse(code=503, msg='parser is not ready', data=health).model_dump(mode='json'),
                )
            return BaseResponse(code=200, msg='success', data=health)

        @app.get('/v1/health')
        def health(self):
            return self.ready()

        @app.get('/v1/internal/parser-url')
        def get_parser_url(self):
            self._lazy_init()
            return BaseResponse(code=200, msg='success', data={'parser_url': self._parser_url})

        def __call__(self, func_name: str, *args, **kwargs):
            return getattr(self, func_name)(*args, **kwargs)

    def __init__(
        self,
        port: Optional[int] = None,
        url: Optional[str] = None,
        parser_url: Optional[str] = None,
        db_config: Optional[Dict[str, Any]] = None,
        parser_db_config: Optional[Dict[str, Any]] = None,
        parser_poll_interval: float = 0.05,
        storage_dir: Optional[str] = None,
        callback_url: Optional[str] = None,
        pythonpath: Optional[str] = None,
        launcher=None,
        enable_scan: bool = False,
        scan_interval: int = 10,
    ):
        super().__init__()
        self._raw_impl = None
        self._storage_dir = storage_dir or os.path.join(os.getcwd(), '.doc_service_uploads')
        self._db_config = db_config or _get_default_db_config('doc_service')
        self._parser_db_config = parser_db_config or _get_default_db_config('doc_service_parser')
        if url:
            self._impl = UrlModule(url=ensure_call_endpoint(url))
        else:
            if not parser_url:
                raise ValueError('parser_url is required; doc_service no longer embeds a mock parsing server')
            self._raw_impl = DocServer._Impl(
                storage_dir=self._storage_dir,
                db_config=self._db_config,
                parser_db_config=self._parser_db_config,
                parser_poll_interval=parser_poll_interval,
                parser_url=parser_url,
                callback_url=callback_url,
                enable_scan=enable_scan,
                scan_interval=scan_interval,
            )
            # DocServer is a lightweight HTTP front-end for doc CRUD; never needs
            # GPU. Default to EmptyLauncher when the caller did not pass one so we
            # don't inherit LAZYLLM_DEFAULT_LAUNCHER (e.g. 'sco' in CI) and try
            # to submit srun jobs for what should be a local Python subprocess.
            import lazyllm as _lazyllm
            effective_launcher = launcher if launcher is not None else _lazyllm.launchers.empty(sync=False)
            self._impl = ServerModule(
                self._raw_impl, port=port, launcher=effective_launcher, pythonpath=pythonpath,
            )

    @staticmethod
    def _register_openapi_routes(openapi_app: 'fastapi.FastAPI', impl: 'DocServer._Impl'):
        def _find_services(cls):
            if '__relay_services__' not in dir(cls):
                return
            if '__relay_services__' in cls.__dict__:
                for (method, path), (name, kw) in cls.__relay_services__.items():
                    if getattr(impl.__class__, name) is getattr(cls, name):
                        route_method = getattr(openapi_app, 'get' if method == 'list' else method)
                        route_method(path, **kw)(getattr(impl, name))
            for base in cls.__bases__:
                _find_services(base)

        app.update()
        _find_services(impl.__class__)

    @classmethod
    def build_openapi_app(cls, title: str = 'LazyLLM DocService API', version: str = '1.0.0'):
        """构建用于导出 OpenAPI 的 FastAPI 应用对象。"""
        openapi_app = fastapi.FastAPI(
            title=title,
            version=version,
            description='OpenAPI schema generated from current DocServer routes.',
        )
        impl = cls._Impl(
            storage_dir=os.path.join(os.getcwd(), '.doc_service_openapi'),
            parser_url='http://127.0.0.1:9966',
        )
        cls._register_openapi_routes(openapi_app, impl)
        app._prepare_openapi(openapi_app)
        return openapi_app

    @classmethod
    def build_openapi_schema(cls, title: str = 'LazyLLM DocService API', version: str = '1.0.0'):
        """生成 doc service 的 OpenAPI schema。"""
        return cls.build_openapi_app(title=title, version=version).openapi()

    @classmethod
    def export_openapi(
        cls,
        output_path: str = DEFAULT_OPENAPI_OUTPUT_PATH,
        title: str = 'LazyLLM DocService API',
        version: str = '1.0.0',
    ):
        """将 doc service 的 OpenAPI schema 导出到文件。"""
        schema = cls.build_openapi_schema(title=title, version=version)
        output_dir = os.path.dirname(output_path)
        if output_dir:
            os.makedirs(output_dir, exist_ok=True)
        with open(output_path, 'w', encoding='utf-8') as fh:
            json.dump(schema, fh, ensure_ascii=False, indent=2, sort_keys=True)
        return output_path

    def start(self):
        result = super().start()
        if self._raw_impl and isinstance(self._impl, ServerModule):
            try:
                callback_url = self._impl._url.rsplit('/', 1)[0] + '/v1/internal/callbacks/tasks'
                self._dispatch('set_runtime_callback_url', callback_url)
            except Exception as exc:
                LOG.warning(f'[DocServer] failed to set runtime callback url: {exc}')
        return result

    def stop(self):
        if self._raw_impl:
            try:
                self._dispatch('stop')
            except Exception as exc:
                LOG.warning(f'[DocServer] stop impl failed: {exc}, {traceback.format_exc()}')
        if isinstance(self._impl, ServerModule):
            self._impl.stop()

    @property
    def url(self):
        return self._impl._url

    @property
    def _url(self):
        return self.url

    @property
    def parser_url(self):
        if self._raw_impl:
            return self._raw_impl._parser_url
        base_url = self.url.rsplit('/', 1)[0]
        try:
            response = requests.get(f'{base_url}/v1/internal/parser-url', timeout=5)
            response.raise_for_status()
            return response.json()['data']['parser_url']
        except (requests.RequestException, KeyError, TypeError, ValueError) as exc:
            LOG.warning(f'[DocServer] failed to resolve remote parser_url from {base_url}: {exc}')
            return None

    @staticmethod
    def _normalize_dispatch_result(result):
        if isinstance(result, fastapi.responses.JSONResponse):
            return json.loads(result.body.decode())
        return result

    def _dispatch(self, method: str, *args, **kwargs):
        if isinstance(self._impl, ServerModule):
            return self._normalize_dispatch_result(self._impl._call(method, *args, **kwargs))
        return self._normalize_dispatch_result(getattr(self._impl, method)(*args, **kwargs))

    # Method-call style wrappers
    def upload(self, request: UploadRequest):
        """通过 ``/v1/docs/upload`` 流程将文件上传到 DocServer 管理的存储目录。

当你希望由 DocServer 保存上传副本时,使用该方法。请求体为 ``UploadRequest``,包含 ``kb_id``
和 ``items``。每个 item 使用 ``file_path`` 作为本地源路径,也可以附带可选的 ``doc_id``、``metadata``。

**Returns:**
    标准 API 响应。``data["items"]`` 中包含接受后的 ``doc_id`` 和异步 ``task_id``。
"""
        return self._dispatch('upload_request', request)

    def add(self, request: AddRequest):
        """通过 ``/v1/docs/add`` 接口添加服务端可直接访问的本地文件。

当文件路径已经对 DocServer 所在机器可见时,使用该方法。请求体为 ``AddRequest``,包含 ``kb_id``
和 ``items``。每个 item 可提供 ``file_path``,以及可选的 ``doc_id``、``metadata``。

**Returns:**
    标准 API 响应。``data["items"]`` 中包含接受后的 ``doc_id`` 和异步 ``task_id``。
"""
        return self._dispatch('add', request)

    def reparse(self, request: ReparseRequest):
        """通过 ``/v1/docs/reparse`` 接口重新解析已有文档。

请求体为 ``ReparseRequest``,包含 ``kb_id`` 和 ``doc_ids``。当元数据或解析配置变更后,
需要为已有文档重新入队解析任务时,可使用该方法。

可通过 ``algo_id`` 指定重解析该算法下的所有节点组,或通过 ``reparse_group``(节点组名称)
指定仅重解析某一个节点组。两个字段互斥,同时传入会触发校验错误。若两者均不传,则使用知识库绑定的
第一个算法,并重解析其所有节点组。
"""
        return self._dispatch('reparse', request)

    def delete(self, request: DeleteRequest):
        """通过 ``/v1/docs/delete`` 接口从知识库中删除文档。

请求体为 ``DeleteRequest``,包含 ``kb_id`` 和 ``doc_ids``。删除是异步操作,因此如果需要最终状态,
应继续通过任务接口跟踪返回的 ``task_id``。

知识库下绑定的所有算法均会被自动处理,无需指定 ``algo_id``。若任意算法的解析任务处于 WORKING
状态,请求会被拒绝并返回 ``E_STATE_CONFLICT``;处于 WAITING 状态的添加任务会在删除前被自动取消。
"""
        return self._dispatch('delete', request)

    def transfer(self, request: TransferRequest):
        """在同一算法下的不同知识库之间转移已解析文档。

请求体为 ``TransferRequest``。每个转移项都必须在目标知识库中提供唯一的 ``target_doc_id``。
当前不支持跨算法 transfer。可选字段 ``target_filename`` 与 ``target_file_path`` 用于覆盖目标文档记录的文件名或文件路径。
"""
        return self._dispatch('transfer', request)

    def patch_metadata(self, request: MetadataPatchRequest):
        """通过 ``/v1/docs/metadata/patch`` 接口更新文档元数据。

请求体为 ``MetadataPatchRequest``,包含 ``kb_id``、``algo_id`` 和 ``items``。每个 item 指向一个文档,
并在 ``patch`` 中携带需要合并的局部元数据。
"""
        return self._dispatch('patch_metadata', request)

    def list_docs(self, **kwargs):
        """分页列出知识库中的文档。"""
        return self._dispatch('list_docs', **kwargs)

    def get_doc(self, doc_id: str):
        """获取单个文档的详细信息。"""
        return self._dispatch('get_doc', doc_id)

    def list_tasks(self, **kwargs):
        """分页列出任务记录。"""
        return self._dispatch('list_tasks', **kwargs)

    def get_tasks_batch(self, task_ids: List[str]):
        """批量获取多个任务记录。"""
        return self._dispatch('get_tasks_batch', TaskBatchRequest(task_ids=task_ids))

    def get_task(self, task_id: str):
        """获取单个任务记录。"""
        return self._dispatch('get_task', task_id)

    def set_runtime_callback_url(self, callback_url: str):
        """更新运行时任务回调 URL。"""
        return self._dispatch('set_runtime_callback_url', callback_url)

    def cancel_task(self, task_id: str):
        """通过 ``/v1/tasks/cancel`` 接口取消一个处于等待中的任务。

Args:
    task_id (str): 要取消的任务 ID。

**Returns:**
    表示任务是否取消成功的标准 API 响应。
"""
        return self._dispatch('cancel_task', TaskCancelRequest(task_id=task_id))

    def list_kbs(self, **kwargs):
        """分页列出知识库。"""
        return self._dispatch('list_kbs', **kwargs)

    def get_kb(self, kb_id: str):
        """获取单个知识库的信息。"""
        return self._dispatch('get_kb', kb_id)

    def list_chunks(self, **kwargs):
        """通过 ``/v1/chunks`` 接口分页查看文档的解析 chunk。

Args:
    kb_id (str): 知识库 ID。
    doc_id (str): 文档 ID。
    group (str): 要查看的节点组名。
    algo_id (str): 算法 ID。
    page (int): 从 1 开始的页码。
    page_size (int): 每页 chunk 数量。
    offset (Optional[int]): 显式偏移量;未传时服务端会根据 ``page`` 和 ``page_size`` 推导。

Returns:
    包含 ``items`` 与 ``total`` 的分页结果。
"""
        return self._dispatch('list_chunks', **kwargs)

    def list_doc_node_groups(self, kb_id: str, doc_id: str):
        """列出指定文档已落库的节点组名称。"""
        return self._dispatch('list_doc_node_groups', kb_id=kb_id, doc_id=doc_id)

    def get_doc_ng_status(self, kb_id: str, doc_id: str):
        """查询指定文档的节点组解析状态。"""
        return self._dispatch('get_doc_ng_status', kb_id=kb_id, doc_id=doc_id)

    def list_algorithms(self):
        """列出可用算法。"""
        return self._dispatch('list_algorithms')

    def get_algorithm_info(self, algo_id: str):
        """获取指定算法的详细信息。"""
        return self._dispatch('get_algorithm_info', AlgorithmInfoRequest(algo_id=algo_id))

    def create_kb(self, kb_id: str, display_name: Optional[str] = None, description: Optional[str] = None,
                  owner_id: Optional[str] = None, meta: Optional[Dict[str, Any]] = None,
                  algo_id: str = '__default__'):
        """创建新的知识库。"""
        return self._dispatch('create_kb', KbCreateRequest(
            kb_id=kb_id, display_name=display_name, description=description,
            owner_id=owner_id, meta=meta, algo_id=algo_id,
        ))

    def update_kb(self, kb_id: str, request: KbUpdateRequest):
        """更新知识库的元信息。"""
        return self._dispatch('update_kb', kb_id, request)

    def delete_kb(self, kb_id: str):
        """删除一个知识库。"""
        return self._dispatch('delete_kb', kb_id)

    def delete_kbs(self, kb_ids: List[str]):
        """批量删除多个知识库。"""
        return self._dispatch('delete_kbs', KbDeleteBatchRequest(kb_ids=kb_ids))

    def unbind_algo(self, kb_id: str, algo_id: str):
        """从知识库解绑一个算法,并异步清理该算法独有节点组的解析数据;与其他算法共享的节点组数据保留。"""
        return self._dispatch('unbind_algo', kb_id, algo_id)

    def ensure_kb_registered(self, kb_id: str, algo_id: Optional[str] = None):
        """Ensure the knowledge base row and algorithm binding exist in the doc service."""
        return self._dispatch('ensure_kb_registered', kb_id, algo_id)

    def enable_scanning(self):
        """Trigger dataset scanning for a local doc service after registrations are ready."""
        return self._dispatch('enable_scanning')

    def set_node_group_lazy_mode(self, group_name: str, lazy_mode: Optional[str] = None):
        """设置指定节点组的懒加载模式。"""
        return self._dispatch('set_node_group_lazy_mode', group_name, lazy_mode)

add(request)

通过 /v1/docs/add 接口添加服务端可直接访问的本地文件。

当文件路径已经对 DocServer 所在机器可见时,使用该方法。请求体为 AddRequest,包含 kb_iditems。每个 item 可提供 file_path,以及可选的 doc_idmetadata

Returns: 标准 API 响应。data["items"] 中包含接受后的 doc_id 和异步 task_id

Source code in lazyllm/tools/rag/doc_service/doc_server.py
    def add(self, request: AddRequest):
        """通过 ``/v1/docs/add`` 接口添加服务端可直接访问的本地文件。

当文件路径已经对 DocServer 所在机器可见时,使用该方法。请求体为 ``AddRequest``,包含 ``kb_id``
和 ``items``。每个 item 可提供 ``file_path``,以及可选的 ``doc_id``、``metadata``。

**Returns:**
    标准 API 响应。``data["items"]`` 中包含接受后的 ``doc_id`` 和异步 ``task_id``。
"""
        return self._dispatch('add', request)

build_openapi_app(title='LazyLLM DocService API', version='1.0.0') classmethod

构建用于导出 OpenAPI 的 FastAPI 应用对象。

Source code in lazyllm/tools/rag/doc_service/doc_server.py
@classmethod
def build_openapi_app(cls, title: str = 'LazyLLM DocService API', version: str = '1.0.0'):
    """构建用于导出 OpenAPI 的 FastAPI 应用对象。"""
    openapi_app = fastapi.FastAPI(
        title=title,
        version=version,
        description='OpenAPI schema generated from current DocServer routes.',
    )
    impl = cls._Impl(
        storage_dir=os.path.join(os.getcwd(), '.doc_service_openapi'),
        parser_url='http://127.0.0.1:9966',
    )
    cls._register_openapi_routes(openapi_app, impl)
    app._prepare_openapi(openapi_app)
    return openapi_app

build_openapi_schema(title='LazyLLM DocService API', version='1.0.0') classmethod

生成 doc service 的 OpenAPI schema。

Source code in lazyllm/tools/rag/doc_service/doc_server.py
@classmethod
def build_openapi_schema(cls, title: str = 'LazyLLM DocService API', version: str = '1.0.0'):
    """生成 doc service 的 OpenAPI schema。"""
    return cls.build_openapi_app(title=title, version=version).openapi()

cancel_task(task_id)

通过 /v1/tasks/cancel 接口取消一个处于等待中的任务。

Parameters:

  • task_id (str) –

    要取消的任务 ID。

Returns: 表示任务是否取消成功的标准 API 响应。

Source code in lazyllm/tools/rag/doc_service/doc_server.py
    def cancel_task(self, task_id: str):
        """通过 ``/v1/tasks/cancel`` 接口取消一个处于等待中的任务。

Args:
    task_id (str): 要取消的任务 ID。

**Returns:**
    表示任务是否取消成功的标准 API 响应。
"""
        return self._dispatch('cancel_task', TaskCancelRequest(task_id=task_id))

create_kb(kb_id, display_name=None, description=None, owner_id=None, meta=None, algo_id='__default__')

创建新的知识库。

Source code in lazyllm/tools/rag/doc_service/doc_server.py
def create_kb(self, kb_id: str, display_name: Optional[str] = None, description: Optional[str] = None,
              owner_id: Optional[str] = None, meta: Optional[Dict[str, Any]] = None,
              algo_id: str = '__default__'):
    """创建新的知识库。"""
    return self._dispatch('create_kb', KbCreateRequest(
        kb_id=kb_id, display_name=display_name, description=description,
        owner_id=owner_id, meta=meta, algo_id=algo_id,
    ))

delete(request)

通过 /v1/docs/delete 接口从知识库中删除文档。

请求体为 DeleteRequest,包含 kb_iddoc_ids。删除是异步操作,因此如果需要最终状态, 应继续通过任务接口跟踪返回的 task_id

知识库下绑定的所有算法均会被自动处理,无需指定 algo_id。若任意算法的解析任务处于 WORKING 状态,请求会被拒绝并返回 E_STATE_CONFLICT;处于 WAITING 状态的添加任务会在删除前被自动取消。

Source code in lazyllm/tools/rag/doc_service/doc_server.py
    def delete(self, request: DeleteRequest):
        """通过 ``/v1/docs/delete`` 接口从知识库中删除文档。

请求体为 ``DeleteRequest``,包含 ``kb_id`` 和 ``doc_ids``。删除是异步操作,因此如果需要最终状态,
应继续通过任务接口跟踪返回的 ``task_id``。

知识库下绑定的所有算法均会被自动处理,无需指定 ``algo_id``。若任意算法的解析任务处于 WORKING
状态,请求会被拒绝并返回 ``E_STATE_CONFLICT``;处于 WAITING 状态的添加任务会在删除前被自动取消。
"""
        return self._dispatch('delete', request)

delete_kb(kb_id)

删除一个知识库。

Source code in lazyllm/tools/rag/doc_service/doc_server.py
def delete_kb(self, kb_id: str):
    """删除一个知识库。"""
    return self._dispatch('delete_kb', kb_id)

delete_kbs(kb_ids)

批量删除多个知识库。

Source code in lazyllm/tools/rag/doc_service/doc_server.py
def delete_kbs(self, kb_ids: List[str]):
    """批量删除多个知识库。"""
    return self._dispatch('delete_kbs', KbDeleteBatchRequest(kb_ids=kb_ids))

enable_scanning()

Trigger dataset scanning for a local doc service after registrations are ready.

Source code in lazyllm/tools/rag/doc_service/doc_server.py
def enable_scanning(self):
    """Trigger dataset scanning for a local doc service after registrations are ready."""
    return self._dispatch('enable_scanning')

ensure_kb_registered(kb_id, algo_id=None)

Ensure the knowledge base row and algorithm binding exist in the doc service.

Source code in lazyllm/tools/rag/doc_service/doc_server.py
def ensure_kb_registered(self, kb_id: str, algo_id: Optional[str] = None):
    """Ensure the knowledge base row and algorithm binding exist in the doc service."""
    return self._dispatch('ensure_kb_registered', kb_id, algo_id)

export_openapi(output_path=DEFAULT_OPENAPI_OUTPUT_PATH, title='LazyLLM DocService API', version='1.0.0') classmethod

将 doc service 的 OpenAPI schema 导出到文件。

Source code in lazyllm/tools/rag/doc_service/doc_server.py
@classmethod
def export_openapi(
    cls,
    output_path: str = DEFAULT_OPENAPI_OUTPUT_PATH,
    title: str = 'LazyLLM DocService API',
    version: str = '1.0.0',
):
    """将 doc service 的 OpenAPI schema 导出到文件。"""
    schema = cls.build_openapi_schema(title=title, version=version)
    output_dir = os.path.dirname(output_path)
    if output_dir:
        os.makedirs(output_dir, exist_ok=True)
    with open(output_path, 'w', encoding='utf-8') as fh:
        json.dump(schema, fh, ensure_ascii=False, indent=2, sort_keys=True)
    return output_path

get_algorithm_info(algo_id)

获取指定算法的详细信息。

Source code in lazyllm/tools/rag/doc_service/doc_server.py
def get_algorithm_info(self, algo_id: str):
    """获取指定算法的详细信息。"""
    return self._dispatch('get_algorithm_info', AlgorithmInfoRequest(algo_id=algo_id))

get_doc(doc_id)

获取单个文档的详细信息。

Source code in lazyllm/tools/rag/doc_service/doc_server.py
def get_doc(self, doc_id: str):
    """获取单个文档的详细信息。"""
    return self._dispatch('get_doc', doc_id)

get_doc_ng_status(kb_id, doc_id)

查询指定文档的节点组解析状态。

Source code in lazyllm/tools/rag/doc_service/doc_server.py
def get_doc_ng_status(self, kb_id: str, doc_id: str):
    """查询指定文档的节点组解析状态。"""
    return self._dispatch('get_doc_ng_status', kb_id=kb_id, doc_id=doc_id)

get_kb(kb_id)

获取单个知识库的信息。

Source code in lazyllm/tools/rag/doc_service/doc_server.py
def get_kb(self, kb_id: str):
    """获取单个知识库的信息。"""
    return self._dispatch('get_kb', kb_id)

get_task(task_id)

获取单个任务记录。

Source code in lazyllm/tools/rag/doc_service/doc_server.py
def get_task(self, task_id: str):
    """获取单个任务记录。"""
    return self._dispatch('get_task', task_id)

get_tasks_batch(task_ids)

批量获取多个任务记录。

Source code in lazyllm/tools/rag/doc_service/doc_server.py
def get_tasks_batch(self, task_ids: List[str]):
    """批量获取多个任务记录。"""
    return self._dispatch('get_tasks_batch', TaskBatchRequest(task_ids=task_ids))

list_algorithms()

列出可用算法。

Source code in lazyllm/tools/rag/doc_service/doc_server.py
def list_algorithms(self):
    """列出可用算法。"""
    return self._dispatch('list_algorithms')

list_chunks(**kwargs)

通过 /v1/chunks 接口分页查看文档的解析 chunk。

Parameters:

  • kb_id (str) –

    知识库 ID。

  • doc_id (str) –

    文档 ID。

  • group (str) –

    要查看的节点组名。

  • algo_id (str) –

    算法 ID。

  • page (int) –

    从 1 开始的页码。

  • page_size (int) –

    每页 chunk 数量。

  • offset (Optional[int]) –

    显式偏移量;未传时服务端会根据 pagepage_size 推导。

Returns:

  • 包含 itemstotal 的分页结果。

Source code in lazyllm/tools/rag/doc_service/doc_server.py
    def list_chunks(self, **kwargs):
        """通过 ``/v1/chunks`` 接口分页查看文档的解析 chunk。

Args:
    kb_id (str): 知识库 ID。
    doc_id (str): 文档 ID。
    group (str): 要查看的节点组名。
    algo_id (str): 算法 ID。
    page (int): 从 1 开始的页码。
    page_size (int): 每页 chunk 数量。
    offset (Optional[int]): 显式偏移量;未传时服务端会根据 ``page`` 和 ``page_size`` 推导。

Returns:
    包含 ``items`` 与 ``total`` 的分页结果。
"""
        return self._dispatch('list_chunks', **kwargs)

list_doc_node_groups(kb_id, doc_id)

列出指定文档已落库的节点组名称。

Source code in lazyllm/tools/rag/doc_service/doc_server.py
def list_doc_node_groups(self, kb_id: str, doc_id: str):
    """列出指定文档已落库的节点组名称。"""
    return self._dispatch('list_doc_node_groups', kb_id=kb_id, doc_id=doc_id)

list_docs(**kwargs)

分页列出知识库中的文档。

Source code in lazyllm/tools/rag/doc_service/doc_server.py
def list_docs(self, **kwargs):
    """分页列出知识库中的文档。"""
    return self._dispatch('list_docs', **kwargs)

list_kbs(**kwargs)

分页列出知识库。

Source code in lazyllm/tools/rag/doc_service/doc_server.py
def list_kbs(self, **kwargs):
    """分页列出知识库。"""
    return self._dispatch('list_kbs', **kwargs)

list_tasks(**kwargs)

分页列出任务记录。

Source code in lazyllm/tools/rag/doc_service/doc_server.py
def list_tasks(self, **kwargs):
    """分页列出任务记录。"""
    return self._dispatch('list_tasks', **kwargs)

patch_metadata(request)

通过 /v1/docs/metadata/patch 接口更新文档元数据。

请求体为 MetadataPatchRequest,包含 kb_idalgo_iditems。每个 item 指向一个文档, 并在 patch 中携带需要合并的局部元数据。

Source code in lazyllm/tools/rag/doc_service/doc_server.py
    def patch_metadata(self, request: MetadataPatchRequest):
        """通过 ``/v1/docs/metadata/patch`` 接口更新文档元数据。

请求体为 ``MetadataPatchRequest``,包含 ``kb_id``、``algo_id`` 和 ``items``。每个 item 指向一个文档,
并在 ``patch`` 中携带需要合并的局部元数据。
"""
        return self._dispatch('patch_metadata', request)

reparse(request)

通过 /v1/docs/reparse 接口重新解析已有文档。

请求体为 ReparseRequest,包含 kb_iddoc_ids。当元数据或解析配置变更后, 需要为已有文档重新入队解析任务时,可使用该方法。

可通过 algo_id 指定重解析该算法下的所有节点组,或通过 reparse_group(节点组名称) 指定仅重解析某一个节点组。两个字段互斥,同时传入会触发校验错误。若两者均不传,则使用知识库绑定的 第一个算法,并重解析其所有节点组。

Source code in lazyllm/tools/rag/doc_service/doc_server.py
    def reparse(self, request: ReparseRequest):
        """通过 ``/v1/docs/reparse`` 接口重新解析已有文档。

请求体为 ``ReparseRequest``,包含 ``kb_id`` 和 ``doc_ids``。当元数据或解析配置变更后,
需要为已有文档重新入队解析任务时,可使用该方法。

可通过 ``algo_id`` 指定重解析该算法下的所有节点组,或通过 ``reparse_group``(节点组名称)
指定仅重解析某一个节点组。两个字段互斥,同时传入会触发校验错误。若两者均不传,则使用知识库绑定的
第一个算法,并重解析其所有节点组。
"""
        return self._dispatch('reparse', request)

set_node_group_lazy_mode(group_name, lazy_mode=None)

设置指定节点组的懒加载模式。

Source code in lazyllm/tools/rag/doc_service/doc_server.py
def set_node_group_lazy_mode(self, group_name: str, lazy_mode: Optional[str] = None):
    """设置指定节点组的懒加载模式。"""
    return self._dispatch('set_node_group_lazy_mode', group_name, lazy_mode)

set_runtime_callback_url(callback_url)

更新运行时任务回调 URL。

Source code in lazyllm/tools/rag/doc_service/doc_server.py
def set_runtime_callback_url(self, callback_url: str):
    """更新运行时任务回调 URL。"""
    return self._dispatch('set_runtime_callback_url', callback_url)

transfer(request)

在同一算法下的不同知识库之间转移已解析文档。

请求体为 TransferRequest。每个转移项都必须在目标知识库中提供唯一的 target_doc_id。 当前不支持跨算法 transfer。可选字段 target_filenametarget_file_path 用于覆盖目标文档记录的文件名或文件路径。

Source code in lazyllm/tools/rag/doc_service/doc_server.py
    def transfer(self, request: TransferRequest):
        """在同一算法下的不同知识库之间转移已解析文档。

请求体为 ``TransferRequest``。每个转移项都必须在目标知识库中提供唯一的 ``target_doc_id``。
当前不支持跨算法 transfer。可选字段 ``target_filename`` 与 ``target_file_path`` 用于覆盖目标文档记录的文件名或文件路径。
"""
        return self._dispatch('transfer', request)

unbind_algo(kb_id, algo_id)

从知识库解绑一个算法,并异步清理该算法独有节点组的解析数据;与其他算法共享的节点组数据保留。

Source code in lazyllm/tools/rag/doc_service/doc_server.py
def unbind_algo(self, kb_id: str, algo_id: str):
    """从知识库解绑一个算法,并异步清理该算法独有节点组的解析数据;与其他算法共享的节点组数据保留。"""
    return self._dispatch('unbind_algo', kb_id, algo_id)

update_kb(kb_id, request)

更新知识库的元信息。

Source code in lazyllm/tools/rag/doc_service/doc_server.py
def update_kb(self, kb_id: str, request: KbUpdateRequest):
    """更新知识库的元信息。"""
    return self._dispatch('update_kb', kb_id, request)

upload(request)

通过 /v1/docs/upload 流程将文件上传到 DocServer 管理的存储目录。

当你希望由 DocServer 保存上传副本时,使用该方法。请求体为 UploadRequest,包含 kb_iditems。每个 item 使用 file_path 作为本地源路径,也可以附带可选的 doc_idmetadata

Returns: 标准 API 响应。data["items"] 中包含接受后的 doc_id 和异步 task_id

Source code in lazyllm/tools/rag/doc_service/doc_server.py
    def upload(self, request: UploadRequest):
        """通过 ``/v1/docs/upload`` 流程将文件上传到 DocServer 管理的存储目录。

当你希望由 DocServer 保存上传副本时,使用该方法。请求体为 ``UploadRequest``,包含 ``kb_id``
和 ``items``。每个 item 使用 ``file_path`` 作为本地源路径,也可以附带可选的 ``doc_id``、``metadata``。

**Returns:**
    标准 API 响应。``data["items"]`` 中包含接受后的 ``doc_id`` 和异步 ``task_id``。
"""
        return self._dispatch('upload_request', request)

lazyllm.tools.rag.doc_service.base.AddFileItem

Bases: BaseModel

Source code in lazyllm/tools/rag/doc_service/base.py
class AddFileItem(BaseModel):
    file_path: str
    doc_id: Optional[str] = None
    metadata: Dict[str, Any] = Field(default_factory=dict)

    @model_validator(mode='after')
    def validate_file_path(self):
        """校验单个文件项中的 ``file_path`` 字段。"""
        if not self.file_path or not self.file_path.strip():
            raise ValueError('file_path is required')
        return self

validate_file_path()

校验单个文件项中的 file_path 字段。

Source code in lazyllm/tools/rag/doc_service/base.py
@model_validator(mode='after')
def validate_file_path(self):
    """校验单个文件项中的 ``file_path`` 字段。"""
    if not self.file_path or not self.file_path.strip():
        raise ValueError('file_path is required')
    return self

lazyllm.tools.rag.doc_service.base.UploadRequest = DocItemsRequest module-attribute

lazyllm.tools.rag.doc_service.base.AddRequest = DocItemsRequest module-attribute

lazyllm.tools.rag.doc_service.base.TransferItem

Bases: BaseModel

Source code in lazyllm/tools/rag/doc_service/base.py
class TransferItem(BaseModel):
    doc_id: str
    target_doc_id: str
    kb_id: str = Field(default='__default__', validation_alias=AliasChoices('kb_id', 'source_kb_id'))
    target_kb_id: str
    target_metadata: Optional[Dict[str, Any]] = None
    target_filename: Optional[str] = None
    target_file_path: Optional[str] = None
    mode: str = 'copy'

    @property
    def source_kb_id(self) -> str:
        return self.kb_id

lazyllm.tools.rag.doc_service.base.TransferRequest

Bases: BaseModel

Source code in lazyllm/tools/rag/doc_service/base.py
class TransferRequest(BaseModel):
    items: List[TransferItem]
    idempotency_key: Optional[str] = None

    @model_validator(mode='after')
    def validate_items(self):
        """校验 transfer 请求中的 item 列表及其基本约束。"""
        if not self.items:
            raise ValueError('items is required')
        return self

validate_items()

校验 transfer 请求中的 item 列表及其基本约束。

Source code in lazyllm/tools/rag/doc_service/base.py
@model_validator(mode='after')
def validate_items(self):
    """校验 transfer 请求中的 item 列表及其基本约束。"""
    if not self.items:
        raise ValueError('items is required')
    return self

lazyllm.tools.rag.data_loaders.DirectoryReader

Read local files with the configured reader registry and return document nodes.

Source code in lazyllm/tools/rag/data_loaders.py
class DirectoryReader:
    """Read local files with the configured reader registry and return document nodes."""

    def __init__(self, input_files: Optional[List[str]], local_readers: Optional[Dict] = None,
                 global_readers: Optional[Dict] = None) -> None:
        """Initialize a directory-backed document reader with local and global reader registries."""
        self._input_files = input_files
        self._local_readers, self._global_readers = local_readers, global_readers

    def _reader_entry_sig(self, reader) -> str:
        if reader is None:
            return '__none__'
        if inspect.isclass(reader):
            return f'{reader.__module__}.{reader.__qualname__}'
        qualname = getattr(reader, '__qualname__', None)
        module = getattr(reader, '__module__', None)
        if qualname and '<lambda>' not in qualname:
            return f'{module}.{qualname}' if module else qualname
        try:
            src = inspect.getsource(reader).strip()
            return '__lambda__::' + hashlib.sha256(src.encode()).hexdigest()[:16]
        except (OSError, TypeError):
            return repr(reader)

    def signature(self) -> str:
        """计算当前读取器配置的指纹哈希,用于检测 reader 注册表变更。

将本地和全局 reader 映射序列化为 JSON 后取 SHA-256 前 16 位十六进制字符串。
当任意 reader 被替换或新增时,返回值会发生变化,可用于判断是否需要重新解析文档。

**Returns:**

- str: 16 位十六进制指纹字符串。
"""
        local_sig = {k: self._reader_entry_sig(v) for k, v in (self._local_readers or {}).items()}
        global_sig = {k: self._reader_entry_sig(v) for k, v in (self._global_readers or {}).items()}
        payload = json.dumps({'local_readers': local_sig, 'global_readers': global_sig}, sort_keys=True)
        return hashlib.sha256(payload.encode()).hexdigest()[:16]

    @once_wrapper
    def _lazy_init(self):
        self._reader = SimpleDirectoryReader(
            file_extractor={**(self._global_readers or {}), **(self._local_readers or {})})

    def load_data(self, input_files: Optional[List[str]] = None, metadatas: Optional[Dict] = None,
                  *, split_nodes_by_type: bool = False) -> List[DocNode]:
        """Load documents from files and optionally split the result by node type."""
        self._lazy_init()
        input_files = input_files or self._input_files
        nodes: Union[List[DocNode], Dict[str, List[DocNode]]] = defaultdict(list) if split_nodes_by_type else []
        for doc in self._reader(input_files=input_files, metadatas=metadatas):
            doc._group = type_mapping.get(type(doc), LAZY_ROOT_NAME)
            nodes[doc._group].append(doc) if split_nodes_by_type else nodes.append(doc)
        if not nodes:
            message = f'No nodes load from path {input_files}, please check your data path.'
            load_error = getattr(self._reader, 'last_load_error', None)
            if load_error is not None:
                raise ValueError(str(load_error) or message) from load_error
            raise ValueError(message)
        LOG.info('DirectoryReader loads data done!')
        return nodes

__init__(input_files, local_readers=None, global_readers=None)

Initialize a directory-backed document reader with local and global reader registries.

Source code in lazyllm/tools/rag/data_loaders.py
def __init__(self, input_files: Optional[List[str]], local_readers: Optional[Dict] = None,
             global_readers: Optional[Dict] = None) -> None:
    """Initialize a directory-backed document reader with local and global reader registries."""
    self._input_files = input_files
    self._local_readers, self._global_readers = local_readers, global_readers

load_data(input_files=None, metadatas=None, *, split_nodes_by_type=False)

Load documents from files and optionally split the result by node type.

Source code in lazyllm/tools/rag/data_loaders.py
def load_data(self, input_files: Optional[List[str]] = None, metadatas: Optional[Dict] = None,
              *, split_nodes_by_type: bool = False) -> List[DocNode]:
    """Load documents from files and optionally split the result by node type."""
    self._lazy_init()
    input_files = input_files or self._input_files
    nodes: Union[List[DocNode], Dict[str, List[DocNode]]] = defaultdict(list) if split_nodes_by_type else []
    for doc in self._reader(input_files=input_files, metadatas=metadatas):
        doc._group = type_mapping.get(type(doc), LAZY_ROOT_NAME)
        nodes[doc._group].append(doc) if split_nodes_by_type else nodes.append(doc)
    if not nodes:
        message = f'No nodes load from path {input_files}, please check your data path.'
        load_error = getattr(self._reader, 'last_load_error', None)
        if load_error is not None:
            raise ValueError(str(load_error) or message) from load_error
        raise ValueError(message)
    LOG.info('DirectoryReader loads data done!')
    return nodes

signature()

计算当前读取器配置的指纹哈希,用于检测 reader 注册表变更。

将本地和全局 reader 映射序列化为 JSON 后取 SHA-256 前 16 位十六进制字符串。 当任意 reader 被替换或新增时,返回值会发生变化,可用于判断是否需要重新解析文档。

Returns:

  • str: 16 位十六进制指纹字符串。
Source code in lazyllm/tools/rag/data_loaders.py
    def signature(self) -> str:
        """计算当前读取器配置的指纹哈希,用于检测 reader 注册表变更。

将本地和全局 reader 映射序列化为 JSON 后取 SHA-256 前 16 位十六进制字符串。
当任意 reader 被替换或新增时,返回值会发生变化,可用于判断是否需要重新解析文档。

**Returns:**

- str: 16 位十六进制指纹字符串。
"""
        local_sig = {k: self._reader_entry_sig(v) for k, v in (self._local_readers or {}).items()}
        global_sig = {k: self._reader_entry_sig(v) for k, v in (self._global_readers or {}).items()}
        payload = json.dumps({'local_readers': local_sig, 'global_readers': global_sig}, sort_keys=True)
        return hashlib.sha256(payload.encode()).hexdigest()[:16]

lazyllm.tools.rag.transform.sentence.SentenceSplitter

Bases: _TextSplitterBase

将句子拆分成指定大小的块。可以指定相邻块之间重合部分的大小。

Parameters:

  • chunk_size (int, default: _UNSET ) –

    拆分之后的块大小

  • chunk_overlap (int, default: _UNSET ) –

    相邻两个块之间重合的内容长度

  • num_workers (int, default: _UNSET ) –

    控制并行处理的线程/进程数量

  • **kwargs

    传递给拆分器的额外参数。

Examples:

>>> import lazyllm
>>> from lazyllm.tools import Document, SentenceSplitter
>>> m = lazyllm.OnlineEmbeddingModule(source="glm")
>>> documents = Document(dataset_path='your_doc_path', embed=m, manager=False)
>>> documents.create_node_group(name="sentences", transform=SentenceSplitter, chunk_size=1024, chunk_overlap=100)
Source code in lazyllm/tools/rag/transform/sentence.py
@cpp_proxy
class SentenceSplitter(_TextSplitterBase):
    """
将句子拆分成指定大小的块。可以指定相邻块之间重合部分的大小。

Args:
    chunk_size (int): 拆分之后的块大小
    chunk_overlap (int): 相邻两个块之间重合的内容长度
    num_workers (int):控制并行处理的线程/进程数量
    **kwargs: 传递给拆分器的额外参数。


Examples:

    >>> import lazyllm
    >>> from lazyllm.tools import Document, SentenceSplitter
    >>> m = lazyllm.OnlineEmbeddingModule(source="glm")
    >>> documents = Document(dataset_path='your_doc_path', embed=m, manager=False)
    >>> documents.create_node_group(name="sentences", transform=SentenceSplitter, chunk_size=1024, chunk_overlap=100)
    """
    def __init__(self, chunk_size: int = _UNSET, chunk_overlap: int = _UNSET, num_workers: int = _UNSET):
        super().__init__(chunk_size=chunk_size, overlap=chunk_overlap, num_workers=num_workers)

    def sig_fields(self) -> Dict:
        return {'chunk_size': self._chunk_size, 'chunk_overlap': self._overlap}

    def _merge(self, splits: List[_Split], chunk_size: int) -> List[str]:
        chunks: List[str] = []
        cur_chunk: List[Tuple[str, int]] = []  # list of (text, length)
        cur_chunk_len = 0
        is_chunk_new = True

        def close_chunk() -> None:
            nonlocal cur_chunk, cur_chunk_len, is_chunk_new

            chunks.append(''.join([text for text, _ in cur_chunk]))
            last_chunk = cur_chunk
            cur_chunk = []
            cur_chunk_len = 0
            is_chunk_new = True

            # Add overlap to the next chunk using the last one first
            overlap_len = 0
            for text, length in reversed(last_chunk):
                if overlap_len + length > self._overlap:
                    break
                cur_chunk.append((text, length))
                overlap_len += length
                cur_chunk_len += length
            cur_chunk.reverse()

        i = 0
        while i < len(splits):
            cur_split = splits[i]
            if cur_split.token_size > chunk_size:
                raise ValueError('Single token exceeded chunk size')
            if cur_chunk_len + cur_split.token_size > chunk_size and not is_chunk_new:
                # if adding split to current chunk exceeds chunk size
                close_chunk()
            else:
                if (
                    cur_split.is_sentence
                    or cur_chunk_len + cur_split.token_size <= chunk_size
                    or is_chunk_new  # new chunk, always add at least one split
                ):
                    # add split to chunk
                    cur_chunk_len += cur_split.token_size
                    cur_chunk.append((cur_split.text, cur_split.token_size))
                    i += 1
                    is_chunk_new = False
                else:
                    close_chunk()

        # handle the last chunk
        if not is_chunk_new:
            chunks.append(''.join([text for text, _ in cur_chunk]))

        # Remove whitespace only chunks and remove leading and trailing whitespace.
        return [stripped_chunk for chunk in chunks if (stripped_chunk := chunk.strip())]

lazyllm.tools.rag.transform.character.CharacterSplitter

Bases: _TextSplitterBase

将文本按字符拆分。

Parameters:

  • chunk_size (int, default: _UNSET ) –

    拆分之后的块大小

  • overlap (int, default: _UNSET ) –

    相邻两个块之间重合的内容长度

  • num_workers (int, default: _UNSET ) –

    控制并行处理的线程/进程数量。

  • separator (str, default: _UNSET ) –

    用于拆分的分隔符。默认为' '。

  • is_separator_regex (bool, default: _UNSET ) –

    是否使用正则表达式作为分隔符。默认为False。

  • keep_separator (bool, default: _UNSET ) –

    是否保留分隔符在拆分后的文本中。默认为False。

  • **kwargs

    传递给拆分器的额外参数。

Examples:

>>> import lazyllm
>>> from lazyllm.tools import Document, CharacterSplitter
>>> m = lazyllm.OnlineEmbeddingModule(source="glm")
>>> documents = Document(dataset_path='your_doc_path', embed=m, manager=False)
>>> documents.create_node_group(name="characters", transform=CharacterSplitter, chunk_size=1024, chunk_overlap=100)
Source code in lazyllm/tools/rag/transform/character.py
class CharacterSplitter(_TextSplitterBase):
    """
将文本按字符拆分。

Args:
    chunk_size (int): 拆分之后的块大小
    overlap (int): 相邻两个块之间重合的内容长度
    num_workers (int): 控制并行处理的线程/进程数量。
    separator (str): 用于拆分的分隔符。默认为' '。
    is_separator_regex (bool): 是否使用正则表达式作为分隔符。默认为False。
    keep_separator (bool): 是否保留分隔符在拆分后的文本中。默认为False。
    **kwargs: 传递给拆分器的额外参数。


Examples:

    >>> import lazyllm
    >>> from lazyllm.tools import Document, CharacterSplitter
    >>> m = lazyllm.OnlineEmbeddingModule(source="glm")
    >>> documents = Document(dataset_path='your_doc_path', embed=m, manager=False)
    >>> documents.create_node_group(name="characters", transform=CharacterSplitter, chunk_size=1024, chunk_overlap=100)
    """
    def __init__(self, chunk_size: int = _UNSET, overlap: int = _UNSET, num_workers: int = _UNSET,
                 separator: str = _UNSET, is_separator_regex: bool = _UNSET, keep_separator: bool = _UNSET, **kwargs):
        super().__init__(chunk_size=chunk_size, overlap=overlap, num_workers=num_workers)
        separator = self._get_param_value('separator', separator, ' ')
        is_separator_regex = self._get_param_value('is_separator_regex', is_separator_regex, False)
        keep_separator = self._get_param_value('keep_separator', keep_separator, False)

        self._separator = separator
        self._is_separator_regex = is_separator_regex
        self._keep_separator = keep_separator
        self._character_split_fns = []
        self._cached_sep_pattern = self._get_separator_pattern(self._separator)
        self._cached_default_split_fns = None

    def sig_fields(self) -> Dict:
        fns_sig = [_callable_sig(fn) for fn in self._character_split_fns] if self._character_split_fns else []
        return {
            'chunk_size': self._chunk_size, 'overlap': self._overlap,
            'separator': self._separator, 'is_separator_regex': self._is_separator_regex,
            'keep_separator': self._keep_separator,
            'split_fns': fns_sig,
        }

    def _split(self, text: str, chunk_size: int) -> List[_Split]:
        token_size = self._token_size(text)
        if token_size <= chunk_size:
            return [_Split(text, is_sentence=True, token_size=token_size)]

        text_splits, is_sentence = self._get_splits_by_fns(text)

        if len(text_splits) == 1 and self._token_size(text_splits[0]) > chunk_size:
            token_splitter = _TokenTextSplitter(chunk_size=chunk_size, overlap=self._overlap)
            token_sub_texts = token_splitter.split_text(text_splits[0], metadata_size=0)
            return [
                _Split(s, is_sentence=is_sentence, token_size=self._token_size(s))
                for s in token_sub_texts
            ]

        results = []
        for segment in text_splits:
            token_size = self._token_size(segment)
            if token_size <= chunk_size:
                results.append(_Split(segment, is_sentence=is_sentence, token_size=token_size))
            else:
                sub_results = self._split(segment, chunk_size=chunk_size)
                results.extend(sub_results)

        return results

    def set_split_fns(self, split_fns: Union[Callable[[str], List[str]], List[Callable[[str], List[str]]]], bind_separator: bool = None):  # noqa: E501
        """
CharacterSplitter有默认的拆分函数,你也可以设置自己的拆分函数。
可以设置多个拆分函数,CharacterSplitter会按顺序使用这些函数,分隔符参数将失效。

Args:
    split_fns (List[Callable[[str], List[str]]]): 要使用的拆分函数列表。


Examples:

    >>> import lazyllm
    >>> from lazyllm.tools import CharacterSplitter
    >>> splitter = CharacterSplitter(separator='
    ')
    >>> splitter.set_split_fns([lambda text: text.split(' '), lambda text: text.split('
    ')])
    >>> text = 'Hello, world!'
    >>> splits = splitter.split_text(text, metadata_size=0)
    >>> print(splits)
    """
        if not isinstance(split_fns, list):
            split_fns = [split_fns]
        self._character_split_fns = []
        for split_fn in split_fns:
            if bind_separator is None:
                sig = inspect.signature(split_fn)
                has_separator = 'separator' in sig.parameters
                should_bind = has_separator
            else:
                should_bind = bind_separator

            if should_bind:
                fn = partial(split_fn, separator=self._separator)
            else:
                fn = split_fn

            self._character_split_fns.append(fn)

    def add_split_fn(self, split_fn: Callable[[str], List[str]], index: Optional[int] = None, bind_separator: bool = None):  # noqa: E501
        """
添加一个拆分函数到CharacterSplitter。

Args:
    split_fn (Callable[[str], List[str]]): 要添加的拆分函数。
    index (Optional[int]): 要添加的拆分函数的位置。默认为最后一个位置。
    bind_separator (bool): 是否将分隔符绑定到拆分函数。默认为False。


Examples:

    >>> import lazyllm
    >>> from lazyllm.tools import CharacterSplitter
    >>> splitter = CharacterSplitter(separator='
    ')
    >>> splitter.add_split_fn(lambda text: text.split(' '), index=0)
    >>> text = 'Hello, world!'
    >>> splits = splitter.split_text(text, metadata_size=0)
    >>> print(splits)
    """
        if bind_separator is None:
            sig = inspect.signature(split_fn)
            has_separator = 'separator' in sig.parameters
            should_bind = has_separator
        else:
            should_bind = bind_separator

        if should_bind:
            fn = partial(split_fn, separator=self._separator)
        else:
            fn = split_fn

        if index is None:
            self._character_split_fns.append(fn)
        else:
            self._character_split_fns.insert(index, fn)

    def clear_split_fns(self):
        """
清除CharacterSplitter的所有拆分函数,并使用默认的拆分函数。


Examples:

    >>> import lazyllm
    >>> from lazyllm.tools import CharacterSplitter
    >>> splitter = CharacterSplitter(separator='
    ')
    >>> splitter.clear_split_fns()
    >>> text = 'Hello, world!'
    >>> splits = splitter.split_text(text, metadata_size=0)
    >>> print(splits)
    """
        self._character_split_fns = []

    def _get_splits_by_fns(self, text: str) -> Tuple[List[str], bool]:
        character_split_fns = self._character_split_fns
        if character_split_fns == []:
            if self._cached_default_split_fns is None:
                self._cached_default_split_fns = [
                    partial(self._default_split, self._cached_sep_pattern),
                    lambda t: t.split(' '),
                    list
                ]
            character_split_fns = self._cached_default_split_fns

        splits = []
        for split_fn in character_split_fns:
            splits = split_fn(text)
            if len(splits) > 1:
                break

        return splits, False

    def _default_split(self, sep_pattern: Union[str, set[str]], text: str) -> List[str]:
        splits = re.split(sep_pattern, text)
        results = []
        if self._keep_separator:
            for i in range(0, len(splits) - 1, 2):
                if i + 1 < len(splits):
                    combined = splits[i] + splits[i + 1]
                    if combined:
                        results.append(combined)
            if len(splits) % 2 == 1 and splits[-1]:
                results.append(splits[-1])
        else:
            results = [split for split in splits if split]
        return results

    def _get_separator_pattern(self, separator: str) -> Union[str, set[str]]:
        lookaround_prefixes = ('(?=', '(?<!', '(?<=', '(?!')
        lookaround_pattern = re.compile(r'^\(\?(?:=|<=|!|<!)')

        is_lookaround = (
            self._is_separator_regex
            and (separator.startswith(lookaround_prefixes) or bool(lookaround_pattern.match(separator)))
        )

        if self._is_separator_regex or is_lookaround:
            sep_pattern = separator
        else:
            needs_escape = any(char in separator for char in r'\.^$*+?{}[]|()')
            sep_pattern = re.escape(separator) if needs_escape else separator

        if self._keep_separator:
            sep_pattern = f'({sep_pattern})'
        else:
            sep_pattern = f'(?:{sep_pattern})'

        return sep_pattern

add_split_fn(split_fn, index=None, bind_separator=None)

添加一个拆分函数到CharacterSplitter。

Parameters:

  • split_fn (Callable[[str], List[str]]) –

    要添加的拆分函数。

  • index (Optional[int], default: None ) –

    要添加的拆分函数的位置。默认为最后一个位置。

  • bind_separator (bool, default: None ) –

    是否将分隔符绑定到拆分函数。默认为False。

Examples:

>>> import lazyllm
>>> from lazyllm.tools import CharacterSplitter
>>> splitter = CharacterSplitter(separator='
')
>>> splitter.add_split_fn(lambda text: text.split(' '), index=0)
>>> text = 'Hello, world!'
>>> splits = splitter.split_text(text, metadata_size=0)
>>> print(splits)
Source code in lazyllm/tools/rag/transform/character.py
    def add_split_fn(self, split_fn: Callable[[str], List[str]], index: Optional[int] = None, bind_separator: bool = None):  # noqa: E501
        """
添加一个拆分函数到CharacterSplitter。

Args:
    split_fn (Callable[[str], List[str]]): 要添加的拆分函数。
    index (Optional[int]): 要添加的拆分函数的位置。默认为最后一个位置。
    bind_separator (bool): 是否将分隔符绑定到拆分函数。默认为False。


Examples:

    >>> import lazyllm
    >>> from lazyllm.tools import CharacterSplitter
    >>> splitter = CharacterSplitter(separator='
    ')
    >>> splitter.add_split_fn(lambda text: text.split(' '), index=0)
    >>> text = 'Hello, world!'
    >>> splits = splitter.split_text(text, metadata_size=0)
    >>> print(splits)
    """
        if bind_separator is None:
            sig = inspect.signature(split_fn)
            has_separator = 'separator' in sig.parameters
            should_bind = has_separator
        else:
            should_bind = bind_separator

        if should_bind:
            fn = partial(split_fn, separator=self._separator)
        else:
            fn = split_fn

        if index is None:
            self._character_split_fns.append(fn)
        else:
            self._character_split_fns.insert(index, fn)

clear_split_fns()

清除CharacterSplitter的所有拆分函数,并使用默认的拆分函数。

Examples:

>>> import lazyllm
>>> from lazyllm.tools import CharacterSplitter
>>> splitter = CharacterSplitter(separator='
')
>>> splitter.clear_split_fns()
>>> text = 'Hello, world!'
>>> splits = splitter.split_text(text, metadata_size=0)
>>> print(splits)
Source code in lazyllm/tools/rag/transform/character.py
    def clear_split_fns(self):
        """
清除CharacterSplitter的所有拆分函数,并使用默认的拆分函数。


Examples:

    >>> import lazyllm
    >>> from lazyllm.tools import CharacterSplitter
    >>> splitter = CharacterSplitter(separator='
    ')
    >>> splitter.clear_split_fns()
    >>> text = 'Hello, world!'
    >>> splits = splitter.split_text(text, metadata_size=0)
    >>> print(splits)
    """
        self._character_split_fns = []

set_split_fns(split_fns, bind_separator=None)

CharacterSplitter有默认的拆分函数,你也可以设置自己的拆分函数。 可以设置多个拆分函数,CharacterSplitter会按顺序使用这些函数,分隔符参数将失效。

Parameters:

  • split_fns (List[Callable[[str], List[str]]]) –

    要使用的拆分函数列表。

Examples:

>>> import lazyllm
>>> from lazyllm.tools import CharacterSplitter
>>> splitter = CharacterSplitter(separator='
')
>>> splitter.set_split_fns([lambda text: text.split(' '), lambda text: text.split('
')])
>>> text = 'Hello, world!'
>>> splits = splitter.split_text(text, metadata_size=0)
>>> print(splits)
Source code in lazyllm/tools/rag/transform/character.py
    def set_split_fns(self, split_fns: Union[Callable[[str], List[str]], List[Callable[[str], List[str]]]], bind_separator: bool = None):  # noqa: E501
        """
CharacterSplitter有默认的拆分函数,你也可以设置自己的拆分函数。
可以设置多个拆分函数,CharacterSplitter会按顺序使用这些函数,分隔符参数将失效。

Args:
    split_fns (List[Callable[[str], List[str]]]): 要使用的拆分函数列表。


Examples:

    >>> import lazyllm
    >>> from lazyllm.tools import CharacterSplitter
    >>> splitter = CharacterSplitter(separator='
    ')
    >>> splitter.set_split_fns([lambda text: text.split(' '), lambda text: text.split('
    ')])
    >>> text = 'Hello, world!'
    >>> splits = splitter.split_text(text, metadata_size=0)
    >>> print(splits)
    """
        if not isinstance(split_fns, list):
            split_fns = [split_fns]
        self._character_split_fns = []
        for split_fn in split_fns:
            if bind_separator is None:
                sig = inspect.signature(split_fn)
                has_separator = 'separator' in sig.parameters
                should_bind = has_separator
            else:
                should_bind = bind_separator

            if should_bind:
                fn = partial(split_fn, separator=self._separator)
            else:
                fn = split_fn

            self._character_split_fns.append(fn)

lazyllm.tools.rag.transform.recursive.RecursiveSplitter

Bases: CharacterSplitter

递归拆分文本。

Parameters:

  • chunk_size (int, default: _UNSET ) –

    拆分之后的块大小

  • overlap (int, default: _UNSET ) –

    相邻两个块之间重合的内容长度

  • num_workers (int, default: _UNSET ) –

    控制并行处理的线程/进程数量。

  • keep_separator (bool, default: _UNSET ) –

    是否保留分隔符在拆分后的文本中。默认为False。

  • is_separator_regex (bool, default: _UNSET ) –

    是否使用正则表达式作为分隔符。默认为False。

  • separators (List[str], default: _UNSET ) –

    用于拆分的分隔符列表。默认为['

', ' ', ' ', '']。如果你想按多个分隔符拆分,可以设置这个参数。

Examples:

>>> import lazyllm
>>> from lazyllm.tools import RecursiveSplitter
>>> splitter = RecursiveSplitter(separators=['

', '
', ' ', ''])
>>> documents = Document(dataset_path='your_doc_path', embed=m, manager=False)
>>> documents.create_node_group(name="recursive", transform=RecursiveSplitter, chunk_size=1024, chunk_overlap=100)
Source code in lazyllm/tools/rag/transform/recursive.py
class RecursiveSplitter(CharacterSplitter):
    """
递归拆分文本。

Args:
    chunk_size (int): 拆分之后的块大小
    overlap (int): 相邻两个块之间重合的内容长度
    num_workers (int):控制并行处理的线程/进程数量。
    keep_separator (bool): 是否保留分隔符在拆分后的文本中。默认为False。
    is_separator_regex (bool): 是否使用正则表达式作为分隔符。默认为False。
    separators (List[str]): 用于拆分的分隔符列表。默认为['

', '
', ' ', '']。如果你想按多个分隔符拆分,可以设置这个参数。


Examples:

    >>> import lazyllm
    >>> from lazyllm.tools import RecursiveSplitter
    >>> splitter = RecursiveSplitter(separators=['

    ', '
    ', ' ', ''])
    >>> documents = Document(dataset_path='your_doc_path', embed=m, manager=False)
    >>> documents.create_node_group(name="recursive", transform=RecursiveSplitter, chunk_size=1024, chunk_overlap=100)
    """
    def __init__(self, chunk_size: int = _UNSET, overlap: int = _UNSET, num_workers: int = _UNSET,
                 keep_separator: bool = _UNSET, is_separator_regex: bool = _UNSET,
                 separators: List[str] = _UNSET, **kwargs):
        super().__init__(chunk_size=chunk_size, overlap=overlap, num_workers=num_workers,
                         keep_separator=keep_separator, is_separator_regex=is_separator_regex)
        separators = self._get_param_value('separators', separators, None)

        self._separators = separators if separators else ['\n\n', '\n', ' ', '']
        self._cached_recursive_split_fns = [
            partial(self._default_split, self._get_separator_pattern(sep))
            for sep in self._separators
        ] + [list]

    def sig_fields(self) -> Dict:
        base = super().sig_fields()
        base.update({'separators': self._separators})
        return base

    def _get_splits_by_fns(self, text: str) -> Tuple[List[str], bool]:
        character_split_fns = self._character_split_fns
        if character_split_fns == []:
            character_split_fns = self._cached_recursive_split_fns
        splits = []
        for split_fn in character_split_fns:
            splits = split_fn(text)
            if len(splits) > 1:
                break

        return splits, False

lazyllm.tools.rag.transform.markdown.MarkdownSplitter

Bases: _TextSplitterBase

递归拆分markdown文本。

Parameters:

  • chunk_size (int, default: _UNSET ) –

    拆分之后的块大小

  • overlap (int, default: _UNSET ) –

    相邻两个块之间重合的内容长度

  • num_workers (int, default: _UNSET ) –

    控制并行处理的线程/进程数量。

  • keep_trace (bool, default: _UNSET ) –

    是否保留markdown文本的追踪。默认为False。

  • keep_headers (bool, default: _UNSET ) –

    是否保留headers在拆分后的文本中。默认为False。

  • keep_lists (bool, default: _UNSET ) –

    是否保留lists在拆分后的文本中。默认为False。

  • keep_code_blocks (bool, default: _UNSET ) –

    是否保留code blocks在拆分后的文本中。默认为False。

  • keep_tables (bool, default: _UNSET ) –

    是否保留tables在拆分后的文本中。默认为False。

  • keep_images (bool, default: _UNSET ) –

    是否保留images在拆分后的文本中。默认为False。

  • keep_links (bool, default: _UNSET ) –

    是否保留links在拆分后的文本中。默认为False。

  • **kwargs

    传递给拆分器的额外参数。

Examples:

>>> import lazyllm
>>> from lazyllm.tools import MarkdownSplitter
>>> documents = Document(dataset_path='your_doc_path', embed=m, manager=False)
>>> documents.create_node_group(name="markdown", transform=MarkdownSplitter,
                                chunk_size=1024, chunk_overlap=100, keep_trace=True, keep_headers=True)
Source code in lazyllm/tools/rag/transform/markdown.py