feat(v3): T29 token 级流式输出(SSE 流式解析 + delta 事件 + 打字机渲染,D10)

- OpenAICompatChat 默认 stream=True:httpx 流式解析 OpenAI chunk,
  tool_calls 碎片按 index 组装(不在正文展示),include_usage 计量;
  失败自动回退非流式一次(已有部分增量输出则如实抛出);
  补 raise_for_status(修复流式 4xx 误入回退的缺陷)
- ToolLoop 增 on_delta(签名探测兼容 2/3 参 chat_fn);_loads_json_object 宽松解析
- DeltaThrottle >=48 字符节流落 delta 事件;单模型与两级模式(含规划者)接线
- 前端:streamText 打字机渲染 + 光标动画;工具/阶段事件到达时清空归档
- 配套修复:AgentView 闭包持有 push 前原始对象导致响应式丢失、过程事件不渲染
- 测试 +7(SSE 解析/碎片组装/回退/部分失败抛出/on_delta/节流/service delta 事件),
  全量 296 passed
This commit is contained in:
tzt
2026-09-01 23:41:45 +08:00
parent ca8add02e2
commit ddef8bec1c
18 changed files with 786 additions and 62 deletions
+277 -33
View File
@@ -1,8 +1,10 @@
"""工具调用内核 —— 让 LLM 以 OpenAI function-calling 协议操作工作区文件。
组成(对齐《实现方案_v4_模型池与工具智能体.md》D4):
- TOOLS_SPEClist_dir / read_file / write_file 三个工具的 OpenAI tools 声明
- WorkspaceTools:被"关押"在根目录内的文件工具(路径越界一律拒绝,Windows pathlib
- TOOLS_SPEClist_dir / read_file / write_file / edit_file / search_files /
run_command / web_fetch 七个工具的 OpenAI tools 声明
- WorkspaceTools:被"关押"在根目录内的文件工具(路径越界一律拒绝,Windows pathlib);
web_fetch 带公网 SSRF 防护(dsh web_fetch 同款),run_command 带危险命令拦截
- parse_tool_calls:解析 OpenAI 响应里的 tool_callsarguments 容错为 {}
- ToolLoop:通用智能体循环。chat_fn 注入(网关传 OpenAI 兼容客户端,测试传假实现),
本模块只负责循环编排:调用 -> 执行工具 -> 回喂结果 -> 直到模型给出最终答复。
@@ -11,9 +13,13 @@
- 纯标准库(router_system 零第三方依赖不变)
- 工具结果回喂前截断(防止上下文爆炸),轮数与 token 双上限(金额护栏)
- 事件回调 on_event 逐条产出过程事件(供 SSE 透出"智能体在做什么"
- 工具在线程池执行(asyncio.to_thread),长命令/大搜索不阻塞事件循环
- 重复同参调用达到阈值回喂警语(防模型原地打转,dsh repeat-reminder 同款)
"""
from __future__ import annotations
import asyncio
import ipaddress
import json
from pathlib import Path
from typing import Any, Awaitable, Callable, Dict, List, Optional
@@ -33,6 +39,17 @@ SEARCH_SKIP_DIRS = {".git", "node_modules", "__pycache__", ".venv", "venv", "dis
# run_command 上限
SHELL_OUTPUT_CHARS = 4000
# web_fetch 上限(SSRF 防护:仅公网 http/https,拒绝私网/环回/链路本地地址)
FETCH_MAX_CHARS = 12000
FETCH_TIMEOUT_S = 15.0
FETCH_MAX_BYTES = 512 * 1024
FETCH_ALLOWED_SCHEMES = ("http", "https")
# NAT64 Well-Known Prefixdsh 同款拒绝项)
_NAT64_PREFIX = ipaddress.ip_network("64:ff9b::/96")
# 重复调用提醒阈值(同一工具 + 同一参数第 N 次起提示模型换策略)
REPEAT_CALL_WARN_AT = 3
# 默认循环上限
DEFAULT_MAX_ROUNDS = 8
@@ -128,6 +145,21 @@ TOOLS_SPEC: List[Dict[str, Any]] = [
},
},
},
{
"type": "function",
"function": {
"name": "web_fetch",
"description": "抓取一个公网 http/https URL 的文本内容(如查文档/接口说明),"
"返回截断后的正文。私网/环回地址会被拒绝;仅当系统开启 allow_net 时可用。",
"parameters": {
"type": "object",
"properties": {
"url": {"type": "string", "description": "要抓取的完整 URLhttp/https"},
},
"required": ["url"],
},
},
},
]
TOOL_NAMES = {t["function"]["name"] for t in TOOLS_SPEC}
@@ -167,6 +199,71 @@ class ToolError(Exception):
"""工具执行失败(路径越界/不存在/参数非法)。"""
def _atomic_write_text(p: Path, text: str) -> None:
"""原子写文本:同目录临时文件 + os.replace(防半截文件;dsh atomic-write 同款)。
Windows 上目标被占用时 os.replace 可能 EPERM:小退避重试一次,
仍失败则退回直接写(保可用性,牺牲原子性)。
"""
import os
import tempfile
import time
tmp = None
try:
with tempfile.NamedTemporaryFile(
"w", encoding="utf-8", dir=str(p.parent),
prefix=p.name + ".", suffix=".tmp", delete=False) as f:
tmp = Path(f.name)
f.write(text)
for attempt in (0, 1):
try:
os.replace(tmp, p)
return
except PermissionError:
if attempt == 0:
time.sleep(0.05)
raise
except PermissionError:
if tmp is not None:
tmp.unlink(missing_ok=True)
p.write_text(text, encoding="utf-8") # 退路:非原子但保可用
except BaseException:
if tmp is not None:
tmp.unlink(missing_ok=True)
raise
# 危险命令模式(大小写不敏感):宁可误拦不可漏拦(用户可换写法绕开误拦项)
_DANGEROUS_PATTERNS = [
(r"\bformat\b\s+[a-z]:", "格式化磁盘"),
(r"\brd\s+/s", "递归删除目录"),
(r"\brmdir\s+/s", "递归删除目录"),
(r"\bdel\s+/[fsmq]", "强制/递归删除"),
(r"\brm\s+(-[a-z]*r[a-z]*f|-[a-z]*f[a-z]*r)\s+[/~]", "递归强制删除根/家目录"),
(r"\bshutdown\b", "关机/重启"),
(r"\bdiskpart\b", "磁盘分区操作"),
(r"\bbcdedit\b", "启动配置修改"),
(r"\breg\s+delete\b", "注册表删除"),
(r"\bvssadmin\s+delete\b", "卷影副本删除"),
(r"\bmkfs\b", "格式化文件系统"),
(r"\bdd\s+if=", "裸磁盘写入"),
(r":\(\)\s*\{.*\};\s*:", "fork 炸弹"),
(r"\bcurl\b[^|]*\|\s*(ba)?sh\b", "下载并执行脚本"),
(r"\bwget\b[^|]*\|\s*(ba)?sh\b", "下载并执行脚本"),
(r"\biwr\b[^|]*\|\s*iex\b", "下载并执行脚本"),
]
def _dangerous_command_reason(command: str) -> str:
"""命中危险命令模式时返回原因,否则返回空串(审批之外的独立防线)。"""
import re
lowered = command.lower()
for pattern, reason in _DANGEROUS_PATTERNS:
if re.search(pattern, lowered):
return reason
return ""
class WorkspaceTools:
"""被限制在根目录内的文件工具(智能体的"")。
@@ -175,11 +272,13 @@ class WorkspaceTools:
"""
def __init__(self, root: str | Path,
allow_shell: bool = False, shell_timeout_s: int = 20):
allow_shell: bool = False, shell_timeout_s: int = 20,
allow_net: bool = True):
self.root = Path(root).resolve()
self.root.mkdir(parents=True, exist_ok=True)
self.allow_shell = bool(allow_shell)
self.shell_timeout_s = max(1, int(shell_timeout_s))
self.allow_net = bool(allow_net)
# ---------- 路径关押 ----------
def resolve(self, rel_path: str) -> Path:
@@ -234,7 +333,7 @@ class WorkspaceTools:
return {"ok": False, "error": f"内容过长(>{MAX_WRITE_CHARS} 字符),拒绝写入"}
p = self.resolve(rel_path)
p.parent.mkdir(parents=True, exist_ok=True)
p.write_text(content, encoding="utf-8")
_atomic_write_text(p, content)
return {"ok": True, "path": rel_path, "bytes_written": len(content.encode("utf-8"))}
def edit_file(self, rel_path: str, old_string: str, new_string: str) -> Dict[str, Any]:
@@ -257,7 +356,7 @@ class WorkspaceTools:
return {"ok": False,
"error": f"old_string 出现 {count} 次(要求唯一);请扩大上下文使其唯一"}
new_text = text.replace(old_string, new_string, 1)
p.write_text(new_text, encoding="utf-8")
_atomic_write_text(p, new_text)
return {
"ok": True, "path": rel_path,
"replaced": 1,
@@ -265,7 +364,8 @@ class WorkspaceTools:
}
def search_files(self, query: str, rel_path: str = "") -> Dict[str, Any]:
"""跨文件文本搜索(跳过依赖/构建目录与二进制大文件,限量返回)。"""
"""跨文件文本搜索(os.walk 修剪依赖/构建目录,限量返回,不跟随符号链接)。"""
import os
if not query:
return {"ok": False, "error": "query 不能为空"}
base = self.resolve(rel_path or "")
@@ -274,40 +374,61 @@ class WorkspaceTools:
matches: List[Dict[str, Any]] = []
scanned = 0
truncated = False
for p in sorted(base.rglob("*")):
if len(matches) >= SEARCH_MAX_MATCHES:
truncated = True
break
if not p.is_file():
continue
rel_parts = p.relative_to(base).parts
if any(part in SEARCH_SKIP_DIRS for part in rel_parts):
continue
try:
if p.stat().st_size > SEARCH_MAX_FILE_BYTES:
continue
text = p.read_text(encoding="utf-8")
except (OSError, UnicodeDecodeError):
continue # 二进制/不可读,跳过
scanned += 1
if scanned > SEARCH_MAX_FILES:
truncated = True
break
for lineno, line in enumerate(text.splitlines(), 1):
def _match_file(p: Path) -> bool:
"""在单文件内找匹配(找到即 True)。"""
nonlocal matches, truncated
for lineno, line in enumerate(p.read_text(encoding="utf-8").splitlines(), 1):
if query in line:
rel = Path(*p.relative_to(self.root).parts).as_posix()
rel = p.relative_to(self.root).as_posix()
matches.append({
"file": rel, "line": lineno,
"text": line.strip()[:300],
})
if len(matches) >= SEARCH_MAX_MATCHES:
truncated = True
return True
return False
for dirpath, dirnames, filenames in os.walk(base, followlinks=False):
# 修剪依赖/构建目录:不进入(rglob 全量物化在大工作区上不可接受)
dirnames[:] = sorted((d for d in dirnames if d not in SEARCH_SKIP_DIRS),
key=str.lower)
if len(matches) >= SEARCH_MAX_MATCHES or scanned > SEARCH_MAX_FILES:
truncated = True
break
for fname in sorted(filenames, key=str.lower):
fpath = Path(dirpath) / fname
try:
if not fpath.is_file():
continue
if fpath.stat().st_size > SEARCH_MAX_FILE_BYTES:
continue
scanned += 1
if scanned > SEARCH_MAX_FILES:
truncated = True
break
if _match_file(fpath):
break
except (OSError, UnicodeDecodeError):
continue # 二进制/不可读/并发删除,跳过
if len(matches) >= SEARCH_MAX_MATCHES:
truncated = True
break
return {"ok": True, "query": query, "matches": matches,
"scanned_files": scanned, "truncated": truncated}
def run_command(self, command: str) -> Dict[str, Any]:
"""在工作区根目录执行 shell 命令(默认关闭,allow_shell 开启后可用)。"""
"""在工作区根目录执行一条 shell 命令(默认关闭,allow_shell 开启后可用)。
安全设计(dsh bash 工具同款语义):
- 危险命令模式先拦截(即使审批通过也拒绝)
- 用显式 shell 解释器的参数列表执行(cmd /c 或 /bin/sh -c),
命令字符串对解释器可见属于功能本体,防护依赖 allow_shell 开关
+ 审批门卫 + 超时熔断 + cwd 关押在本工作区
"""
import os
import subprocess
if not self.allow_shell:
return {"ok": False,
"error": "run_command 未启用(系统设置 allow_shell 为关)。"
@@ -315,11 +436,17 @@ class WorkspaceTools:
command = (command or "").strip()
if not command:
return {"ok": False, "error": "command 不能为空"}
import subprocess
creationflags = 0x08000000 if __import__("os").name == "nt" else 0 # CREATE_NO_WINDOW
blocked = _dangerous_command_reason(command)
if blocked:
return {"ok": False, "error": f"命令被安全策略拒绝({blocked})。请换一种不具破坏性的做法。"}
if os.name == "nt":
argv = [os.environ.get("COMSPEC", "cmd.exe"), "/c", command]
else:
argv = ["/bin/sh", "-c", command]
creationflags = 0x08000000 if os.name == "nt" else 0 # CREATE_NO_WINDOW
try:
proc = subprocess.run(
command, shell=True, cwd=str(self.root), capture_output=True,
argv, cwd=str(self.root), capture_output=True,
timeout=self.shell_timeout_s, creationflags=creationflags,
)
out = (proc.stdout or b"").decode("utf-8", errors="replace")
@@ -334,6 +461,72 @@ class WorkspaceTools:
except OSError as e:
return {"ok": False, "error": f"命令执行失败: {type(e).__name__}: {e}"}
def web_fetch(self, url: str) -> Dict[str, Any]:
"""抓取公网 URL 文本(SSRF 防护:拒绝非 http(s)、私网/环回/链路本地/NAT64 目标)。
校验流程对齐 dsh web_fetch:解析 DNS -> 全部地址必须公网 -> 才发起请求;
响应限量(字节/字符)、超时熔断、二进制嗅探拒绝。
"""
import socket
import urllib.error
import urllib.parse
import urllib.request
if not self.allow_net:
return {"ok": False, "error": "web_fetch 未启用(系统设置 allow_net 为关)。"}
raw = (url or "").strip()
try:
parsed = urllib.parse.urlsplit(raw)
except ValueError:
return {"ok": False, "error": f"URL 无法解析: {raw[:120]}"}
if parsed.scheme.lower() not in FETCH_ALLOWED_SCHEMES:
return {"ok": False, "error": f"仅允许 http/https URL(收到 {parsed.scheme or ''}"}
host = parsed.hostname or ""
if not host:
return {"ok": False, "error": "URL 缺少主机名"}
try:
port = parsed.port
except ValueError:
return {"ok": False, "error": "URL 端口非法"}
# DNS 解析后逐一校验:任何私网/环回/链路本地/保留/NAT64 地址都拒绝
try:
infos = socket.getaddrinfo(host, port or (443 if parsed.scheme == "https" else 80),
proto=socket.IPPROTO_TCP)
except socket.gaierror as e:
return {"ok": False, "error": f"域名解析失败: {host} ({e})"}
for info in infos:
ip = info[4][0]
try:
addr = ipaddress.ip_address(ip.split("%")[0]) # 剥 zone id
except ValueError:
return {"ok": False, "error": f"解析出非法地址: {ip}"}
if (addr.is_private or addr.is_loopback or addr.is_link_local
or addr.is_reserved or addr.is_multicast or addr.is_unspecified
or (addr.version == 6 and addr in _NAT64_PREFIX)):
return {"ok": False,
"error": f"目标地址 {addr} 属于内网/保留段,已被 SSRF 防护拒绝"}
try:
req = urllib.request.Request(raw, headers={"User-Agent": "router-agent/1.0"})
with urllib.request.urlopen(req, timeout=FETCH_TIMEOUT_S) as resp:
body = resp.read(FETCH_MAX_BYTES + 1)
charset = resp.headers.get_content_charset() or "utf-8"
except urllib.error.HTTPError as e:
return {"ok": False, "error": f"HTTP {e.code}: {e.reason}"}
except (urllib.error.URLError, OSError, ValueError) as e:
return {"ok": False, "error": f"抓取失败: {type(e).__name__}: {e}"}
if len(body) > FETCH_MAX_BYTES:
return {"ok": False, "error": f"响应超过 {FETCH_MAX_BYTES // 1024}KB 上限,拒绝处理"}
if b"\x00" in body[:512]:
return {"ok": False, "error": "非文本内容(检测到二进制),拒绝处理"}
try:
text = body.decode(charset, errors="replace")
except LookupError:
text = body.decode("utf-8", errors="replace")
truncated = len(text) > FETCH_MAX_CHARS
return {"ok": True, "url": raw,
"content": text[:FETCH_MAX_CHARS], "truncated": truncated,
"total_chars": len(text)}
# ---------- 统一执行入口 ----------
def execute(self, name: str, arguments: Dict[str, Any]) -> Dict[str, Any]:
"""按名字执行工具;任何异常折叠为 {"ok": False, "error": ...}。"""
@@ -355,6 +548,8 @@ class WorkspaceTools:
str(arguments.get("query", "")), str(arguments.get("path", "")))
if name == "run_command":
return self.run_command(str(arguments.get("command", "")))
if name == "web_fetch":
return self.web_fetch(str(arguments.get("url", "")))
return {"ok": False, "error": f"未知工具: {name}"}
except ToolError as e:
return {"ok": False, "error": str(e)}
@@ -362,6 +557,28 @@ class WorkspaceTools:
return {"ok": False, "error": f"文件系统错误: {type(e).__name__}: {e}"}
def _loads_json_object(text: str) -> Dict[str, Any]:
"""宽松解析 JSON 对象:剥除 markdown 围栏、取首个 {...};失败返回 {}"""
t = (text or "").strip()
fence = "`" * 3
t = t.replace(fence + "json", fence).replace(fence, "").strip()
try:
obj = json.loads(t)
if isinstance(obj, dict):
return obj
except json.JSONDecodeError:
pass
start, end = t.find("{"), t.rfind("}")
if start != -1 and end > start:
try:
obj = json.loads(t[start:end + 1])
if isinstance(obj, dict):
return obj
except json.JSONDecodeError:
pass
return {}
def parse_tool_calls(message: Dict[str, Any]) -> List[Dict[str, Any]]:
"""从 OpenAI 响应的 message 解析 tool_calls。
@@ -415,6 +632,7 @@ class ToolLoop:
result_preview_chars: int = MAX_RESULT_CHARS,
emit_final: bool = True,
approval_hook: Optional[Callable[[str, Dict[str, Any]], Awaitable[bool]]] = None,
on_delta: Optional[Callable[[str], None]] = None,
):
self.tools = tools
self.chat_fn = chat_fn
@@ -425,6 +643,11 @@ class ToolLoop:
self.emit_final = emit_final # 两级模式内层循环置 False,由外层统一收尾
# 审批门卫(D9):执行工具前调用,返回 False = 用户拒绝(可选;缺省跳过审批)
self.approval_hook = approval_hook
# 流式增量回调(D10,可选):chat_fn 支持 on_delta 参数时逐段转发模型文本
self.on_delta = on_delta
self._accepts_delta: Optional[bool] = None
# 重复调用计数(dsh repeat-tool-reminder 同款提醒,防模型原地打转)
self._call_counts: Dict[tuple, int] = {}
def _emit(self, ev: Dict[str, Any]) -> None:
if ev.get("type") == "final" and not self.emit_final:
@@ -438,6 +661,18 @@ class ToolLoop:
def _total_tokens(self, usage: Dict[str, int]) -> int:
return int(usage.get("prompt_tokens", 0)) + int(usage.get("completion_tokens", 0))
def _invoke_chat(self, messages: List[Dict[str, Any]]) -> Awaitable[Dict[str, Any]]:
"""调用 chat_fn;其签名支持 on_delta 时才传入(对旧假实现向后兼容)。"""
if self._accepts_delta is None:
try:
import inspect
self._accepts_delta = len(inspect.signature(self.chat_fn).parameters) >= 3
except (TypeError, ValueError):
self._accepts_delta = False
if self.on_delta is not None and self._accepts_delta:
return self.chat_fn(messages, TOOLS_SPEC, self.on_delta)
return self.chat_fn(messages, TOOLS_SPEC)
async def run(self, task: str, system: str = "",
history: Optional[List[Dict[str, Any]]] = None) -> Dict[str, Any]:
"""执行任务直到模型给出最终答复或触顶。
@@ -458,7 +693,7 @@ class ToolLoop:
for round_no in range(1, self.max_rounds + 1):
self._emit({"type": "round", "round": round_no})
try:
resp = await self.chat_fn(messages, TOOLS_SPEC)
resp = await self._invoke_chat(messages)
except Exception as e:
self._emit({"type": "final", "round": round_no, "reason": "error",
"error": f"{type(e).__name__}: {e}"})
@@ -521,10 +756,19 @@ class ToolLoop:
messages.append({"role": "tool", "tool_call_id": c["id"],
"content": preview})
continue
result = self.tools.execute(c["name"], c["arguments"])
result = await asyncio.to_thread(self.tools.execute, c["name"], c["arguments"])
preview = json.dumps(result, ensure_ascii=False)
if len(preview) > self.result_preview_chars:
preview = preview[:self.result_preview_chars] + "…(截断)"
# 重复调用提醒:同一工具同一参数第 N 次起,回喂时附警语促使换策略
key = (c["name"], json.dumps(c["arguments"], sort_keys=True, ensure_ascii=False))
seen = self._call_counts.get(key, 0) + 1
self._call_counts[key] = seen
if seen >= REPEAT_CALL_WARN_AT:
preview += (f"\n[系统提示] 该工具已第 {seen} 次以完全相同的参数调用。"
"重复同样的调用不会带来新信息;请改变做法或直接给出最终答复。")
self._emit({"type": "repeat_warning", "round": round_no,
"name": c["name"], "count": seen})
self._emit({"type": "tool_result", "round": round_no,
"name": c["name"], "ok": bool(result.get("ok")),
"preview": preview})