code.py 12 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275
  1. #!/usr/bin/env python3
  2. """
  3. s04: 钩子系统 — 把扩展逻辑从循环中移出,挂到钩子上。
  4. 用户输入问题
  5. ┌──────────────────┐
  6. │ UserPromptSubmit │ ── LLM 调用前触发 trigger_hooks()
  7. └────────┬─────────┘
  8. ┌────────────┐ ┌─────────────────────────────┐
  9. │ messages │────▶│ LLM (stop_reason=tool_use?)│
  10. └────────────┘ │ 否 ──▶ Stop hooks ──▶ 退出 │
  11. │ 是 ──▶ tool_use block ──┐ │
  12. └────────────────────────────┘ │
  13. ┌──────────────────┐
  14. │ trigger_hooks() │
  15. │ PreToolUse: │
  16. │ permission_hook │
  17. │ log_hook │
  18. └───────┬──────────┘
  19. │ (未被拦截)
  20. ┌───────▼──────────┐
  21. │ TOOL_HANDLERS[x] │
  22. └───────┬──────────┘
  23. ┌───────▼──────────┐
  24. │ trigger_hooks() │
  25. │ PostToolUse: │
  26. │ large_output │
  27. └───────┬──────────┘
  28. results ──▶ 回到 messages
  29. """
  30. import os, subprocess
  31. from pathlib import Path
  32. try:
  33. import readline
  34. readline.parse_and_bind('set bind-tty-special-chars off')
  35. readline.parse_and_bind('set input-meta on')
  36. readline.parse_and_bind('set output-meta on')
  37. readline.parse_and_bind('set convert-meta off')
  38. except ImportError:
  39. pass
  40. from anthropic import Anthropic
  41. from dotenv import load_dotenv
  42. load_dotenv(override=True)
  43. if os.getenv("ANTHROPIC_BASE_URL"):
  44. os.environ.pop("ANTHROPIC_AUTH_TOKEN", None)
  45. WORKDIR = Path.cwd()
  46. client = Anthropic(base_url=os.getenv("ANTHROPIC_BASE_URL"))
  47. MODEL = os.environ["MODEL_ID"]
  48. SYSTEM = f"你是位于 {WORKDIR}. 使用工具解决任务。直接行动,不要只解释。"
  49. # ═══════════════════════════════════════════════════════════
  50. # 来自 s02-s03 : 工具实现
  51. # ═══════════════════════════════════════════════════════════
  52. def run_bash(command: str) -> str:
  53. try:
  54. r = subprocess.run(command, shell=True, cwd=WORKDIR,
  55. capture_output=True, text=True, timeout=120)
  56. out = (r.stdout + r.stderr).strip()
  57. return out[:50000] if out else "(无输出)"
  58. except subprocess.TimeoutExpired:
  59. return "错误:执行超时(120 秒)"
  60. def run_read(path: str, limit: int | None = None) -> str:
  61. try:
  62. file_path = (WORKDIR / path).resolve()
  63. lines = file_path.read_text().splitlines()
  64. if limit and limit < len(lines):
  65. lines = lines[:limit] + [f"... ({len(lines) - limit} 行更多内容)"]
  66. return "\n".join(lines)
  67. except Exception as e:
  68. return f"错误:{e}"
  69. def run_write(path: str, content: str) -> str:
  70. try:
  71. file_path = (WORKDIR / path).resolve()
  72. file_path.parent.mkdir(parents=True, exist_ok=True)
  73. file_path.write_text(content)
  74. return f"已写入 {len(content)} 字节到 {path}"
  75. except Exception as e:
  76. return f"错误:{e}"
  77. def run_edit(path: str, old_text: str, new_text: str) -> str:
  78. try:
  79. file_path = (WORKDIR / path).resolve()
  80. text = file_path.read_text()
  81. if old_text not in text:
  82. return f"错误:在文件中未找到目标文本:{path}"
  83. file_path.write_text(text.replace(old_text, new_text, 1))
  84. return f"已编辑 {path}"
  85. except Exception as e:
  86. return f"错误:{e}"
  87. def run_glob(pattern: str) -> str:
  88. import glob as g
  89. try:
  90. results = []
  91. for match in g.glob(pattern, root_dir=WORKDIR):
  92. if (WORKDIR / match).resolve().is_relative_to(WORKDIR):
  93. results.append(match)
  94. return "\n".join(results) if results else "(无匹配)"
  95. except Exception as e:
  96. return f"错误:{e}"
  97. TOOLS = [
  98. {"name": "bash", "description": "运行一条 shell 命令。",
  99. "input_schema": {"type": "object", "properties": {"command": {"type": "string"}}, "required": ["command"]}},
  100. {"name": "read_file", "description": "读取文件内容。",
  101. "input_schema": {"type": "object", "properties": {"path": {"type": "string"}, "limit": {"type": "integer"}}, "required": ["path"]}},
  102. {"name": "write_file", "description": "向文件写入内容。",
  103. "input_schema": {"type": "object", "properties": {"path": {"type": "string"}, "content": {"type": "string"}}, "required": ["path", "content"]}},
  104. {"name": "edit_file", "description": "在文件中替换一次完全匹配的文本。",
  105. "input_schema": {"type": "object", "properties": {"path": {"type": "string"}, "old_text": {"type": "string"}, "new_text": {"type": "string"}}, "required": ["path", "old_text", "new_text"]}},
  106. {"name": "glob", "description": "查找匹配 glob 模式的文件。",
  107. "input_schema": {"type": "object", "properties": {"pattern": {"type": "string"}}, "required": ["pattern"]}},
  108. ]
  109. TOOL_HANDLERS = {
  110. "bash": run_bash, "read_file": run_read, "write_file": run_write,
  111. "edit_file": run_edit, "glob": run_glob,
  112. }
  113. # ═══════════════════════════════════════════════════════════
  114. # 新增于 s04: 钩子系统(s03 权限逻辑现在通过钩子实现)
  115. # ═══════════════════════════════════════════════════════════
  116. HOOKS = {"UserPromptSubmit": [], "PreToolUse": [], "PostToolUse": [], "Stop": []}
  117. def register_hook(event: str, callback):
  118. HOOKS[event].append(callback)
  119. def trigger_hooks(event: str, *args):
  120. for callback in HOOKS[event]:
  121. result = callback(*args)
  122. if result is not None: # 教学快捷方式:拦截这个工具调用
  123. return result
  124. return None
  125. # s03 权限检查逻辑,现在封装成钩子
  126. DENY_LIST = ["rm -rf /", "sudo", "shutdown", "reboot", "mkfs", "dd if="]
  127. DESTRUCTIVE = ["rm ", "> /etc/", "chmod 777"]
  128. def permission_hook(block):
  129. """PreToolUse:这里承载从 s03 迁移过来的 check_permission() 逻辑。"""
  130. if block.name == "bash":
  131. for pattern in DENY_LIST:
  132. if pattern in block.input.get("command", ""):
  133. print(f"\n\033[31m⛔ 已拦截:'{pattern}'\033[0m")
  134. return "被拒绝列表拒绝授权"
  135. for kw in DESTRUCTIVE:
  136. if kw in block.input.get("command", ""):
  137. print(f"\n\033[33m⚠ 可能具有破坏性的命令\033[0m")
  138. print(f" 工具:{block.name}({block.input})")
  139. choice = input(" 是否允许?[y/N] ").strip().lower()
  140. if choice not in ("y", "yes"):
  141. return "用户拒绝授权"
  142. if block.name in ("read_file", "write_file", "edit_file"):
  143. path = block.input.get("path", "")
  144. if not (WORKDIR / path).resolve().is_relative_to(WORKDIR):
  145. print(f"\n\033[33m⚠ 访问工作区外部路径\033[0m")
  146. print(f" 工具:{block.name}({block.input})")
  147. choice = input(" 是否允许?[y/N] ").strip().lower()
  148. if choice not in ("y", "yes"):
  149. return "用户拒绝授权"
  150. return None
  151. def log_hook(block):
  152. """PreToolUse:记录每一次工具调用。"""
  153. args_preview = str(list(block.input.values())[:2])[:60]
  154. print(f"\033[90m[钩子] {block.name}({args_preview})\033[0m")
  155. return None
  156. def large_output_hook(block, output):
  157. """PostToolUse:对大型输出给出提醒。"""
  158. if len(str(output)) > 100000:
  159. print(f"\033[33m[钩子] ⚠ 来自以下工具的大输出:{block.name}: {len(str(output))} 个字符\033[0m")
  160. return None
  161. # 用户提交提示词钩子:在用户输入到达 LLM 前记录它
  162. def context_inject_hook(query: str):
  163. print(f"\033[90m[钩子] UserPromptSubmit: 工作目录:{WORKDIR}\033[0m")
  164. return None
  165. # 停止钩子:在循环即将退出时打印摘要
  166. def summary_hook(messages: list):
  167. tool_count = sum(1 for m in messages
  168. for b in (m.get("content") if isinstance(m.get("content"), list) else [])
  169. if isinstance(b, dict) and b.get("type") == "tool_result")
  170. print(f"\033[90m[钩子] Stop:会话使用了 {tool_count} 次工具调用\033[0m")
  171. return None
  172. register_hook("UserPromptSubmit", context_inject_hook)
  173. register_hook("PreToolUse", permission_hook)
  174. register_hook("PreToolUse", log_hook)
  175. register_hook("PostToolUse", large_output_hook)
  176. register_hook("Stop", summary_hook)
  177. # ═══════════════════════════════════════════════════════════
  178. # agent_loop — 与 s03 结构相同,但没有硬编码检查
  179. # s03: if not check_permission(block): ...
  180. # s04: if trigger_hooks("PreToolUse", block): ...
  181. # ═══════════════════════════════════════════════════════════
  182. def agent_loop(messages: list):
  183. while True:
  184. response = client.messages.create(
  185. model=MODEL, system=SYSTEM, messages=messages,
  186. tools=TOOLS, max_tokens=8000,
  187. )
  188. messages.append({"role": "assistant", "content": response.content})
  189. if response.stop_reason != "tool_use":
  190. force = trigger_hooks("Stop", messages)
  191. if force:
  192. messages.append({"role": "user", "content": force})
  193. continue
  194. return
  195. results = []
  196. for block in response.content:
  197. if block.type != "tool_use":
  198. continue
  199. # s04 变化: Hook 替代硬编码的 check_permission()
  200. blocked = trigger_hooks("PreToolUse", block)
  201. if blocked:
  202. results.append({"type": "tool_result", "tool_use_id": block.id,
  203. "content": str(blocked)})
  204. continue
  205. handler = TOOL_HANDLERS.get(block.name)
  206. output = handler(**block.input) if handler else f"未知工具:{block.name}"
  207. trigger_hooks("PostToolUse", block, output) # s04: 后置钩子
  208. results.append({"type": "tool_result", "tool_use_id": block.id, "content": output})
  209. messages.append({"role": "user", "content": results})
  210. if __name__ == "__main__":
  211. print("s04: 钩子系统 — 扩展逻辑挂到钩子上,循环保持干净")
  212. print("输入问题后按回车。输入 q 退出。\n")
  213. history = []
  214. while True:
  215. try:
  216. query = input("\033[36ms04 >> \033[0m")
  217. except (EOFError, KeyboardInterrupt):
  218. break
  219. if query.strip().lower() in ("q", "exit", ""):
  220. break
  221. trigger_hooks("UserPromptSubmit", query)
  222. history.append({"role": "user", "content": query})
  223. agent_loop(history)
  224. for block in history[-1]["content"]:
  225. if getattr(block, "type", None) == "text":
  226. print(block.text)
  227. print()