nju-coding-agent-harness 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (67) hide show
  1. harness/__init__.py +0 -0
  2. harness/agent.py +334 -0
  3. harness/config.py +45 -0
  4. harness/credentials.py +86 -0
  5. harness/fake_llm.py +34 -0
  6. harness/guardrails.py +46 -0
  7. harness/hooks.py +57 -0
  8. harness/llm.py +135 -0
  9. harness/main.py +400 -0
  10. harness/mcp.py +253 -0
  11. harness/memory.py +125 -0
  12. harness/policy.py +101 -0
  13. harness/registry.py +322 -0
  14. harness/sandbox.py +147 -0
  15. harness/state.py +56 -0
  16. harness/tests/__init__.py +0 -0
  17. harness/tests/conftest.py +14 -0
  18. harness/tests/fixtures/fake_mcp_server.py +64 -0
  19. harness/tests/mechanism_demo/__init__.py +1 -0
  20. harness/tests/mechanism_demo/demo_1_guardrail_deny.py +65 -0
  21. harness/tests/mechanism_demo/demo_2_feedback_change.py +85 -0
  22. harness/tests/mechanism_demo/demo_3_hitl_trace.py +68 -0
  23. harness/tests/test_acceptance_matrix.py +434 -0
  24. harness/tests/test_agent_context.py +101 -0
  25. harness/tests/test_agent_core.py +99 -0
  26. harness/tests/test_agent_end.py +53 -0
  27. harness/tests/test_agent_feedback.py +59 -0
  28. harness/tests/test_config.py +20 -0
  29. harness/tests/test_credentials.py +75 -0
  30. harness/tests/test_docs.py +41 -0
  31. harness/tests/test_guardrails.py +24 -0
  32. harness/tests/test_hooks.py +33 -0
  33. harness/tests/test_llm.py +98 -0
  34. harness/tests/test_mcp.py +136 -0
  35. harness/tests/test_mechanism_demo.py +31 -0
  36. harness/tests/test_memory.py +49 -0
  37. harness/tests/test_perf_smoke.py +40 -0
  38. harness/tests/test_policy.py +64 -0
  39. harness/tests/test_registry.py +172 -0
  40. harness/tests/test_repl.py +320 -0
  41. harness/tests/test_sandbox.py +96 -0
  42. harness/tests/test_security_scan.py +29 -0
  43. harness/tests/test_skeleton.py +4 -0
  44. harness/tests/test_state.py +42 -0
  45. harness/tests/test_tools_ask.py +53 -0
  46. harness/tests/test_tools_bash.py +49 -0
  47. harness/tests/test_tools_files.py +52 -0
  48. harness/tests/test_tools_memory.py +58 -0
  49. harness/tests/test_tools_notes.py +25 -0
  50. harness/tests/test_tools_skills.py +59 -0
  51. harness/tests/test_tools_subagent.py +146 -0
  52. harness/tests/test_tools_web.py +48 -0
  53. harness/tools/__init__.py +0 -0
  54. harness/tools/ask.py +46 -0
  55. harness/tools/bash.py +38 -0
  56. harness/tools/files.py +125 -0
  57. harness/tools/memory.py +91 -0
  58. harness/tools/notes.py +56 -0
  59. harness/tools/search.py +133 -0
  60. harness/tools/skills.py +133 -0
  61. harness/tools/subagent.py +110 -0
  62. harness/tools/web.py +56 -0
  63. harness/transcript.py +32 -0
  64. nju_coding_agent_harness-0.1.0.dist-info/METADATA +11 -0
  65. nju_coding_agent_harness-0.1.0.dist-info/RECORD +67 -0
  66. nju_coding_agent_harness-0.1.0.dist-info/WHEEL +5 -0
  67. nju_coding_agent_harness-0.1.0.dist-info/top_level.txt +1 -0
harness/__init__.py ADDED
File without changes
harness/agent.py ADDED
@@ -0,0 +1,334 @@
1
+ from __future__ import annotations
2
+
3
+ import json
4
+ from dataclasses import dataclass, field
5
+ from datetime import datetime
6
+ from pathlib import Path
7
+ from typing import TYPE_CHECKING, Any, Callable
8
+
9
+ from harness.guardrails import evaluate
10
+ from harness.registry import (
11
+ Context,
12
+ ToolResult,
13
+ build_request_tools,
14
+ validate_args,
15
+ )
16
+
17
+ if TYPE_CHECKING:
18
+ from harness.config import Config
19
+ from harness.fake_llm import FakeLLM
20
+ from harness.hooks import HookBus
21
+ from harness.llm import LLM
22
+ from harness.memory import MemoryStore
23
+ from harness.policy import Policy
24
+ from harness.sandbox import Sandbox
25
+ from harness.state import StateMachine
26
+
27
+ SYSTEM_PROMPT = (
28
+ "你是编码代理助手。可调用工具完成任务:"
29
+ "先规划,再执行,最后给出最终答案。"
30
+ )
31
+
32
+ _ASK_OPTIONS = ["y", "n", "always_allow", "never_allow"]
33
+
34
+
35
+ @dataclass
36
+ class AgentResult:
37
+ text: str = ""
38
+ steps_used: int = 0
39
+ tool_results: list[dict] = field(default_factory=list)
40
+ policy_changes: list[dict] = field(default_factory=list)
41
+ messages: list[dict] = field(default_factory=list)
42
+ failed_sequence: int = 0
43
+ transcript_path: str | None = None
44
+
45
+
46
+ class Agent:
47
+ def __init__(
48
+ self,
49
+ llm: LLM,
50
+ registry: dict[str, Any],
51
+ sandbox: Sandbox,
52
+ hooks: HookBus,
53
+ policy: Policy,
54
+ state: StateMachine,
55
+ memory: MemoryStore | None,
56
+ config: Config,
57
+ ask_callback: Callable | None = None,
58
+ on_text: Callable[[str], None] | None = None,
59
+ ):
60
+ self.llm = llm
61
+ self.registry = registry
62
+ self.sandbox = sandbox
63
+ self.hooks = hooks
64
+ self.policy = policy
65
+ self.state = state
66
+ self.memory = memory
67
+ self.config = config
68
+ self.ask_callback = ask_callback
69
+ self.on_text = on_text
70
+ self._compress_calls = 0
71
+ self.warnings: list[str] = []
72
+ self._tool_calls: list[dict] = []
73
+
74
+ def context_for_tool(self) -> Context:
75
+ return Context(
76
+ workspace=self.config.workspace,
77
+ sandbox=self.sandbox,
78
+ hooks=self.hooks,
79
+ policy=self.policy,
80
+ state=self.state,
81
+ memory=self.memory,
82
+ config=self.config,
83
+ ask_callback=self.ask_callback,
84
+ llm=self.llm,
85
+ registry=self.registry,
86
+ )
87
+
88
+ def pipeline(self, call: dict, ctx: Context) -> ToolResult:
89
+ name = call["name"]
90
+ args = call["arguments"]
91
+ verdict = evaluate(self.policy.rules, name, args)
92
+ if verdict.action == "deny":
93
+ return ToolResult(status="error", error=f"guardrail denied: {verdict.reason}")
94
+ if verdict.action == "ask":
95
+ self.state.fire("approval_needed", "guardrail")
96
+ answer = self._ask(verdict.matched_rule, verdict.reason)
97
+ self.policy.apply_answer(verdict.matched_rule, answer)
98
+ if self.state.state == "awaiting_user":
99
+ self.state.fire("user_answered", "user")
100
+ if answer in ("n", "never_allow"):
101
+ return ToolResult(status="error", error=f"guardrail denied: {verdict.reason}")
102
+ if self.state.state == "awaiting_user":
103
+ self.state.fire("user_answered", "user")
104
+ args, ok = self.hooks.pre_tool_use(name, args)
105
+ if not ok:
106
+ return ToolResult(status="error", error=f"pre_tool_use hook rejected {name}")
107
+ self.state.fire("tool_requested", "loop")
108
+ tool = self.registry.get(name)
109
+ if tool is None:
110
+ result = ToolResult(status="error", error=f"unknown tool: {name}")
111
+ else:
112
+ err = validate_args(tool.parameters, args, self.config.workspace)
113
+ if err is not None:
114
+ result = ToolResult(status="error", error=f"参数错误: {err}")
115
+ else:
116
+ try:
117
+ result = tool.handler(args, ctx)
118
+ except Exception as exc:
119
+ result = ToolResult(
120
+ status="error",
121
+ error=f"工具 {name} 异常:{type(exc).__name__}: {exc}",
122
+ )
123
+ self.state.fire("tool_finished", "loop")
124
+ self.hooks.post_tool_use(name, args, result)
125
+ return result
126
+
127
+ def _emit_text(self, text: str) -> None:
128
+ if text and self.on_text is not None:
129
+ self.on_text(text)
130
+
131
+ def _ask(self, rule, reason: str) -> str:
132
+ question = f"是否允许执行该操作?\n规则: {rule.pattern}\n原因: {reason}"
133
+ if self.ask_callback is None:
134
+ return "n"
135
+ try:
136
+ answer = self.ask_callback(question, list(_ASK_OPTIONS))
137
+ except Exception:
138
+ return "n"
139
+ if answer not in _ASK_OPTIONS:
140
+ return "n"
141
+ return answer
142
+
143
+ def _check_budget(self, messages: list[dict]) -> bool:
144
+ total = sum(len(m.get("content", "")) / 4 for m in messages)
145
+ return total > self.config.max_budget_tokens
146
+
147
+ def _compress(self, messages: list[dict]) -> list[dict]:
148
+ keep = self.config.compression_keep_turns
149
+ oldest = messages[:-keep]
150
+ if not oldest:
151
+ return messages
152
+ self._compress_calls += 1
153
+ if self._compress_calls > self.config.compression_max_rounds:
154
+ return self._drop_oldest(messages, keep)
155
+ try:
156
+ summary = self.llm.complete(
157
+ [
158
+ {
159
+ "role": "system",
160
+ "content": (
161
+ "请将以下较早回合总结为简洁摘要,"
162
+ "保留关键事实、决定与结果:"
163
+ ),
164
+ },
165
+ {
166
+ "role": "user",
167
+ "content": json.dumps(oldest, ensure_ascii=False),
168
+ },
169
+ ],
170
+ tools=[],
171
+ )
172
+ except Exception:
173
+ return self._drop_oldest(messages, keep)
174
+ text = (summary.text or "").strip()
175
+ if not text:
176
+ return self._drop_oldest(messages, keep)
177
+ return [{"role": "system", "content": f"[summary] {text}"}] + messages[-keep:]
178
+
179
+ def _drop_oldest(self, messages: list[dict], keep: int) -> list[dict]:
180
+ window = messages[-keep:]
181
+ if window and window[0]["role"] not in ("user", "system"):
182
+ for m in reversed(messages[:-keep]):
183
+ if m["role"] in ("user", "system"):
184
+ return [m] + window
185
+ return window
186
+
187
+ def run(self, task: str) -> AgentResult:
188
+ result = AgentResult()
189
+ messages = [
190
+ {"role": "system", "content": SYSTEM_PROMPT},
191
+ {"role": "user", "content": task},
192
+ ]
193
+ if self.memory is not None:
194
+ for chunk in self.memory.top_k_chunks(task):
195
+ messages.append(
196
+ {"role": "system", "content": f"[memory] {chunk['chunk']}"}
197
+ )
198
+ call_uid = 0
199
+ fail_seq = 0
200
+ fail_tool: str | None = None
201
+ max_fail_seq = 0
202
+ self._compress_calls = 0
203
+ self._tool_calls = []
204
+ self.state.fire("task_submitted", "loop")
205
+ while result.steps_used < self.config.max_steps:
206
+ if self._check_budget(messages):
207
+ messages = self._compress(messages)
208
+ response = self.llm.complete(messages, build_request_tools(self.registry))
209
+ result.steps_used += 1
210
+ if not response.tool_calls:
211
+ final = response.text or "任务完成"
212
+ self._emit_text(final)
213
+ messages.append({"role": "assistant", "content": final})
214
+ result.text = final
215
+ return self._finish(result, messages, max_fail_seq)
216
+ self._emit_text(response.text)
217
+ assistant_call = []
218
+ for i, call in enumerate(response.tool_calls):
219
+ assistant_call.append({
220
+ "id": f"call_{call_uid}",
221
+ "type": "function",
222
+ "function": {
223
+ "name": call["name"],
224
+ "arguments": json.dumps(call["arguments"], ensure_ascii=False),
225
+ },
226
+ })
227
+ call_uid += 1
228
+ messages.append({
229
+ "role": "assistant",
230
+ "content": response.text,
231
+ "tool_calls": assistant_call,
232
+ })
233
+ for i, call in enumerate(response.tool_calls):
234
+ tool_id = f"call_{call_uid - len(response.tool_calls) + i}"
235
+ tool_result = self.pipeline(call, self.context_for_tool())
236
+ result.tool_results.append(tool_result)
237
+ self._tool_calls.append({"name": call["name"], "arguments": call["arguments"]})
238
+ messages.append({
239
+ "role": "tool",
240
+ "tool_call_id": tool_id,
241
+ "name": call["name"],
242
+ "content": json.dumps(
243
+ self._result_to_dict(tool_result), ensure_ascii=False
244
+ ),
245
+ })
246
+ norm = self._result_to_dict(tool_result)
247
+ failed = norm.get("status") != "success" or bool(norm.get("error"))
248
+ if failed:
249
+ if call["name"] == fail_tool:
250
+ fail_seq += 1
251
+ else:
252
+ fail_seq = 1
253
+ fail_tool = call["name"]
254
+ max_fail_seq = max(max_fail_seq, fail_seq)
255
+ if fail_seq >= self.config.failure_budget:
256
+ final = (
257
+ f"连续失败 {fail_seq} 次(工具 {fail_tool}),"
258
+ f"超过失败预算 {self.config.failure_budget},停止重试。"
259
+ )
260
+ self._emit_text(final)
261
+ messages.append({"role": "assistant", "content": final})
262
+ result.text = final
263
+ return self._finish(result, messages, max_fail_seq)
264
+ else:
265
+ fail_seq = 0
266
+ fail_tool = None
267
+ final = f"达到步数上限 {self.config.max_steps},任务终止,未挂死。"
268
+ self._emit_text(final)
269
+ messages.append({"role": "assistant", "content": final})
270
+ result.text = final
271
+ return self._finish(result, messages, max_fail_seq)
272
+
273
+ def _finish(self, result: AgentResult, messages: list[dict], max_fail_seq: int) -> AgentResult:
274
+ result.messages = messages
275
+ self.messages = messages
276
+ result.failed_sequence = max_fail_seq
277
+ self.state.fire("final_answer", "loop")
278
+ self._finalize(result)
279
+ return result
280
+
281
+ def _finalize(self, result: AgentResult) -> None:
282
+ self.hooks.session_data["tool_calls"] = list(self._tool_calls)
283
+ self.hooks.session_data["policy_changes"] = self.policy.changes()
284
+ self.hooks.session_end(self.messages)
285
+ if self.hooks.transcript_dir is not None:
286
+ files = [p for p in Path(self.hooks.transcript_dir).glob("*.json")]
287
+ if files:
288
+ result.transcript_path = str(
289
+ max(files, key=lambda p: p.stat().st_mtime)
290
+ )
291
+ if self.memory is not None and self.llm is not None:
292
+ self._consolidate()
293
+
294
+ def _consolidate(self) -> None:
295
+ try:
296
+ response = self.llm.complete(
297
+ [
298
+ {
299
+ "role": "system",
300
+ "content": (
301
+ "请总结本次会话的关键事实、决策与工具结果,"
302
+ "作为长期记忆保存:"
303
+ ),
304
+ }
305
+ ]
306
+ + self.messages,
307
+ tools=[],
308
+ )
309
+ except Exception as exc:
310
+ self.warnings.append(f"memory consolidation failed: {exc}")
311
+ return
312
+ summary = (response.text or "").strip()
313
+ if not summary:
314
+ self.warnings.append("memory consolidation skipped: empty summary")
315
+ return
316
+ try:
317
+ self.memory.save(
318
+ f"session-summary-{datetime.now():%Y%m%d-%H%M%S}", summary
319
+ )
320
+ except Exception as exc:
321
+ self.warnings.append(f"memory save failed: {exc}")
322
+
323
+ @staticmethod
324
+ def _result_to_dict(tool_result) -> dict:
325
+ if isinstance(tool_result, ToolResult):
326
+ return {
327
+ "status": tool_result.status,
328
+ "output": tool_result.output,
329
+ "error": tool_result.error,
330
+ "exit_code": tool_result.exit_code,
331
+ }
332
+ if isinstance(tool_result, dict):
333
+ return tool_result
334
+ return {"status": "success", "output": str(tool_result)}
harness/config.py ADDED
@@ -0,0 +1,45 @@
1
+ from __future__ import annotations
2
+
3
+ import tomllib
4
+ import warnings
5
+ from dataclasses import dataclass, field, fields
6
+ from pathlib import Path
7
+
8
+
9
+ @dataclass
10
+ class Config:
11
+ model: str = "deepseek-chat"
12
+ base_url: str = "https://api.deepseek.com"
13
+ max_steps: int = 50
14
+ failure_budget: int = 3
15
+ tool_timeout: int = 30
16
+ memory_top_k: int = 2
17
+ max_budget_tokens: int = 6000
18
+ compression_keep_turns: int = 10
19
+ compression_max_rounds: int = 3
20
+ workspace: Path = field(default_factory=Path.cwd)
21
+ max_output_bytes: int = 51200
22
+ mcp_servers: list[dict] = field(default_factory=list)
23
+
24
+ @classmethod
25
+ def load(cls, path: Path | None = None) -> Config:
26
+ cfg = cls()
27
+ if path is None:
28
+ return cfg
29
+ path = Path(path)
30
+ if not path.exists():
31
+ return cfg
32
+ try:
33
+ with path.open("rb") as f:
34
+ data = tomllib.load(f)
35
+ except Exception as exc:
36
+ warnings.warn(f"config 解析失败,使用默认值: {exc}")
37
+ return cfg
38
+ known = {f.name for f in fields(cls)}
39
+ for key, value in data.items():
40
+ if key not in known:
41
+ continue
42
+ if key == "workspace":
43
+ value = Path(value)
44
+ setattr(cfg, key, value)
45
+ return cfg
harness/credentials.py ADDED
@@ -0,0 +1,86 @@
1
+ from __future__ import annotations
2
+
3
+ import getpass
4
+ from pathlib import Path
5
+
6
+ _ENV_KEY = "DEEPSEEK_API_KEY"
7
+
8
+
9
+ class CredentialStore:
10
+ def __init__(self, service: str = "coding-agent-harness", env_file: str = ".env", keyring_backend=None):
11
+ self.service = service
12
+ self.env_file = Path(env_file)
13
+ self._keyring = keyring_backend
14
+ if self._keyring is None:
15
+ try:
16
+ import keyring
17
+ self._keyring = keyring.get_keyring()
18
+ except Exception:
19
+ self._keyring = None
20
+
21
+ def _keyring_get(self, user: str) -> str | None:
22
+ if self._keyring is None:
23
+ return None
24
+ try:
25
+ return self._keyring.get_password(self.service, user)
26
+ except Exception:
27
+ return None
28
+
29
+ def _env_value(self) -> str | None:
30
+ try:
31
+ if not self.env_file.exists():
32
+ return None
33
+ for line in self.env_file.read_text(encoding="utf-8").splitlines():
34
+ line = line.strip()
35
+ if line.startswith(_ENV_KEY + "="):
36
+ value = line[len(_ENV_KEY) + 1:].strip()
37
+ if value:
38
+ return value
39
+ except Exception:
40
+ return None
41
+ return None
42
+
43
+ def get(self) -> str | None:
44
+ key = self._keyring_get("api_key")
45
+ if key:
46
+ return key
47
+ return self._env_value()
48
+
49
+ def set(self, key: str) -> None:
50
+ if self._keyring is None:
51
+ raise RuntimeError("keyring 不可用(无凭据服务),无法保存 API Key")
52
+ try:
53
+ self._keyring.set_password(self.service, "api_key", key)
54
+ except Exception:
55
+ raise RuntimeError("keyring 写入失败,无法保存 API Key") from None
56
+
57
+ def clear(self) -> None:
58
+ if self._keyring is None:
59
+ raise RuntimeError("keyring 不可用(无凭据服务),无法删除 API Key")
60
+ try:
61
+ self._keyring.delete_password(self.service, "api_key")
62
+ except Exception:
63
+ raise RuntimeError("keyring 删除失败,无法删除 API Key") from None
64
+
65
+ def verified_at(self) -> str | None:
66
+ return self._keyring_get("verified_at")
67
+
68
+ def status(self) -> dict:
69
+ if self._keyring_get("api_key"):
70
+ source = "keyring"
71
+ elif self._env_value():
72
+ source = "env"
73
+ else:
74
+ source = None
75
+ return {
76
+ "configured": source is not None,
77
+ "source": source,
78
+ "verified_at": self.verified_at(),
79
+ }
80
+
81
+
82
+ def wizard_enter_key() -> str:
83
+ key = getpass.getpass("请粘贴 API Key(输入不可见): ")
84
+ if not key.strip():
85
+ raise ValueError("API Key 不能为空")
86
+ return key
harness/fake_llm.py ADDED
@@ -0,0 +1,34 @@
1
+ from __future__ import annotations
2
+
3
+ from dataclasses import dataclass
4
+
5
+ from harness.llm import LLM, LLMResult
6
+
7
+
8
+ @dataclass
9
+ class FakeTurn:
10
+ text: str = ""
11
+ tool_calls: list[dict] | None = None
12
+ usage_approx: int = 10
13
+
14
+
15
+ class FakeLLM(LLM):
16
+ """脚本化 LLM:按序重放 turns,耗尽后重放最后一个;绝无网络。"""
17
+
18
+ def __init__(self, turns: list[FakeTurn]):
19
+ self.turns = list(turns)
20
+ self.calls = 0
21
+ self.turn_index = 0
22
+
23
+ def complete(self, messages: list[dict], tools: list[dict]) -> LLMResult:
24
+ self.calls += 1
25
+ if not self.turns:
26
+ return LLMResult(text="", tool_calls=[], usage={"approx_tokens": 0})
27
+ index = min(self.turn_index, len(self.turns) - 1)
28
+ self.turn_index += 1
29
+ turn = self.turns[index]
30
+ return LLMResult(
31
+ text=turn.text,
32
+ tool_calls=turn.tool_calls or [],
33
+ usage={"approx_tokens": turn.usage_approx},
34
+ )
harness/guardrails.py ADDED
@@ -0,0 +1,46 @@
1
+ import json
2
+ import re
3
+ from dataclasses import dataclass
4
+
5
+
6
+ @dataclass
7
+ class Rule:
8
+ pattern: str
9
+ action: str
10
+ source: str
11
+
12
+
13
+ @dataclass
14
+ class Verdict:
15
+ action: str
16
+ matched_rule: Rule | None
17
+ reason: str
18
+
19
+
20
+ def pattern_matches(pattern: str, tool_name: str, args: dict) -> bool:
21
+ if ":" in pattern:
22
+ tool, _, regex = pattern.partition(":")
23
+ if tool != tool_name:
24
+ return False
25
+ return re.search(regex, json.dumps(args, ensure_ascii=False)) is not None
26
+ return pattern == tool_name
27
+
28
+
29
+ def default_rules() -> list[Rule]:
30
+ return [
31
+ Rule(r"bash:rm -rf\s+[\\/]?(/|[A-Z]:[\\/]|[A-Z](\b|:)|etc\b|boot\b|bin\b)",
32
+ "deny", "builtin"),
33
+ Rule(r"bash:.*:\(\)\s*\{.*:.*\};", "deny", "builtin"),
34
+ Rule(r"bash:format.*", "deny", "builtin"),
35
+ Rule(r"bash:del /f.*", "deny", "builtin"),
36
+ ]
37
+
38
+
39
+ def evaluate(rules: list[Rule], tool_name: str, args: dict) -> Verdict:
40
+ matched = None
41
+ for rule in rules:
42
+ if pattern_matches(rule.pattern, tool_name, args):
43
+ matched = rule
44
+ if matched is None:
45
+ return Verdict("allow", None, "no rule matched")
46
+ return Verdict(matched.action, matched, f"matched {matched.pattern}")
harness/hooks.py ADDED
@@ -0,0 +1,57 @@
1
+ from datetime import datetime
2
+ from typing import Callable
3
+
4
+ from harness.transcript import default_session_end_hook
5
+
6
+
7
+ class HookBus:
8
+ def __init__(self, transcript_dir=None):
9
+ self.transcript_dir = transcript_dir
10
+ self.session_data: dict = {"tool_calls": [], "policy_changes": []}
11
+ self._hooks: dict[str, list[Callable]] = {}
12
+ self._records: list[dict] = []
13
+ self.errors: list[str] = []
14
+ if transcript_dir is not None:
15
+ self.register("session_end", default_session_end_hook(transcript_dir, self.session_data))
16
+
17
+ def register(self, name: str, fn: Callable) -> None:
18
+ self._hooks.setdefault(name, []).append(fn)
19
+
20
+ def _record(self, hook_name: str, tool_name, args, result) -> None:
21
+ self._records.append({
22
+ "hook_name": hook_name,
23
+ "tool_name": tool_name,
24
+ "args": args,
25
+ "result": result,
26
+ "timestamp": datetime.now().isoformat(),
27
+ })
28
+
29
+ def pre_tool_use(self, tool_name: str, args: dict) -> tuple[dict, bool]:
30
+ ok = True
31
+ for hook in self._hooks.get("pre", []):
32
+ try:
33
+ args, flag = hook(tool_name, args)
34
+ ok = ok and flag
35
+ self._record("pre", tool_name, args, None)
36
+ except Exception as exc:
37
+ self.errors.append(f"pre_tool_use({tool_name}): {exc}")
38
+ return args, ok
39
+
40
+ def post_tool_use(self, tool_name: str, args: dict, result) -> None:
41
+ for hook in self._hooks.get("post", []):
42
+ try:
43
+ hook(tool_name, args, result)
44
+ self._record("post", tool_name, args, result)
45
+ except Exception as exc:
46
+ self.errors.append(f"post_tool_use({tool_name}): {exc}")
47
+
48
+ def session_end(self, messages: list[dict]) -> None:
49
+ for hook in self._hooks.get("session_end", []):
50
+ try:
51
+ hook(messages)
52
+ self._record("session_end", None, None, None)
53
+ except Exception as exc:
54
+ self.errors.append(f"session_end: {exc}")
55
+
56
+ def records(self) -> list[dict]:
57
+ return self._records