chore: import original project baseline
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.
This commit is contained in:
@@ -0,0 +1,482 @@
|
||||
"""
|
||||
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"
|
||||
Reference in New Issue
Block a user