engine.py 15 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377
  1. """Engine: orchestrates the full Drill run lifecycle."""
  2. from __future__ import annotations
  3. import json
  4. import os
  5. import re
  6. import subprocess
  7. import time
  8. from dataclasses import dataclass, field
  9. from datetime import datetime
  10. from pathlib import Path
  11. from typing import Any
  12. import yaml
  13. from drill.actor import Actor
  14. from drill.assertions import AssertionResult, run_verify_assertions
  15. from drill.backend import load_backend
  16. from drill.normalizer import (
  17. NORMALIZERS,
  18. collect_new_logs,
  19. filter_codex_logs_by_cwd,
  20. snapshot_log_dir,
  21. )
  22. from drill.session import TmuxSession
  23. from drill.setup import run_assertions, run_helpers
  24. from drill.verifier import Verifier
  25. @dataclass
  26. class VerifyConfig:
  27. criteria: list[str] = field(default_factory=list)
  28. assertions: list[str] = field(default_factory=list)
  29. observe: bool = False
  30. @dataclass
  31. class ScenarioConfig:
  32. scenario: str
  33. description: str
  34. user_posture: str
  35. setup: dict[str, Any]
  36. turns: list[dict[str, Any]]
  37. limits: dict[str, Any]
  38. verify: VerifyConfig
  39. @classmethod
  40. def from_yaml(cls, path: Path) -> ScenarioConfig:
  41. with open(path) as f:
  42. data = yaml.safe_load(f)
  43. verify_data = data.get("verify", {})
  44. return cls(
  45. scenario=data["scenario"],
  46. description=data.get("description", ""),
  47. user_posture=data.get("user_posture", "naive"),
  48. setup=data.get("setup", {}),
  49. turns=data.get("turns", []),
  50. limits=data.get("limits", {"max_turns": 20, "turn_timeout": 120}),
  51. verify=VerifyConfig(
  52. criteria=verify_data.get("criteria", []),
  53. assertions=verify_data.get("assertions", []),
  54. observe=verify_data.get("observe", False),
  55. ),
  56. )
  57. @dataclass
  58. class RunResult:
  59. scenario: str
  60. backend: str
  61. timestamp: str
  62. session_log: str
  63. filesystem_json: str
  64. tool_calls_jsonl: str
  65. verdict_json: str
  66. meta: dict[str, Any]
  67. def save_artifacts(self, output_dir: Path) -> None:
  68. output_dir.mkdir(parents=True, exist_ok=True)
  69. (output_dir / "session.log").write_text(self.session_log)
  70. (output_dir / "filesystem.json").write_text(self.filesystem_json)
  71. (output_dir / "tool_calls.jsonl").write_text(self.tool_calls_jsonl)
  72. def save_verdict(self, output_dir: Path) -> None:
  73. output_dir.mkdir(parents=True, exist_ok=True)
  74. (output_dir / "verdict.json").write_text(self.verdict_json)
  75. (output_dir / "meta.json").write_text(json.dumps(self.meta, indent=2))
  76. def save(self, output_dir: Path) -> None:
  77. self.save_artifacts(output_dir)
  78. self.save_verdict(output_dir)
  79. def snapshot_filesystem(workdir: Path) -> str:
  80. files: list[str] = []
  81. for f in sorted(workdir.rglob("*")):
  82. if ".git" in f.parts:
  83. continue
  84. if f.is_file():
  85. files.append(str(f.relative_to(workdir)))
  86. git_status = _git_cmd(workdir, ["git", "status", "--short"])
  87. branch = _git_cmd(workdir, ["git", "branch", "--show-current"])
  88. worktree_list = _git_cmd(workdir, ["git", "worktree", "list"])
  89. return json.dumps(
  90. {
  91. "files": files,
  92. "git_status": git_status,
  93. "branch": branch,
  94. "worktree_list": worktree_list,
  95. },
  96. indent=2,
  97. )
  98. class Engine:
  99. def __init__(
  100. self,
  101. scenario_path: Path,
  102. backend_name: str,
  103. backends_dir: Path,
  104. fixtures_dir: Path,
  105. results_dir: Path,
  106. ) -> None:
  107. self.scenario = ScenarioConfig.from_yaml(scenario_path)
  108. self.backend = load_backend(backend_name, backends_dir)
  109. self.fixtures_dir = fixtures_dir
  110. self.results_dir = results_dir
  111. def run(self, *, output_dir: Path | None = None, run_suffix: str = "") -> RunResult:
  112. start_time = time.time()
  113. timestamp = datetime.now().strftime("%Y-%m-%dT%H-%M-%S")
  114. self.backend.validate_env()
  115. workdir = Path(f"/tmp/drill-{self.scenario.scenario}-{timestamp}{run_suffix}")
  116. self._setup(workdir)
  117. actual_workdir = workdir
  118. override = self.scenario.setup.get("workdir_override")
  119. if override:
  120. resolved = override.replace("${WORKDIR_NAME}", workdir.name)
  121. actual_workdir = (workdir / resolved).resolve()
  122. # Run assertions in the actual workdir (after override)
  123. assertions = self.scenario.setup.get("assertions", [])
  124. if assertions:
  125. run_assertions(assertions, actual_workdir)
  126. session_name = f"drill-{self.scenario.scenario}-{timestamp}{run_suffix}"
  127. session = TmuxSession(name=session_name, cols=self.backend.cols, rows=self.backend.rows)
  128. log_dir = self._resolve_log_dir(actual_workdir)
  129. log_snapshot = snapshot_log_dir(log_dir) if log_dir else set()
  130. session_log, actor_turns = self._run_session(session, actual_workdir)
  131. filesystem_json = snapshot_filesystem(actual_workdir)
  132. tool_calls = self._collect_tool_calls(log_dir, log_snapshot, actual_workdir)
  133. tool_calls_jsonl = "\n".join(json.dumps(tc) for tc in tool_calls)
  134. # Write artifacts to disk before assertions (assertions read from disk)
  135. if output_dir is None:
  136. output_dir = self.results_dir / self.scenario.scenario / self.backend.name / timestamp
  137. output_dir.mkdir(parents=True, exist_ok=True)
  138. (output_dir / "session.log").write_text(session_log)
  139. (output_dir / "filesystem.json").write_text(filesystem_json)
  140. (output_dir / "tool_calls.jsonl").write_text(tool_calls_jsonl)
  141. # Run deterministic assertions
  142. assertion_results: list[AssertionResult] = []
  143. if self.scenario.verify.assertions:
  144. if not tool_calls_jsonl.strip():
  145. assertion_results = [
  146. AssertionResult(
  147. command="<pre-check>",
  148. passed=False,
  149. exit_code=1,
  150. stdout="",
  151. stderr="tool_calls.jsonl is empty — session may have crashed",
  152. )
  153. ]
  154. else:
  155. assertion_results = run_verify_assertions(
  156. self.scenario.verify.assertions,
  157. output_dir,
  158. actual_workdir,
  159. )
  160. # Run LLM verifier
  161. verifier = Verifier()
  162. verdict = verifier.verify(
  163. session_log=session_log,
  164. filesystem_json=filesystem_json,
  165. tool_calls_jsonl=tool_calls_jsonl,
  166. criteria=self.scenario.verify.criteria,
  167. )
  168. # Merge assertion results into verdict
  169. for ar in assertion_results:
  170. verdict.criteria.append(ar.to_criterion_result())
  171. duration = time.time() - start_time
  172. meta: dict[str, Any] = {
  173. "scenario": self.scenario.scenario,
  174. "backend": self.backend.name,
  175. "backend_model": self.backend.model,
  176. "user_posture": self.scenario.user_posture,
  177. "timestamp": timestamp,
  178. "duration_seconds": round(duration, 1),
  179. "actor_turns": actor_turns,
  180. "actor_model": "claude-sonnet-4-6",
  181. "verifier_model": "claude-sonnet-4-6",
  182. }
  183. result = RunResult(
  184. scenario=self.scenario.scenario,
  185. backend=self.backend.name,
  186. timestamp=timestamp,
  187. session_log=session_log,
  188. filesystem_json=filesystem_json,
  189. tool_calls_jsonl=tool_calls_jsonl,
  190. verdict_json=verdict.model_dump_json(indent=2),
  191. meta=meta,
  192. )
  193. # Write verdict + meta (artifacts already on disk)
  194. (output_dir / "verdict.json").write_text(result.verdict_json)
  195. (output_dir / "meta.json").write_text(json.dumps(result.meta, indent=2))
  196. return result
  197. def _setup(self, workdir: Path) -> None:
  198. # Scenario helpers first (create_base_repo needs to run before anything else)
  199. helpers = self.scenario.setup.get("helpers", [])
  200. run_helpers(helpers, workdir, self.fixtures_dir)
  201. # Backend pre_run hooks after (e.g., codex symlink needs workdir to exist)
  202. hooks_needing_superpowers_root = {"symlink_superpowers", "link_gemini_extension"}
  203. for hook_name in self.backend.hooks.get("pre_run", []):
  204. from setup_helpers import HELPER_REGISTRY
  205. hook = HELPER_REGISTRY.get(hook_name)
  206. if hook and hook_name in hooks_needing_superpowers_root:
  207. hook(workdir, os.environ["SUPERPOWERS_ROOT"]) # ty: ignore[invalid-argument-type, too-many-positional-arguments, missing-argument]
  208. elif hook:
  209. hook(workdir) # ty: ignore[invalid-argument-type, missing-argument]
  210. def _run_session(self, session: TmuxSession, workdir: Path) -> tuple[str, int]:
  211. session.create()
  212. try:
  213. cmd = self.backend.build_command(str(workdir))
  214. session.launch(cmd, str(workdir))
  215. self._wait_for_ready(session, timeout=self.backend.startup_timeout)
  216. actor = Actor()
  217. intents = [t["intent"] for t in self.scenario.turns]
  218. actor.build_system_prompt(posture=self.scenario.user_posture, intents=intents)
  219. max_turns = self.scenario.limits.get("max_turns", 20)
  220. turn_timeout = self.backend.turn_timeout or self.scenario.limits.get(
  221. "turn_timeout", 120
  222. )
  223. all_captures: list[str] = []
  224. turn_count = 0
  225. for turn in range(max_turns):
  226. self._wait_for_ready(session, timeout=turn_timeout)
  227. capture = session.capture()
  228. all_captures.append(f"=== Turn {turn + 1} ===\n{capture}")
  229. actor.append_capture(f"Terminal output:\n{capture}")
  230. action = actor.decide()
  231. turn_count += 1
  232. if action.action == "done" or action.action == "stuck":
  233. break
  234. elif action.action == "type":
  235. session.send_keys(action.text or "")
  236. elif action.action == "key":
  237. session.send_special_key(action.key or "")
  238. final_capture = session.capture()
  239. all_captures.append(f"=== Final ===\n{final_capture}")
  240. if self.backend.shutdown.startswith("<<KEY:"):
  241. key = self.backend.shutdown[6:-2]
  242. session.send_special_key(key)
  243. else:
  244. session.send_keys(self.backend.shutdown)
  245. time.sleep(3)
  246. return "\n".join(all_captures), turn_count
  247. finally:
  248. session.kill()
  249. def _wait_for_ready(self, session: TmuxSession, timeout: float) -> None:
  250. """Wait until the agent's terminal is ready for Actor input.
  251. Returns when the terminal is quiescent AND matches the backend's
  252. ready pattern. If the backend's busy pattern matches (spinner
  253. visible, "Thinking...", timer counting), the deadline is extended
  254. by small increments up to `max_busy_seconds` total. This prevents
  255. the Actor from interrupting long-running subagent work (multi-file
  256. implementation, parallel dispatch, etc.).
  257. Exits silently if the final deadline (timeout + busy extensions)
  258. passes without reaching a ready state.
  259. """
  260. quiescence = self.backend.quiescence_seconds
  261. max_busy_extension = float(self.backend.max_busy_seconds)
  262. start = time.time()
  263. deadline = start + timeout
  264. total_busy_extended = 0.0
  265. last_output: str = ""
  266. stable_since: float | None = None
  267. while time.time() < deadline:
  268. current = session.capture()
  269. lines = current.strip().split("\n")
  270. is_busy = any(self.backend.is_busy_line(line) for line in lines)
  271. # If the agent is actively busy, extend the deadline so we
  272. # don't time out mid-subagent-work. Extensions are capped at
  273. # max_busy_seconds total across all extensions combined.
  274. if is_busy:
  275. remaining_budget = max_busy_extension - total_busy_extended
  276. if remaining_budget > 0:
  277. # Ensure we have at least 30 more seconds of headroom.
  278. needed = 30.0 - (deadline - time.time())
  279. if needed > 0:
  280. grant = min(needed, remaining_budget)
  281. deadline += grant
  282. total_busy_extended += grant
  283. # Strip animated elements so they don't reset the quiescence timer:
  284. # - Time counters: "Thinking... (4m 1s)" or "(esc to cancel, 4m 1s)"
  285. # - Braille spinner characters that rotate every frame
  286. normalized = re.sub(r"\((?:esc to cancel, )?(?:\d+[hms]\s*)+\)", "(…)", current)
  287. normalized = re.sub(r"[⠇⠏⠋⠙⠹⠸⠼⠴⠦⠧⠶⠾⠽⠻⠿]", "·", normalized)
  288. if normalized != last_output:
  289. last_output = normalized
  290. stable_since = time.time()
  291. elif stable_since and (time.time() - stable_since) >= quiescence:
  292. if is_busy:
  293. stable_since = None # Reset — agent is still working
  294. elif any(self.backend.is_ready_line(line) for line in lines):
  295. return
  296. time.sleep(0.5)
  297. def _resolve_log_dir(self, workdir: Path) -> Path | None:
  298. """Resolve the log directory for the given backend and workdir.
  299. Claude Code stores logs at ~/.claude/projects/<encoded-path>/
  300. where the path is the real workdir with / replaced by -.
  301. Codex stores logs at ~/.codex/sessions/.
  302. """
  303. if self.backend.family == "claude":
  304. real_workdir = workdir.resolve()
  305. encoded = str(real_workdir).replace("/", "-")
  306. log_dir = Path.home() / ".claude" / "projects" / encoded
  307. return log_dir
  308. elif self.backend.family == "codex":
  309. # Codex stores at ~/.codex/sessions/YYYY/MM/DD/rollout-*.jsonl
  310. return Path.home() / ".codex" / "sessions"
  311. elif self.backend.family == "gemini":
  312. # Gemini stores at ~/.gemini/tmp/<project-name>/chats/session-*.json
  313. # Project name is the workdir basename, lowercased
  314. project = workdir.resolve().name.lower()
  315. return Path.home() / ".gemini" / "tmp" / project
  316. pattern = self.backend.session_logs.get("pattern", "")
  317. if not pattern:
  318. return None
  319. expanded = os.path.expanduser(pattern)
  320. parts = expanded.split("*")[0].rstrip("/")
  321. return Path(parts)
  322. def _collect_tool_calls(
  323. self, log_dir: Path | None, snapshot: set[str], workdir: Path
  324. ) -> list[dict[str, Any]]:
  325. if log_dir is None:
  326. return []
  327. new_files = collect_new_logs(log_dir, snapshot)
  328. if self.backend.family == "codex":
  329. new_files = filter_codex_logs_by_cwd(new_files, str(workdir.resolve()))
  330. normalizer = NORMALIZERS.get(self.backend.family)
  331. if not normalizer:
  332. return []
  333. results: list[dict[str, Any]] = []
  334. for log_file in new_files:
  335. results.extend(normalizer(log_file.read_text()))
  336. return results
  337. def _git_cmd(workdir: Path, cmd: list[str]) -> str:
  338. result = subprocess.run(cmd, cwd=workdir, capture_output=True, text=True)
  339. return result.stdout.strip()