429 lines
16 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

#!/usr/bin/env python3
"""pi 框架适配器:把任务包派发为 pi 子代理并归一其 JSON 事件流。
这是全仓唯一直接调用 Agent 框架二进制的位置(架构门禁
tests/architecture/test_import_boundaries.py 白名单)。适配器只做三件事:
构造 argv(角色 prompt 注入 + 工具白名单 + 隔离上下文)、逐行消费框架事件流、
把事件归一转发给 TraceWriter。它不含任何业务决策:补证、重写、下一步做什么
全部属于框架里的模型,不属于本模块。
执行策略(provider/model/thinking)由派发方给定并如实记账;框架把模型模式解析为
完整模型 ID,匹配口径见 agent_trace.model_ids_match。超时用看门狗线程杀进程:
阻塞读 stdout 不会自己抛超时,挂死的框架进程必须被强制终止才能失败关闭。
"""
from __future__ import annotations
import json
import subprocess
import threading
import time
from dataclasses import dataclass, field
from typing import Any, Callable, Iterable, Iterator, Mapping, Sequence
from agent_task import TaskPackage
from agent_trace import AgentTraceWriter
DEFAULT_FRAMEWORK = "pi"
DEFAULT_PI_BIN = "pi"
# 事件流单行上限:防御性截断,正常 JSONL 行远小于此。
MAX_STREAM_LINE_BYTES = 8 * 1024 * 1024
class FrameworkError(RuntimeError):
"""框架执行失败:超时、非零退出或事件流不可解析。"""
def __init__(
self,
error_code: str,
message: str,
*,
outcome: "AgentStreamOutcome | None" = None,
) -> None:
super().__init__(message)
self.error_code = error_code
self.outcome = outcome
@dataclass(frozen=True)
class ExecutionPolicy:
"""框架侧执行策略;provider 与 model 必须由调用方显式传入。"""
provider: str | None = None
model: str | None = None
thinking: str | None = None
pi_bin: str = DEFAULT_PI_BIN
cwd: str | None = None
def __post_init__(self) -> None:
if not isinstance(self.provider, str) or not self.provider.strip():
raise ValueError("provider 必须显式传入")
if not isinstance(self.model, str) or not self.model.strip():
raise ValueError("model 必须显式传入")
if self.thinking is not None and self.thinking not in {
"off", "minimal", "low", "medium", "high", "xhigh", "max"
}:
raise ValueError("thinking 不受支持")
@property
def framework(self) -> str:
return DEFAULT_FRAMEWORK
@property
def requested_model_id(self) -> str:
"""账本口径的显式请求模型。"""
return f"{self.provider}/{self.model}"
@dataclass(frozen=True)
class ModelCall:
"""一个模型回合的账本投影材料(usage 为框架归一后的原始字典)。"""
actual_model_id: str
provider: str | None
usage: Mapping[str, Any]
stop_reason: str | None
cost_usd: float | None
@dataclass
class ToolCallRecord:
tool_call_id: str
name: str
is_error: bool = False
@dataclass
class AgentStreamOutcome:
"""框架子代理一次执行的客观结果(不含业务判断)。"""
session_id: str | None = None
exit_code: int | None = None
timed_out: bool = False
final_text: str | None = None
model_calls: list[ModelCall] = field(default_factory=list)
tool_calls: list[ToolCallRecord] = field(default_factory=list)
turns: int = 0
parse_error_lines: int = 0
duration_ms: int = 0
def build_pi_argv(package: TaskPackage, policy: ExecutionPolicy) -> list[str]:
"""构造 pi 子代理 argv:system prompt 注入、工具白名单、上下文隔离。"""
argv = [policy.pi_bin, "--print", "--mode", "json", "--no-session"]
argv += ["--provider", policy.provider, "--model", policy.model]
if policy.thinking:
argv += ["--thinking", policy.thinking]
# 上下文隔离:不加载项目 AGENTS.md/skills/extensions,角色合同全部来自任务包。
argv += ["--no-context-files", "--no-skills", "--no-extensions", "--no-approve"]
allowlist = package.spec.tool_allowlist
if allowlist:
argv += ["--tools", ",".join(allowlist)]
else:
argv += ["--no-tools"]
argv += ["--system-prompt", package.system_prompt, package.user_message]
return argv
def _message_text(message: Mapping[str, Any]) -> str:
"""提取消息中的全部文本块(跳过 thinking/tool_call 块)。"""
parts: list[str] = []
for block in message.get("content") or []:
if isinstance(block, Mapping) and block.get("type") == "text":
parts.append(str(block.get("text") or ""))
return "".join(parts)
def _qualified_model_id(provider: Any, model: Any) -> str:
"""把 pi 分开的 provider/model 字段合成账本要求的完整模型 ID。"""
model_id = str(model or "").strip()
provider_id = str(provider or "").strip()
if not model_id or "/" in model_id or not provider_id:
return model_id
return f"{provider_id}/{model_id}"
def _usage_cost(usage: Mapping[str, Any] | None) -> float | None:
"""读取框架报告的单回合成本;供应商未定价(0/缺失)记 None,不伪造。"""
if not isinstance(usage, Mapping):
return None
cost = usage.get("cost")
if isinstance(cost, Mapping):
total = cost.get("total")
try:
return float(total) if total and float(total) > 0 else None
except (TypeError, ValueError):
return None
return None
class _SubprocessStream:
"""把 Popen stdout 包装成字节行迭代器;看门狗超时杀进程,stderr 丢弃防管道死锁。"""
def __init__(self, proc: subprocess.Popen, timeout_seconds: float) -> None:
self._proc = proc
self.exit_code: int | None = None
self.timed_out = False
self._watchdog = threading.Timer(
max(timeout_seconds, 0.1),
self._kill,
)
self._watchdog.daemon = True
self._watchdog.start()
def _kill(self) -> None:
if self._proc.poll() is None:
self.timed_out = True
self._proc.kill()
def __iter__(self) -> Iterator[bytes]:
assert self._proc.stdout is not None
try:
for raw_line in self._proc.stdout:
if len(raw_line) > MAX_STREAM_LINE_BYTES:
if self._proc.poll() is None:
self._proc.kill()
raise FrameworkError("STREAM_LINE_TOO_LARGE", "事件流单行超限")
yield raw_line
finally:
self._watchdog.cancel()
self.exit_code = self._proc.wait()
self._proc.stdout.close()
def close(self) -> None:
self._watchdog.cancel()
if self._proc.poll() is None:
self._proc.kill()
self._proc.wait()
class PiAgentRunner:
"""启动 pi 子代理、消费事件流并转发归一事件。"""
def __init__(self, launcher: Callable[..., Iterable[bytes]] | None = None) -> None:
# launcher(argv, timeout, cwd) -> 字节行迭代器(带 exit_code 属性);测试注入假 pi。
self._launcher = launcher
def _launch(self, argv: Sequence[str], timeout: float, cwd: str | None) -> Iterable[bytes]:
if self._launcher is not None:
return self._launcher(argv, timeout, cwd)
proc = subprocess.Popen(
list(argv),
stdout=subprocess.PIPE,
stderr=subprocess.DEVNULL,
cwd=cwd,
)
return _SubprocessStream(proc, timeout)
def run(
self,
package: TaskPackage,
policy: ExecutionPolicy,
sink: AgentTraceWriter,
*,
timeout_seconds: float,
raw_sink: Callable[[bytes], None] | None = None,
) -> AgentStreamOutcome:
"""执行一次框架派发;框架层异常抛 FrameworkError(业务校验在派发器)。
raw_sink 逐行接收框架原始事件流字节(转录 tap),供派发器固定全量原始证据。
"""
argv = build_pi_argv(package, policy)
outcome = AgentStreamOutcome()
started = time.monotonic()
sink.emit(
"agent.started",
status="ok",
requested_model_id=policy.requested_model_id,
details={"framework": policy.framework, "thinking": policy.thinking},
)
try:
stream = self._launch(argv, timeout_seconds, policy.cwd)
except OSError as exc:
raise FrameworkError(
"FRAMEWORK_START_FAILED", f"框架进程启动失败: {type(exc).__name__}"
) from exc
final_message: Mapping[str, Any] | None = None
allowed_tools = frozenset(package.spec.tool_allowlist)
stream_error: FrameworkError | None = None
try:
try:
for raw_line in stream:
if raw_sink is not None:
try:
raw_sink(raw_line)
except OSError as exc:
raise FrameworkError(
"TRANSCRIPT_WRITE_FAILED", f"框架转录写入失败: {type(exc).__name__}"
) from exc
line = raw_line.decode("utf-8", errors="replace").strip()
if not line:
continue
try:
event = json.loads(line)
except json.JSONDecodeError:
outcome.parse_error_lines += 1
continue
if not isinstance(event, Mapping):
outcome.parse_error_lines += 1
continue
self._consume(event, sink, policy, outcome, allowed_tools)
if event.get("type") == "agent_end":
messages = event.get("messages") or []
for message in reversed(messages):
if isinstance(message, Mapping) and message.get("role") == "assistant":
final_message = message
break
except FrameworkError as exc:
stream_error = exc
finally:
closer = getattr(stream, "close", None)
if callable(closer):
closer()
outcome.duration_ms = int((time.monotonic() - started) * 1000)
outcome.exit_code = getattr(stream, "exit_code", None)
outcome.timed_out = bool(getattr(stream, "timed_out", False))
outcome.final_text = _message_text(final_message) if final_message else None
if stream_error is not None:
stream_error.outcome = outcome
raise stream_error
if outcome.timed_out:
raise FrameworkError(
"FRAMEWORK_TIMEOUT", f"框架执行超时(>{timeout_seconds}s)", outcome=outcome
)
if outcome.exit_code != 0:
raise FrameworkError(
"FRAMEWORK_EXIT_NONZERO", f"框架进程退出码 {outcome.exit_code}", outcome=outcome
)
if outcome.parse_error_lines:
raise FrameworkError(
"STREAM_PARSE_ERROR",
f"事件流有 {outcome.parse_error_lines} 行不可解析",
outcome=outcome,
)
if not outcome.model_calls:
raise FrameworkError("NO_MODEL_RESPONSE", "事件流未含任何模型回合", outcome=outcome)
if outcome.final_text is None or not outcome.final_text.strip():
raise FrameworkError("EMPTY_FINAL_MESSAGE", "框架未返回最终文本", outcome=outcome)
sink.emit(
"agent.completed",
status="ok",
requested_model_id=policy.requested_model_id,
actual_model_id=outcome.model_calls[-1].actual_model_id,
details={
"sessionId": outcome.session_id,
"turns": outcome.turns,
"modelCalls": len(outcome.model_calls),
"toolCalls": len(outcome.tool_calls),
"durationMs": outcome.duration_ms,
},
)
return outcome
def _consume(
self,
event: Mapping[str, Any],
sink: AgentTraceWriter,
policy: ExecutionPolicy,
outcome: AgentStreamOutcome,
allowed_tools: frozenset[str],
) -> None:
"""把单个框架事件归一转发;未知事件类型静默忽略(框架可演进)。"""
kind = event.get("type")
if kind == "session":
outcome.session_id = str(event.get("id") or "") or None
elif kind == "turn_start":
outcome.turns += 1
elif kind == "message_end":
message = event.get("message") or {}
if message.get("role") != "assistant":
return
usage = message.get("usage") or {}
actual_model_id = _qualified_model_id(message.get("provider"), message.get("model"))
if not actual_model_id:
raise FrameworkError("MODEL_ID_MISSING", "模型回合缺 provider/model 身份")
call = ModelCall(
actual_model_id=actual_model_id,
provider=message.get("provider"),
usage=usage,
stop_reason=message.get("stopReason"),
cost_usd=_usage_cost(usage),
)
outcome.model_calls.append(call)
model_failed = call.stop_reason in {"error", "aborted"}
sink.emit(
"model.completed",
status="error" if model_failed else "ok",
requested_model_id=policy.requested_model_id,
actual_model_id=call.actual_model_id,
usage=usage,
cost_usd=call.cost_usd,
details={
"stopReason": call.stop_reason,
"provider": call.provider,
"sessionId": outcome.session_id,
"turn": outcome.turns,
"errorMessage": str(message.get("errorMessage") or "")[:256] if model_failed else None,
},
)
if model_failed:
raise FrameworkError(
"MODEL_TURN_FAILED",
f"模型回合结束状态为 {call.stop_reason}",
)
elif kind == "tool_execution_start":
name = str(event.get("toolName") or "")
tool_call_id = str(event.get("toolCallId") or "")
if not tool_call_id:
raise FrameworkError("TOOL_EVENT_INVALID", "工具开始事件缺 toolCallId")
if name not in allowed_tools:
raise FrameworkError("TOOL_NOT_ALLOWED", f"框架执行了未授权工具: {name or '<empty>'}")
outcome.tool_calls.append(ToolCallRecord(tool_call_id=tool_call_id, name=name))
sink.emit(
"tool.started",
status="ok",
tool_name=name,
details={"toolCallId": tool_call_id},
)
elif kind == "tool_execution_end":
is_error = bool(event.get("isError"))
name = str(event.get("toolName") or "")
tool_call_id = str(event.get("toolCallId") or "")
if not tool_call_id:
raise FrameworkError("TOOL_EVENT_INVALID", "工具结束事件缺 toolCallId")
if name not in allowed_tools:
raise FrameworkError("TOOL_NOT_ALLOWED", f"框架执行了未授权工具: {name or '<empty>'}")
matching = next(
(call for call in reversed(outcome.tool_calls) if call.tool_call_id == tool_call_id),
None,
)
if matching is None:
raise FrameworkError("TOOL_EVENT_INVALID", "工具结束事件缺对应开始事件")
matching.is_error = is_error
sink.emit(
"tool.completed",
status="error" if is_error else "ok",
tool_name=name,
details={"toolCallId": tool_call_id},
)
__all__ = [
"AgentStreamOutcome",
"DEFAULT_PI_BIN",
"ExecutionPolicy",
"FrameworkError",
"ModelCall",
"PiAgentRunner",
"ToolCallRecord",
"build_pi_argv",
]