research_workflow.py 8.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237
  1. """
  2. agent67 ResearchWorkflow 工具 — 一键科研流程
  3. 将 research skill pack 的 4 阶段 pipeline 封装为 agent67 的一个工具。
  4. agent67 调用一次 ResearchWorkflow,内部自动串联:
  5. 读论文 → 复现 → 实验 → 记录
  6. Lambda 语义:
  7. ResearchWorkflow = λpaper. notebook(experiment(reproduce(read(paper))))
  8. 用法 (agent67 工具调用):
  9. {"action": "ResearchWorkflow", "input": {"paper": "Attention Is All You Need"}}
  10. {"action": "ResearchWorkflow", "input": {"paper": "https://arxiv.org/abs/1706.03762"}}
  11. {"action": "ResearchWorkflow", "input": {"paper": "transformer", "workspace": "./research/transformer"}}
  12. """
  13. from __future__ import annotations
  14. import json
  15. import os
  16. import sys
  17. import time
  18. from pathlib import Path
  19. PROJECT_ROOT = Path(__file__).resolve().parent.parent.parent.parent
  20. sys.path.insert(0, str(PROJECT_ROOT))
  21. from lambdagent.core import Context
  22. def research_workflow(input_str: str) -> str:
  23. """
  24. 执行完整科研流程: 读论文→复现→实验→记录
  25. 输入格式:
  26. JSON: {"paper": "论文标题/URL", "workspace": "保存路径(可选)", "skip": ["reproduce"](可选)}
  27. 纯文本: 论文标题或URL
  28. """
  29. # 解析输入
  30. if isinstance(input_str, str):
  31. try:
  32. data = json.loads(input_str)
  33. paper = data.get("paper", data.get("query", input_str))
  34. workspace = data.get("workspace", "")
  35. skip_phases = data.get("skip", [])
  36. except (json.JSONDecodeError, AttributeError):
  37. paper = input_str.strip()
  38. workspace = ""
  39. skip_phases = []
  40. else:
  41. paper = str(input_str)
  42. workspace = ""
  43. skip_phases = []
  44. if not paper:
  45. return "[ResearchWorkflow] 请提供论文标题或URL"
  46. # 设置工作空间
  47. if not workspace:
  48. safe_name = "".join(c if c.isalnum() or c in "-_" else "_" for c in paper[:50])
  49. workspace = f"./research/{safe_name}"
  50. os.makedirs(workspace, exist_ok=True)
  51. results = {}
  52. t0 = time.time()
  53. # ── Phase 1: 读论文 ──
  54. if "read" not in skip_phases:
  55. print(f" 📖 [1/4] 读论文: {paper[:60]}...")
  56. try:
  57. reader = _get_skill("paper-reader")
  58. if reader:
  59. paper_info = reader.apply(paper, Context())
  60. results["paper_info"] = str(paper_info)
  61. # 保存到工作空间
  62. _save_phase(workspace, "01_paper_reading", paper_info)
  63. print(f" ✅ [1/4] 论文分析完成")
  64. else:
  65. # fallback: 简化版读论文 (用内置工具)
  66. results["paper_info"] = _fallback_read_paper(paper)
  67. print(f" ✅ [1/4] 论文分析完成 (fallback)")
  68. except Exception as e:
  69. results["paper_info"] = f"[读论文失败: {e}]"
  70. print(f" ⚠️ [1/4] 读论文失败: {e}")
  71. else:
  72. results["paper_info"] = "[跳过]"
  73. print(f" ⏭️ [1/4] 跳过读论文")
  74. # ── Phase 2: 复现 ──
  75. if "reproduce" not in skip_phases:
  76. print(f" 🔬 [2/4] 代码复现...")
  77. try:
  78. reproducer = _get_skill("code-reproducer")
  79. if reproducer:
  80. reproduction = reproducer.apply(results["paper_info"], Context())
  81. results["reproduction"] = str(reproduction)
  82. _save_phase(workspace, "02_reproduction", reproduction)
  83. print(f" ✅ [2/4] 复现完成")
  84. else:
  85. results["reproduction"] = "[code-reproducer skill 未注册]"
  86. print(f" ⚠️ [2/4] 复现跳过 (skill 未注册)")
  87. except Exception as e:
  88. results["reproduction"] = f"[复现失败: {e}]"
  89. print(f" ⚠️ [2/4] 复现失败: {e}")
  90. else:
  91. results["reproduction"] = "[跳过]"
  92. print(f" ⏭️ [2/4] 跳过复现")
  93. # ── Phase 3: 实验 ──
  94. if "experiment" not in skip_phases:
  95. print(f" 🧪 [3/4] 实验...")
  96. try:
  97. experimenter = _get_skill("experimenter")
  98. if experimenter:
  99. experiment = experimenter.apply(
  100. json.dumps({"paper": results["paper_info"], "reproduction": results["reproduction"]},
  101. ensure_ascii=False),
  102. Context()
  103. )
  104. results["experiment"] = str(experiment)
  105. _save_phase(workspace, "03_experiment", experiment)
  106. print(f" ✅ [3/4] 实验完成")
  107. else:
  108. results["experiment"] = "[experimenter skill 未注册]"
  109. print(f" ⚠️ [3/4] 实验跳过 (skill 未注册)")
  110. except Exception as e:
  111. results["experiment"] = f"[实验失败: {e}]"
  112. print(f" ⚠️ [3/4] 实验失败: {e}")
  113. else:
  114. results["experiment"] = "[跳过]"
  115. print(f" ⏭️ [3/4] 跳过实验")
  116. # ── Phase 4: 记录 ──
  117. if "record" not in skip_phases:
  118. print(f" 📝 [4/4] 记录笔记...")
  119. try:
  120. notebook = _get_skill("lab-notebook")
  121. if notebook:
  122. record = notebook.apply(
  123. json.dumps(results, ensure_ascii=False),
  124. Context()
  125. )
  126. results["record"] = str(record)
  127. _save_phase(workspace, "04_notebook", record)
  128. print(f" ✅ [4/4] 笔记完成")
  129. else:
  130. # fallback: 自己写笔记
  131. results["record"] = _fallback_write_notebook(workspace, results)
  132. print(f" ✅ [4/4] 笔记完成 (fallback)")
  133. except Exception as e:
  134. results["record"] = f"[记录失败: {e}]"
  135. print(f" ⚠️ [4/4] 记录失败: {e}")
  136. else:
  137. results["record"] = "[跳过]"
  138. print(f" ⏭️ [4/4] 跳过记录")
  139. elapsed = time.time() - t0
  140. # 汇总
  141. summary = {
  142. "paper": paper,
  143. "workspace": workspace,
  144. "phases_completed": [k for k, v in results.items() if "[跳过]" not in str(v) and "[失败" not in str(v)],
  145. "elapsed_seconds": round(elapsed, 1),
  146. "results_preview": {k: str(v)[:200] for k, v in results.items()},
  147. }
  148. return json.dumps(summary, ensure_ascii=False, indent=2)
  149. def _get_skill(name: str):
  150. """从 SkillRegistry 获取 skill,如不存在则尝试注册"""
  151. try:
  152. from lambdagent.skills import SkillRegistry
  153. registry = SkillRegistry()
  154. skill = registry.get(name)
  155. if skill is None:
  156. # 尝试注册 research skill pack
  157. from lambdagent.skillpacks.research import register_all
  158. register_all()
  159. skill = registry.get(name)
  160. return skill
  161. except ImportError:
  162. return None
  163. def _save_phase(workspace: str, phase_name: str, content) -> None:
  164. """保存阶段结果到工作空间"""
  165. phase_dir = os.path.join(workspace, phase_name)
  166. os.makedirs(phase_dir, exist_ok=True)
  167. output_path = os.path.join(phase_dir, "output.json")
  168. try:
  169. with open(output_path, "w", encoding="utf-8") as f:
  170. if isinstance(content, str):
  171. try:
  172. data = json.loads(content)
  173. json.dump(data, f, ensure_ascii=False, indent=2)
  174. except json.JSONDecodeError:
  175. f.write(content)
  176. else:
  177. json.dump(str(content), f, ensure_ascii=False, indent=2)
  178. except Exception:
  179. pass
  180. def _fallback_read_paper(paper: str) -> str:
  181. """当 paper-reader skill 不可用时的 fallback"""
  182. return json.dumps({
  183. "title": paper,
  184. "status": "需要手动搜索",
  185. "hint": "请用 WebSearch 搜索论文,然后 WebFetch 获取内容",
  186. }, ensure_ascii=False)
  187. def _fallback_write_notebook(workspace: str, results: dict) -> str:
  188. """当 lab-notebook skill 不可用时的 fallback"""
  189. notebook_path = os.path.join(workspace, "research_notes.md")
  190. content = f"""# 研究笔记
  191. ## 论文信息
  192. {results.get('paper_info', '无')}
  193. ## 复现结果
  194. {results.get('reproduction', '无')}
  195. ## 实验结果
  196. {results.get('experiment', '无')}
  197. ---
  198. *由 ResearchWorkflow 自动生成*
  199. """
  200. try:
  201. with open(notebook_path, "w", encoding="utf-8") as f:
  202. f.write(content)
  203. return f"笔记已保存到: {notebook_path}"
  204. except Exception as e:
  205. return f"[保存失败: {e}]"