"""Tests for reme common steps with only local dependencies.""" # pylint: disable=protected-access import asyncio import os import tempfile import warnings from reme.components.agent_wrapper import BaseAgentWrapper from reme.components.application_context import ApplicationContext from reme.components.file_store import LocalFileStore from reme.schema import FileFrontMatter, FileLink, FileNode, TraverseGraph from reme.steps.common.add import AddStep from reme.steps.common.health_check import _file_graph_status from reme.steps.common.llm_demo import LLMDemoStep from reme.steps.common.python_execute import PythonExecuteStep from reme.steps.index import traverse as traverse_mod warnings.filterwarnings("ignore", category=DeprecationWarning, module="jieba") warnings.filterwarnings("ignore", category=DeprecationWarning, module="pkg_resources") class _temp_chdir: """chdir to path for the duration of the block; restore on exit.""" def __init__(self, path): self.path = path self._old = None def __enter__(self): self._old = os.getcwd() os.chdir(self.path) return self def __exit__(self, *exc): os.chdir(self._old) def _run(coro): """Run an async coroutine on a fresh isolated event loop.""" asyncio.run(coro) def _node( path: str, links: list[tuple[str, str | None]] | None = None, *, name: str = "", description: str = "", ) -> FileNode: """Build a FileNode with (target_path, target_anchor) outgoing edges.""" return FileNode( path=path, st_mtime=1.0, links=[ FileLink(source_path=path, target_path=target, target_anchor=anchor) for target, anchor in (links or []) ], front_matter=FileFrontMatter(name=name, description=description), ) async def _make_store(nodes: list[FileNode]) -> LocalFileStore: """LocalFileStore seeded with the given graph nodes (no files on disk).""" store = LocalFileStore(name="t", embedding_store="") await store.start() if nodes: await store.file_graph.upsert_nodes(nodes) return store def _edges(step) -> list[dict]: return step.context.response.answer["edges"] def test_add_step_coerces_numeric_inputs(): """add accepts numeric strings as numbers, not string concatenation.""" async def run(): step = AddStep() resp = await step(a="1", b="2.5") assert resp.success is True assert resp.answer == "3.5" assert resp.metadata["result"] == 3.5 print("✓ test_add_step_coerces_numeric_inputs passed") _run(run()) def test_add_step_rejects_invalid_inputs(): """invalid add arguments should return a failed response instead of throwing or concatenating.""" async def run(): step = AddStep() resp = await step(a="one", b=2) assert resp.success is False assert "Invalid add arguments" in resp.answer print("✓ test_add_step_rejects_invalid_inputs passed") _run(run()) def test_python_execute_step_returns_printed_stdout(tmp_path): """python_execute returns stdout as answer and runs under workspace_dir.""" async def run(): app_context = ApplicationContext(workspace_dir=str(tmp_path)) step = PythonExecuteStep(app_context=app_context) resp = await step(code="from pathlib import Path\nprint(Path.cwd().name)") assert resp.success is True assert resp.answer == f"{tmp_path.name}\n" assert resp.metadata["returncode"] == 0 assert resp.metadata["stderr"] == "" print("✓ test_python_execute_step_returns_printed_stdout passed") _run(run()) def test_python_execute_step_reports_stderr_on_failure(): """python_execute captures traceback stderr instead of throwing.""" async def run(): step = PythonExecuteStep() resp = await step(code='raise RuntimeError("boom")') assert resp.success is False assert resp.metadata["returncode"] != 0 assert "RuntimeError: boom" in resp.answer assert "RuntimeError: boom" in resp.metadata["stderr"] print("✓ test_python_execute_step_reports_stderr_on_failure passed") _run(run()) def test_python_execute_step_times_out(): """python_execute converts subprocess timeout into a failed response.""" async def run(): step = PythonExecuteStep() resp = await step(code="import time\ntime.sleep(1)", timeout=0.01) assert resp.success is False assert resp.answer == "Python execution timed out after 0.01s" assert resp.metadata["timeout"] == 0.01 print("✓ test_python_execute_step_times_out passed") _run(run()) class _FakeAgentWrapper(BaseAgentWrapper): """Capture reply kwargs without calling a real model.""" def __init__(self): super().__init__() self.last_kwargs = None async def reply(self, inputs, **kwargs) -> dict: self.last_kwargs = kwargs return {"result": "ok"} def test_llm_demo_always_registers_add_tool(): """LLM demo always passes the add job as a tool.""" async def run(): wrapper = _FakeAgentWrapper() step = LLMDemoStep() resp = await step(query="hello", agent_wrapper=wrapper) assert resp.success is True assert wrapper.last_kwargs["job_tools"] == ["add"] assert "job_tools" not in resp.metadata print("✓ test_llm_demo_always_registers_add_tool passed") _run(run()) def test_file_graph_health_reports_neo4j_cached_counts(): """Neo4j file graph health should not be reported as an empty local graph.""" class FakeNeo4jGraph: """Minimal Neo4j graph stub with cached health counters.""" is_started = True _driver = object() _uri = "bolt://example" _database = "neo4j" _n_nodes = 3 _n_edges = 4 _n_virtual = 1 status = _file_graph_status(FakeNeo4jGraph()) assert status["n_nodes"] == 3 assert status["n_edges"] == 4 assert status["n_virtual"] == 1 print("✓ test_file_graph_health_reports_neo4j_cached_counts passed") # =========================================================================== # Direct unit tests: TraverseStep # (LocalFileStore, no HTTP server — BFS over wikilink edges) # =========================================================================== def test_traverse_forward_depth_1(): """depth=1 forward returns direct outbound neighbors.""" async def run(): with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp): store = await _make_store( [ _node("a.md", [("b.md", None), ("c.md", "intro")]), _node("b.md"), _node("c.md"), ], ) step = traverse_mod.TraverseStep(file_store=store) await step(path="a.md", direction="forward", depth=1) results = _edges(step) paths = {r["target"] for r in results} assert paths == {"b.md", "c.md"} c_edge = next(r for r in results if r["target"] == "c.md") assert c_edge["target_anchor"] == "intro" assert c_edge["source"] == "a.md" assert c_edge["depth"] == 1 await store.close() print("✓ test_traverse_forward_depth_1 passed") asyncio.run(run()) def test_traverse_normalizes_windows_seed_path(): """Windows-style seeds match the graph's portable POSIX path keys.""" async def run(): with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp): store = await _make_store( [ _node("topics/a.md", [("topics/b.md", None)]), _node("topics/b.md"), ], ) step = traverse_mod.TraverseStep(file_store=store) await step(path=r"topics\a.md", direction="forward", depth=1) assert [edge["target"] for edge in _edges(step)] == ["topics/b.md"] await store.close() asyncio.run(run()) def test_traverse_backward_returns_inlinks(): """direction=backward walks inbound edges.""" async def run(): with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp): store = await _make_store( [ _node("a.md", [("b.md", None)]), _node("c.md", [("b.md", None)]), _node("b.md"), ], ) step = traverse_mod.TraverseStep(file_store=store) await step(path="b.md", direction="backward", depth=1) results = _edges(step) assert {(r["source"], r["target"]) for r in results} == { ("a.md", "b.md"), ("c.md", "b.md"), } await store.close() print("✓ test_traverse_backward_returns_inlinks passed") asyncio.run(run()) def test_traverse_depth_2_expands(): """depth=2 traverses one hop beyond direct neighbors.""" async def run(): with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp): store = await _make_store( [ _node("a.md", [("b.md", None)]), _node("b.md", [("c.md", None)]), _node("c.md"), ], ) step = traverse_mod.TraverseStep(file_store=store) await step(path="a.md", direction="forward", depth=2) results = _edges(step) depth_map = {r["target"]: r["depth"] for r in results} assert depth_map.get("b.md") == 1 assert depth_map.get("c.md") == 2 await store.close() print("✓ test_traverse_depth_2_expands passed") asyncio.run(run()) def test_traverse_short_seed_yields_empty(): """A short (not relative to the workspace) seed isn't resolved anymore — BFS simply finds no edges from a path that doesn't match any graph node.""" async def run(): with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp): store = await _make_store( [ _node("topics/Bob.md"), _node("people/Bob.md"), ], ) step = traverse_mod.TraverseStep(file_store=store) await step(path="Bob", direction="forward", depth=1) payload = _edges(step) # No error, just empty results because "Bob" isn't a graph key. assert payload == [] await store.close() print("✓ test_traverse_short_seed_yields_empty passed") asyncio.run(run()) def test_traverse_not_found_seed(): """A seed not in the graph returns an empty list (no error).""" async def run(): with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp): store = await _make_store([_node("a.md")]) step = traverse_mod.TraverseStep(file_store=store) await step(path="topics/ghost.md", direction="forward", depth=1) payload = _edges(step) assert payload == [] await store.close() print("✓ test_traverse_not_found_seed passed") asyncio.run(run()) def test_traverse_both_directions(): """direction=both walks out- and in-bound; depth=1 returns one hop in each direction.""" async def run(): with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp): store = await _make_store( [ _node("upstream.md", [("center.md", None)]), _node("center.md", [("downstream.md", None)]), _node("downstream.md"), ], ) step = traverse_mod.TraverseStep(file_store=store) await step(path="center.md", direction="both", depth=1) results = _edges(step) assert {(r["source"], r["target"]) for r in results} == { ("upstream.md", "center.md"), ("center.md", "downstream.md"), } await store.close() print("✓ test_traverse_both_directions passed") asyncio.run(run()) def test_traverse_both_preserves_reciprocal_edge_directions(): """Opposite wikilinks remain two distinct directed graph edges.""" async def run(): with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp): store = await _make_store( [ _node("a.md", [("b.md", None)]), _node("b.md", [("a.md", None)]), ], ) step = traverse_mod.TraverseStep(file_store=store) response = await step(path="a.md", direction="both", depth=1) graph = TraverseGraph.model_validate(response.answer) assert {(edge.source, edge.target) for edge in graph.edges} == { ("a.md", "b.md"), ("b.md", "a.md"), } assert response.metadata == {} await store.close() asyncio.run(run()) def test_traverse_returns_frontmatter_and_unresolved_nodes(): """Graph nodes expose labels and distinguish indexed files from dangling targets.""" async def run(): with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp): store = await _make_store( [ _node( "a.md", [("missing.md", "details")], name="Alpha", description="Root note", ), ], ) step = traverse_mod.TraverseStep(file_store=store) response = await step(path="a.md", direction="forward", depth=1) graph = TraverseGraph.model_validate(response.answer) nodes = {node.path: node for node in graph.nodes} assert graph.version == 1 assert graph.seeds == ["a.md"] assert nodes["a.md"].model_dump() == { "id": "a.md", "path": "a.md", "name": "Alpha", "description": "Root note", "depth": 0, "indexed": True, } assert nodes["missing.md"].indexed is False assert nodes["missing.md"].depth == 1 assert graph.edges[0].target_anchor == "details" await store.close() asyncio.run(run()) def test_traverse_depth_zero_returns_only_seed_nodes(): """A zero-hop traversal is valid and emits no edges.""" async def run(): with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp): store = await _make_store([_node("a.md", [("b.md", None)]), _node("b.md")]) step = traverse_mod.TraverseStep(file_store=store) response = await step(path="a.md", direction="forward", depth=0) graph = TraverseGraph.model_validate(response.answer) assert [node.path for node in graph.nodes] == ["a.md"] assert graph.edges == [] await store.close() asyncio.run(run()) if __name__ == "__main__": print("\n=== traverse step tests ===") test_traverse_forward_depth_1() test_traverse_backward_returns_inlinks() test_traverse_depth_2_expands() test_traverse_short_seed_yields_empty() test_traverse_not_found_seed() test_traverse_both_directions() test_traverse_both_preserves_reciprocal_edge_directions() test_traverse_returns_frontmatter_and_unresolved_nodes() test_traverse_depth_zero_returns_only_seed_nodes() print("\n所有测试通过!")