Import the pre-repair source tree as the history baseline. Runtime data (data/), virtualenvs, bytecode caches and logs are gitignored so local secrets and user state stay out of the repo.
483 lines
22 KiB
Python
483 lines
22 KiB
Python
"""
|
||
core/agent/loop.py
|
||
==================
|
||
🌟 pi 核心循环的 Python 1:1 移植 —— 整个框架的心脏
|
||
|
||
对照 pi-main 源码(packages/agent/src/agent-loop.ts, 796 行):
|
||
runAgentLoop() → run_loop(agent, new_message, ...)
|
||
runLoop() → _run_loop() (行 163-278)
|
||
streamAssistantResponse() → _stream_turn() (行 281-380)
|
||
executeToolCalls() → execute_tool_calls() (行 413-427)
|
||
executeToolCallsParallel/Sequential → _execute_parallel/_execute_sequential
|
||
prepareToolCall() → tools.prepare_tool_calls
|
||
failToolCallsFromTruncatedMessage → tools.fail_tool_calls_from_truncated_message
|
||
shouldTerminateToolBatch → _should_terminate_batch (行 561-563)
|
||
|
||
pi runLoop 的结构(本文件逐行对应):
|
||
emit agent_start
|
||
pending = getSteeringMessages() # 循环开始前先取一次
|
||
outer: while True:
|
||
inner: while hasMoreToolCalls or pending:
|
||
turn_start(首轮不发,首轮 turn_start 在 runAgentLoop 入口发)
|
||
注入 pending(message_start/end → 入 context + newMessages)
|
||
assistant = streamAssistantResponse(...)
|
||
if stopReason in (error, aborted): turn_end + agent_end + return
|
||
toolCalls = assistant 的 toolCall 块
|
||
if toolCalls:
|
||
length → failToolCallsFromTruncatedMessage(不执行!)
|
||
否则 → executeToolCalls(sequential/parallel 二选一)
|
||
hasMoreToolCalls = !batch.terminate
|
||
toolResults 入 context + newMessages
|
||
turn_end
|
||
prepareNextTurn 钩子(可换 context/model)
|
||
shouldStopAfterTurn 钩子 → agent_end + return
|
||
pending = getSteeringMessages() # 每轮结束取一次
|
||
followUps = getFollowUpMessages()
|
||
if followUps: pending = followUps; continue # 外层续跑
|
||
break
|
||
agent_end
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
from concurrent.futures import ThreadPoolExecutor
|
||
from dataclasses import dataclass, field
|
||
from typing import Any, Callable, Dict, List, Optional, Tuple
|
||
|
||
from .context import clamp_max_tokens_to_context
|
||
from .tools import (PreparedToolCall, execute_tool_call, fail_tool_calls_from_truncated_message,
|
||
prepare_tool_calls)
|
||
from .types import (AgentConfig, AgentError, AgentEvent, AgentMessage,
|
||
AgentTool, AgentToolResult, AbortSignal, AssistantMessageEvent,
|
||
RunResult, ToolCall, new_id)
|
||
|
||
|
||
# ======================================================================
|
||
# 事件发射辅助
|
||
# ======================================================================
|
||
def _emit(agent, event: AgentEvent):
|
||
agent._emit(event)
|
||
|
||
|
||
def _message_events(agent, msg: AgentMessage):
|
||
_emit(agent, AgentEvent(type="message_start", message=msg))
|
||
_emit(agent, AgentEvent(type="message_end", message=msg))
|
||
|
||
|
||
def _tool_result_message(finalized: Dict[str, Any]) -> AgentMessage:
|
||
"""对照 createToolResultMessage"""
|
||
tc: ToolCall = finalized["tool_call"]
|
||
result: AgentToolResult = finalized["result"]
|
||
return AgentMessage(
|
||
role="toolResult",
|
||
tool_call_id=tc.id,
|
||
tool_name=tc.name,
|
||
content=result.content,
|
||
is_error=finalized.get("is_error", False) or result.is_error,
|
||
)
|
||
|
||
|
||
# ======================================================================
|
||
# 流式生成一轮助手消息 —— 对照 streamAssistantResponse (行 281-380)
|
||
# ======================================================================
|
||
def _stream_turn(agent, current_context: List[AgentMessage],
|
||
config: AgentConfig, signal: AbortSignal,
|
||
stream_fn: Callable) -> Tuple[AgentMessage, Optional[AgentError]]:
|
||
"""
|
||
返回 (assistant_message, error)。
|
||
error 非 None 时 message.stop_reason == "error"。
|
||
"""
|
||
# 🌟 pi 同款:流开始前检查中止 → 立即产出 aborted 消息
|
||
if signal.aborted:
|
||
msg = AgentMessage(role="assistant", stop_reason="aborted")
|
||
agent.state.messages.append(msg)
|
||
_message_events(agent, msg)
|
||
return msg, None
|
||
|
||
# 🌟 系统提示词注入(对照 pi:systemPrompt 放在每次 API 请求头部,
|
||
# 不进 state.messages、不占压缩/历史)
|
||
if config.system_prompt:
|
||
api_context: List[AgentMessage] = [
|
||
AgentMessage(role="system", content=config.system_prompt)
|
||
] + current_context
|
||
else:
|
||
api_context = current_context
|
||
|
||
# 🌟 输出预算钳制(对照 simple-options clampMaxTokensToContext)
|
||
# 🆕 P2: 把 system + 工具 schema 传给估算器(对照 pi 传完整 Context),
|
||
# 仅在无 usage 锚点分支生效;否则新会话首轮会多预留 ~system+tools 的预算。
|
||
clamped = clamp_max_tokens_to_context(
|
||
config.model, current_context,
|
||
system_prompt=config.system_prompt or "", tools=config.tools)
|
||
if clamped is None:
|
||
err = AgentError(
|
||
message="上下文溢出:估算输入 token 已超过模型窗口(需要压缩)",
|
||
kind="overflow", recoverable=True)
|
||
msg = AgentMessage(role="assistant", stop_reason="error",
|
||
error_message=err.message)
|
||
agent.state.messages.append(msg)
|
||
_message_events(agent, msg)
|
||
return msg, err
|
||
max_tokens, _input_tokens = clamped
|
||
|
||
msg = AgentMessage(role="assistant")
|
||
agent.state.streaming_message = msg
|
||
agent.state.streaming_delta = {}
|
||
_emit(agent, AgentEvent(type="message_start", message=msg))
|
||
|
||
raw_tc: Dict[int, Dict[str, str]] = {} # 中止时用于保留部分工具调用(仅展示)
|
||
got_final = False
|
||
error: Optional[AgentError] = None
|
||
|
||
try:
|
||
for kind, payload in stream_fn(api_context, config.model, signal,
|
||
max_tokens, config.tools):
|
||
if kind == "event":
|
||
ev = payload
|
||
if ev.type == "text_delta":
|
||
msg.content = (msg.content or "") + ev.text
|
||
agent.state.streaming_delta["text"] = ev.text
|
||
_emit(agent, AgentEvent(type="message_update", message=msg,
|
||
assistant_message_event=ev))
|
||
elif ev.type == "thinking_delta":
|
||
msg.reasoning += ev.text
|
||
agent.state.streaming_delta["thinking"] = ev.text
|
||
_emit(agent, AgentEvent(type="message_update", message=msg,
|
||
assistant_message_event=ev))
|
||
elif ev.type == "toolcall_delta":
|
||
slot = raw_tc.setdefault(ev.tool_call_index,
|
||
{"id": "", "name": "", "args": ""})
|
||
if ev.tool_call_field == "id":
|
||
slot["id"] = ev.tool_call_delta
|
||
elif ev.tool_call_field == "name":
|
||
slot["name"] += ev.tool_call_delta
|
||
else:
|
||
slot["args"] += ev.tool_call_delta
|
||
_emit(agent, AgentEvent(type="message_update", message=msg,
|
||
assistant_message_event=ev))
|
||
else: # final
|
||
final: AgentMessage = payload
|
||
got_final = True
|
||
# 采用权威 final:usage / stop_reason / 解析好的 tool_calls
|
||
msg.usage = final.usage
|
||
msg.stop_reason = final.stop_reason
|
||
msg.tool_calls = final.tool_calls
|
||
if not final.content and msg.content:
|
||
pass # 以循环侧累积为准(二者一致)
|
||
else:
|
||
msg.content = final.content if final.content else msg.content
|
||
msg.reasoning = final.reasoning or msg.reasoning
|
||
except AgentError as e:
|
||
error = e
|
||
msg.stop_reason = "error"
|
||
msg.error_message = e.message
|
||
_emit(agent, AgentEvent(type="message_update", message=msg,
|
||
assistant_message_event=
|
||
AssistantMessageEvent.error(e.message)))
|
||
except Exception as e:
|
||
from .stream_fn import classify_error
|
||
error = classify_error(e)
|
||
msg.stop_reason = "error"
|
||
msg.error_message = error.message
|
||
|
||
if not got_final and error is None:
|
||
# 生成器中途结束且无异常 = 中止(流被关闭,pi 同款语义)
|
||
msg.stop_reason = "aborted"
|
||
# 保留部分工具调用(仅用于 UI 展示;中止后不会执行,也不会进入下次上下文)
|
||
for idx in sorted(raw_tc.keys()):
|
||
slot = raw_tc[idx]
|
||
if slot.get("name"):
|
||
msg.tool_calls.append(ToolCall(id=slot.get("id") or new_id("call"),
|
||
name=slot["name"],
|
||
raw_arguments=slot.get("args", "")))
|
||
if error is None and msg.stop_reason not in ("stop", "length", "aborted"):
|
||
msg.stop_reason = "stop"
|
||
|
||
# 🌟 兜底(haocode 扩展,pi 无此层):不支持 tools API 的供应商/模型会把
|
||
# 工具调用用文字"演"出来(如 <bash>ls</bash>)。识别单参工具 bash/read
|
||
# 转成真 tool_call 继续执行;write/edit 多参歧义大不兜底。
|
||
if error is None and not msg.tool_calls and msg.content:
|
||
from .tools import parse_text_tool_calls
|
||
cleaned, txt_calls = parse_text_tool_calls(msg.content)
|
||
if txt_calls:
|
||
msg.tool_calls = txt_calls
|
||
|
||
agent.state.streaming_message = None
|
||
agent.state.streaming_delta = {}
|
||
agent.state.messages.append(msg)
|
||
_emit(agent, AgentEvent(type="message_end", message=msg))
|
||
return msg, error
|
||
|
||
|
||
# ======================================================================
|
||
# 工具批量执行 —— 对照 executeToolCalls (行 413-427)
|
||
# ======================================================================
|
||
def _should_terminate_batch(finalized: List[Dict[str, Any]]) -> bool:
|
||
"""对照 shouldTerminateToolBatch: 全部结果都请求 terminate 才终止"""
|
||
return len(finalized) > 0 and all(
|
||
f["result"].terminate for f in finalized)
|
||
|
||
|
||
def _run_prepared_in_pool(prep: PreparedToolCall, assistant: AgentMessage,
|
||
config: AgentConfig, signal: AbortSignal,
|
||
agent) -> Dict[str, Any]:
|
||
"""线程池内执行单个工具(对照 parallel 版的 async 闭包)"""
|
||
on_update = None
|
||
on_timer = None
|
||
if agent is not None:
|
||
def on_update(partial: str):
|
||
_emit(agent, AgentEvent(type="tool_execution_update",
|
||
tool_call=prep.tool_call, arg=partial))
|
||
def on_timer(elapsed_i: int, timeout_i: int):
|
||
# 🆕 bash 运行中每秒滴一次 → 前端气泡读秒 N/Ts
|
||
_emit(agent, AgentEvent(type="tool_execution_timer",
|
||
tool_call=prep.tool_call,
|
||
arg=(int(elapsed_i), int(timeout_i))))
|
||
result = execute_tool_call(prep, assistant, config, signal, on_update,
|
||
on_timer)
|
||
finalized = {"tool_call": prep.tool_call, "result": result,
|
||
"is_error": result.is_error}
|
||
_emit(agent, AgentEvent(type="tool_execution_end",
|
||
tool_call=prep.tool_call, result=result,
|
||
is_error=result.is_error))
|
||
return finalized
|
||
|
||
|
||
def _execute_parallel(current_context, assistant, config, signal, agent,
|
||
prepared: List[PreparedToolCall]) -> Dict[str, Any]:
|
||
"""对照 executeToolCallsParallel:start 事件串行发、准备串行、执行并发、
|
||
结果消息按原始顺序产出"""
|
||
finalized: List[Dict[str, Any]] = []
|
||
pending_futures: List = []
|
||
max_workers = max(1, min(8, len(prepared)))
|
||
with ThreadPoolExecutor(max_workers=max_workers,
|
||
thread_name_prefix="tool") as pool:
|
||
for prep in prepared:
|
||
_emit(agent, AgentEvent(type="tool_execution_start",
|
||
tool_call=prep.tool_call,
|
||
arg=str(prep.args) if prep.args else ""))
|
||
if prep.error:
|
||
# 准备失败(未知工具/参数非法)→ 立即结果(对照 kind:"immediate")
|
||
result = AgentToolResult.text(prep.error, is_error=True)
|
||
fin = {"tool_call": prep.tool_call, "result": result,
|
||
"is_error": True}
|
||
_emit(agent, AgentEvent(type="tool_execution_end",
|
||
tool_call=prep.tool_call, result=result,
|
||
is_error=True))
|
||
finalized.append(fin)
|
||
if signal.aborted:
|
||
break
|
||
continue
|
||
fut = pool.submit(_run_prepared_in_pool, prep, assistant, config,
|
||
signal, agent)
|
||
pending_futures.append(fut)
|
||
if signal.aborted:
|
||
# 已提交的任务仍会完成(它们内部检查 signal),不再追加
|
||
pass
|
||
for fut in pending_futures:
|
||
finalized.append(fut.result())
|
||
|
||
messages: List[AgentMessage] = []
|
||
for fin in finalized:
|
||
tr = _tool_result_message(fin)
|
||
_message_events(agent, tr)
|
||
messages.append(tr)
|
||
return {"messages": messages, "terminate": _should_terminate_batch(finalized)}
|
||
|
||
|
||
def _execute_sequential(current_context, assistant, config, signal, agent,
|
||
prepared: List[PreparedToolCall]) -> Dict[str, Any]:
|
||
"""对照 executeToolCallsSequential:一次一个,完成一个再下一个"""
|
||
finalized: List[Dict[str, Any]] = []
|
||
messages: List[AgentMessage] = []
|
||
for prep in prepared:
|
||
_emit(agent, AgentEvent(type="tool_execution_start",
|
||
tool_call=prep.tool_call,
|
||
arg=str(prep.args) if prep.args else ""))
|
||
if not prep.error:
|
||
fin = _run_prepared_in_pool(prep, assistant, config, signal, agent)
|
||
else:
|
||
result = AgentToolResult.text(prep.error, is_error=True)
|
||
fin = {"tool_call": prep.tool_call, "result": result, "is_error": True}
|
||
_emit(agent, AgentEvent(type="tool_execution_end",
|
||
tool_call=prep.tool_call, result=result,
|
||
is_error=True))
|
||
finalized.append(fin)
|
||
tr = _tool_result_message(fin)
|
||
_message_events(agent, tr)
|
||
messages.append(tr)
|
||
if signal.aborted:
|
||
break
|
||
return {"messages": messages, "terminate": _should_terminate_batch(finalized)}
|
||
|
||
|
||
def execute_tool_calls(current_context, assistant: AgentMessage,
|
||
config: AgentConfig, signal: AbortSignal, agent) -> Dict[str, Any]:
|
||
"""对照 executeToolCalls: 批次里有 sequential 工具 → 整批串行"""
|
||
prepared = prepare_tool_calls(assistant, config.tools)
|
||
has_sequential = any(
|
||
(p.tool and p.tool.execution_mode == "sequential") for p in prepared
|
||
if not p.error)
|
||
if config.tool_execution == "sequential" or has_sequential:
|
||
return _execute_sequential(current_context, assistant, config, signal,
|
||
agent, prepared)
|
||
return _execute_parallel(current_context, assistant, config, signal, agent,
|
||
prepared)
|
||
|
||
|
||
# ======================================================================
|
||
# 主循环 —— 对照 runLoop (行 163-278)
|
||
# ======================================================================
|
||
def run_loop(agent, new_message: Optional[AgentMessage], signal: AbortSignal,
|
||
stream_fn: Callable) -> RunResult:
|
||
config = agent.config
|
||
new_messages: List[AgentMessage] = []
|
||
|
||
# 入口校验(对照 runAgentLoopContinue 的前置检查由 Agent 层负责)
|
||
_emit(agent, AgentEvent(type="agent_start"))
|
||
_emit(agent, AgentEvent(type="turn_start")) # 首轮 turn_start(对照行 142/170)
|
||
|
||
# context 准备 + transformContext 钩子
|
||
base = list(agent.state.messages)
|
||
if config.transform_context:
|
||
try:
|
||
base = config.transform_context(base) or base
|
||
except Exception:
|
||
pass
|
||
current_context = list(base)
|
||
|
||
first_turn = True
|
||
pending: List[AgentMessage] = agent._take_steering()
|
||
|
||
while True: # outer loop
|
||
has_more_tool_calls = True
|
||
|
||
while has_more_tool_calls or pending: # inner loop
|
||
if not first_turn:
|
||
_emit(agent, AgentEvent(type="turn_start"))
|
||
else:
|
||
first_turn = False
|
||
|
||
# ---- 注入 pending 消息(steering / followUp)----
|
||
if pending:
|
||
for m in pending:
|
||
_message_events(agent, m)
|
||
current_context.append(m)
|
||
new_messages.append(m)
|
||
agent.state.messages.append(m)
|
||
pending = []
|
||
|
||
# ---- 🆕 轮中主动压缩检查(haocode 增强,偏离 pi 1:1)----
|
||
# 单条工具输出可能把上下文顶出窗口;发下一次请求前主动检查
|
||
#(与轮首 should_compact 同公式)。compact_fn 返回新列表
|
||
#(发生了压缩)→ 同步循环局部上下文。
|
||
if config.compact_fn is not None and agent.state.messages:
|
||
try:
|
||
compacted = config.compact_fn(list(agent.state.messages))
|
||
if compacted is not None:
|
||
agent.state.messages = compacted
|
||
current_context = list(compacted)
|
||
except Exception:
|
||
pass
|
||
|
||
# ---- 流式生成助手消息 ----
|
||
assistant, error = _stream_turn(agent, current_context, config,
|
||
signal, stream_fn)
|
||
new_messages.append(assistant)
|
||
current_context.append(assistant)
|
||
|
||
if assistant.stop_reason in ("error", "aborted"):
|
||
agent.state.error = error
|
||
_emit(agent, AgentEvent(type="turn_end", message=assistant))
|
||
_emit(agent, AgentEvent(type="agent_end",
|
||
stop_reason=assistant.stop_reason,
|
||
error=error,
|
||
messages=new_messages))
|
||
agent._finish_run(new_messages, assistant.stop_reason, error)
|
||
return RunResult(stop_reason=assistant.stop_reason, error=error,
|
||
message_count=len(new_messages))
|
||
|
||
# ---- 工具调用 ----
|
||
tool_results: List[AgentMessage] = []
|
||
has_more_tool_calls = False
|
||
if assistant.tool_calls:
|
||
if assistant.stop_reason == "length":
|
||
# 🌟 截断保护(对照行 212-213):参数可能残缺,一律失败不执行
|
||
fail_msgs = fail_tool_calls_from_truncated_message(
|
||
assistant, "length")
|
||
# 事件流与正常执行对齐
|
||
for tc, tr in zip(assistant.tool_calls, fail_msgs):
|
||
_emit(agent, AgentEvent(type="tool_execution_start",
|
||
tool_call=tc))
|
||
res = AgentToolResult.text(
|
||
f'工具调用 "{tc.name}" 未执行:输出达到 token 上限,'
|
||
f'参数可能被截断。请用完整参数重新发起。', is_error=True)
|
||
_emit(agent, AgentEvent(type="tool_execution_end",
|
||
tool_call=tc, result=res,
|
||
is_error=True))
|
||
_message_events(agent, tr)
|
||
batch = {"messages": fail_msgs, "terminate": False}
|
||
else:
|
||
batch = execute_tool_calls(current_context, assistant, config,
|
||
signal, agent)
|
||
tool_results = batch["messages"]
|
||
has_more_tool_calls = not batch["terminate"]
|
||
for r in tool_results:
|
||
current_context.append(r)
|
||
new_messages.append(r)
|
||
agent.state.messages.append(r)
|
||
|
||
_emit(agent, AgentEvent(type="turn_end", message=assistant))
|
||
|
||
# ---- prepareNextTurn 钩子 ----
|
||
if config.prepare_next_turn:
|
||
try:
|
||
snap = config.prepare_next_turn({
|
||
"message": assistant, "tool_results": tool_results,
|
||
"context": current_context, "new_messages": new_messages,
|
||
})
|
||
if snap:
|
||
current_context = snap.get("context") or current_context
|
||
except Exception:
|
||
pass
|
||
|
||
# ---- shouldStopAfterTurn 钩子 ----
|
||
if config.should_stop_after_turn:
|
||
try:
|
||
if config.should_stop_after_turn({
|
||
"message": assistant, "tool_results": tool_results,
|
||
"context": current_context,
|
||
"new_messages": new_messages}):
|
||
stop = _last_assistant_stop(new_messages)
|
||
_emit(agent, AgentEvent(type="agent_end",
|
||
stop_reason=stop,
|
||
messages=new_messages))
|
||
agent._finish_run(new_messages, stop, None)
|
||
return RunResult(stop_reason=stop,
|
||
message_count=len(new_messages))
|
||
except Exception:
|
||
pass
|
||
|
||
# ---- 每轮结束取 steering ----
|
||
pending = agent._take_steering()
|
||
|
||
# ---- 外层:followUp ----
|
||
follow_ups = agent._take_follow_ups()
|
||
if follow_ups:
|
||
pending = follow_ups
|
||
continue
|
||
break
|
||
|
||
stop = _last_assistant_stop(new_messages)
|
||
_emit(agent, AgentEvent(type="agent_end", stop_reason=stop,
|
||
messages=new_messages))
|
||
agent._finish_run(new_messages, stop, None)
|
||
return RunResult(stop_reason=stop, message_count=len(new_messages))
|
||
|
||
|
||
def _last_assistant_stop(new_messages: List[AgentMessage]) -> str:
|
||
"""对照 pi: agent_end 不携带 stopReason,会话层从最后一条 assistant 消息读取"""
|
||
for m in reversed(new_messages):
|
||
if m.role == "assistant":
|
||
return m.stop_reason or "stop"
|
||
return "stop"
|