from pathlib import Path import pytest from langchain_core.messages import AIMessage from gcode.graph import Agent from gcode.safety import SafetyGate from gcode.tools.fs import make_fs_tools class ScriptedLLM: def __init__(self, responses: list[AIMessage]) -> None: self.responses = list(responses) self.calls = 0 def bind_tools(self, tools): return self async def ainvoke(self, messages, **kwargs): if self.calls >= len(self.responses): return AIMessage(content="(script exhausted)") msg = self.responses[self.calls] self.calls += 1 return msg async def _always(kind: str, detail: str) -> bool: return True async def _never(kind: str, detail: str) -> bool: return False @pytest.mark.asyncio async def test_one_tool_call_then_end(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.chdir(tmp_path) (tmp_path / "README.md").write_text("hello gcode\n", encoding="utf-8") safety = SafetyGate(_always) tools = make_fs_tools(safety, tmp_path) llm = ScriptedLLM( [ AIMessage( content="", tool_calls=[ { "name": "read_file", "args": {"path": "README.md"}, "id": "call_1", "type": "tool_call", } ], ), AIMessage(content="README 只有一句 hello gcode。"), ] ) agent = Agent(llm, tools, cwd=tmp_path) messages = await agent.invoke_turn("请读 README") texts = [getattr(m, "content", "") for m in messages] assert any("hello gcode" in str(t) for t in texts) assert any("README 只有一句" in str(t) for t in texts) assert llm.calls == 2 @pytest.mark.asyncio async def test_replace_denied_file_unchanged(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.chdir(tmp_path) target = tmp_path / "app.py" target.write_text("print(1)\n", encoding="utf-8") safety = SafetyGate(_never) tools = make_fs_tools(safety, tmp_path) llm = ScriptedLLM( [ AIMessage( content="", tool_calls=[ { "name": "replace_in_file", "args": { "path": "app.py", "old_text": "print(1)", "new_text": "print(2)", }, "id": "call_w", "type": "tool_call", } ], ), AIMessage(content="用户拒绝了写入,我改用其它办法。"), ] ) agent = Agent(llm, tools, cwd=tmp_path) messages = await agent.invoke_turn("把 print(1) 改成 print(2)") joined = "\n".join(str(getattr(m, "content", "")) for m in messages) assert "用户拒绝本次写入" in joined assert target.read_text(encoding="utf-8") == "print(1)\n" assert llm.calls == 2