465 lines
18 KiB
Python
465 lines
18 KiB
Python
#!/usr/bin/env python3
|
||
"""pi 框架适配器:把任务包派发为 pi 子代理并归一其 JSON 事件流。
|
||
|
||
这是全仓唯一直接调用 Agent 框架二进制的位置(架构门禁
|
||
tests/architecture/test_import_boundaries.py 白名单)。适配器只做三件事:
|
||
构造 argv(角色 prompt 注入 + 工具白名单 + 隔离上下文)、逐行消费框架事件流、
|
||
把事件归一转发给注入的 TraceSink。它不含任何业务决策:补证、重写、下一步做什么
|
||
全部属于框架里的模型,不属于本模块。
|
||
|
||
执行策略(provider/model/thinking)由派发方给定并如实记账;框架把模型模式解析为
|
||
完整模型 ID;业务侧负责把 requested/actual 结果投影到自己的证据账本。超时用看门狗线程杀进程:
|
||
阻塞读 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, Protocol, Sequence
|
||
|
||
from framework.primitives.execution import FrameworkExecutionRequest
|
||
|
||
|
||
class TraceSink(Protocol):
|
||
"""Muse 注入的事件接收端;框架不拥有其持久化实现。"""
|
||
|
||
def emit(self, event_type: str, **kwargs: Any) -> int: ...
|
||
|
||
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
|
||
# 会话复用:同一章补证/修改继续原会话;两者同给或同缺。
|
||
session_id: str | None = None
|
||
session_dir: str | None = None
|
||
# 显式加载的框架扩展(工具 server);--no-extensions 只禁自动发现,-e 照常加载。
|
||
extension_path: 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 不受支持")
|
||
if bool(self.session_id) != bool(self.session_dir):
|
||
raise ValueError("session_id 与 session_dir 必须同时给出")
|
||
|
||
@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
|
||
args: Mapping[str, Any] | None = None
|
||
|
||
|
||
@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
|
||
unknown_event_types: list[str] = field(default_factory=list)
|
||
duration_ms: int = 0
|
||
|
||
|
||
def build_pi_argv(request: FrameworkExecutionRequest, policy: ExecutionPolicy) -> list[str]:
|
||
"""构造 Pi argv:只消费通用执行请求,不读取 Muse 业务合同。"""
|
||
|
||
argv = [policy.pi_bin, "--print", "--mode", "json"]
|
||
if policy.session_id:
|
||
argv += ["--session-id", policy.session_id, "--session-dir", policy.session_dir or ""]
|
||
else:
|
||
argv.append("--no-session")
|
||
argv += ["--provider", policy.provider, "--model", policy.model]
|
||
if policy.thinking:
|
||
argv += ["--thinking", policy.thinking]
|
||
# 上下文隔离:不加载项目 AGENTS.md/skills/extensions,角色合同全部来自任务包。
|
||
# --no-extensions 只禁自动发现,-e 显式传入的工具 server 扩展不受影响。
|
||
argv += ["--no-context-files", "--no-skills", "--no-extensions", "--no-approve"]
|
||
if policy.extension_path:
|
||
argv += ["-e", policy.extension_path]
|
||
allowlist = request.tool_allowlist
|
||
if allowlist:
|
||
argv += ["--tools", ",".join(allowlist)]
|
||
else:
|
||
argv += ["--no-tools"]
|
||
argv += ["--system-prompt", request.system_prompt, request.user_content]
|
||
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,
|
||
request: FrameworkExecutionRequest,
|
||
policy: ExecutionPolicy,
|
||
sink: TraceSink,
|
||
*,
|
||
timeout_seconds: float,
|
||
raw_sink: Callable[[bytes], None] | None = None,
|
||
) -> AgentStreamOutcome:
|
||
"""执行一次框架派发;框架层异常抛 FrameworkError(业务校验在派发器)。
|
||
|
||
raw_sink 逐行接收框架原始事件流字节(转录 tap),供派发器固定全量原始证据。
|
||
"""
|
||
|
||
argv = build_pi_argv(request, 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(request.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: TraceSink,
|
||
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>'}")
|
||
raw_args = event.get("args")
|
||
args = raw_args if isinstance(raw_args, Mapping) else None
|
||
outcome.tool_calls.append(
|
||
ToolCallRecord(tool_call_id=tool_call_id, name=name, args=args)
|
||
)
|
||
# 依赖清单材料:工具读取参数进事件账本(截断,避免撑爆事件行)。
|
||
args_summary: str | None = None
|
||
if args is not None:
|
||
args_summary = json.dumps(dict(args), ensure_ascii=False, sort_keys=True)
|
||
if len(args_summary) > 512:
|
||
args_summary = args_summary[:512] + "…"
|
||
sink.emit(
|
||
"tool.started",
|
||
status="ok",
|
||
tool_name=name,
|
||
details={"toolCallId": tool_call_id, "args": args_summary},
|
||
)
|
||
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},
|
||
)
|
||
else:
|
||
# 原始 JSONL 已由 raw_sink 保存;显式记录未知类型,避免把轨迹完整性误报为已知闭集。
|
||
unknown = str(kind or "<missing>")
|
||
outcome.unknown_event_types.append(unknown)
|
||
|
||
|
||
__all__ = [
|
||
"AgentStreamOutcome",
|
||
"DEFAULT_PI_BIN",
|
||
"ExecutionPolicy",
|
||
"FrameworkError",
|
||
"ModelCall",
|
||
"PiAgentRunner",
|
||
"ToolCallRecord",
|
||
"build_pi_argv",
|
||
"TraceSink",
|
||
]
|