141 lines
4.6 KiB
Python
141 lines
4.6 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""pbl_agent_runtime 挂载入口(M4a/M4b)。
|
||
|
||
load_pbl_agent_runtime(env=None):
|
||
1. ensure_tables 建 4 表(幂等)
|
||
2. **挂载即执行 self_check()**:22=13+9 数量不符、或写保护域存在启用写工具、
|
||
或 6 探针裁决不符 → 抛 RuntimeError,应用启动失败(fail-closed,绝不带病上线)
|
||
3. 注册 tools / enabled_tools / disabled_tools / api(judge/call_tool/verdict_stats/self_check)
|
||
4. 返回 default_verdict(= 'DENY')
|
||
|
||
对齐 docs/01-design/agent-tool-contract.md(13 启用 / 9 禁用,fail-closed)。
|
||
"""
|
||
|
||
from . import tables as _tables
|
||
from . import tool_registry as _registry
|
||
from . import verdict as _verdict
|
||
|
||
MODULE = "pbl_agent_runtime"
|
||
EXPECTED_TOTAL = 22
|
||
EXPECTED_ENABLED = 13
|
||
EXPECTED_DISABLED = 9
|
||
|
||
|
||
def ensure_tables(env=None, sor=None):
|
||
"""幂等建 4 表。返回 (ok_list, bad_list)。"""
|
||
return _tables.ensure_tables(env=env, sor=sor)
|
||
|
||
|
||
def self_check(env=None, run_probes=True):
|
||
"""模块自检:数量契约 + 写保护域 + 裁决探针。返回 (all_ok, msgs)。"""
|
||
msgs = []
|
||
all_ok = True
|
||
|
||
tools = _registry.all_tools()
|
||
enabled = [t for t in tools if t.get("enabled")]
|
||
disabled = [t for t in tools if not t.get("enabled")]
|
||
if len(tools) != EXPECTED_TOTAL:
|
||
all_ok = False
|
||
msgs.append("工具总数=%d 应为 %d" % (len(tools), EXPECTED_TOTAL))
|
||
if len(enabled) != EXPECTED_ENABLED:
|
||
all_ok = False
|
||
msgs.append("启用工具=%d 应为 %d" % (len(enabled), EXPECTED_ENABLED))
|
||
if len(disabled) != EXPECTED_DISABLED:
|
||
all_ok = False
|
||
msgs.append("禁用工具=%d 应为 %d" % (len(disabled), EXPECTED_DISABLED))
|
||
msgs.append("工具契约:%d = %d 启用 + %d 禁用" % (len(tools), len(enabled), len(disabled)))
|
||
|
||
reg_ok, reg_msgs = _registry.self_check()
|
||
if not reg_ok:
|
||
all_ok = False
|
||
msgs.extend(reg_msgs)
|
||
|
||
tbl_ok, tbl_msgs = _tables.self_check()
|
||
if not tbl_ok:
|
||
all_ok = False
|
||
msgs.extend(tbl_msgs)
|
||
|
||
if run_probes:
|
||
v_ok, v_msgs = _verdict.runtime_self_check()
|
||
if not v_ok:
|
||
all_ok = False
|
||
msgs.extend(v_msgs)
|
||
|
||
if _verdict.DEFAULT_VERDICT != "DENY":
|
||
all_ok = False
|
||
msgs.append("DEFAULT_VERDICT=%r 必须为 'DENY'(fail-closed)"
|
||
% _verdict.DEFAULT_VERDICT)
|
||
|
||
if all_ok:
|
||
msgs.append("SELF_CHECK %s: PASS %d/%d" % (MODULE, len(tools), len(tools)))
|
||
return all_ok, msgs
|
||
|
||
|
||
def api():
|
||
"""对外 API 契约(供应用/其它模块调用)。"""
|
||
return {
|
||
"judge": _verdict.judge,
|
||
"call_tool": _verdict.call_tool,
|
||
"verdict_stats": _verdict.verdict_stats,
|
||
"self_check": self_check,
|
||
"get_tool": _registry.get_tool,
|
||
"all_tools": _registry.all_tools,
|
||
"ensure_tables": ensure_tables,
|
||
"POLICY_VERSION": _verdict.POLICY_VERSION,
|
||
"MAX_CALLS_PER_RUN": _verdict.MAX_CALLS_PER_RUN,
|
||
"PblError": _verdict.PblError,
|
||
}
|
||
|
||
|
||
def load_pbl_agent_runtime(env=None):
|
||
"""挂载入口:建表 → 自检(不过即抛)→ 注册契约 → 返回 default_verdict。"""
|
||
srv = env
|
||
if srv is None:
|
||
try:
|
||
from ahserver.serverenv import ServerEnv
|
||
srv = ServerEnv()
|
||
except Exception: # noqa: BLE001
|
||
srv = None
|
||
|
||
if srv is not None:
|
||
ensure_tables(srv)
|
||
|
||
ok, msgs = self_check()
|
||
for m in msgs:
|
||
_log(srv, m)
|
||
if not ok:
|
||
raise RuntimeError(
|
||
"%s self_check FAILED(fail-closed,拒绝启动):%s"
|
||
% (MODULE, "; ".join([m for m in msgs if "FAIL" in m or "应为" in m or "必须" in m][:6]))
|
||
)
|
||
|
||
if srv is not None:
|
||
setattr(srv, "pbl_agent_tools", _registry.all_tools())
|
||
setattr(srv, "pbl_agent_enabled_tools", [t["name"] for t in _registry.ENABLED_TOOLS])
|
||
setattr(srv, "pbl_agent_disabled_tools", [t["name"] for t in _registry.DISABLED_TOOLS])
|
||
setattr(srv, "pbl_agent_api", api())
|
||
modules = getattr(srv, "modules", None)
|
||
if isinstance(modules, list) and MODULE not in modules:
|
||
modules.append(MODULE)
|
||
|
||
return _verdict.DEFAULT_VERDICT
|
||
|
||
|
||
def _log(srv, msg):
|
||
logger = getattr(srv, "logger", None) if srv is not None else None
|
||
if logger is not None and hasattr(logger, "info"):
|
||
try:
|
||
logger.info("[%s] %s" % (MODULE, msg))
|
||
return
|
||
except Exception: # noqa: BLE001
|
||
pass
|
||
print("[%s] %s" % (MODULE, msg))
|
||
|
||
|
||
if __name__ == "__main__":
|
||
_ok, _msgs = self_check(run_probes=True)
|
||
for _m in _msgs:
|
||
print(_m)
|
||
print("RESULT: %s" % ("PASS" if _ok else "FAIL"))
|
||
raise SystemExit(0 if _ok else 1)
|