207 lines
11 KiB
Python
207 lines
11 KiB
Python
"""Offline CLI integration: python3 tests/cli_flow.py (cargo build -p grokboy first)."""
|
|
import json
|
|
import os
|
|
from pathlib import Path
|
|
import subprocess
|
|
import tempfile
|
|
import threading
|
|
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
|
|
|
BINARY = Path(__file__).resolve().parents[1] / "target/debug/grokboy"
|
|
|
|
|
|
def tool(name, args):
|
|
return {"role": "assistant", "tool_calls": [{"id": "call", "type": "function", "function": {
|
|
"name": name, "arguments": json.dumps(args)}}]}
|
|
|
|
|
|
def response(message, finish=None):
|
|
return {"choices": [{"message": message, "finish_reason": finish or (
|
|
"tool_calls" if message.get("tool_calls") else "stop")}]}
|
|
|
|
|
|
def paired(messages):
|
|
for i, message in enumerate(messages):
|
|
for offset, call in enumerate(message.get("tool_calls", [])):
|
|
result = messages[i + offset + 1]
|
|
assert result["role"] == "tool" and result["tool_call_id"] == call["id"]
|
|
|
|
|
|
class Provider(BaseHTTPRequestHandler):
|
|
def do_POST(self):
|
|
body = json.loads(self.rfile.read(int(self.headers["Content-Length"])))
|
|
self.server.requests.append(body)
|
|
try:
|
|
paired(body["messages"])
|
|
if getattr(self.server, "callback", None):
|
|
reply = self.server.callback(body)
|
|
else:
|
|
reply = self.server.replies.pop(0)
|
|
content_type = "application/json"
|
|
if body.get("stream"):
|
|
choices = []
|
|
for index, choice in enumerate(reply.get("choices", [])):
|
|
delta = dict(choice.get("message", {}))
|
|
if "tool_calls" in delta:
|
|
delta["tool_calls"] = [dict(call, index=i) for i, call in enumerate(delta["tool_calls"])]
|
|
choices.append({"index":index,"delta":delta,"finish_reason":choice.get("finish_reason")})
|
|
encoded = b"data: " + json.dumps({"choices":choices}).encode() + b"\n\ndata: [DONE]\n\n"
|
|
content_type = "text/event-stream"
|
|
else:
|
|
encoded = json.dumps(reply).encode()
|
|
self.send_response(200)
|
|
except Exception as exc:
|
|
content_type = "application/json"
|
|
self.server.errors.append(str(exc))
|
|
encoded = b'{"error":{"message":"unexpected request or unpaired history"}}'
|
|
self.send_response(500)
|
|
self.send_header("Content-Type", content_type)
|
|
self.send_header("Content-Length", str(len(encoded)))
|
|
self.end_headers()
|
|
self.wfile.write(encoded)
|
|
|
|
def log_message(self, *_):
|
|
pass
|
|
|
|
|
|
def main():
|
|
with tempfile.TemporaryDirectory(prefix="grokboy-cli-") as temp, ThreadingHTTPServer(("127.0.0.1", 0), Provider) as server:
|
|
server.requests, server.replies, server.errors = [], [], []
|
|
thread = threading.Thread(target=server.serve_forever, daemon=True)
|
|
thread.start()
|
|
env = {**os.environ, "GROKBOY_API_KEY": "offline-test", "GROKBOY_MODEL": "mock",
|
|
"GROKBOY_BASE_URL": f"http://127.0.0.1:{server.server_port}/v1",
|
|
"GROKBOY_SESSIONS_DIR": str(Path(temp) / "sessions"),
|
|
"GROKBOY_MAX_ROUNDS": "1", "GROKBOY_MAX_ROUNDS_TOTAL": "6",
|
|
"GROKBOY_CONTEXT_CHARS": "100000", "GROKBOY_PROGRESS": "1"}
|
|
|
|
def run(replies, args, verdict, input_text=None):
|
|
server.requests.clear()
|
|
server.replies = list(replies)
|
|
result = subprocess.run([str(BINARY), *args], cwd=temp, env=env, input=input_text,
|
|
capture_output=True, text=True, timeout=15)
|
|
assert result.returncode == (0 if verdict in ("done", "answer", "waiting") or args[0] == "agent" else 1), result.stderr
|
|
assert not server.replies and not server.errors, (server.replies, server.errors, result.stderr)
|
|
sessions = list((Path(temp) / "sessions").glob("*.json"))
|
|
session = json.loads(max(sessions, key=lambda p: p.stat().st_mtime_ns).read_text())
|
|
assert session["last_verdict"] == verdict, session
|
|
paired(session["messages"])
|
|
assert len(server.requests) == len(replies), "hidden summary request"
|
|
return session, result
|
|
|
|
try:
|
|
run([response(tool("external_write_file", {"path": "note.txt", "content": "ok"})),
|
|
response(tool("external_read_file", {"path": "note.txt"})),
|
|
response(tool("report_done", {"message": "verified note.txt"}))], ["run", "write and verify"], "done")
|
|
assert Path(temp, "note.txt").read_text() == "ok"
|
|
print("PASS write -> read -> done across progress intervals")
|
|
|
|
env["GROKBOY_MAX_ROUNDS_TOTAL"] = "1"
|
|
saved, _ = run([response(tool("external_read_file", {"path": "note.txt"}))], ["run", "inspect"], "budget_exhausted")
|
|
run([response({"role": "assistant", "content": "continued"})], ["run", "--session", saved["id"], "continue"], "answer")
|
|
print("PASS hard request budget -> saved session -> resume")
|
|
|
|
env["GROKBOY_MAX_ROUNDS_TOTAL"] = "6"
|
|
saved, _ = run([response(tool("external_read_file", {"path": "note.txt"}))] * 3, ["run", "repeat"], "blocked")
|
|
run([response({"role": "assistant", "content": "recovered"})], ["run", "--session", saved["id"], "change approach"], "answer")
|
|
print("PASS repeated observations -> paired history -> resume")
|
|
|
|
run([response(tool("external_write_file", {"path": "bad.txt", "content": "bad"}), "length")], ["run", "truncated"], "failed")
|
|
assert not Path(temp, "bad.txt").exists()
|
|
print("PASS truncated tool response never executes")
|
|
|
|
run([{"error": {"message": "mock provider failure"}}, response({"role": "assistant", "content": "retry ok"})],
|
|
["agent"], "answer", "first\nretry\n/exit\n")
|
|
print("PASS provider failure keeps REPL alive for next input")
|
|
|
|
env["GROKBOY_MAX_ROUNDS_TOTAL"] = "8"
|
|
session, result = run([
|
|
response(tool("send_message", {"type": "text", "content": "hello visible"})),
|
|
response({"role": "assistant", "content": "scratchpad must not leak"}),
|
|
], ["run", "say hi"], "answer")
|
|
assert "hello visible" in result.stdout
|
|
assert "scratchpad must not leak" not in result.stdout
|
|
assert session["last_message"] == "hello visible"
|
|
print("PASS send_message then no-tools is answer on stdout")
|
|
|
|
child_n = {"n": 0}
|
|
parent_n = {"n": 0}
|
|
sent = {"n": False}
|
|
|
|
def revival(req):
|
|
msgs = req["messages"]
|
|
child = any("You are a background subagent" in (m.get("content") or "") for m in msgs)
|
|
if child:
|
|
child_n["n"] += 1
|
|
if child_n["n"] == 1:
|
|
return response(tool("external_write_file", {"path": "from-child.txt", "content": "child-ok"}))
|
|
return response({"role": "assistant", "content": "wrote from-child.txt"})
|
|
parent_n["n"] += 1
|
|
if parent_n["n"] == 1:
|
|
return response(tool("spawn_subagent", {"goal": "write from-child.txt", "title": "writer"}))
|
|
text = " ".join((m.get("content") or "") for m in msgs)
|
|
if "background subagent has finished" in text:
|
|
if not sent["n"]:
|
|
sent["n"] = True
|
|
return response(tool("send_message", {"type": "text", "content": "child finished"}))
|
|
return response({"role": "assistant", "content": "scratch"})
|
|
return response({"role": "assistant", "content": "scratch"})
|
|
|
|
server.callback = revival
|
|
server.requests.clear()
|
|
server.replies = []
|
|
result = subprocess.run([str(BINARY), "run", "delegate write"], cwd=temp, env=env,
|
|
capture_output=True, text=True, timeout=20)
|
|
server.callback = None
|
|
assert result.returncode == 0, result.stderr
|
|
assert not server.errors, server.errors
|
|
session = json.loads(max((Path(temp) / "sessions").glob("*.json"),
|
|
key=lambda p: p.stat().st_mtime_ns).read_text())
|
|
paired(session["messages"])
|
|
assert session["last_verdict"] == "answer", session
|
|
assert Path(temp, "from-child.txt").read_text() == "child-ok"
|
|
assert any("background subagent has finished" in (m.get("content") or "") for m in session["messages"])
|
|
assert "child finished" in result.stdout
|
|
assert "scratch" not in result.stdout, result.stdout
|
|
assert parent_n["n"] < 8 and 2 <= child_n["n"] <= 4, (parent_n, child_n, result.stderr[-2000:])
|
|
print("PASS spawn_subagent yield + revival on real CLI")
|
|
|
|
sent["n"] = False
|
|
parent_n["n"] = 0
|
|
|
|
def cmd_revival(req):
|
|
msgs = req["messages"]
|
|
parent_n["n"] += 1
|
|
if parent_n["n"] == 1:
|
|
return response(tool("external_exec_command", {"cmd": "printf revived-cli", "block_until_ms": 0}))
|
|
text = " ".join((m.get("content") or "") for m in msgs)
|
|
if "background command has finished" in text:
|
|
if not sent["n"]:
|
|
sent["n"] = True
|
|
return response(tool("send_message", {"type": "text", "content": "command done"}))
|
|
return response({"role": "assistant", "content": "scratch"})
|
|
return response({"role": "assistant", "content": "scratch"})
|
|
|
|
server.callback = cmd_revival
|
|
server.requests.clear()
|
|
result = subprocess.run([str(BINARY), "run", "run cmd"], cwd=temp, env=env,
|
|
capture_output=True, text=True, timeout=20)
|
|
server.callback = None
|
|
assert result.returncode == 0, result.stderr
|
|
assert not server.errors, server.errors
|
|
session = json.loads(max((Path(temp) / "sessions").glob("*.json"),
|
|
key=lambda p: p.stat().st_mtime_ns).read_text())
|
|
paired(session["messages"])
|
|
assert session["last_verdict"] == "answer", session
|
|
assert any("background command has finished" in (m.get("content") or "") for m in session["messages"])
|
|
assert "command done" in result.stdout
|
|
assert "scratch" not in result.stdout, result.stdout
|
|
print("PASS exec_command block_until_ms 0 yield + revival on real CLI")
|
|
finally:
|
|
server.shutdown()
|
|
thread.join(timeout=5)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|