feat(memory): P1 长期记忆升级 —— 异步攒批 Consolidate + 软删 + importance/last_seen #1

Merged
Blizzard merged 181 commits from feat/wails3 into main 2026-07-17 01:12:32 +00:00
31 changed files with 1731 additions and 107 deletions
Showing only changes of commit 17955f0088 - Show all commits
+1 -1
View File
@@ -10,7 +10,7 @@
- [x] **Phase A · 地基**`llm.Pool` 换 Eino ChatModel 组件(commit d84b1ec,验收通过) - [x] **Phase A · 地基**`llm.Pool` 换 Eino ChatModel 组件(commit d84b1ec,验收通过)
- [x] **Phase B · 质变**MCP 工具→`InvokableTool` + ReAct agent(模型自主调工具,验收 7/7 命中) - [x] **Phase B · 质变**MCP 工具→`InvokableTool` + ReAct agent(模型自主调工具,验收 7/7 命中)
- [x] **Phase C · 编排归一**:✅ 全图 `DSL→compose.Graph` 编译器(全节点 + branch + DAG 并行调度)+ callbacks→ExecEvent 归一,等价回归通过(EINO_COMPOSE 灰度开关,默认关;过渡期后 graph.go 退役) - [x] **Phase C · 编排归一**:✅ 全图 `DSL→compose.Graph` 编译器(全节点 + branch + DAG 并行调度)+ callbacks→ExecEvent 归一,等价回归通过(EINO_COMPOSE 灰度开关,默认关;过渡期后 graph.go 退役)
- [~] **Phase D · 状态化执行**:✅ 任务生命周期 FSM(已完成)/ HITL 中断恢复 / ⬜ 多智能体(按场景) - [~] **Phase D · 状态化执行**:✅ 任务生命周期 FSM / HITL 人工审批中断(审批节点暂停→waiting→批准回 running / 拒绝→rejectedNATS 决定回传 + AckWait 续租,全栈含 Studio 审批节点 + 运行抽屉批准条)/ ⬜ 多智能体(按场景)
**组件化补完(A**:检索 → `ragRetriever`(`components/retriever.Retriever``eino_components.go`);提示词 → `buildMessages` 改用 `prompt.FromMessages`+`MessagesPlaceholder`;工具 → `mcpTool`(`InvokableTool`,模型自主调用)。至此终态架构 8 层中 模型/工具/检索/提示词/编排/智能体(单)/可观测 均已 Eino 组件化;剩 人机交互(中断恢复,按场景)。 **组件化补完(A**:检索 → `ragRetriever`(`components/retriever.Retriever``eino_components.go`);提示词 → `buildMessages` 改用 `prompt.FromMessages`+`MessagesPlaceholder`;工具 → `mcpTool`(`InvokableTool`,模型自主调用)。至此终态架构 8 层中 模型/工具/检索/提示词/编排/智能体(单)/可观测 均已 Eino 组件化;剩 人机交互(中断恢复,按场景)。
**自主 agent 工具集动态化**:agent 工具集不再硬编码——mcp-go 注册表(单一事实源)每个工具声明 `agent/params/inject``list_tools` 上报,dispatcher `agentTools()` 动态发现并建 `InvokableTool`、运行时注入 `user_id/session_id/kb`(不暴露给模型)。**加工具只改 mcp-go 注册表,dispatcher 零改动。** 当前暴露 4 个(wiki_search/recall_user_memory/remember_user_fact/history_get);实测模型自主调用新暴露的 remember_user_fact 成功(参数自生成、user_id 服务端注入)。 **自主 agent 工具集动态化**:agent 工具集不再硬编码——mcp-go 注册表(单一事实源)每个工具声明 `agent/params/inject``list_tools` 上报,dispatcher `agentTools()` 动态发现并建 `InvokableTool`、运行时注入 `user_id/session_id/kb`(不暴露给模型)。**加工具只改 mcp-go 注册表,dispatcher 零改动。** 当前暴露 4 个(wiki_search/recall_user_memory/remember_user_fact/history_get);实测模型自主调用新暴露的 remember_user_fact 成功(参数自生成、user_id 服务端注入)。
+20 -1
View File
@@ -14,7 +14,7 @@ import { Placeholder } from "./views/Placeholder";
import { CommandPalette, type Command } from "./components/CommandPalette"; import { CommandPalette, type Command } from "./components/CommandPalette";
import { UpdateBanner } from "./components/UpdateBanner"; import { UpdateBanner } from "./components/UpdateBanner";
import { Login } from "./views/Login"; import { Login } from "./views/Login";
import { submitTask, streamTokens, streamExec, authMe, logout, type Identity, type AuthUser } from "./lib/api"; import { submitTask, streamTokens, streamExec, taskStatus, authMe, logout, type Identity, type AuthUser } from "./lib/api";
import type { TaskDsl } from "./lib/dsl"; import type { TaskDsl } from "./lib/dsl";
import { emptyRun, type RunState } from "./lib/run"; import { emptyRun, type RunState } from "./lib/run";
import { ToastProvider } from "./ui"; import { ToastProvider } from "./ui";
@@ -48,6 +48,13 @@ export default function App() {
const closeRef = useRef<(() => void) | null>(null); const closeRef = useRef<(() => void) | null>(null);
const execCloseRef = useRef<(() => void) | null>(null); const execCloseRef = useRef<(() => void) | null>(null);
const pollRef = useRef<number | null>(null);
const stopPoll = () => {
if (pollRef.current != null) {
window.clearInterval(pollRef.current);
pollRef.current = null;
}
};
// 全局 ⌘K / Ctrl+K 唤起命令面板(键盘优先工作站入口)。 // 全局 ⌘K / Ctrl+K 唤起命令面板(键盘优先工作站入口)。
useEffect(() => { useEffect(() => {
@@ -98,6 +105,7 @@ export default function App() {
async (dsl: TaskDsl) => { async (dsl: TaskDsl) => {
closeRef.current?.(); closeRef.current?.();
execCloseRef.current?.(); execCloseRef.current?.();
stopPoll();
const t0 = Date.now(); const t0 = Date.now();
setRun({ phase: "submitting", output: "", events: [{ t: 0, label: "提交任务" }], exec: [] }); setRun({ phase: "submitting", output: "", events: [{ t: 0, label: "提交任务" }], exec: [] });
try { try {
@@ -109,6 +117,17 @@ export default function App() {
taskId, taskId,
events: [...r.events, { t: Date.now() - t0, label: `已发布 ${taskId}` }], events: [...r.events, { t: Date.now() - t0, label: `已发布 ${taskId}` }],
})); }));
// 轮询后端任务状态:可靠捕获 waiting(审批中断)—— exec 实时事件会抢跑,状态落 PG 不会丢。
const terminal = new Set(["done", "failed", "timeout", "rejected"]);
pollRef.current = window.setInterval(async () => {
try {
const s = await taskStatus(taskId);
setRun((r) => (r.taskId === taskId ? { ...r, lifecycle: s.status, detail: s.detail } : r));
if (terminal.has(s.status)) stopPoll();
} catch {
/* 忽略瞬时失败,下个 tick 再试 */
}
}, 1500);
// 执行轨迹(运行·观测):与 token 流并行订阅,逐节点点亮。 // 执行轨迹(运行·观测):与 token 流并行订阅,逐节点点亮。
execCloseRef.current = streamExec( execCloseRef.current = streamExec(
taskId, taskId,
@@ -0,0 +1,126 @@
import type { ChartSpec } from "../lib/chartspec";
// ChartView:把 chart spec 自绘成 SVGbar / line / pie),无第三方图表依赖。
const PALETTE = ["#60a5fa", "#34d399", "#fbbf24", "#f87171", "#a78bfa", "#22d3ee", "#fb923c", "#4ade80"];
export function ChartView({ spec }: { spec: ChartSpec }) {
return (
<figure className="my-2 rounded-lg border border-line bg-ink-950/50 p-3">
{spec.title && <figcaption className="mb-2 text-[12px] font-medium text-slate-200">{spec.title}</figcaption>}
{spec.type === "pie" ? <Pie spec={spec} /> : <BarLine spec={spec} />}
{spec.type !== "pie" && spec.series.length > 1 && <Legend names={spec.series.map((s, i) => s.name || `系列${i + 1}`)} />}
</figure>
);
}
function Legend({ names }: { names: string[] }) {
return (
<div className="mt-2 flex flex-wrap gap-3">
{names.map((n, i) => (
<span key={i} className="flex items-center gap-1 text-[10px] text-slate-400">
<span className="inline-block h-2 w-2 rounded-sm" style={{ background: PALETTE[i % PALETTE.length] }} />
{n}
</span>
))}
</div>
);
}
// BarLine:柱状图 / 折线图(共用坐标系)。
function BarLine({ spec }: { spec: ChartSpec }) {
const W = 480, H = 220, padL = 40, padB = 28, padT = 8, padR = 8;
const plotW = W - padL - padR, plotH = H - padT - padB;
const all = spec.series.flatMap((s) => s.data);
const max = Math.max(1, ...all);
const min = Math.min(0, ...all);
const span = max - min || 1;
const y = (v: number) => padT + plotH - ((v - min) / span) * plotH;
const n = spec.labels.length;
const slot = plotW / n;
return (
<svg viewBox={`0 0 ${W} ${H}`} className="w-full" role="img" aria-label={spec.title || "图表"}>
{/* y 轴基准线 */}
<line x1={padL} y1={y(min)} x2={W - padR} y2={y(min)} stroke="#334155" strokeWidth={1} />
<text x={padL - 6} y={y(max)} fill="#64748b" fontSize={9} textAnchor="end">{fmt(max)}</text>
<text x={padL - 6} y={y(min) + 3} fill="#64748b" fontSize={9} textAnchor="end">{fmt(min)}</text>
{spec.type === "bar"
? spec.series.map((s, si) =>
s.data.map((v, i) => {
const bw = (slot * 0.7) / spec.series.length;
const x = padL + i * slot + slot * 0.15 + si * bw;
return (
<rect key={`${si}-${i}`} x={x} y={Math.min(y(v), y(0))} width={bw} height={Math.abs(y(v) - y(0))}
fill={PALETTE[si % PALETTE.length]} rx={1}>
<title>{`${spec.labels[i]}: ${v}`}</title>
</rect>
);
}),
)
: spec.series.map((s, si) => {
const pts = s.data.map((v, i) => `${padL + i * slot + slot / 2},${y(v)}`).join(" ");
return (
<g key={si}>
<polyline points={pts} fill="none" stroke={PALETTE[si % PALETTE.length]} strokeWidth={2} />
{s.data.map((v, i) => (
<circle key={i} cx={padL + i * slot + slot / 2} cy={y(v)} r={2.5} fill={PALETTE[si % PALETTE.length]}>
<title>{`${spec.labels[i]}: ${v}`}</title>
</circle>
))}
</g>
);
})}
{/* x 轴标签 */}
{spec.labels.map((lb, i) => (
<text key={i} x={padL + i * slot + slot / 2} y={H - padB + 14} fill="#64748b" fontSize={9} textAnchor="middle">
{lb.length > 6 ? lb.slice(0, 6) + "…" : lb}
</text>
))}
</svg>
);
}
// Pie:饼图(用第一条系列)。
function Pie({ spec }: { spec: ChartSpec }) {
const data = spec.series[0]?.data ?? [];
const total = data.reduce((a, b) => a + Math.max(0, b), 0) || 1;
const cx = 110, cy = 110, r = 90;
let acc = -Math.PI / 2; // 从 12 点方向起
const arcs = data.map((v, i) => {
const frac = Math.max(0, v) / total;
const a0 = acc;
const a1 = acc + frac * Math.PI * 2;
acc = a1;
const large = a1 - a0 > Math.PI ? 1 : 0;
const x0 = cx + r * Math.cos(a0), y0 = cy + r * Math.sin(a0);
const x1 = cx + r * Math.cos(a1), y1 = cy + r * Math.sin(a1);
const d = `M${cx},${cy} L${x0.toFixed(2)},${y0.toFixed(2)} A${r},${r} 0 ${large} 1 ${x1.toFixed(2)},${y1.toFixed(2)} Z`;
return { d, color: PALETTE[i % PALETTE.length], label: spec.labels[i], pct: Math.round(frac * 100) };
});
return (
<div className="flex items-center gap-4">
<svg viewBox="0 0 220 220" className="h-44 w-44 shrink-0" role="img" aria-label={spec.title || "饼图"}>
{arcs.map((a, i) => (
<path key={i} d={a.d} fill={a.color} stroke="#0b1220" strokeWidth={1}>
<title>{`${a.label}: ${data[i]} (${a.pct}%)`}</title>
</path>
))}
</svg>
<div className="flex flex-col gap-1">
{arcs.map((a, i) => (
<span key={i} className="flex items-center gap-1.5 text-[10px] text-slate-400">
<span className="inline-block h-2 w-2 rounded-sm" style={{ background: a.color }} />
{a.label} · {a.pct}%
</span>
))}
</div>
</div>
);
}
function fmt(v: number): string {
if (Math.abs(v) >= 1000) return (v / 1000).toFixed(1) + "k";
return String(Math.round(v * 100) / 100);
}
@@ -37,6 +37,7 @@ function StatusDot({ status }: { status: NodeTrace["status"] }) {
if (status === "running") return <Loader2 className="h-4 w-4 animate-spin text-accent-400" strokeWidth={2.4} />; if (status === "running") return <Loader2 className="h-4 w-4 animate-spin text-accent-400" strokeWidth={2.4} />;
if (status === "done") return <CheckCircle2 className="h-4 w-4 text-success" strokeWidth={2.2} />; if (status === "done") return <CheckCircle2 className="h-4 w-4 text-success" strokeWidth={2.2} />;
if (status === "error") return <XCircle className="h-4 w-4 text-danger" strokeWidth={2.2} />; if (status === "error") return <XCircle className="h-4 w-4 text-danger" strokeWidth={2.2} />;
if (status === "waiting") return <Loader2 className="h-4 w-4 animate-spin text-amber-400" strokeWidth={2.4} />; // 待审批
return <Circle className="h-4 w-4 text-slate-600" strokeWidth={2} />; return <Circle className="h-4 w-4 text-slate-600" strokeWidth={2} />;
} }
+21
View File
@@ -109,6 +109,27 @@ export async function submitTask(dsl: TaskDsl, id: Identity): Promise<string> {
return data.task_id; return data.task_id;
} }
// taskStatus: GET /api/v1/tasks/:id —— 任务生命周期状态(submitted/running/waiting/done/failed/timeout/rejected)。
// 轮询它来可靠捕获 waiting(审批中断)—— 不依赖易抢跑的实时 exec 事件。
export async function taskStatus(taskId: string): Promise<{ status: string; detail: string }> {
const res = guard401(await fetch(`${GATEWAY}/api/v1/tasks/${taskId}`, { headers: bearer() }));
if (!res.ok) throw new Error(`status failed: ${res.status}`);
const d = (await res.json()) as { status?: string; detail?: string };
return { status: d.status ?? "", detail: d.detail ?? "" };
}
// approveTask: POST /api/v1/tasks/:id/approve —— HITL 人工审批决定(批准放行 / 拒绝中止)。
export async function approveTask(taskId: string, approved: boolean, opts?: { node?: string; note?: string }): Promise<void> {
const res = guard401(
await fetch(`${GATEWAY}/api/v1/tasks/${taskId}/approve`, {
method: "POST",
headers: { "Content-Type": "application/json", ...bearer() },
body: JSON.stringify({ approved, node: opts?.node ?? "", note: opts?.note ?? "" }),
}),
);
if (!res.ok) throw new Error(`approve failed: ${res.status} ${await res.text()}`);
}
// streamTokens: 订阅 SSE /api/v1/tasks/:id/stream,逐 token 回调,done 收尾。 // streamTokens: 订阅 SSE /api/v1/tasks/:id/stream,逐 token 回调,done 收尾。
// 返回关闭函数。注意 EventSource 无法带请求头,但流按 task_id 寻址,无需身份头。 // 返回关闭函数。注意 EventSource 无法带请求头,但流按 task_id 寻址,无需身份头。
export function streamTokens( export function streamTokens(
@@ -0,0 +1,47 @@
import { describe, it, expect } from "vitest";
import { extractChartBlocks, isChartSpec, hasChart } from "./chartspec";
const spec = { type: "bar", title: "销量", labels: ["Q1", "Q2"], series: [{ name: "销量", data: [120, 180] }] };
describe("isChartSpec", () => {
it("合法 spec 通过", () => expect(isChartSpec(spec)).toBe(true));
it.each([
{},
{ type: "x", labels: ["a"], series: [{ data: [1] }] },
{ type: "bar", labels: [], series: [{ data: [] }] },
{ type: "bar", labels: ["a"], series: [] },
{ type: "bar", labels: ["a"], series: [{ name: "x" }] }, // 无 data
])("非法 spec 拒绝 %#", (bad) => expect(isChartSpec(bad)).toBe(false));
});
describe("extractChartBlocks", () => {
it("纯文本 → 单个 text 段", () => {
const segs = extractChartBlocks("你好世界");
expect(segs).toEqual([{ kind: "text", text: "你好世界" }]);
});
it("文本 + chart 块 + 文本 → 三段且顺序正确", () => {
const out = "看图:\n```chart\n" + JSON.stringify(spec) + "\n```\n以上。";
const segs = extractChartBlocks(out);
expect(segs.map((s) => s.kind)).toEqual(["text", "chart", "text"]);
expect(segs[1].kind === "chart" && segs[1].spec.title).toBe("销量");
});
it("非法 JSON 的 chart 块 → 回退为文本,不丢内容", () => {
const out = "```chart\n{坏的}\n```";
const segs = extractChartBlocks(out);
expect(segs).toHaveLength(1);
expect(segs[0].kind).toBe("text");
});
it("两个 chart 块都解析", () => {
const blk = "```chart\n" + JSON.stringify(spec) + "\n```";
const segs = extractChartBlocks(blk + "\n中间\n" + blk);
expect(segs.filter((s) => s.kind === "chart")).toHaveLength(2);
});
});
describe("hasChart", () => {
it("含 chart 围栏 → true", () => expect(hasChart("a\n```chart\n{}\n```")).toBe(true));
it("不含 → false", () => expect(hasChart("```json\n{}\n```")).toBe(false));
});
@@ -0,0 +1,60 @@
// 图表 spec 与「从模型输出里抽取 ```chart 代码块」的解析。
// 约定:chart 工具返回图表 JSONagent 在答复里用 ```chart 围栏原样包裹;前端据此渲染 SVG。
export interface ChartSeries {
name?: string;
data: number[];
}
export interface ChartSpec {
type: "bar" | "line" | "pie";
title?: string;
labels: string[];
series: ChartSeries[];
}
// Segment:把一段输出拆成「普通文本」与「图表」交错的有序片段,供输出区分别渲染。
export type Segment = { kind: "text"; text: string } | { kind: "chart"; spec: ChartSpec };
const chartFence = /```chart\s*\n([\s\S]*?)```/g;
// isChartSpec 校验解析出的对象是否是合法图表 spec(防脏数据炸渲染)。
export function isChartSpec(v: unknown): v is ChartSpec {
if (!v || typeof v !== "object") return false;
const o = v as Record<string, unknown>;
if (o.type !== "bar" && o.type !== "line" && o.type !== "pie") return false;
if (!Array.isArray(o.labels) || o.labels.length === 0) return false;
if (!Array.isArray(o.series) || o.series.length === 0) return false;
return (o.series as unknown[]).every(
(s) => s && typeof s === "object" && Array.isArray((s as ChartSeries).data),
);
}
// extractChartBlocks 把输出文本拆为 text / chart 片段(保持原顺序)。
// ```chart 块解析失败或非法 → 回退为普通文本,绝不丢内容。
export function extractChartBlocks(input: string): Segment[] {
if (!input) return [];
const segs: Segment[] = [];
let last = 0;
for (let m = chartFence.exec(input); m; m = chartFence.exec(input)) {
if (m.index > last) segs.push({ kind: "text", text: input.slice(last, m.index) });
let parsed: unknown;
try {
parsed = JSON.parse(m[1].trim());
} catch {
parsed = null;
}
if (isChartSpec(parsed)) segs.push({ kind: "chart", spec: parsed });
else segs.push({ kind: "text", text: m[0] }); // 解析失败:原样当文本,不丢
last = m.index + m[0].length;
}
chartFence.lastIndex = 0;
if (last < input.length) segs.push({ kind: "text", text: input.slice(last) });
return segs;
}
// hasChart 快速判断输出里是否含图表块(决定是否走分段渲染)。
export function hasChart(input: string): boolean {
const has = /```chart\s*\n/.test(input);
return has;
}
+33 -1
View File
@@ -1,6 +1,6 @@
import { describe, it, expect } from "vitest"; import { describe, it, expect } from "vitest";
import type { ExecEvent } from "./api"; import type { ExecEvent } from "./api";
import { deriveNodes } from "./run"; import { deriveNodes, pendingApproval } from "./run";
let seq = 0; let seq = 0;
function ev(node: string, phase: string, extra: Partial<ExecEvent> = {}): ExecEvent { function ev(node: string, phase: string, extra: Partial<ExecEvent> = {}): ExecEvent {
@@ -48,4 +48,36 @@ describe("deriveNodesExecEvent 流 → 节点轨迹)", () => {
const [n] = deriveNodes([ev("a", "start", { label: "旧" }), ev("a", "end", { label: "新" })]); const [n] = deriveNodes([ev("a", "start", { label: "旧" }), ev("a", "end", { label: "新" })]);
expect(n.label).toBe("新"); expect(n.label).toBe("新");
}); });
it("await 事件 → 节点状态 waiting", () => {
const [n] = deriveNodes([ev("ap", "await", { kind: "approval", label: "审批" })]);
expect(n.status).toBe("waiting");
});
});
describe("pendingApprovalHITL 待审批中断)", () => {
it("await 未收口 → 返回待审批", () => {
const p = pendingApproval([ev("ap", "await", { kind: "approval", label: "高危审批", detail: "摘要" })]);
expect(p).toEqual({ node: "ap", title: "高危审批", summary: "摘要" });
});
it("await 后 end(同节点)→ 已决,返回 null", () => {
const p = pendingApproval([
ev("ap", "await", { kind: "approval" }),
ev("ap", "end", { kind: "approval" }),
]);
expect(p).toBeNull();
});
it("await 后 error(超时/拒绝)→ 已决,返回 null", () => {
const p = pendingApproval([
ev("ap", "await", { kind: "approval" }),
ev("ap", "error", { kind: "approval" }),
]);
expect(p).toBeNull();
});
it("无审批事件 → null", () => {
expect(pendingApproval([ev("a", "start"), ev("a", "end")])).toBeNull();
});
}); });
+28 -1
View File
@@ -15,13 +15,15 @@ export interface RunState {
events: RunEvent[]; events: RunEvent[];
exec: ExecEvent[]; // 后端回流的节点级执行轨迹(运行·观测) exec: ExecEvent[]; // 后端回流的节点级执行轨迹(运行·观测)
error?: string; error?: string;
lifecycle?: string; // 轮询到的后端任务状态(waiting 时弹审批;done/rejected 等终态)
detail?: string; // 状态附带说明(如审批标题)
} }
export const emptyRun: RunState = { phase: "idle", output: "", events: [], exec: [] }; export const emptyRun: RunState = { phase: "idle", output: "", events: [], exec: [] };
// ---- 执行轨迹派生:把扁平 ExecEvent 流归并为按节点聚合的轨迹 ---- // ---- 执行轨迹派生:把扁平 ExecEvent 流归并为按节点聚合的轨迹 ----
export type NodeStatus = "running" | "done" | "error" | "info"; export type NodeStatus = "running" | "done" | "error" | "info" | "waiting";
export interface NodeTrace { export interface NodeTrace {
node: string; node: string;
@@ -34,6 +36,28 @@ export interface NodeTrace {
order: number; order: number;
} }
// PendingApproval 是一个待人工审批的中断(HITL):审批节点已发 await 事件、尚未被 end/error 收口。
export interface PendingApproval {
node: string;
title: string;
summary: string;
}
// pendingApproval 从执行事件流里找出当前待审批的中断:取最后一个 kind=approval & phase=await
// 且其后该节点没有 end/error(未被批准/拒绝收口)的事件。无则 null。
export function pendingApproval(events: ExecEvent[]): PendingApproval | null {
let pending: PendingApproval | null = null;
for (const e of events) {
if (e.kind !== "approval") continue;
if (e.phase === "await") {
pending = { node: e.node, title: e.label || "人工审批", summary: e.detail || "" };
} else if (e.phase === "end" || e.phase === "error") {
if (pending && pending.node === e.node) pending = null; // 已决,收口
}
}
return pending;
}
// deriveNodes 把事件流按 node 归并:start→runningend→done(带耗时)error→errorinfo→点事件/附注。 // deriveNodes 把事件流按 node 归并:start→runningend→done(带耗时)error→errorinfo→点事件/附注。
export function deriveNodes(events: ExecEvent[]): NodeTrace[] { export function deriveNodes(events: ExecEvent[]): NodeTrace[] {
const map = new Map<string, NodeTrace>(); const map = new Map<string, NodeTrace>();
@@ -50,6 +74,9 @@ export function deriveNodes(events: ExecEvent[]): NodeTrace[] {
case "start": case "start":
if (n.status !== "done" && n.status !== "error") n.status = "running"; if (n.status !== "done" && n.status !== "error") n.status = "running";
break; break;
case "await": // HITL 审批节点暂停,等人工决定
n.status = "waiting";
break;
case "end": case "end":
n.status = "done"; n.status = "done";
n.ms = e.ms; n.ms = e.ms;
@@ -1,7 +1,10 @@
import { useState } from "react"; import { useState } from "react";
import { ChevronDown, ChevronUp, Wrench } from "lucide-react"; import { ChevronDown, ChevronUp, Wrench, ShieldCheck, Check, X } from "lucide-react";
import { deriveNodes, type RunState } from "../lib/run"; import { deriveNodes, pendingApproval, type RunState } from "../lib/run";
import { ExecTrace } from "../components/ExecTrace"; import { ExecTrace } from "../components/ExecTrace";
import { ChartView } from "../components/ChartView";
import { extractChartBlocks, hasChart } from "../lib/chartspec";
import { approveTask } from "../lib/api";
import { Tabs, Badge, cn, type TabDef } from "../ui"; import { Tabs, Badge, cn, type TabDef } from "../ui";
type Tab = "output" | "trace" | "tools" | "cite" | "eval"; type Tab = "output" | "trace" | "tools" | "cite" | "eval";
@@ -26,8 +29,15 @@ export function BottomDrawer({ run }: { run: RunState }) {
const statusText = const statusText =
run.phase === "streaming" ? "流式中…" : run.phase === "done" ? "完成 ✓" : run.phase === "error" ? `${run.error ?? "出错"}` : run.phase === "submitting" ? "提交中…" : "就绪"; run.phase === "streaming" ? "流式中…" : run.phase === "done" ? "完成 ✓" : run.phase === "error" ? `${run.error ?? "出错"}` : run.phase === "submitting" ? "提交中…" : "就绪";
// 审批条触发以「后端状态 waiting」为准(可靠,落 PG);exec 的 await 事件仅用于丰富摘要(可能抢跑丢失)。
const approval =
run.taskId && run.lifecycle === "waiting"
? pendingApproval(run.exec) ?? { node: "", title: run.detail || "人工审批", summary: "" }
: null;
return ( return (
<div className="shrink-0 border-t border-line bg-ink-900"> <div className="shrink-0 border-t border-line bg-ink-900">
{approval && run.taskId && <ApprovalBar taskId={run.taskId} node={approval.node} title={approval.title} summary={approval.summary} />}
<div className="flex items-center border-b border-line px-2"> <div className="flex items-center border-b border-line px-2">
<Tabs <Tabs
tabs={tabs} tabs={tabs}
@@ -45,11 +55,7 @@ export function BottomDrawer({ run }: { run: RunState }) {
</div> </div>
{open && ( {open && (
<div className="h-44 overflow-auto p-3 text-xs"> <div className="h-44 overflow-auto p-3 text-xs">
{tab === "output" && ( {tab === "output" && <OutputView output={run.output} />}
<pre className="whitespace-pre-wrap font-mono leading-relaxed text-emerald-300">
{run.output || "在编排页搭图 → 运行,模型注入画像与历史后流式作答,token 在此呈现。"}
</pre>
)}
{tab === "trace" && <ExecTrace events={run.exec} phase={run.phase} />} {tab === "trace" && <ExecTrace events={run.exec} phase={run.phase} />}
{tab === "tools" && <ToolCalls run={run} />} {tab === "tools" && <ToolCalls run={run} />}
{tab === "cite" && <p className="text-slate-600">RAG + + </p>} {tab === "cite" && <p className="text-slate-600">RAG + + </p>}
@@ -60,6 +66,89 @@ export function BottomDrawer({ run }: { run: RunState }) {
); );
} }
// ApprovalBarHITL 人工审批中断条。任务停在审批节点时常驻顶部,展示待审摘要 + 批准/拒绝。
// 决定经 approveTask 发回;dispatcher 续跑后 SSE 会推来 end/error 事件,本条随 pendingApproval 归 null 自动消失。
function ApprovalBar({ taskId, node, title, summary }: { taskId: string; node: string; title: string; summary: string }) {
const [note, setNote] = useState("");
const [busy, setBusy] = useState<"approve" | "reject" | null>(null);
const [err, setErr] = useState("");
const decide = async (approved: boolean) => {
setBusy(approved ? "approve" : "reject");
setErr("");
try {
await approveTask(taskId, approved, { node, note });
} catch (e) {
setErr((e as Error).message);
setBusy(null);
}
};
return (
<div className="flex flex-col gap-2 border-b border-amber-500/30 bg-amber-500/10 px-3 py-2">
<div className="flex items-center gap-2 text-[12px] text-amber-200">
<ShieldCheck className="h-4 w-4 text-amber-400" strokeWidth={2.2} />
<span className="font-semibold"></span>
<span className="text-amber-300/80">{title}</span>
<Badge tone="warn"></Badge>
</div>
{summary && <pre className="max-h-20 overflow-auto whitespace-pre-wrap font-mono text-[11px] leading-relaxed text-amber-100/80">{summary}</pre>}
<div className="flex items-center gap-2">
<input
value={note}
onChange={(e) => setNote(e.target.value)}
placeholder="备注(可选,拒绝原因等)"
className="min-w-0 flex-1 rounded border border-line bg-ink-950/60 px-2 py-1 text-[11px] text-slate-200 placeholder:text-slate-600 focus:border-amber-500/50 focus:outline-none"
/>
<button
onClick={() => decide(true)}
disabled={busy !== null}
className="flex items-center gap-1 rounded bg-success/20 px-2.5 py-1 text-[11px] font-medium text-success hover:bg-success/30 disabled:opacity-50"
>
<Check className="h-3.5 w-3.5" /> {busy === "approve" ? "提交中…" : "批准"}
</button>
<button
onClick={() => decide(false)}
disabled={busy !== null}
className="flex items-center gap-1 rounded bg-danger/20 px-2.5 py-1 text-[11px] font-medium text-danger hover:bg-danger/30 disabled:opacity-50"
>
<X className="h-3.5 w-3.5" /> {busy === "reject" ? "提交中…" : "拒绝"}
</button>
</div>
{err && <p className="text-[11px] text-danger">{err}</p>}
</div>
);
}
// OutputView:渲染模型输出。含 ```chart 块时分段渲染(文本 + SVG 图表),否则纯文本。
function OutputView({ output }: { output: string }) {
if (!output) {
return (
<pre className="whitespace-pre-wrap font-mono leading-relaxed text-emerald-300">
token
</pre>
);
}
if (!hasChart(output)) {
return <pre className="whitespace-pre-wrap font-mono leading-relaxed text-emerald-300">{output}</pre>;
}
return (
<div>
{extractChartBlocks(output).map((seg, i) =>
seg.kind === "chart" ? (
<ChartView key={i} spec={seg.spec} />
) : (
seg.text.trim() && (
<pre key={i} className="whitespace-pre-wrap font-mono leading-relaxed text-emerald-300">
{seg.text}
</pre>
)
),
)}
</div>
);
}
// ToolCalls:从执行事件里筛出工具调用节点,逐条展示入参 → 产出 + 耗时/状态。 // ToolCalls:从执行事件里筛出工具调用节点,逐条展示入参 → 产出 + 耗时/状态。
function ToolCalls({ run }: { run: RunState }) { function ToolCalls({ run }: { run: RunState }) {
const tools = deriveNodes(run.exec).filter((n) => n.kind === "tool"); const tools = deriveNodes(run.exec).filter((n) => n.kind === "tool");
@@ -102,6 +102,18 @@ export const NODE_KINDS: Record<string, NodeKind> = {
fields: [{ key: "condition", label: "条件", type: "text", placeholder: "score > 0.8" }], fields: [{ key: "condition", label: "条件", type: "text", placeholder: "score > 0.8" }],
defaults: { condition: "" }, defaults: { condition: "" },
}, },
approval: {
kind: "approval",
label: "人工审批",
accent: "border-l-orange-500",
badge: "bg-orange-100 text-orange-700",
desc: "HITL:暂停等人工批准",
fields: [
{ key: "title", label: "审批标题", type: "text", placeholder: "如:高危操作审批" },
{ key: "prompt", label: "审批说明", type: "textarea", placeholder: "向审批人说明待执行的操作…" },
],
defaults: { title: "人工审批", prompt: "请审批是否继续执行后续步骤" },
},
map: { map: {
kind: "map", kind: "map",
label: "并行 / Map", label: "并行 / Map",
@@ -150,6 +162,7 @@ export const NODE_ORDER = [
"tool", "tool",
"memory", "memory",
"branch", "branch",
"approval",
"map", "map",
"aggregate", "aggregate",
"render", "render",
+3 -2
View File
@@ -45,8 +45,9 @@ func main() {
} }
go sub.FetchModelConfigWithRetry(context.Background(), pool.SetConfig) go sub.FetchModelConfigWithRetry(context.Background(), pool.SetConfig)
// sub 同时作为 Token 回流(TokenSink)、MCP 工具调用(ToolCaller)、执行事件(ExecSink)与任务状态回写(StatusSink)出口。 // sub 同时作为 Token 回流(TokenSink)、MCP 工具调用(ToolCaller)、执行事件(ExecSink)
orch, err := eino.NewOrchestrator(pool, breaker, eval, sub, sub, sub, sub) // 任务状态回写(StatusSink)与 HITL 审批等待(ApprovalWaiter)出口。
orch, err := eino.NewOrchestrator(pool, breaker, eval, sub, sub, sub, sub, sub)
if err != nil { if err != nil {
log.Fatalf("[dispatcher] build eino graph: %v", err) log.Fatalf("[dispatcher] build eino graph: %v", err)
} }
+10 -1
View File
@@ -48,9 +48,18 @@ func buildMessages(ctx context.Context, rc *RunCtx) ([]*schema.Message, error) {
sys.WriteString("\n\n以下是前序协作 agent 的产出,请在此基础上继续完成你的部分(不要重头再来):\n") sys.WriteString("\n\n以下是前序协作 agent 的产出,请在此基础上继续完成你的部分(不要重头再来):\n")
sys.WriteString(strings.Join(rc.Upstream, "\n---\n")) sys.WriteString(strings.Join(rc.Upstream, "\n---\n"))
} }
// 防御:剔除空 content 的历史消息。OpenAI 兼容 API 拒绝「content 与 tool_calls 都为空」的
// assistant 消息(400 Invalid assistant message),一条脏历史会毒化整段会话的后续请求。
history := make([]*schema.Message, 0, len(rc.History))
for _, m := range rc.History {
if m == nil || (strings.TrimSpace(m.Content) == "" && len(m.ToolCalls) == 0) {
continue
}
history = append(history, m)
}
return chatTemplate.Format(ctx, map[string]any{ return chatTemplate.Format(ctx, map[string]any{
"system": sys.String(), "system": sys.String(),
"history": rc.History, "history": history,
"query": rc.Query, "query": rc.Query,
}) })
} }
+64 -1
View File
@@ -32,6 +32,8 @@ type board struct {
sections []reportSection // map 并行 fan-out 产出的分项成稿(供 render 多章渲染) sections []reportSection // map 并行 fan-out 产出的分项成稿(供 render 多章渲染)
answer string // 当前成稿(多 agent 协作时 = 最近一个 agent 的产出 = 成品) answer string // 当前成稿(多 agent 协作时 = 最近一个 agent 的产出 = 成品)
agentOut []string // 各上游 agent 的产出(按序),注入下游 agent 上下文以实现接力协作 agentOut []string // 各上游 agent 的产出(按序),注入下游 agent 上下文以实现接力协作
rejected bool // HITL 审批节点拒绝/超时 → 置位,runGraph 中止并返回 errRejected
fatalErr error // agent 节点 LLM 调用失败 → 置位,runGraph 中止并上抛 → 任务判 failed(而非 done-空)
} }
// runGraph 按 DSL 图的真实拓扑与连线执行(替代旧的线性拍平 compileFlow)。 // runGraph 按 DSL 图的真实拓扑与连线执行(替代旧的线性拍平 compileFlow)。
@@ -55,7 +57,7 @@ func (o *Orchestrator) runGraph(ctx context.Context, t *contract.Task, tr *execT
b.profile = o.fetchMemory(ctx, b.uid, b.query) b.profile = o.fetchMemory(ctx, b.uid, b.query)
b.history = o.fetchHistory(ctx, b.sid) b.history = o.fetchHistory(ctx, b.sid)
o.runConversation(ctx, t.ID, b, plan.System, tr, "agent") o.runConversation(ctx, t.ID, b, plan.System, tr, "agent")
return b.answer, nil return b.answer, b.fatalErr // 模型失败 → 上抛判 failed
} }
// 建邻接与入度(只认两端都存在的边)。保留整条边以便 branch 按 true/false 标签选路。 // 建邻接与入度(只认两端都存在的边)。保留整条边以便 branch 按 true/false 标签选路。
@@ -144,6 +146,8 @@ func (o *Orchestrator) runGraph(ctx context.Context, t *contract.Task, tr *execT
o.renderNode(nctx, t.ID, n, b, tr) o.renderNode(nctx, t.ID, n, b, tr)
case "branch": case "branch":
propagate = o.branchNode(n, b, outE[n.ID], nodeByID, tr) propagate = o.branchNode(n, b, outE[n.ID], nodeByID, tr)
case "approval":
propagate = o.approvalNode(nctx, t.ID, n, b, tr, outE[n.ID])
case "map": case "map":
o.mapNode(nctx, t.ID, n, b, tr) o.mapNode(nctx, t.ID, n, b, tr)
case "output": case "output":
@@ -152,15 +156,28 @@ func (o *Orchestrator) runGraph(ctx context.Context, t *contract.Task, tr *execT
tr.info(n.Kind+":"+n.ID, "system", labelOf(n, n.Kind), "未识别节点,跳过") tr.info(n.Kind+":"+n.ID, "system", labelOf(n, n.Kind), "未识别节点,跳过")
} }
nspan.End() nspan.End()
if b.rejected || b.fatalErr != nil {
break // 审批拒绝/超时 或 模型失败:中止后续节点
}
for _, tgt := range propagate { for _, tgt := range propagate {
active[tgt] = true active[tgt] = true
} }
} }
if b.rejected {
return b.answer, errRejected // 合法终态,Handle 据此判 rejected 并优雅收尾
}
if b.fatalErr != nil {
return b.answer, b.fatalErr // 上抛 → Handle 判 failed(带原因)
}
// 图里无 agent 节点(纯工具/检索图)也要出一段模型答复,否则没有输出。 // 图里无 agent 节点(纯工具/检索图)也要出一段模型答复,否则没有输出。
if b.answer == "" { if b.answer == "" {
o.runConversation(ctx, t.ID, b, plan.System, tr, "agent") o.runConversation(ctx, t.ID, b, plan.System, tr, "agent")
} }
if b.fatalErr != nil {
return b.answer, b.fatalErr
}
return b.answer, nil return b.answer, nil
} }
@@ -273,6 +290,11 @@ func (o *Orchestrator) runAgent(ctx context.Context, taskID string, b *board, sy
} }
if err != nil { if err != nil {
tr.emit(node, "model", "error", "模型流式推理", err.Error(), time.Since(t0).Milliseconds()) tr.emit(node, "model", "error", "模型流式推理", err.Error(), time.Since(t0).Milliseconds())
// 未产出任何 token 即失败 → 标记致命错,让任务判 failed(暴露原因,便于监控告警),
// 而非静默 done-空。已流出部分 token 的中断也算失败(结果不完整)。
if b.fatalErr == nil {
b.fatalErr = fmt.Errorf("agent 模型推理失败: %w", err)
}
return return
} }
if redacted > 0 { if redacted > 0 {
@@ -369,6 +391,47 @@ func (o *Orchestrator) branchNode(n dsl.Node, b *board, outs []dsl.Edge, byID ma
return chosen return chosen
} }
// approvalNode 是 HITL 人工审批中断:执行到此暂停,把待审摘要推给 UI(exec 事件 kind=approval/phase=await
// + 状态 waiting),阻塞等人工批准/拒绝(带超时,安全默认拒绝)。
// 批准 → 状态回 running 并放行下游;拒绝/超时 → 置 b.rejected 中止全图。返回应激活的下游(拒绝=空)。
func (o *Orchestrator) approvalNode(ctx context.Context, taskID string, n dsl.Node, b *board, tr *execTracer, outs []dsl.Edge) []string {
title := firstNonEmpty(cstr(n.Config, "title"), labelOf(n, "人工审批"))
prompt := firstNonEmpty(cstr(n.Config, "prompt"), "请审批是否继续执行后续步骤")
summary := prompt
if b.answer != "" { // 带上当前产出预览,便于审批人判断
summary = prompt + "\n—— 当前产出预览 ——\n" + truncate(b.answer, 400)
}
// 未接审批通道(单测/降级)→ 自动放行,避免无人应答卡死。
if o.approval == nil {
tr.info("approval:"+n.ID, "approval", title, "未接审批通道,自动放行")
return targetsOf(outs)
}
// 暂停:发待审事件(UI 据 kind=approval & phase=await 弹批准/拒绝)+ 置任务 waiting。
tr.emit("approval:"+n.ID, "approval", "await", title, summary, 0)
o.setStatus(taskID, contract.TaskWaiting, title)
dec, err := o.approval.WaitApproval(ctx, taskID, approvalTimeout)
switch {
case err != nil: // 超时 / ctx 取消 → 安全默认拒绝
b.rejected = true
b.answer = "❌ 审批超时未决,已自动拒绝:" + title
tr.emit("approval:"+n.ID, "approval", "error", title, "审批超时,自动拒绝", 0)
return nil
case !dec.Approved:
b.rejected = true
note := firstNonEmpty(dec.Note, "审批人拒绝")
b.answer = "❌ 已被拒绝:" + note
tr.emit("approval:"+n.ID, "approval", "end", title, "拒绝:"+note, 0)
return nil
default: // 批准 → 恢复执行,放行下游
o.setStatus(taskID, contract.TaskRunning, "审批通过,继续执行")
tr.emit("approval:"+n.ID, "approval", "end", title, "批准:"+firstNonEmpty(dec.Note, "放行"), 0)
return targetsOf(outs)
}
}
// targetsOf 取一组边的目标节点 ID(保持顺序)。 // targetsOf 取一组边的目标节点 ID(保持顺序)。
func targetsOf(edges []dsl.Edge) []string { func targetsOf(edges []dsl.Edge) []string {
out := make([]string, 0, len(edges)) out := make([]string, 0, len(edges))
@@ -8,6 +8,7 @@ import (
"fmt" "fmt"
"log" "log"
"log/slog" "log/slog"
"strings"
"sync" "sync"
"time" "time"
@@ -40,6 +41,15 @@ type StatusSink interface {
PublishTaskStatus(taskID, status, detail string) error PublishTaskStatus(taskID, status, detail string) error
} }
// ApprovalWaiter 阻塞等待审批节点的人工决定(由 NATS bus 实现;可为 nil → 审批节点自动放行)。
type ApprovalWaiter interface {
WaitApproval(ctx context.Context, taskID string, timeout time.Duration) (*contract.ApprovalDecision, error)
}
// errRejected 是审批节点拒绝(或超时)时图执行返回的哨兵错误:它是合法终态而非故障,
// Handle 据此判 rejected 并优雅收尾(不计熔断失败)。
var errRejected = errors.New("approval rejected")
// LLM 是编排所需的语言模型能力(生产由 *llm.Pool 实现)。抽成接口便于测试注入假模型。 // LLM 是编排所需的语言模型能力(生产由 *llm.Pool 实现)。抽成接口便于测试注入假模型。
type LLM interface { type LLM interface {
Ready() bool Ready() bool
@@ -56,17 +66,22 @@ type LLM interface {
const toolCallTimeout = 3 * time.Second const toolCallTimeout = 3 * time.Second
// taskExecTimeout 是单个任务整体执行上限;超时即判 timeout(状态机),避免无限期"运行中"。 // taskExecTimeout 是单个任务整体执行上限;超时即判 timeout(状态机),避免无限期"运行中"。
const taskExecTimeout = 3 * time.Minute // 含 HITL 审批等待预算(approvalTimeout+ 常规图执行;须 < bus 消费者 AckWait(15min) 以免重投。
const taskExecTimeout = 10 * time.Minute
// approvalTimeout 是单个审批节点等待人工决定的上限;超时安全默认拒绝(fail-safe)。
const approvalTimeout = 5 * time.Minute
// Orchestrator 把每个 DSL 任务动态编译为 Eino 图并执行(记忆召回 → 工具节点 → 注入 → 流式)。 // Orchestrator 把每个 DSL 任务动态编译为 Eino 图并执行(记忆召回 → 工具节点 → 注入 → 流式)。
type Orchestrator struct { type Orchestrator struct {
pool LLM pool LLM
breaker *harness.CircuitBreaker breaker *harness.CircuitBreaker
eval *harness.Evaluator eval *harness.Evaluator
sink TokenSink sink TokenSink
tools ToolCaller tools ToolCaller
exec ExecSink exec ExecSink
status StatusSink // 任务生命周期状态回写(可为 nil) status StatusSink // 任务生命周期状态回写(可为 nil)
approval ApprovalWaiter // HITL 审批等待(可为 nil → 审批节点自动放行)
turnMu sync.Mutex // 保护 turns(攒批计数,多任务 goroutine 共享) turnMu sync.Mutex // 保护 turns(攒批计数,多任务 goroutine 共享)
turns map[string]int // sessionID → 累计轮次,用于每 N 轮触发 consolidate turns map[string]int // sessionID → 累计轮次,用于每 N 轮触发 consolidate
@@ -74,9 +89,9 @@ type Orchestrator struct {
// NewOrchestrator 持有依赖;图按任务的 DSL 在 Handle 内动态编译。 // NewOrchestrator 持有依赖;图按任务的 DSL 在 Handle 内动态编译。
// exec 为执行可视化事件出口(可为 nil,则不发轨迹事件);eval 为自动化评测(可为 nil); // exec 为执行可视化事件出口(可为 nil,则不发轨迹事件);eval 为自动化评测(可为 nil);
// status 为任务生命周期状态回写出口(可为 nil)。 // status 为任务生命周期状态回写出口(可为 nil);approval 为 HITL 审批等待(可为 nil)。
func NewOrchestrator(pool LLM, breaker *harness.CircuitBreaker, eval *harness.Evaluator, sink TokenSink, tools ToolCaller, exec ExecSink, status StatusSink) (*Orchestrator, error) { func NewOrchestrator(pool LLM, breaker *harness.CircuitBreaker, eval *harness.Evaluator, sink TokenSink, tools ToolCaller, exec ExecSink, status StatusSink, approval ApprovalWaiter) (*Orchestrator, error) {
return &Orchestrator{pool: pool, breaker: breaker, eval: eval, sink: sink, tools: tools, exec: exec, status: status}, nil return &Orchestrator{pool: pool, breaker: breaker, eval: eval, sink: sink, tools: tools, exec: exec, status: status, approval: approval}, nil
} }
// setStatus 回写一次任务状态流转(status 为 nil 时静默跳过)。 // setStatus 回写一次任务状态流转(status 为 nil 时静默跳过)。
@@ -144,6 +159,17 @@ func (o *Orchestrator) Handle(ctx context.Context, t *contract.Task) error {
// 按 DSL 图执行:compose.GraphEINO_COMPOSE=1)或自研 graph.go(默认);agent 节点流式回流 token。 // 按 DSL 图执行:compose.GraphEINO_COMPOSE=1)或自研 graph.go(默认);agent 节点流式回流 token。
answer, err := o.executeGraph(tctx, t, tr) answer, err := o.executeGraph(tctx, t, tr)
if errors.Is(err, errRejected) {
// HITL 拒绝:合法终态,非故障。收尾流 + 置 rejected,不计熔断、不重投。
slog.InfoContext(ctx, "task rejected by approval", "task_id", t.ID)
if answer != "" {
_ = o.sink.PublishToken(t.ID, []byte(answer))
}
_ = o.sink.CompleteStream(t.ID)
o.breaker.Report(true) // 拒绝是人为决策,不算后端失败
o.setStatus(t.ID, contract.TaskRejected, truncate(answer, 120))
return nil
}
if err != nil { if err != nil {
span.RecordError(err) span.RecordError(err)
span.SetStatus(codes.Error, err.Error()) span.SetStatus(codes.Error, err.Error())
@@ -245,6 +271,11 @@ const consolidateEveryTurns = 3
// memorize 写回阶段(异步、离热路径):落短期历史;每 N 轮做一次记忆对账(consolidate)。 // memorize 写回阶段(异步、离热路径):落短期历史;每 N 轮做一次记忆对账(consolidate)。
func (o *Orchestrator) memorize(t *contract.Task, answer string) { func (o *Orchestrator) memorize(t *contract.Task, answer string) {
// 空答复(LLM 失败/降级/拒绝)不落历史:空 assistant 消息会被 LLM API 拒绝(400),
// 一旦写入会毒化该会话后续所有请求。失败的一轮干脆不留痕。
if strings.TrimSpace(answer) == "" {
return
}
uid, _ := t.Meta[contract.MetaUserID].(string) uid, _ := t.Meta[contract.MetaUserID].(string)
sid, _ := t.Meta[contract.MetaSessionID].(string) sid, _ := t.Meta[contract.MetaSessionID].(string)
if sid != "" && o.tools != nil { if sid != "" && o.tools != nil {
@@ -81,6 +81,11 @@ func (s *Subscriber) PublishTaskStatus(taskID, status, detail string) error {
}) })
} }
// WaitApproval 让 Subscriber 满足 eino.ApprovalWaiter,阻塞等待审批节点的人工决定。
func (s *Subscriber) WaitApproval(ctx context.Context, taskID string, timeout time.Duration) (*contract.ApprovalDecision, error) {
return s.inner.WaitApproval(ctx, taskID, timeout)
}
// RequestModelConfig 向控制面(Gateway)取当前激活的对话模型配置。 // RequestModelConfig 向控制面(Gateway)取当前激活的对话模型配置。
func (s *Subscriber) RequestModelConfig(ctx context.Context) (*contract.ModelConfig, error) { func (s *Subscriber) RequestModelConfig(ctx context.Context) (*contract.ModelConfig, error) {
return s.inner.RequestConfig(ctx, contract.ConfigKindChat) return s.inner.RequestConfig(ctx, contract.ConfigKindChat)
@@ -89,6 +89,34 @@ func (h *Handler) TaskStatus(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"task_id": id, "status": status, "detail": detail}) c.JSON(http.StatusOK, gin.H{"task_id": id, "status": status, "detail": detail})
} }
// ApproveTask: POST /api/v1/tasks/:id/approve {approved, node?, note?} —— 人工审批决定(HITL)。
// 把决定经 NATS 发给 dispatcher,解除审批节点的阻塞(批准放行 / 拒绝中止)。
func (h *Handler) ApproveTask(c *gin.Context) {
id := c.Param("id")
var body struct {
Approved bool `json:"approved"`
Node string `json:"node"`
Note string `json:"note"`
}
if err := c.ShouldBindJSON(&body); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
// 仅在任务确为等待审批时受理(幂等:重复/迟到的决定不报错,dispatcher 侧已只取首条)。
if status, _ := h.db.GetTaskStatus(c.Request.Context(), id); status != contract.TaskWaiting {
c.JSON(http.StatusConflict, gin.H{"error": "任务当前非待审批状态", "status": status})
return
}
if err := h.bus.PublishApproval(&contract.ApprovalDecision{
TaskID: id, Node: body.Node, Approved: body.Approved, Note: body.Note,
By: userID(c), TS: time.Now().UnixMilli(),
}); err != nil {
c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"task_id": id, "approved": body.Approved})
}
// StreamTask: 以 SSE 把 Token Stream 推给客户端。 // StreamTask: 以 SSE 把 Token Stream 推给客户端。
// 优先从 Redis Stream 读(可回放 + 断点续传,根治连晚/重连丢 token);Redis 降级时回退 live NATS。 // 优先从 Redis Stream 读(可回放 + 断点续传,根治连晚/重连丢 token);Redis 降级时回退 live NATS。
func (h *Handler) StreamTask(c *gin.Context) { func (h *Handler) StreamTask(c *gin.Context) {
@@ -64,6 +64,11 @@ func (b *Bus) SubscribeTaskStatus(onEvent func(*contract.TaskStatusEvent)) (func
return b.inner.SubscribeTaskStatus(onEvent) return b.inner.SubscribeTaskStatus(onEvent)
} }
// PublishApproval 把一次人工审批决定发给 dispatcher(解除审批节点阻塞)。
func (b *Bus) PublishApproval(dec *contract.ApprovalDecision) error {
return b.inner.PublishApproval(dec)
}
// ServeConfig 让网关作为配置控制面,响应某 kind 的配置请求。 // ServeConfig 让网关作为配置控制面,响应某 kind 的配置请求。
func (b *Bus) ServeConfig(kind string, provide func() *contract.ModelConfig) (func() error, error) { func (b *Bus) ServeConfig(kind string, provide func() *contract.ModelConfig) (func() error, error) {
return b.inner.ServeConfig(kind, provide) return b.inner.ServeConfig(kind, provide)
+20 -19
View File
@@ -50,25 +50,26 @@ func New(db *store.Postgres, cache *store.Redis, bus *nats.Bus, blobStore *blob.
// —— 受保护:owner 作用域业务,必须携带有效 JWT —— // —— 受保护:owner 作用域业务,必须携带有效 JWT ——
p := api.Group("", middleware.RequireAuth()) p := api.Group("", middleware.RequireAuth())
{ {
p.POST("/tasks", h.SubmitTask) // 解析 DSL 并 Publish 到 NATS(带已验证 uid p.POST("/tasks", h.SubmitTask) // 解析 DSL 并 Publish 到 NATS(带已验证 uid
p.GET("/tasks/:id", h.TaskStatus) // 任务生命周期状态(UI 轮询 submitted/running/done/failed/timeout p.GET("/tasks/:id", h.TaskStatus) // 任务生命周期状态(UI 轮询 submitted/running/done/failed/timeout/waiting/rejected
p.PUT("/memory", h.SetMemory) // 偏好记忆登记(→ mcp-go memory_upsert p.POST("/tasks/:id/approve", h.ApproveTask) // HITL 人工审批决定(批准/拒绝
p.GET("/memory", h.ListMemory) // 列出当前用户偏好记忆面板 p.PUT("/memory", h.SetMemory) // 偏好记忆登记(→ mcp-go memory_upsert
p.DELETE("/memory", h.DeleteMemory) // 软删一条偏好(?key= p.GET("/memory", h.ListMemory) // 列出当前用户偏好(记忆面板
p.GET("/kb/list", h.KbList) // 当前用户的知识库列表(owner 隔离 p.DELETE("/memory", h.DeleteMemory) // 软删一条偏好(?key=
p.POST("/kb/create", h.KbCreate) // 新建知识库 p.GET("/kb/list", h.KbList) // 当前用户的知识库列表(owner 隔离)
p.POST("/kb/ingest", h.KbIngest) // 文本入 p.POST("/kb/create", h.KbCreate) // 新建知识
p.POST("/kb/ingest_file", h.KbIngestFile) // 文入库 p.POST("/kb/ingest", h.KbIngest) // 文入库
p.POST("/kb/search", h.KbSearch) // 检索台 p.POST("/kb/ingest_file", h.KbIngestFile) // 文件入库
p.GET("/kb/vault", h.KbVault) // 文库列表 p.POST("/kb/search", h.KbSearch) // 检索台
p.GET("/kb/doc", h.KbDoc) // 取单篇文档 p.GET("/kb/vault", h.KbVault) // 文库列表
p.GET("/kb/links", h.KbLinks) // 某库双链 p.GET("/kb/doc", h.KbDoc) // 取单篇文档
p.POST("/kb/note", h.KbSaveNote) // 新建/编辑笔记 p.GET("/kb/links", h.KbLinks) // 某库双链
p.GET("/kb/graph", h.KbGraph) // 知识图谱三元组 p.POST("/kb/note", h.KbSaveNote) // 新建/编辑笔记
p.GET("/agents", h.AgentList) // 我的编排列表(owner 隔离) p.GET("/kb/graph", h.KbGraph) // 知识图谱三元组
p.POST("/agents", h.AgentSave) // 保存/更新编排 p.GET("/agents", h.AgentList) // 我的编排列表(owner 隔离)
p.DELETE("/agents", h.AgentDelete) // 删除编排 p.POST("/agents", h.AgentSave) // 保存/更新编排
p.POST("/reports", h.GenerateReport) // 报告生成 p.DELETE("/agents", h.AgentDelete) // 删除编排
p.POST("/reports", h.GenerateReport) // 报告生成
p.GET("/billing", h.Billing) p.GET("/billing", h.Billing)
} }
+24
View File
@@ -5,6 +5,9 @@ import (
"context" "context"
"errors" "errors"
"log" "log"
"os"
"strconv"
"time"
"gorm.io/driver/postgres" "gorm.io/driver/postgres"
"gorm.io/gorm" "gorm.io/gorm"
@@ -13,6 +16,26 @@ import (
"github.com/sundynix/sundynix-shared/contract" "github.com/sundynix/sundynix-shared/contract"
) )
// envInt 读正整数环境变量,缺省回退 def。
func envInt(key string, def int) int {
if v := os.Getenv(key); v != "" {
if n, err := strconv.Atoi(v); err == nil && n > 0 {
return n
}
}
return def
}
// tunePool 给连接池设上限:高并发下不至于无限开连接打爆 PG(max_connections 默认 100)。
// 各服务默认 25,可经 DB_MAX_OPEN_CONNS / DB_MAX_IDLE_CONNS 调整。
func tunePool(db *gorm.DB) {
if sqlDB, err := db.DB(); err == nil {
sqlDB.SetMaxOpenConns(envInt("DB_MAX_OPEN_CONNS", 25))
sqlDB.SetMaxIdleConns(envInt("DB_MAX_IDLE_CONNS", 5))
sqlDB.SetConnMaxLifetime(time.Hour)
}
}
// errStoreDisabled 表示 Postgres 处于降级(未连接)模式,写操作无法进行。 // errStoreDisabled 表示 Postgres 处于降级(未连接)模式,写操作无法进行。
var errStoreDisabled = errors.New("postgres store disabled") var errStoreDisabled = errors.New("postgres store disabled")
@@ -36,6 +59,7 @@ func OpenPostgres(dsn string) *Postgres {
log.Printf("[store] postgres 不可用,降级运行(不持久化): %v", err) log.Printf("[store] postgres 不可用,降级运行(不持久化): %v", err)
return &Postgres{} return &Postgres{}
} }
tunePool(db) // 连接池上限,防高并发打爆 PG
// 一次性迁移:旧表用整型自增 id,与新雪花字符串 id 不兼容(AutoMigrate 不改主键类型)。 // 一次性迁移:旧表用整型自增 id,与新雪花字符串 id 不兼容(AutoMigrate 不改主键类型)。
// 备份模型密钥(唯一不可再生的数据) → 重建全部表 → 回灌模型。其余为可重建的测试数据。 // 备份模型密钥(唯一不可再生的数据) → 重建全部表 → 回灌模型。其余为可重建的测试数据。
migrateLegacyIntIDs(db) migrateLegacyIntIDs(db)
+1 -1
View File
@@ -90,7 +90,7 @@ func main() {
go b.RequestConfigWithRetry(ctx, contract.ConfigKindEmbedding, applyEmbed) go b.RequestConfigWithRetry(ctx, contract.ConfigKindEmbedding, applyEmbed)
go b.RequestConfigWithRetry(ctx, contract.ConfigKindChat, applyChat) go b.RequestConfigWithRetry(ctx, contract.ConfigKindChat, applyChat)
gw := mcp.NewGateway(b, engine, mem, hist, ragEngine) gw := mcp.NewGateway(b, engine, mem, hist, ragEngine, pgDSN)
log.Println("[mcp_go] serving MCP over sundynix.tools.go.* (Ctrl-C to quit)") log.Println("[mcp_go] serving MCP over sundynix.tools.go.* (Ctrl-C to quit)")
if err := gw.Serve(ctx); err != nil && err != context.Canceled { if err := gw.Serve(ctx); err != nil && err != context.Canceled {
+60
View File
@@ -0,0 +1,60 @@
package mcp
import (
"encoding/json"
"fmt"
"context"
"github.com/sundynix/sundynix-shared/contract"
)
// chartSeries 是一条数据系列。
type chartSeries struct {
Name string `json:"name,omitempty"`
Data []float64 `json:"data"`
}
// chartSpec 是图表的结构化规范(工具产出,前端据此渲染 SVG,不在后端出图)。
type chartSpec struct {
Type string `json:"type"` // bar / line / pie
Title string `json:"title,omitempty"` //
Labels []string `json:"labels"` // x 轴/扇区标签
Series []chartSeries `json:"series"` // 一条或多条数据系列(pie 取第一条)
}
// chart 工具:只校验并返回规范化图表 JSON(渲染交前端)。职责单一、零图片传输。
// 返回内容即一段 chart JSON;工具说明会指示 agent 在最终答复里用 ```chart 围栏原样包裹它。
func (g *Gateway) chart(_ context.Context, call *contract.ToolCall) *contract.ToolResult {
// 用 JSON round-trip 把 args 收进类型化结构(args 里 labels/series 是 []any,手解繁琐)。
raw, _ := json.Marshal(call.Args)
var in chartSpec
if err := json.Unmarshal(raw, &in); err != nil {
return &contract.ToolResult{OK: false, Error: "chart: 参数解析失败 —— " + err.Error()}
}
switch in.Type {
case "bar", "line", "pie":
case "":
in.Type = "bar"
default:
return &contract.ToolResult{OK: false, Error: "chart: type 仅支持 bar / line / pie"}
}
if len(in.Labels) == 0 {
return &contract.ToolResult{OK: false, Error: "chart: labels 必填"}
}
if len(in.Series) == 0 || len(in.Series[0].Data) == 0 {
return &contract.ToolResult{OK: false, Error: "chart: series 至少一条且 data 非空"}
}
for i, s := range in.Series {
if len(s.Data) != len(in.Labels) {
return &contract.ToolResult{OK: false,
Error: fmt.Sprintf("chart: 第 %d 条系列 data 长度(%d) 与 labels 长度(%d) 不一致", i+1, len(s.Data), len(in.Labels))}
}
}
if in.Type == "pie" {
in.Series = in.Series[:1] // pie 只用第一条系列
}
out, _ := json.Marshal(in)
// 提示 agent:把这段 JSON 用 ```chart 围栏原样放进最终答复,前端会渲染成图。
return &contract.ToolResult{OK: true, Content: string(out)}
}
+58 -7
View File
@@ -3,6 +3,7 @@ package mcp
import ( import (
"context" "context"
"database/sql"
"encoding/json" "encoding/json"
"fmt" "fmt"
"log" "log"
@@ -10,6 +11,7 @@ import (
"path/filepath" "path/filepath"
"sort" "sort"
"strings" "strings"
"sync"
"time" "time"
sharedbus "github.com/sundynix/sundynix-shared/bus" sharedbus "github.com/sundynix/sundynix-shared/bus"
@@ -30,6 +32,11 @@ type Gateway struct {
history *history.Store history *history.Store
rag *rag.Engine rag *rag.Engine
tools map[string]toolDef // 工具注册表:唯一事实源,dispatch 与 list_tools 共用,杜绝漂移 tools map[string]toolDef // 工具注册表:唯一事实源,dispatch 与 list_tools 共用,杜绝漂移
pgDSN string // 平台 PG DSNsql_query 兜底库;SQL_QUERY_DSN 未设时用它)
sqlOnce sync.Once // sql_query 只读连接懒连一次
sqlDB *sql.DB //
sqlDBErr error //
} }
// paramSpec 是一个工具参数的声明(供自主 agent 据此生成调用入参)。 // paramSpec 是一个工具参数的声明(供自主 agent 据此生成调用入参)。
@@ -53,8 +60,8 @@ type toolDef struct {
handler func(context.Context, *contract.ToolCall) *contract.ToolResult handler func(context.Context, *contract.ToolCall) *contract.ToolResult
} }
func NewGateway(b *sharedbus.Bus, s *search.Hybrid, m *memory.Store, h *history.Store, r *rag.Engine) *Gateway { func NewGateway(b *sharedbus.Bus, s *search.Hybrid, m *memory.Store, h *history.Store, r *rag.Engine, pgDSN string) *Gateway {
g := &Gateway{bus: b, search: s, memory: m, history: h, rag: r} g := &Gateway{bus: b, search: s, memory: m, history: h, rag: r, pgDSN: pgDSN}
g.tools = g.buildRegistry() g.tools = g.buildRegistry()
return g return g
} }
@@ -81,7 +88,7 @@ func (g *Gateway) buildRegistry() map[string]toolDef {
// —— 暴露给自主 agent 的工具(带参数 schema / 注入声明)—— // —— 暴露给自主 agent 的工具(带参数 schema / 注入声明)——
"wiki_search": { "wiki_search": {
cn: "知识检索", desc: "检索知识库,返回与查询最相关的资料片段。需要外部知识/事实依据时调用。", cn: "知识检索", desc: "检索知识库,返回与查询最相关的资料片段。需要外部知识/事实依据时调用。",
agent: true, agent: true,
params: []paramSpec{{Name: "q", Type: "string", Desc: "检索查询语句", Required: true}}, params: []paramSpec{{Name: "q", Type: "string", Desc: "检索查询语句", Required: true}},
inject: []string{"kb"}, handler: g.wikiSearch, inject: []string{"kb"}, handler: g.wikiSearch,
}, },
@@ -102,6 +109,50 @@ func (g *Gateway) buildRegistry() map[string]toolDef {
cn: "历史召回", desc: "取当前会话最近多轮对话,用于理解上下文。", cn: "历史召回", desc: "取当前会话最近多轮对话,用于理解上下文。",
agent: true, inject: []string{"session_id"}, handler: g.historyGet, agent: true, inject: []string{"session_id"}, handler: g.historyGet,
}, },
"web_search": {
cn: "联网搜索", desc: "联网搜索,返回最新网页结果(标题/链接/摘要)。需要实时/最新信息或外部事实时调用。",
agent: true,
params: []paramSpec{
{Name: "q", Type: "string", Desc: "搜索关键词", Required: true},
{Name: "topK", Type: "integer", Desc: "返回结果条数(默认 5,最多 10)"},
},
handler: g.webSearch,
},
"web_fetch": {
cn: "网页抓取", desc: "抓取一个网页 URL 并提取正文文本。需要读取某个链接的内容时调用。",
agent: true,
params: []paramSpec{{Name: "url", Type: "string", Desc: "要抓取的网页 URL", Required: true}},
handler: g.webFetch,
},
"calculator": {
cn: "计算器", desc: "精确计算数学表达式(+ - * / % ^ 与括号)。涉及算术/数值计算时调用,不要心算。",
agent: true,
params: []paramSpec{{Name: "expr", Type: "string", Desc: "数学表达式,如 (3+4)*2^3", Required: true}},
handler: g.calculator,
},
"current_datetime": {
cn: "当前时间", desc: "获取当前日期与时间(含星期)。需要“现在/今天几号/星期几”等时间信息时调用。",
agent: true,
params: []paramSpec{{Name: "tz", Type: "string", Desc: "可选时区,如 Asia/Shanghai;缺省服务器本地时区"}},
handler: g.currentDatetime,
},
"sql_query": {
cn: "SQL查询", desc: "对数据库执行只读 SQL 查询(仅 SELECT/WITH),返回结果表。需要查业务数据/统计时调用。",
agent: true,
params: []paramSpec{{Name: "sql", Type: "string", Desc: "只读 SQL,如 SELECT count(*) FROM sundynix_task", Required: true}},
handler: g.sqlQuery,
},
"chart": {
cn: "图表", desc: "把数据生成图表。返回图表 JSON——请在最终答复中用 ```chart 代码块原样包裹该 JSON,前端会渲染成图。需要可视化数据分布/趋势时调用。",
agent: true,
params: []paramSpec{
{Name: "type", Type: "string", Desc: "图表类型:bar / line / pie", Required: true},
{Name: "title", Type: "string", Desc: "图表标题"},
{Name: "labels", Type: "array", Desc: "x 轴/扇区标签数组,如 [\"Q1\",\"Q2\"]", Required: true},
{Name: "series", Type: "array", Desc: "数据系列数组,如 [{\"name\":\"销量\",\"data\":[120,180]}]", Required: true},
},
handler: g.chart,
},
// —— 仅内部/流水线/管理用,不暴露给自主 agent —— // —— 仅内部/流水线/管理用,不暴露给自主 agent ——
"kb_ingest": {cn: "知识入库", desc: "文本切块 → 向量化 → 写入 Milvus / Bleve", handler: g.kbIngest}, "kb_ingest": {cn: "知识入库", desc: "文本切块 → 向量化 → 写入 Milvus / Bleve", handler: g.kbIngest},
@@ -147,10 +198,10 @@ func (g *Gateway) listTools() *contract.ToolResult {
Name string `json:"name"` Name string `json:"name"`
CN string `json:"cn"` CN string `json:"cn"`
Desc string `json:"desc"` Desc string `json:"desc"`
Agent bool `json:"agent_exposed"` // 是否给自主 agent Agent bool `json:"agent_exposed"` // 是否给自主 agent
AgentName string `json:"agent_name,omitempty"`// 模型可见名(空=name AgentName string `json:"agent_name,omitempty"` // 模型可见名(空=name
Params []paramSpec `json:"params,omitempty"` // 模型可填参数 Params []paramSpec `json:"params,omitempty"` // 模型可填参数
Inject []string `json:"inject,omitempty"` // 服务端注入参数(不暴露给模型) Inject []string `json:"inject,omitempty"` // 服务端注入参数(不暴露给模型)
} }
out := make([]info, 0, len(g.tools)) out := make([]info, 0, len(g.tools))
for name, td := range g.tools { for name, td := range g.tools {
+168
View File
@@ -0,0 +1,168 @@
package mcp
import (
"context"
"database/sql"
"fmt"
"os"
"strings"
"time"
"github.com/sundynix/sundynix-shared/contract"
"gorm.io/driver/postgres"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)
const (
sqlQueryMaxRows = 100 // 结果行上限,防超大结果撑爆上下文
sqlQueryTimeout = 10 * time.Second // 单次查询超时
sqlCellMaxRunes = 200 // 单元格文本上限
)
// sqlQueryDB 懒连只读查询库:优先 SQL_QUERY_DSN(生产应指向专用只读库 / 只读账号),
// 未配置则回退本服务已解析的平台 PG DSN(g.pgDSN,仅供开发;线上务必单配,避免把平台库直接暴露给 agent)。
func (g *Gateway) sqlQueryDB() (*sql.DB, error) {
g.sqlOnce.Do(func() {
dsn := strings.TrimSpace(os.Getenv("SQL_QUERY_DSN"))
if dsn == "" {
dsn = g.pgDSN
}
if dsn == "" {
g.sqlDBErr = fmt.Errorf("未配置 SQL_QUERY_DSN")
return
}
gdb, err := gorm.Open(postgres.New(postgres.Config{DSN: dsn}), &gorm.Config{
Logger: logger.Default.LogMode(logger.Silent),
})
if err != nil {
g.sqlDBErr = err
return
}
db, err := gdb.DB()
if err != nil {
g.sqlDBErr = err
return
}
db.SetMaxOpenConns(8) // 只读查询独立小池,不挤占主连接
db.SetMaxIdleConns(2)
db.SetConnMaxLifetime(time.Hour)
g.sqlDB = db
})
return g.sqlDB, g.sqlDBErr
}
// sqlQuery 对数据库执行只读 SQL。三重防护:① 仅允许单条 SELECT/WITH 语句(拒多语句/写/DDL);
// ② 在「只读事务」中执行(Postgres 引擎级强制只读,是真正的兜底);③ 行数 + 超时 + 单元格上限。
func (g *Gateway) sqlQuery(ctx context.Context, call *contract.ToolCall) *contract.ToolResult {
q := strings.TrimSpace(fmt.Sprint(call.Args["sql"]))
if q == "" || q == "<nil>" {
return &contract.ToolResult{OK: false, Error: "sql_query: sql 必填"}
}
if reason, ok := validateReadOnlySQL(q); !ok {
return &contract.ToolResult{OK: false, Error: "sql_query: " + reason}
}
db, err := g.sqlQueryDB()
if err != nil {
return &contract.ToolResult{OK: false, Error: "sql_query: " + err.Error()}
}
cctx, cancel := context.WithTimeout(ctx, sqlQueryTimeout)
defer cancel()
// 只读事务:即便上面的语句校验被绕过,Postgres 也会在引擎层拒绝任何写操作。
tx, err := db.BeginTx(cctx, &sql.TxOptions{ReadOnly: true})
if err != nil {
return &contract.ToolResult{OK: false, Error: "sql_query: " + err.Error()}
}
defer func() { _ = tx.Rollback() }()
rows, err := tx.QueryContext(cctx, q)
if err != nil {
return &contract.ToolResult{OK: false, Error: "sql_query: " + err.Error()}
}
defer rows.Close()
cols, _ := rows.Columns()
var b strings.Builder
b.WriteString(strings.Join(cols, " | ") + "\n")
b.WriteString(strings.Repeat("-", len(strings.Join(cols, " | "))) + "\n")
n := 0
for rows.Next() {
if n >= sqlQueryMaxRows {
break
}
cells := make([]any, len(cols))
ptrs := make([]any, len(cols))
for i := range cells {
ptrs[i] = &cells[i]
}
if err := rows.Scan(ptrs...); err != nil {
return &contract.ToolResult{OK: false, Error: "sql_query: scan " + err.Error()}
}
strs := make([]string, len(cols))
for i, c := range cells {
strs[i] = truncateRunes(cellToString(c), sqlCellMaxRunes)
}
b.WriteString(strings.Join(strs, " | ") + "\n")
n++
}
if err := rows.Err(); err != nil {
return &contract.ToolResult{OK: false, Error: "sql_query: " + err.Error()}
}
b.WriteString(fmt.Sprintf("(%d 行%s)", n, map[bool]string{true: ",已截断至上限"}[n >= sqlQueryMaxRows]))
return &contract.ToolResult{OK: true, Content: b.String()}
}
// validateReadOnlySQL 静态校验:单条语句、以 SELECT/WITH 开头、不含写/DDL 关键字。
// 与只读事务双保险(这层挡明显误用,事务层是硬保证)。
func validateReadOnlySQL(q string) (reason string, ok bool) {
s := strings.TrimSpace(q)
s = strings.TrimSuffix(s, ";")
if strings.Contains(s, ";") {
return "只允许单条语句", false
}
low := strings.ToLower(s)
if !strings.HasPrefix(low, "select") && !strings.HasPrefix(low, "with") {
return "只允许 SELECT / WITH 查询", false
}
// 词边界匹配写/DDL 关键字(避免误伤列名如 created_at)。
for _, kw := range []string{"insert", "update", "delete", "drop", "alter", "create",
"truncate", "grant", "revoke", "comment", "copy", "merge", "call", "do "} {
if containsWord(low, kw) {
return "检测到写/DDL 关键字「" + strings.TrimSpace(kw) + "」,仅允许只读查询", false
}
}
return "", true
}
func containsWord(s, word string) bool {
for i := 0; ; {
idx := strings.Index(s[i:], word)
if idx < 0 {
return false
}
j := i + idx
before := j == 0 || !isWordChar(rune(s[j-1]))
after := j+len(word) >= len(s) || !isWordChar(rune(s[j+len(word)]))
if before && after {
return true
}
i = j + len(word)
}
}
func isWordChar(r rune) bool {
return r == '_' || (r >= 'a' && r <= 'z') || (r >= 'A' && r <= 'Z') || (r >= '0' && r <= '9')
}
func cellToString(v any) string {
switch x := v.(type) {
case nil:
return "NULL"
case []byte:
return string(x)
case time.Time:
return x.Format("2006-01-02 15:04:05")
default:
return fmt.Sprint(x)
}
}
+406
View File
@@ -0,0 +1,406 @@
package mcp
import (
"context"
"fmt"
"html"
"io"
"math"
"net/http"
"net/url"
"os"
"regexp"
"strconv"
"strings"
"time"
"github.com/sundynix/sundynix-shared/contract"
)
// ===== current_datetime:当前日期时间(LLM 不知道"现在"=====
var weekdaysCN = []string{"周日", "周一", "周二", "周三", "周四", "周五", "周六"}
func (g *Gateway) currentDatetime(_ context.Context, call *contract.ToolCall) *contract.ToolResult {
loc := time.Local
if tz, _ := call.Args["tz"].(string); strings.TrimSpace(tz) != "" {
if l, err := time.LoadLocation(strings.TrimSpace(tz)); err == nil {
loc = l
}
}
now := time.Now().In(loc)
out := fmt.Sprintf("%s %s%s,时区 %sUnix %d",
now.Format("2006-01-02"), now.Format("15:04:05"),
weekdaysCN[int(now.Weekday())], now.Format("MST-07:00"), now.Unix())
return &contract.ToolResult{OK: true, Content: out}
}
// ===== calculator:安全表达式求值(+ - * / % ^ 括号、一元负号)=====
func (g *Gateway) calculator(_ context.Context, call *contract.ToolCall) *contract.ToolResult {
expr := strings.TrimSpace(fmt.Sprint(call.Args["expr"]))
if expr == "" || expr == "<nil>" {
return &contract.ToolResult{OK: false, Error: "calculator: expr 必填"}
}
v, err := evalArith(expr)
if err != nil {
return &contract.ToolResult{OK: false, Error: "calculator: " + err.Error()}
}
return &contract.ToolResult{OK: true, Content: strconv.FormatFloat(v, 'g', -1, 64)}
}
// evalArith 用调度场算法把中缀表达式转 RPN 再求值。仅支持数值与 + - * / % ^ ( ),杜绝任意代码执行。
func evalArith(s string) (float64, error) {
toks, err := tokenizeArith(s)
if err != nil {
return 0, err
}
prec := map[string]int{"+": 1, "-": 1, "*": 2, "/": 2, "%": 2, "^": 3, "u-": 4}
rightAssoc := map[string]bool{"^": true, "u-": true}
var output, ops []string
for i, t := range toks {
switch {
case isNumber(t):
output = append(output, t)
case t == "(":
ops = append(ops, t)
case t == ")":
for len(ops) > 0 && ops[len(ops)-1] != "(" {
output = append(output, ops[len(ops)-1])
ops = ops[:len(ops)-1]
}
if len(ops) == 0 {
return 0, fmt.Errorf("括号不匹配")
}
ops = ops[:len(ops)-1] // 弹出 "("
default: // 运算符
op := t
// 一元负号:在表达式开头或运算符/左括号之后的 "-"。
if op == "-" && (i == 0 || isOperator(toks[i-1]) || toks[i-1] == "(") {
op = "u-"
}
for len(ops) > 0 {
top := ops[len(ops)-1]
if top == "(" {
break
}
if prec[top] > prec[op] || (prec[top] == prec[op] && !rightAssoc[op]) {
output = append(output, top)
ops = ops[:len(ops)-1]
} else {
break
}
}
ops = append(ops, op)
}
}
for len(ops) > 0 {
if ops[len(ops)-1] == "(" {
return 0, fmt.Errorf("括号不匹配")
}
output = append(output, ops[len(ops)-1])
ops = ops[:len(ops)-1]
}
return evalRPN(output)
}
func evalRPN(rpn []string) (float64, error) {
var st []float64
pop := func() (float64, error) {
if len(st) == 0 {
return 0, fmt.Errorf("表达式非法")
}
v := st[len(st)-1]
st = st[:len(st)-1]
return v, nil
}
for _, t := range rpn {
if isNumber(t) {
f, _ := strconv.ParseFloat(t, 64)
st = append(st, f)
continue
}
if t == "u-" {
a, err := pop()
if err != nil {
return 0, err
}
st = append(st, -a)
continue
}
b, err := pop()
if err != nil {
return 0, err
}
a, err := pop()
if err != nil {
return 0, err
}
switch t {
case "+":
st = append(st, a+b)
case "-":
st = append(st, a-b)
case "*":
st = append(st, a*b)
case "/":
if b == 0 {
return 0, fmt.Errorf("除以零")
}
st = append(st, a/b)
case "%":
st = append(st, math.Mod(a, b))
case "^":
st = append(st, math.Pow(a, b))
default:
return 0, fmt.Errorf("未知运算符 %q", t)
}
}
if len(st) != 1 {
return 0, fmt.Errorf("表达式非法")
}
return st[0], nil
}
func tokenizeArith(s string) ([]string, error) {
var toks []string
r := []rune(s)
for i := 0; i < len(r); {
c := r[i]
switch {
case c == ' ' || c == '\t':
i++
case strings.ContainsRune("+-*/%^()", c):
toks = append(toks, string(c))
i++
case (c >= '0' && c <= '9') || c == '.':
j := i
for j < len(r) && ((r[j] >= '0' && r[j] <= '9') || r[j] == '.' || r[j] == 'e' || r[j] == 'E' ||
((r[j] == '+' || r[j] == '-') && j > i && (r[j-1] == 'e' || r[j-1] == 'E'))) {
j++
}
num := string(r[i:j])
if _, err := strconv.ParseFloat(num, 64); err != nil {
return nil, fmt.Errorf("非法数字 %q", num)
}
toks = append(toks, num)
i = j
default:
return nil, fmt.Errorf("非法字符 %q(仅支持数字与 + - * / %% ^ ( )", string(c))
}
}
if len(toks) == 0 {
return nil, fmt.Errorf("空表达式")
}
return toks, nil
}
func isNumber(t string) bool { _, err := strconv.ParseFloat(t, 64); return err == nil }
func isOperator(t string) bool {
return t == "+" || t == "-" || t == "*" || t == "/" || t == "%" || t == "^"
}
// ===== web_fetch:抓取网页并提取正文文本 =====
var (
// RE2 不支持反向引用,逐标签枚举其配对闭合。
reScriptStyle = regexp.MustCompile(`(?is)<script\b[^>]*>.*?</script\s*>|<style\b[^>]*>.*?</style\s*>|<head\b[^>]*>.*?</head\s*>|<noscript\b[^>]*>.*?</noscript\s*>`)
reTag = regexp.MustCompile(`(?s)<[^>]+>`)
reWS = regexp.MustCompile(`[ \t\x{00a0}]+`)
reBlankLines = regexp.MustCompile(`\n\s*\n\s*\n+`)
)
const webFetchMaxText = 8000 // 提取正文上限(rune),避免塞爆上下文
func (g *Gateway) webFetch(ctx context.Context, call *contract.ToolCall) *contract.ToolResult {
raw := strings.TrimSpace(fmt.Sprint(call.Args["url"]))
if raw == "" || raw == "<nil>" {
return &contract.ToolResult{OK: false, Error: "web_fetch: url 必填"}
}
if reason, ok := validateExternalURL(raw, extAllowlist()); !ok {
return &contract.ToolResult{OK: false, Error: "web_fetch: URL 被拦截 —— " + reason}
}
body, status, err := httpGet(ctx, raw)
if err != nil {
return &contract.ToolResult{OK: false, Error: "web_fetch: " + err.Error()}
}
text := htmlToText(body)
if rs := []rune(text); len(rs) > webFetchMaxText {
text = string(rs[:webFetchMaxText]) + "\n…(已截断)"
}
return &contract.ToolResult{OK: true, Content: fmt.Sprintf("URL: %s (HTTP %d)\n\n%s", raw, status, text)}
}
// htmlToText 把 HTML 粗提取为可读文本:去脚本/样式 → 去标签 → 解实体 → 收敛空白。
func htmlToText(h string) string {
h = reScriptStyle.ReplaceAllString(h, " ")
h = regexp.MustCompile(`(?i)<\s*(br|/p|/div|/li|/h[1-6]|/tr)\s*/?>`).ReplaceAllString(h, "\n")
h = reTag.ReplaceAllString(h, "")
h = html.UnescapeString(h)
h = reWS.ReplaceAllString(h, " ")
h = reBlankLines.ReplaceAllString(h, "\n\n")
return strings.TrimSpace(h)
}
// ===== web_search:联网搜索(Tavily 有 key 则用,否则 DuckDuckGo HTML 兜底)=====
const webSearchDefaultTopK = 5
func (g *Gateway) webSearch(ctx context.Context, call *contract.ToolCall) *contract.ToolResult {
q := strings.TrimSpace(fmt.Sprint(call.Args["q"]))
if q == "" || q == "<nil>" {
return &contract.ToolResult{OK: false, Error: "web_search: q 必填"}
}
topK := webSearchDefaultTopK
if n, ok := toInt(call.Args["topK"]); ok && n > 0 && n <= 10 {
topK = n
}
if key := strings.TrimSpace(os.Getenv("TAVILY_API_KEY")); key != "" {
if out, err := tavilySearch(ctx, key, q, topK); err == nil {
return &contract.ToolResult{OK: true, Content: out}
}
// Tavily 失败 → 落 DDG 兜底
}
out, err := ddgSearch(ctx, q, topK)
if err != nil {
return &contract.ToolResult{OK: false, Error: "web_search: " + err.Error()}
}
return &contract.ToolResult{OK: true, Content: out}
}
var (
reDDGLink = regexp.MustCompile(`(?is)<a[^>]+class="result__a"[^>]+href="([^"]+)"[^>]*>(.*?)</a>`)
reDDGSnippet = regexp.MustCompile(`(?is)<a[^>]+class="result__snippet"[^>]*>(.*?)</a>`)
)
// ddgSearch 抓 DuckDuckGo HTML 版结果(免 key)。SERP 结构变动时尽力解析,解析不到则返回提示。
func ddgSearch(ctx context.Context, q string, topK int) (string, error) {
u := "https://html.duckduckgo.com/html/?q=" + url.QueryEscape(q)
body, _, err := httpGet(ctx, u)
if err != nil {
return "", err
}
links := reDDGLink.FindAllStringSubmatch(body, -1)
snips := reDDGSnippet.FindAllStringSubmatch(body, -1)
if len(links) == 0 {
return "", fmt.Errorf("未解析到结果(DDG 可能限流或改版)")
}
var b strings.Builder
for i := 0; i < len(links) && i < topK; i++ {
title := strings.TrimSpace(htmlToText(links[i][2]))
href := ddgUnwrap(links[i][1])
snippet := ""
if i < len(snips) {
snippet = strings.TrimSpace(htmlToText(snips[i][1]))
}
fmt.Fprintf(&b, "%d. %s\n %s\n %s\n", i+1, title, href, snippet)
}
return strings.TrimSpace(b.String()), nil
}
// ddgUnwrap 还原 DDG 的跳转链接 /l/?uddg=<编码真实 url>。
func ddgUnwrap(href string) string {
if i := strings.Index(href, "uddg="); i >= 0 {
raw := href[i+len("uddg="):]
if amp := strings.IndexByte(raw, '&'); amp >= 0 {
raw = raw[:amp]
}
if dec, err := url.QueryUnescape(raw); err == nil {
return dec
}
}
if strings.HasPrefix(href, "//") {
return "https:" + href
}
return href
}
// tavilySearch 走 Tavily Search API(需 TAVILY_API_KEY)—— 返回干净的标题/URL/摘要。
func tavilySearch(ctx context.Context, key, q string, topK int) (string, error) {
payload := fmt.Sprintf(`{"api_key":%q,"query":%q,"max_results":%d}`, key, q, topK)
req, _ := http.NewRequestWithContext(ctx, http.MethodPost, "https://api.tavily.com/search", strings.NewReader(payload))
req.Header.Set("Content-Type", "application/json")
resp, err := (&http.Client{Timeout: extTimeout}).Do(req)
if err != nil {
return "", err
}
defer resp.Body.Close()
if resp.StatusCode >= 400 {
return "", fmt.Errorf("tavily HTTP %d", resp.StatusCode)
}
data, _ := io.ReadAll(io.LimitReader(resp.Body, extMaxBytes))
// 轻量提取(不引 JSON 结构体):用正则挑 title/url/content。
reItem := regexp.MustCompile(`(?is)"title"\s*:\s*"(.*?)".*?"url"\s*:\s*"(.*?)".*?"content"\s*:\s*"(.*?)"`)
items := reItem.FindAllStringSubmatch(string(data), -1)
if len(items) == 0 {
return "", fmt.Errorf("tavily 无结果")
}
var b strings.Builder
for i, m := range items {
if i >= topK {
break
}
fmt.Fprintf(&b, "%d. %s\n %s\n %s\n", i+1, jsonUnesc(m[1]), jsonUnesc(m[2]), truncateRunes(jsonUnesc(m[3]), 240))
}
return strings.TrimSpace(b.String()), nil
}
// ===== 公共小工具 =====
func httpGet(ctx context.Context, raw string) (body string, status int, err error) {
allow := extAllowlist()
req, err := http.NewRequestWithContext(ctx, http.MethodGet, raw, nil)
if err != nil {
return "", 0, err
}
// 带个常见 UA,部分站点对空 UA 返回 403。
req.Header.Set("User-Agent", "Mozilla/5.0 (compatible; sundynix-agentix/1.0)")
client := &http.Client{
Timeout: extTimeout,
CheckRedirect: func(r *http.Request, via []*http.Request) error {
if len(via) >= 3 {
return fmt.Errorf("重定向过多")
}
if reason, ok := validateExternalURL(r.URL.String(), allow); !ok {
return fmt.Errorf("重定向被拦截:%s", reason)
}
return nil
},
}
resp, err := client.Do(req)
if err != nil {
return "", 0, err
}
defer resp.Body.Close()
data, _ := io.ReadAll(io.LimitReader(resp.Body, extMaxBytes))
return string(data), resp.StatusCode, nil
}
func toInt(v any) (int, bool) {
switch n := v.(type) {
case float64:
return int(n), true
case int:
return n, true
case string:
if i, err := strconv.Atoi(strings.TrimSpace(n)); err == nil {
return i, true
}
}
return 0, false
}
// truncateRunes 按 rune 截断(中文安全),超出加省略号。
func truncateRunes(s string, n int) string {
r := []rune(s)
if len(r) <= n {
return s
}
return string(r[:n]) + "…"
}
// jsonUnesc 还原 JSON 字符串里的常见转义(够用即可,不做完整解码)。
func jsonUnesc(s string) string {
r := strings.NewReplacer(`\"`, `"`, `\\`, `\`, `\n`, " ", `\t`, " ", `\/`, "/")
return strings.TrimSpace(r.Replace(s))
}
@@ -0,0 +1,87 @@
package mcp
import (
"math"
"strings"
"testing"
)
func TestEvalArith(t *testing.T) {
cases := map[string]float64{
"1+2*3": 7,
"(1+2)*3": 9,
"2^3^2": 512, // 右结合
"-3+5": 2,
"10/4": 2.5,
"10%3": 1,
"2*(3+4)-5": 9,
"-(2+3)*2": -10,
"3.5e1+0.5": 35.5,
}
for in, want := range cases {
got, err := evalArith(in)
if err != nil {
t.Errorf("evalArith(%q) 报错: %v", in, err)
continue
}
if math.Abs(got-want) > 1e-9 {
t.Errorf("evalArith(%q)=%v want %v", in, got, want)
}
}
}
func TestEvalArithErrors(t *testing.T) {
for _, in := range []string{"1/0", "1+", "(1+2", "1+2)", "rm -rf /", "import os", ""} {
if _, err := evalArith(in); err == nil {
t.Errorf("evalArith(%q) 应报错但没有", in)
}
}
}
func TestHTMLToText(t *testing.T) {
in := `<html><head><title>x</title></head><body><script>alert(1)</script>
<h1>标题</h1><p>第一段&amp;符号</p><style>.a{}</style><div>第二段</div></body></html>`
got := htmlToText(in)
if strings.Contains(got, "alert") || strings.Contains(got, ".a{}") {
t.Fatalf("脚本/样式未剔除: %q", got)
}
if !strings.Contains(got, "标题") || !strings.Contains(got, "第一段&符号") || !strings.Contains(got, "第二段") {
t.Fatalf("正文/实体提取不正确: %q", got)
}
}
func TestDDGUnwrap(t *testing.T) {
got := ddgUnwrap("//duckduckgo.com/l/?uddg=https%3A%2F%2Fexample.com%2Fa&rut=x")
if got != "https://example.com/a" {
t.Fatalf("ddgUnwrap 解码错误: %q", got)
}
}
func TestValidateReadOnlySQL(t *testing.T) {
okCases := []string{
"SELECT 1",
"select count(*) from sundynix_task",
"WITH x AS (SELECT 1) SELECT * FROM x",
" select * from t where created_at > now() ", // created/update 作为列名一部分不应误伤
}
for _, q := range okCases {
if _, ok := validateReadOnlySQL(q); !ok {
t.Errorf("应通过: %q", q)
}
}
badCases := []string{
"insert into t values (1)",
"update t set a=1",
"delete from t",
"drop table t",
"select 1; drop table t",
"truncate t",
"SELECT 1; SELECT 2",
"grant all on t to x",
}
for _, q := range badCases {
if _, ok := validateReadOnlySQL(q); ok {
t.Errorf("应拒绝: %q", q)
}
}
}
+14
View File
@@ -7,7 +7,9 @@ import (
"fmt" "fmt"
"log" "log"
"math" "math"
"os"
"sort" "sort"
"strconv"
"strings" "strings"
"time" "time"
@@ -69,6 +71,18 @@ func Open(dsn string) *Store {
log.Printf("[memory] postgres 不可用,记忆降级(召回为空): %v", err) log.Printf("[memory] postgres 不可用,记忆降级(召回为空): %v", err)
return &Store{} return &Store{}
} }
// 连接池上限,防高并发打爆 PGmax_connections 默认 100);默认 25,可经 DB_MAX_OPEN_CONNS 调。
if sqlDB, derr := db.DB(); derr == nil {
maxOpen := 25
if v := os.Getenv("DB_MAX_OPEN_CONNS"); v != "" {
if n, perr := strconv.Atoi(v); perr == nil && n > 0 {
maxOpen = n
}
}
sqlDB.SetMaxOpenConns(maxOpen)
sqlDB.SetMaxIdleConns(5)
sqlDB.SetConnMaxLifetime(time.Hour)
}
// 一次性迁移:旧表用复合主键 (user_id,key) 无 id/时间戳,与雪花规约不兼容。 // 一次性迁移:旧表用复合主键 (user_id,key) 无 id/时间戳,与雪花规约不兼容。
migrateLegacyProfile(db) migrateLegacyProfile(db)
if err := db.AutoMigrate(&Profile{}); err != nil { if err := db.AutoMigrate(&Profile{}); err != nil {
@@ -47,6 +47,8 @@ class McpGateway:
self.interpreter = CodeInterpreter() # Docker 隔离沙箱 self.interpreter = CodeInterpreter() # Docker 隔离沙箱
self._nc: nats.NATS | None = None self._nc: nats.NATS | None = None
self._tools: dict[str, callable] = {} self._tools: dict[str, callable] = {}
# 单实例工具并发上限(垂直扩);队列组多副本再叠加水平扩。
self._sem = asyncio.Semaphore(int(os.getenv("MCP_TOOL_CONCURRENCY", "16")))
async def register_tools(self) -> None: async def register_tools(self) -> None:
"""注册工具名 → 处理协程。""" """注册工具名 → 处理协程。"""
@@ -77,25 +79,30 @@ class McpGateway:
await asyncio.Event().wait() await asyncio.Event().wait()
async def _on_call(self, msg) -> None: async def _on_call(self, msg) -> None:
"""解析 ToolCall → 路由到工具 → Respond ToolResult""" """立即派发到独立 task 并发处理,不阻塞订阅消息循环(nats-py 单订阅回调本是串行 await)"""
try: asyncio.create_task(self._handle(msg))
req = json.loads(msg.data)
except Exception as e: # noqa: BLE001 async def _handle(self, msg) -> None:
await self._reply(msg, ok=False, error=f"bad tool call: {e}") """解析 ToolCall → 路由到工具 → Respond ToolResult(限并发)。"""
return async with self._sem:
tool = req.get("tool", "") try:
task_id = req.get("task_id", "") req = json.loads(msg.data)
args = req.get("args") or {} except Exception as e: # noqa: BLE001
log.info("[mcp_py] tool=%s task=%s args=%s", tool, task_id, args) await self._reply(msg, ok=False, error=f"bad tool call: {e}")
fn = self._tools.get(tool) return
if fn is None: tool = req.get("tool", "")
await self._reply(msg, ok=False, error=f"unknown tool: {tool}") task_id = req.get("task_id", "")
return args = req.get("args") or {}
try: log.info("[mcp_py] tool=%s task=%s args=%s", tool, task_id, args)
content = await fn(args) fn = self._tools.get(tool)
await self._reply(msg, ok=True, content=content) if fn is None:
except Exception as e: # noqa: BLE001 await self._reply(msg, ok=False, error=f"unknown tool: {tool}")
await self._reply(msg, ok=False, error=str(e)) return
try:
content = await fn(args)
await self._reply(msg, ok=True, content=content)
except Exception as e: # noqa: BLE001
await self._reply(msg, ok=False, error=str(e))
@staticmethod @staticmethod
async def _reply(msg, *, ok: bool, content: str = "", error: str = "") -> None: async def _reply(msg, *, ok: bool, content: str = "", error: str = "") -> None:
+130 -33
View File
@@ -6,6 +6,9 @@ import (
"context" "context"
"encoding/json" "encoding/json"
"fmt" "fmt"
"log"
"os"
"strconv"
"time" "time"
"github.com/nats-io/nats.go" "github.com/nats-io/nats.go"
@@ -230,27 +233,48 @@ func (b *Bus) CallTool(ctx context.Context, subject string, call *contract.ToolC
// ToolHandler 处理一次工具调用并返回结果。 // ToolHandler 处理一次工具调用并返回结果。
type ToolHandler func(ctx context.Context, call *contract.ToolCall) *contract.ToolResult type ToolHandler func(ctx context.Context, call *contract.ToolCall) *contract.ToolResult
// ServeTool 以队列组订阅工具主题(可用通配 sundynix.tools.go.>), // toolConcurrency 返回单实例工具并发处理上限(env MCP_TOOL_CONCURRENCY,默认 16)。
// 对每个请求调用 h 并 Respond,队列组内多副本自动负载均衡。 func toolConcurrency() int {
// 返回的 unsub 用于退订。 if v := os.Getenv("MCP_TOOL_CONCURRENCY"); v != "" {
if n, err := strconv.Atoi(v); err == nil && n > 0 {
return n
}
}
return 16
}
// ServeTool 以队列组订阅工具主题(可用通配 sundynix.tools.go.>),队列组内多副本自动负载均衡(水平扩)。
// 单实例内每个请求分发到独立 goroutine 并发处理(垂直扩,上限 MCP_TOOL_CONCURRENCY)——
// 回调立即返回不阻塞 NATS 投递;信号量约束同时执行的 handler 数。返回的 unsub 用于退订。
func (b *Bus) ServeTool(subject, queue string, h ToolHandler) (unsub func() error, err error) { func (b *Bus) ServeTool(subject, queue string, h ToolHandler) (unsub func() error, err error) {
sem := make(chan struct{}, toolConcurrency())
sub, err := b.nc.QueueSubscribe(subject, queue, func(m *nats.Msg) { sub, err := b.nc.QueueSubscribe(subject, queue, func(m *nats.Msg) {
var call contract.ToolCall // 立即起 goroutine 并返回,让 NATS 继续投递下一条(核心 NATS 单订阅回调本是串行)。
if err := json.Unmarshal(m.Data, &call); err != nil { go func() {
respond(m, &contract.ToolResult{OK: false, Error: "bad tool call: " + err.Error()}) sem <- struct{}{} // 限并发:超额则在此短暂排队
return defer func() {
} <-sem
// 还原上游链路并开服务端 span(成为 dispatcher tool.call span 的子节点)。 if r := recover(); r != nil { // handler panic → 回错误结果,调用方不必干等超时
mctx := extractTrace(context.Background(), m.Header) respond(m, &contract.ToolResult{OK: false, Error: fmt.Sprintf("tool panic: %v", r)})
mctx, span := tracer().Start(mctx, "tool.serve "+call.Tool, }
trace.WithSpanKind(trace.SpanKindServer), }()
trace.WithAttributes(attribute.String("sundynix.tool", call.Tool))) var call contract.ToolCall
res := h(mctx, &call) if err := json.Unmarshal(m.Data, &call); err != nil {
if res != nil { respond(m, &contract.ToolResult{OK: false, Error: "bad tool call: " + err.Error()})
span.SetAttributes(attribute.Bool("sundynix.tool.ok", res.OK)) return
} }
span.End() // 还原上游链路并开服务端 span(成为 dispatcher tool.call span 的子节点)。
respond(m, res) mctx := extractTrace(context.Background(), m.Header)
mctx, span := tracer().Start(mctx, "tool.serve "+call.Tool,
trace.WithSpanKind(trace.SpanKindServer),
trace.WithAttributes(attribute.String("sundynix.tool", call.Tool)))
defer span.End()
res := h(mctx, &call)
if res != nil {
span.SetAttributes(attribute.Bool("sundynix.tool.ok", res.OK))
}
respond(m, res)
}()
}) })
if err != nil { if err != nil {
return nil, fmt.Errorf("serve tool %s: %w", subject, err) return nil, fmt.Errorf("serve tool %s: %w", subject, err)
@@ -314,6 +338,48 @@ func (b *Bus) SubscribeTaskStatus(onEvent func(*contract.TaskStatusEvent)) (unsu
return sub.Unsubscribe, nil return sub.Unsubscribe, nil
} }
// ---- 人工审批(HITLcore NATS pub-sub----
// PublishApproval 广播一次人工审批决定(网关在收到 UI 的批准/拒绝后调用)。
func (b *Bus) PublishApproval(dec *contract.ApprovalDecision) error {
data, err := json.Marshal(dec)
if err != nil {
return err
}
return b.nc.Publish(contract.ApprovalSubject(dec.TaskID), data)
}
// WaitApproval 阻塞等待某任务的人工审批决定,直到收到、ctx 取消或超时。
// dispatcher 在审批节点调用:先订阅再等待(订阅早于决定到达,避免错过)。
// 超时返回 (nil, error) —— 调用方据安全默认(拒绝)处理。
func (b *Bus) WaitApproval(ctx context.Context, taskID string, timeout time.Duration) (*contract.ApprovalDecision, error) {
ch := make(chan *contract.ApprovalDecision, 1)
sub, err := b.nc.Subscribe(contract.ApprovalSubject(taskID), func(m *nats.Msg) {
var dec contract.ApprovalDecision
if json.Unmarshal(m.Data, &dec) == nil {
select {
case ch <- &dec:
default: // 已收到一条,丢弃后续重复决定
}
}
})
if err != nil {
return nil, fmt.Errorf("subscribe approval: %w", err)
}
defer func() { _ = sub.Unsubscribe() }()
timer := time.NewTimer(timeout)
defer timer.Stop()
select {
case dec := <-ch:
return dec, nil
case <-timer.C:
return nil, fmt.Errorf("approval timeout after %s", timeout)
case <-ctx.Done():
return nil, ctx.Err()
}
}
// ---- 配置控制面(core NATS request-reply + broadcast---- // ---- 配置控制面(core NATS request-reply + broadcast----
// RequestConfig 向控制面(Gateway)请求某 kind 当前激活配置(chat/embedding)。 // RequestConfig 向控制面(Gateway)请求某 kind 当前激活配置(chat/embedding)。
@@ -399,38 +465,69 @@ func (b *Bus) SubscribeConfigUpdated(kind string, onUpdate func(*contract.ModelC
// TaskHandler 处理一个消费到的任务。 // TaskHandler 处理一个消费到的任务。
type TaskHandler func(ctx context.Context, t *contract.Task) error type TaskHandler func(ctx context.Context, t *contract.Task) error
// taskConcurrency 返回任务并发处理上限(env DISPATCHER_CONCURRENCY,默认 8)。
func taskConcurrency() int {
if v := os.Getenv("DISPATCHER_CONCURRENCY"); v != "" {
if n, err := strconv.Atoi(v); err == nil && n > 0 {
return n
}
}
return 8
}
// ConsumeTasks 在持久消费者上消费任务,队列组内负载均衡。 // ConsumeTasks 在持久消费者上消费任务,队列组内负载均衡。
// 返回的 stop 函数用于优雅停止消费 // 每个任务分发到独立 worker goroutine 并发执行——一个慢任务/HITL 待审不再阻塞后续任务
// 并发上限由信号量 + 消费者 MaxAckPending 双重约束(背压)。返回的 stop 用于优雅停止消费。
func (b *Bus) ConsumeTasks(ctx context.Context, h TaskHandler) (stop func(), err error) { func (b *Bus) ConsumeTasks(ctx context.Context, h TaskHandler) (stop func(), err error) {
concurrency := taskConcurrency()
cons, err := b.js.CreateOrUpdateConsumer(ctx, contract.StreamTasks, jetstream.ConsumerConfig{ cons, err := b.js.CreateOrUpdateConsumer(ctx, contract.StreamTasks, jetstream.ConsumerConfig{
Durable: contract.ConsumerDurable, Durable: contract.ConsumerDurable,
AckPolicy: jetstream.AckExplicitPolicy, AckPolicy: jetstream.AckExplicitPolicy,
FilterSubject: contract.SubjectTasksAll, FilterSubject: contract.SubjectTasksAll,
// HITL:审批节点会让 Handle 阻塞等人工决定(最长约 5 分钟),
// AckWait 必须覆盖「审批等待 + 图执行」总时长,否则消息在途未 ack 会被重投成重复任务。
AckWait: 15 * time.Minute,
// 在途未 ack 上限 = 并发上限:服务端不会下发超过本节点同时能处理的量(背压)。
MaxAckPending: concurrency,
}) })
if err != nil { if err != nil {
return nil, fmt.Errorf("create consumer: %w", err) return nil, fmt.Errorf("create consumer: %w", err)
} }
sem := make(chan struct{}, concurrency) // 限并发:最多 N 个任务同时执行
cc, err := cons.Consume(func(msg jetstream.Msg) { cc, err := cons.Consume(func(msg jetstream.Msg) {
t, err := contract.Unmarshal(msg.Data()) t, err := contract.Unmarshal(msg.Data())
if err != nil { if err != nil {
_ = msg.Term() // 脏数据,丢弃不重投 _ = msg.Term() // 脏数据,丢弃不重投
return return
} }
// 从消息头还原上游链路,开消费 span(成为 gateway 发布 span 的子节点)。 // 背压:并发已满则在此等空位;正在关停则留消息不 ack(稍后重投)。
mctx := extractTrace(ctx, nats.Header(msg.Headers())) select {
mctx, span := tracer().Start(mctx, "nats.consume task", case sem <- struct{}{}:
trace.WithSpanKind(trace.SpanKindConsumer), case <-ctx.Done():
trace.WithAttributes(attribute.String("sundynix.task_id", t.ID)))
herr := h(mctx, t)
if herr != nil {
span.RecordError(herr)
}
span.End()
if herr != nil {
_ = msg.NakWithDelay(time.Second) // 处理失败,延迟重投
return return
} }
_ = msg.Ack() go func() {
defer func() {
<-sem // 释放并发额度
if r := recover(); r != nil {
// 任务处理 panic:丢弃不重投(避免崩溃循环),记录后继续。
log.Printf("[bus] task %s handler panic: %v", t.ID, r)
_ = msg.Term()
}
}()
// 从消息头还原上游链路,开消费 span(成为 gateway 发布 span 的子节点)。
mctx := extractTrace(ctx, nats.Header(msg.Headers()))
mctx, span := tracer().Start(mctx, "nats.consume task",
trace.WithSpanKind(trace.SpanKindConsumer),
trace.WithAttributes(attribute.String("sundynix.task_id", t.ID)))
defer span.End()
if herr := h(mctx, t); herr != nil {
span.RecordError(herr)
_ = msg.NakWithDelay(time.Second) // 处理失败,延迟重投
return
}
_ = msg.Ack()
}()
}) })
if err != nil { if err != nil {
return nil, fmt.Errorf("consume: %w", err) return nil, fmt.Errorf("consume: %w", err)
+112
View File
@@ -194,3 +194,115 @@ func TestTokenStreamRoundTrip(t *testing.T) {
t.Fatal("timeout: 未收到流结束信号") t.Fatal("timeout: 未收到流结束信号")
} }
} }
// TestConcurrentConsume 验证并发消费:一个长时间阻塞的任务(模拟 HITL 待审)
// 不应阻塞后续任务——旧的串行消费下本测试会超时失败。
func TestConcurrentConsume(t *testing.T) {
url := startEmbeddedNATS(t)
gw, err := bus.Connect(url)
if err != nil {
t.Fatalf("gateway connect: %v", err)
}
defer gw.Close()
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
defer cancel()
if err := gw.EnsureTaskStream(ctx); err != nil {
t.Fatalf("ensure stream: %v", err)
}
dp, err := bus.Connect(url)
if err != nil {
t.Fatalf("dispatcher connect: %v", err)
}
defer dp.Close()
blockA := make(chan struct{})
doneB := make(chan string, 1)
stop, err := dp.ConsumeTasks(ctx, func(_ context.Context, task *contract.Task) error {
switch task.ID {
case "A":
<-blockA // 模拟审批长阻塞
case "B":
doneB <- task.ID
}
return nil
})
if err != nil {
t.Fatalf("consume: %v", err)
}
defer stop()
defer close(blockA)
// 先发 A 并等它被消费、卡在 handler 里;再发 B。
if _, err := gw.PublishTask(ctx, &contract.Task{ID: "A", Graph: json.RawMessage(`{}`)}); err != nil {
t.Fatalf("publish A: %v", err)
}
time.Sleep(400 * time.Millisecond)
if _, err := gw.PublishTask(ctx, &contract.Task{ID: "B", Graph: json.RawMessage(`{}`)}); err != nil {
t.Fatalf("publish B: %v", err)
}
select {
case id := <-doneB:
if id != "B" {
t.Fatalf("unexpected task done: %q", id)
}
t.Log("✓ 并发消费生效:A 仍阻塞时 B 已完成")
case <-time.After(4 * time.Second):
t.Fatal("B 被阻塞的 A 卡住——并发消费未生效(仍是串行)")
}
}
// TestConcurrentToolServe 验证单实例工具服务并发:一个慢工具调用不阻塞另一个快调用。
// 旧的串行 ServeTool 下本测试会超时失败。
func TestConcurrentToolServe(t *testing.T) {
url := startEmbeddedNATS(t)
srv, err := bus.Connect(url)
if err != nil {
t.Fatalf("server connect: %v", err)
}
defer srv.Close()
blockSlow := make(chan struct{})
unsub, err := srv.ServeTool(contract.SubjectToolsGoAll, contract.QueueToolsGo,
func(_ context.Context, call *contract.ToolCall) *contract.ToolResult {
if call.Tool == "slow" {
<-blockSlow // 模拟慢工具(如 RAG 嵌入)
}
return &contract.ToolResult{OK: true, Content: call.Tool}
})
if err != nil {
t.Fatalf("serve tool: %v", err)
}
defer unsub()
defer close(blockSlow)
cli, err := bus.Connect(url)
if err != nil {
t.Fatalf("client connect: %v", err)
}
defer cli.Close()
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
// 先发 slow(会阻塞在 handler 里),再发 fast——fast 应在 slow 仍阻塞时返回。
go func() {
_, _ = cli.CallTool(ctx, contract.ToolSubjectGo("slow"), &contract.ToolCall{Tool: "slow"})
}()
time.Sleep(300 * time.Millisecond)
fctx, fcancel := context.WithTimeout(ctx, 3*time.Second)
defer fcancel()
res, err := cli.CallTool(fctx, contract.ToolSubjectGo("fast"), &contract.ToolCall{Tool: "fast"})
if err != nil {
t.Fatalf("fast 工具被慢工具阻塞——单实例并发未生效: %v", err)
}
if res.Content != "fast" {
t.Fatalf("unexpected result: %q", res.Content)
}
t.Log("✓ 单实例工具并发生效:slow 仍阻塞时 fast 已返回")
}
+22 -2
View File
@@ -33,26 +33,46 @@ const (
// core NATS pub-sub(状态是幂等覆盖,丢一条由下一条纠正,无需持久化)。 // core NATS pub-sub(状态是幂等覆盖,丢一条由下一条纠正,无需持久化)。
// 注意:必须在 sundynix.tasks.> 之外,否则会被任务流捕获成"幽灵任务"自我放大。 // 注意:必须在 sundynix.tasks.> 之外,否则会被任务流捕获成"幽灵任务"自我放大。
SubjectTaskStatus = "sundynix.status.task" SubjectTaskStatus = "sundynix.status.task"
// 人工审批(HITL)决定回传前缀:实际 sundynix.approval.<task_id>。
// core NATS pub-subUI 点批准/拒绝 → 网关发到此 → dispatcher 解除审批节点的阻塞。
SubjectApproval = "sundynix.approval"
) )
// 任务生命周期状态机:submitted(网关建任务)→ runningdispatcher 开跑) // 任务生命周期状态机:submitted(网关建任务)→ runningdispatcher 开跑)
// → done / failed / timeoutdispatcher 收尾)。 // → done / failed / timeoutdispatcher 收尾)。
// HITL:执行到审批节点 → waiting(等人工决定)→ 批准回 running / 拒绝→rejected。
const ( const (
TaskSubmitted = "submitted" TaskSubmitted = "submitted"
TaskRunning = "running" TaskRunning = "running"
TaskDone = "done" TaskDone = "done"
TaskFailed = "failed" TaskFailed = "failed"
TaskTimeout = "timeout" TaskTimeout = "timeout"
TaskWaiting = "waiting" // 等待人工审批
TaskRejected = "rejected" // 人工拒绝(或审批超时,安全默认拒绝)
) )
// TaskStatusEvent 是一次任务状态流转事件(经 SubjectTaskStatus 回流给网关)。 // TaskStatusEvent 是一次任务状态流转事件(经 SubjectTaskStatus 回流给网关)。
type TaskStatusEvent struct { type TaskStatusEvent struct {
TaskID string `json:"task_id"` TaskID string `json:"task_id"`
Status string `json:"status"` // running / done / failed / timeout Status string `json:"status"` // running / done / failed / timeout / waiting / rejected
Detail string `json:"detail,omitempty"` // 失败原因等 Detail string `json:"detail,omitempty"` // 失败原因 / 审批摘要
TS int64 `json:"ts"` // unix 毫秒 TS int64 `json:"ts"` // unix 毫秒
} }
// ApprovalSubject 返回某任务的人工审批决定回传主题。
func ApprovalSubject(id string) string { return SubjectApproval + "." + id }
// ApprovalDecision 是一次人工审批的结果(UI → 网关 → dispatcher)。
type ApprovalDecision struct {
TaskID string `json:"task_id"`
Node string `json:"node,omitempty"` // 审批节点 id(可空:按 task 维度兜底匹配)
Approved bool `json:"approved"` // true=批准放行,false=拒绝中止
Note string `json:"note,omitempty"` // 审批人备注
By string `json:"by,omitempty"` // 审批人(用户 id
TS int64 `json:"ts"` // unix 毫秒
}
const ( const (
// MetaUserID 是 Task.Meta 中承载已登录用户标识的键(用于偏好记忆召回)。 // MetaUserID 是 Task.Meta 中承载已登录用户标识的键(用于偏好记忆召回)。
MetaUserID = "user_id" MetaUserID = "user_id"