Newer
Older
navi-1 / tests / unit / tools / test_loader.py
"""Loading user tools from tools/ — the reload must see the file, not a cache.

The loader used to run the file through importlib's source loader, so the
bytecode cache in __pycache__ decided the result. Its key is (mtime in whole
seconds, file size): an edit of the same length written in the same second as
the previous load re-ran the OLD code, and reload_tools reported a successful
reload while the previous version of the tool stayed live.
"""

import textwrap
from pathlib import Path

from navi.core.registry import ToolRegistry
from navi.tools._internal.loader import load_tools_from_dir


def _write(path: Path, version: str) -> None:
    """Write a module-level tool whose source length does not depend on *version*."""
    path.write_text(
        textwrap.dedent(
            f'''
            name = "probe"
            description = "{version}"
            parameters = {{"type": "object", "properties": {{}}}}

            async def execute(params):
                return "{version}"
            '''
        ).lstrip()
    )


def test_reload_sees_an_edit_written_in_the_same_second(tmp_path):
    registry = ToolRegistry()
    tool_file = tmp_path / "probe.py"

    _write(tool_file, "v1")
    assert [t.name for t in registry.reload_user_tools(str(tmp_path)).loaded] == ["probe"]
    assert registry.get("probe").description == "v1"

    # Same length, written the same second — the stale-bytecode window.
    _write(tool_file, "v2")
    registry.reload_user_tools(str(tmp_path))
    assert registry.get("probe").description == "v2"


def test_reload_replaces_the_module_level_execute(tmp_path):
    """The new body must run, not just the new description show."""
    import asyncio

    registry = ToolRegistry()
    tool_file = tmp_path / "probe.py"

    _write(tool_file, "v1")
    registry.reload_user_tools(str(tmp_path))
    assert asyncio.run(registry.get("probe").execute({})).output == "v1"

    _write(tool_file, "v2")
    registry.reload_user_tools(str(tmp_path))
    assert asyncio.run(registry.get("probe").execute({})).output == "v2"


def test_class_based_tools_still_load(tmp_path):
    (tmp_path / "cls_tool.py").write_text(
        textwrap.dedent(
            '''
            from navi.tools._internal.base import Tool, ToolResult


            class ClsTool(Tool):
                name = "cls_tool"
                description = "class-based"
                parameters = {"type": "object", "properties": {}}

                async def execute(self, params):
                    return ToolResult(success=True, output="cls")
            '''
        ).lstrip()
    )
    result = load_tools_from_dir(str(tmp_path))
    assert [t.name for t in result.loaded] == ["cls_tool"]
    assert result.errors == {}


def test_class_based_tool_may_take_ctx(tmp_path):
    """The signature every built-in uses — it must not read as "wrong signature"."""
    (tmp_path / "ctx_tool.py").write_text(
        textwrap.dedent(
            '''
            from navi.tools._internal.base import Tool, ToolResult


            class CtxTool(Tool):
                name = "ctx_tool"
                description = "takes ctx"
                parameters = {"type": "object", "properties": {}}

                async def execute(self, params, ctx=None):
                    return ToolResult(success=True, output="ctx")
            '''
        ).lstrip()
    )
    result = load_tools_from_dir(str(tmp_path))
    assert [t.name for t in result.loaded] == ["ctx_tool"]
    assert result.errors == {}


def test_a_broken_file_does_not_hide_the_others(tmp_path):
    (tmp_path / "broken.py").write_text("this is not python (")
    _write(tmp_path / "probe.py", "v1")

    result = load_tools_from_dir(str(tmp_path))
    assert [t.name for t in result.loaded] == ["probe"]
    assert list(result.errors) == ["broken.py"]
    assert "SyntaxError" in result.errors["broken.py"]


def test_a_missing_directory_is_not_an_error(tmp_path):
    result = load_tools_from_dir(str(tmp_path / "nope"))
    assert result.loaded == []
    assert result.errors == {}


def test_underscored_files_are_ignored(tmp_path):
    _write(tmp_path / "_private.py", "v1")
    result = load_tools_from_dir(str(tmp_path))
    assert result.loaded == [] and result.errors == {}