diff --git a/.env.example b/.env.example index 7cd85bf..3f43b5b 100644 --- a/.env.example +++ b/.env.example @@ -29,6 +29,13 @@ LLM_BASE_URL= # 要操作本仓时显式写成仓根。pnpm dev:api 在沙箱存在时默认用沙箱。 GIT_REPO= +# 插件目录。每个子目录放一个与目录同名的二进制。未设则不拉插件进程。 +# 示例:编译 plugin/redact 后设 PLUGIN_DIR=<仓根>/plugin +# PLUGIN_DIR= + +# 单次插件 RPC 超时(默认 10s) +# PLUGIN_RPC_TIMEOUT=10s + # 本机 Codex CLI 可执行文件(默认 codex)。未安装时主服务仍可启动,/codex/status 会说明原因。 CODEX_BIN=codex diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 1f12e76..685faaf 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -24,15 +24,20 @@ jobs: - uses: actions/setup-go@v6 with: go-version-file: server/go.mod - cache-dependency-path: server/go.sum + cache-dependency-path: | + server/go.sum + plugin/example/go.sum + plugin/redact/go.sum - name: Lint run: | - test -z "$(gofmt -l .)" || { echo "gofmt needed:"; gofmt -l .; exit 1; } + test -z "$(gofmt -l . ../plugin)" || { echo "gofmt needed:"; gofmt -l . ../plugin; exit 1; } go vet ./... - name: Test - run: go test ./... + run: | + go test ./... + (cd ../plugin/redact && go test .) - name: Codex coverage run: | diff --git a/.gitignore b/.gitignore index b5e6d1c..3a8676f 100644 --- a/.gitignore +++ b/.gitignore @@ -5,6 +5,8 @@ *.db-shm *.db-wal tmp/ +/plugin/*/* +!/plugin/*/*.* .DS_Store node_modules .pnpm-store diff --git a/AGENTS.md b/AGENTS.md index 70039c8..2117e90 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -12,7 +12,8 @@ Agent Loop 已闭环:用户发文本、装上下文、调模型、产出文字 - 本机 Codex app-server 生命周期、内存排队/问票/SSE 放在 `server/internal/codex`。不新增 Codex 业务表;凡官方 API 能读到的都不入库。 - Markdown 记忆(热层目录+专题)与 context message 索引(冷层按工作区 FTS)放在 `server/internal/agent/memory`;不放 `pkg/memory`。memory 不 import 父包 `internal/agent`,不定义 Tool。 - 具体工具定义放在 `server/internal/agent/tools`。工具名、入参/出参、schema、权限和编排都在本包;Execute 若要调外部能力,只通过 `Ports` 里的接口。Runtime `New` 时由 `cmd/server` 注入 `Ports` 的具体实现,再 `Register`。每个工具只定义入参/出参结构体,执行用 `encoding/json`,schema 从类型推断。`tools` 可 import `memory`,不 import 父包 `internal/agent`。 -- Agent 通用无状态逻辑放在 `server/pkg/agent`:类型、token 统计、提示词、上下文、Tool 抽象(不含具体工具定义)、Agent 配置、模型调用。 +- Agent 通用无状态逻辑放在 `server/pkg/agent`:类型、token 统计、提示词、上下文、Tool 抽象(不含具体工具定义)、Agent 配置、模型调用。六个口的信封放在 `pkg/agent/seam`。 +- 插件 SDK 与 proto 放在 `server/pkg/plugin`;宿主放在 `server/internal/pluginhost`。两者都不进 `pkg/agent`,也不知道主循环内部状态机。可装载的插件放在仓根 `plugin/`(子目录名即插件名)。`plugin/example` 是作者拷贝模板:不订阅、不改正文、不换向、不登记方法。不进 `pkg/plugin`。插件共享参数用 `PluginContext`,不进模型、不复用 Hidden。 - Git CLI 操作放在 `server/pkg/git`:无状态,不写产品流程;Handler 直接调用。不进 `pkg/agent`。 - Codex 协议与领域类型放在 `server/pkg/codex`:看板的子模块,JSONL 客户端给 `internal/codex` 调用;不查库、不 spawn CLI。不进 `pkg/agent`。 - 进程内事件总线放在 `server/internal/events`。 @@ -42,4 +43,10 @@ Agent Loop 已闭环:用户发文本、装上下文、调模型、产出文字 - 前端三层不得反依赖:`core` 不依赖 React / Next / DOM / `process.env`;`ui` 不依赖 `core`;`views` 不 import `next/*`;`apps/web` 只做路由与平台装配。 - 前端按业务域拆模块,不要 `src/`:`core` / `views` 用同名域目录(现有 `chat` / `git` / `codex`);`ui` 只用 `components` / `lib` / `styles`。新业务再建目录,不预建空文件夹。 +## 注释规则 + +- 每个函数、方法正上方必须有一行注释,写清它做什么。Go 用文档注释(`// Name ...`),TypeScript 同等要求。不要用注释复述函数名或参数列表。 +- 结构体 / 接口里,单看字段名读不懂含义或取值约定的属性必须加字段注释;`ID`、`Name`、`Content` 这类自明字段不必硬加。 +- 改现有代码时顺手补上缺的注释;sqlc / proto 生成文件不要手改。 + 当需求变更没有明显的代码归属时,先依据 `docs/architecture.md` 对其分类,再开始编写代码。 diff --git a/docs/architecture.md b/docs/architecture.md index e2a9fc2..af8b373 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -55,6 +55,8 @@ CodeDock/ │ ├── ui/ # 无业务语义;components / lib / styles,不要 src/ │ └── views/ # 组合层;按业务域拆(现有 chat/ git/ codex/),不要 src/ ├── docs/ +├── plugin/example/ # 插件拷贝模板;不改正文、不换向 +├── plugin/redact/ # 脱敏插件;PLUGIN_DIR 指到 plugin/ ├── data/ # 运行时文件(sqlite 等),gitignore ├── server/ │ ├── cmd/server/ # 服务启动、配置、Router 和依赖装配 @@ -66,12 +68,15 @@ CodeDock/ │ │ │ └── tools/ # 具体工具定义:ping、memory_*、编码八工具、plan_* │ │ ├── codex/ # 本机 app-server 生命周期与内存排队/问票/SSE │ │ ├── events/ # 进程内事件总线 +│ │ ├── pluginhost/ # go-plugin 宿主:拉进程、Dispatch、Host 白名单 │ │ ├── config/ │ │ ├── logger/ │ │ ├── errors/ │ │ └── util/ │ ├── pkg/ │ │ ├── agent/ # 全部通用无状态逻辑,含模型调用与 Tool 抽象 +│ │ │ └── seam/ # Envelope / Dispatcher / 六个口的类型常量 +│ │ ├── plugin/ # 插件 SDK 与 proto;作者只 import 这个包 │ │ ├── git/ # 无状态 Git CLI 操作,供 Handler 直接调用 │ │ ├── codex/ # 看板的 Codex 子模块:协议客户端与领域类型 │ │ └── db/ # Client 与 sqlc 生成代码 @@ -90,6 +95,7 @@ cmd/server -> internal/codex -> internal/agent -> internal/events + -> internal/pluginhost -> pkg/db internal/handler @@ -127,10 +133,26 @@ internal/agent/tools -> pkg/db/sqlite.Queries 不 import 父包 internal/agent +internal/pluginhost + -> pkg/plugin + -> pkg/agent / pkg/agent/seam / pkg/agent/tool + -> internal/agent/memory + -> pkg/db/sqlite.Queries + -> internal/events + 不 import 父包 internal/agent + +pkg/plugin + -> pkg/agent / pkg/agent/seam / pkg/agent/tool + -> pkg/plugin/proto + 不依赖 handler、internal、sqlc + 插件作者只 import 这个包 + pkg/agent 不依赖 handler、internal、sqlc 不持有包级状态,不查库 + 不知道 gRPC / go-plugin Tool 包只含接口、Registry、Dispatch,不含具体工具定义 + seam 是叶子包:信封与六个口,agent 与 tool 都能 import pkg/git 不依赖 handler、internal、sqlc @@ -217,6 +239,16 @@ Handler 直接依赖 `*sqlite.Queries`,不经过 Store 接口。Git 带 `sessi 每个工具只定义入参/出参结构体;执行用 `encoding/json`,给模型的 schema 由 `jsonschema.For` 从类型推断。发给模型的是注册表全量工具;`Profile.Tools.Names` 是可执行绑定。模式规则由 `Build` 注入一条 developer 消息;发给网关时紧跟底座 system,不改底座正文。审批流水线:工具默认+参数校验 → 本 Agent `Names`(未绑定 deny,yolo / 已批准都不能抬)→ Agent `Effects` → `approval`(manual/auto/yolo);每层只审上一层的 `ask`。一批待批工具对应一条审批,一次提交审完再流转。不 import 父包 `internal/agent`。测试用 Tool 可留在测试文件。 +### 插件 + +主循环在六个口把当前数据递给 `seam.Dispatcher`:`agent/input`、`agent/pre-step`、`agent/request`、`llm/stream`、`tools/pre-execute`、`tools/post-execute`。没有 Dispatcher 时原样通过。作者侧每个口是一对入参/回包结构体(`OnAgentInput` 的 `AgentInput` / `AgentInputResult` 等),换向是回包字段,不解信封。 + +多个插件按子目录名排序依次改同一份载荷。回包类型不变则继续;类型变成 `input/handled` / `run/blocked` / `tools/denied` / `tools/ask` 则换方向。同一条链上问过的插件记入 `Seen`,不会再问自己。插件之间的参数走 `PluginContext`(信封 `Context`),由宿主按会话/Run 暂存,不进模型、不进消息表。 + +插件跑在独立进程里,经 go-plugin gRPC 通信。宿主白名单:`Emit`(不能发六个口的同名事件)、`RegisterMethod`、记忆读写、`Complete`、`AppendNotice`。超时或进程挂了:拦截口按否决,只改数据的口保留原样。`assistant.delta` 不发给插件。已批准的工具不再拦一次。换插件二进制要重启服务。未设 `PLUGIN_DIR` 不拉进程。 + +详见 [plugin.md](plugin.md)。 + ### `pkg/git` 无状态 Git CLI:`Open` / `Status`(`SiteState` 整局)/ Diff / 图 / 暂存提交 / reset / revert / 推拉 / remote / 分支 / worktree / `stash create` 副本 / 冲突读写。不进 `pkg/agent`,不写 HTTP 或产品流程。Workspace / Branch / Undo / 说明 / Agent 快照的产品组合在 Handler。 @@ -296,7 +328,7 @@ Worker ## 配置 -`LLM_PROVIDER`(`openai` | `fake`,默认 `fake`)、`LLM_MODEL`、`LLM_API_KEY`、`LLM_BASE_URL`。`GIT_REPO` 指向本地仓库根,未设则用进程 cwd(不向上找 `.git`)。未设 `DB_DSN` 时 SQLite 写仓根 `data/codedock.db`,不写 `server/`。`CODEX_BIN` 为本机 Codex CLI(默认 `codex`)。Handler 创建 Run 时写入 `RunConfigSnapshot`,后续 Turn 只读快照。 +`LLM_PROVIDER`(`openai` | `fake`,默认 `fake`)、`LLM_MODEL`、`LLM_API_KEY`、`LLM_BASE_URL`。`GIT_REPO` 指向本地仓库根,未设则用进程 cwd(不向上找 `.git`)。未设 `DB_DSN` 时 SQLite 写仓根 `data/codedock.db`,不写 `server/`。`PLUGIN_DIR` 指向插件根目录,未设则不拉插件进程;`PLUGIN_RPC_TIMEOUT` 默认 `10s`。`CODEX_BIN` 为本机 Codex CLI(默认 `codex`)。Handler 创建 Run 时写入 `RunConfigSnapshot`,后续 Turn 只读快照。 HTTP 出站领域对象使用 snake_case JSON。Router 只对本地回环 Origin 放行 CORS,便于本机 Web 直连 `:8080`。Web 用 `NEXT_PUBLIC_API_BASE`(默认 `http://localhost:8080`)和 `NEXT_PUBLIC_USER_ID`(默认 `local`)。 diff --git a/docs/plugin.md b/docs/plugin.md new file mode 100644 index 0000000..359a1d2 --- /dev/null +++ b/docs/plugin.md @@ -0,0 +1,76 @@ +# 编写 CodeDock 插件 + +外来程序可以在对话的六个口上改数据或换方向,也可以给模型登记方法。插件跑在独立进程里,通过 gRPC 和宿主说话。 + +没有 `PLUGIN_DIR` 时,服务不拉任何插件,主循环与现在完全一样。 + +## 从 example 模板开始 + +`plugin/example` 是作者拷贝的模板:不订阅、不改正文、不换向、不登记方法。装进 `PLUGIN_DIR` 也不会改对话。复制本目录,改 `go.mod` 模块名、`Manifest.Name` 和策略,再在 `Bootstrap` 里填 `Subscriptions`。作者只 import `codedock/pkg/plugin`。`go.mod` 用 `replace` 指到本仓 `server/`。 + +可装载的插件放在仓根 `plugin/`,子目录名即插件名。`plugin/redact` 是脱敏插件:不拦工具,只在 input / request / post-execute 把秘密换成占位符。六个口要哪些字段,以 `codedock/pkg/plugin` 里的结构体为准(`AgentInput`、`AgentInputResult` 等),跳进类型就能看到。 + +本地编译后: + +```sh +(cd plugin/example && go build -o example .) +(cd plugin/redact && go build -o redact .) +PLUGIN_DIR=$PWD/plugin pnpm run dev +``` + +正文或工具回包含 `AKIA…` / `API_KEY=…` 时,落库和发给模型的是占位符。 + +目录约定: + +```text +plugin/ # PLUGIN_DIR 指这里 + example/ + example # 拷贝模板,不改对话 + redact/ + redact # 与子目录同名的二进制 +``` + +## 六个口 + +每个口是一对 SDK 结构体,实现对应方法即可。没实现的口原样通过。字段含义写在结构体上。 + +| 口 | 方法 | 入参 | 回包 | +| --- | --- | --- | --- | +| `agent/input` | `OnAgentInput` | `AgentInput`(正文、模式) | `AgentInputResult`;`Handle()` 不建 Run | +| `agent/pre-step` | `OnAgentPreStep` | `AgentPreStep`(系统提示、隐藏消息) | `AgentPreStepResult`;`Block()` 取消本轮 | +| `agent/request` | `OnAgentRequest` | `AgentRequest` | `AgentRequestResult`;只能 `Reply()` | +| `llm/stream` | `OnLLMStream` | `LLMStream`(请求头、请求体) | `LLMStreamResult`;只能 `Reply()`;fake 模型不插 | +| `tools/pre-execute` | `OnToolPreExecute` | `ToolPreExecute`(工具调用) | `ToolPreExecuteResult`;`Deny()` 当失败;`AskApproval()` 进审批 | +| `tools/post-execute` | `OnToolPostExecute` | `ToolPostExecute`(工具结果) | `ToolPostExecuteResult`;只能 `Reply()` | + +账本通知走 `OnLedgerNotify`,没有换向。多个插件按子目录名排序,后一个看到前一个改完的结果。同一条链上问过的插件记在 `Seen` 里,不会再问自己。 + +跨口、跨插件传参数用各口上的 `Context`(`Set("example.xxx", v)` / `Get`)。这是宿主暂存的 JSON 对象,不进模型、不进消息表、不换向。键建议 `插件名.字段`。`agent/input` 时按会话挂;建 Run 后迁到该 Run。`Handle()` 或 Run 终态会清掉。上限 8KB,超了保留上一份。进程重启即丢。不要把协议塞进 `Hidden`。 + +隐藏提示用 `sdk.HiddenText("...")` 加进 `AgentPreStep.Hidden`。已批准但还没执行的工具不再走 `OnToolPreExecute`。流式增量 `assistant.delta` 不发给插件。 + +## Host 白名单 + +`Bootstrap` 拿到的 `Host` 只能做这些事: + +- `Emit`:另发一条与当前口无关的事件。不能发六个口的同名事件。 +- `RegisterMethod`:给模型加方法。不能覆盖 `ping`、`memory_read`、`memory_write`、`memory_search`。 +- `MemoryGet` / `MemoryUpsert`:按会话读写一篇专题记忆。 +- `Complete`:自己打一次模型,不进当前助手流。 +- `AppendNotice`:写一条用户看得见的 system 消息(本期只落库,不实时推送)。 + +## 失败与超时 + +`PLUGIN_RPC_TIMEOUT` 默认 `10s`。超时或进程挂了: + +- `agent/input`、`agent/pre-step`、`tools/pre-execute` 按否决处理 +- `agent/request`、`llm/stream`、`tools/post-execute` 保留原数据 + +进程崩了不会自动拉起。换二进制要重启服务。 + +## 环境变量 + +| 变量 | 含义 | +| --- | --- | +| `PLUGIN_DIR` | 插件根目录。每个子目录一个常驻进程。未设则不加载。 | +| `PLUGIN_RPC_TIMEOUT` | 单次 RPC 超时,如 `10s`。 | diff --git a/packages/core/chat/reducer.test.ts b/packages/core/chat/reducer.test.ts index 0ac92d3..831506b 100644 --- a/packages/core/chat/reducer.test.ts +++ b/packages/core/chat/reducer.test.ts @@ -10,6 +10,7 @@ import { applyEvent, applyLocalCancel, applyOptimisticUser, + dropOptimisticUser, decisionsForApproval, emptyState, hydrate, @@ -508,6 +509,13 @@ test("applyLocalCancel clears the executing run", () => { assert.equal(optimistic.queued, false); }); +test("dropOptimisticUser removes a handled local bubble", () => { + let state = applyOptimisticUser(emptyState(), { runId: "local:1", text: "/skip" }); + assert.equal(state.items.length, 1); + state = dropOptimisticUser(state, "local:1"); + assert.equal(state.items.length, 0); +}); + test("optimistic user is replaced when run.created arrives", () => { let state = applyOptimisticUser(emptyState(), { runId: "r1", text: "hi" }); assert.equal(state.items[0]?.kind === "user" && state.items[0].messageId, "pending:r1"); diff --git a/packages/core/chat/types.ts b/packages/core/chat/types.ts index e051c55..5c97ac8 100644 --- a/packages/core/chat/types.ts +++ b/packages/core/chat/types.ts @@ -237,7 +237,8 @@ export interface StartRunRequest { export interface StartRunResponse { session_id: string; - run_id: string; + run_id?: string; + handled?: boolean; } export interface DecideApprovalRequest { diff --git a/packages/views/chat/chat-page.tsx b/packages/views/chat/chat-page.tsx index db37354..ecdc2ac 100644 --- a/packages/views/chat/chat-page.tsx +++ b/packages/views/chat/chat-page.tsx @@ -3,7 +3,7 @@ import type { ApprovalMode, Session, TimelineItem, WorkMode } from "@codedock/core/chat"; import type { Session as CodexSession } from "@codedock/core/codex"; import { Button } from "@codedock/ui"; -import { useMemo, useState, useEffect, type ReactNode } from "react"; +import { useEffect, useMemo, useState, type ReactNode } from "react"; import { CodexPane } from "../codex/codex-pane.tsx"; import { useCodexSessionList } from "../codex/hooks/use-session-list.ts"; @@ -37,6 +37,7 @@ export type ChatPageProps = { headerActions?: ReactNode; }; +// ChatPage 组合会话侧栏、时间线与输入条。 export function ChatPage({ sessionId, engine, @@ -62,6 +63,8 @@ export function ChatPage({ const [composerError, setComposerError] = useState(null); const [workspaceDraft, setWorkspaceDraft] = useState(""); const [pickingWorkspace, setPickingWorkspace] = useState(false); + + // 上次目录只在本机 localStorage,等 hydration 后再读,避免 SSR 文本对不上。 useEffect(() => { setWorkspaceDraft(readLastWorkspace()); }, []); @@ -308,6 +311,7 @@ export function ChatPage({ ); } +// mergeSessions 把 Agent 与 Codex 会话按更新时间合成侧栏列表,同引擎同 ID 只留更新的一条。 function mergeSessions(agent: Session[], codex: CodexSession[]): SidebarSession[] { const mapped: SidebarSession[] = [ ...agent.map((session) => ({ ...session, engine: "agent" as const })), @@ -324,6 +328,7 @@ function mergeSessions(agent: Session[], codex: CodexSession[]): SidebarSession[ return [...seen.values()].sort((left, right) => (left.updated_at < right.updated_at ? 1 : -1)); } +// asSidebarSession 把 Codex 会话收成侧栏条目,目录用 cwd,标题优先 title。 function asSidebarSession(session: CodexSession): SidebarSession { return { id: session.id, @@ -341,6 +346,7 @@ function asSidebarSession(session: CodexSession): SidebarSession { }; } +// stampToIso 把秒或毫秒时间戳收成 ISO 字符串,无效值用纪元。 function stampToIso(value?: number): string { if (!value) { return new Date(0).toISOString(); diff --git a/packages/views/chat/hooks/use-session-timeline.ts b/packages/views/chat/hooks/use-session-timeline.ts index d9addac..5b4353b 100644 --- a/packages/views/chat/hooks/use-session-timeline.ts +++ b/packages/views/chat/hooks/use-session-timeline.ts @@ -262,7 +262,17 @@ export function useSessionTimeline(sessionId: string | undefined) { return next; }); try { - await client.startRun(sessionId, { content, mode, approval }); + const started = await client.startRun(sessionId, { content, mode, approval }); + if (started.handled) { + setState((current) => { + if (sessionRef.current !== sessionId) { + return current; + } + const next = dropOptimisticUser(current, pendingRunId); + cacheSet(sessionId, next); + return next; + }); + } setRecoverableRunId(null); setError(null); } catch (err) { diff --git a/plugin/example/go.mod b/plugin/example/go.mod new file mode 100644 index 0000000..2bf7a33 --- /dev/null +++ b/plugin/example/go.mod @@ -0,0 +1,25 @@ +module codedock-example + +go 1.26.5 + +require codedock v0.0.0 + +require ( + github.com/fatih/color v1.13.0 // indirect + github.com/golang/protobuf v1.5.4 // indirect + github.com/google/uuid v1.6.0 // indirect + github.com/hashicorp/go-hclog v1.6.3 // indirect + github.com/hashicorp/go-plugin v1.8.0 // indirect + github.com/hashicorp/yamux v0.1.2 // indirect + github.com/mattn/go-colorable v0.1.12 // indirect + github.com/mattn/go-isatty v0.0.24 // indirect + github.com/oklog/run v1.1.0 // indirect + golang.org/x/net v0.58.0 // indirect + golang.org/x/sys v0.47.0 // indirect + golang.org/x/text v0.41.0 // indirect + google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa // indirect + google.golang.org/grpc v1.83.2 // indirect + google.golang.org/protobuf v1.36.12 // indirect +) + +replace codedock => ../../server diff --git a/plugin/example/go.sum b/plugin/example/go.sum new file mode 100644 index 0000000..724238c --- /dev/null +++ b/plugin/example/go.sum @@ -0,0 +1,75 @@ +github.com/bufbuild/protocompile v0.14.1 h1:iA73zAf/fyljNjQKwYzUHD6AD4R8KMasmwa/FBatYVw= +github.com/bufbuild/protocompile v0.14.1/go.mod h1:ppVdAIhbr2H8asPk6k4pY7t9zB1OU5DoEw9xY/FUi1c= +github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs= +github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= +github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= +github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/fatih/color v1.13.0 h1:8LOYc1KYPPmyKMuN8QV2DNRWNbLo6LZ0iLs8+mlH53w= +github.com/fatih/color v1.13.0/go.mod h1:kLAiJbzzSOZDVNGyDpeOxJ47H46qBXwg5ILebYFFOfk= +github.com/go-logr/logr v1.4.3 h1:CjnDlHq8ikf6E492q6eKboGOC0T8CDaOvkHCIg8idEI= +github.com/go-logr/logr v1.4.3/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY= +github.com/go-logr/stdr v1.2.2 h1:hSWxHoqTgW2S2qGc0LTAI563KZ5YKYRhT3MFKZMbjag= +github.com/go-logr/stdr v1.2.2/go.mod h1:mMo/vtBO5dYbehREoey6XUKy/eSumjCCveDpRre4VKE= +github.com/golang/protobuf v1.5.4 h1:i7eJL8qZTpSEXOPTxNKhASYpMn+8e5Q6AdndVa1dWek= +github.com/golang/protobuf v1.5.4/go.mod h1:lnTiLA8Wa4RWRcIUkrtSVa5nRhsEGBg48fD6rSs7xps= +github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= +github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= +github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= +github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +github.com/hashicorp/go-hclog v1.6.3 h1:Qr2kF+eVWjTiYmU7Y31tYlP1h0q/X3Nl3tPGdaB11/k= +github.com/hashicorp/go-hclog v1.6.3/go.mod h1:W4Qnvbt70Wk/zYJryRzDRU/4r0kIg0PVHBcfoyhpF5M= +github.com/hashicorp/go-plugin v1.8.0 h1:ie8S6RRY8RvB2usYZv+AAZ/wBvx2AU5p5QeP5j/FORs= +github.com/hashicorp/go-plugin v1.8.0/go.mod h1:BExt6KEaIYx804z8k4gRzRLEvxKVb+kn0NMcihqOqb8= +github.com/hashicorp/yamux v0.1.2 h1:XtB8kyFOyHXYVFnwT5C3+Bdo8gArse7j2AQ0DA0Uey8= +github.com/hashicorp/yamux v0.1.2/go.mod h1:C+zze2n6e/7wshOZep2A70/aQU6QBRWJO/G6FT1wIns= +github.com/jhump/protoreflect v1.17.0 h1:qOEr613fac2lOuTgWN4tPAtLL7fUSbuJL5X5XumQh94= +github.com/jhump/protoreflect v1.17.0/go.mod h1:h9+vUUL38jiBzck8ck+6G/aeMX8Z4QUY/NiJPwPNi+8= +github.com/mattn/go-colorable v0.1.9/go.mod h1:u6P/XSegPjTcexA+o6vUJrdnUu04hMope9wVRipJSqc= +github.com/mattn/go-colorable v0.1.12 h1:jF+Du6AlPIjs2BiUiQlKOX0rt3SujHxPnksPKZbaA40= +github.com/mattn/go-colorable v0.1.12/go.mod h1:u5H1YNBxpqRaxsYJYSkiCWKzEfiAb1Gb520KVy5xxl4= +github.com/mattn/go-isatty v0.0.12/go.mod h1:cbi8OIDigv2wuxKPP5vlRcQ1OAZbq2CE4Kysco4FUpU= +github.com/mattn/go-isatty v0.0.14/go.mod h1:7GGIvUiUoEMVVmxf/4nioHXj79iQHKdU27kJ6hsGG94= +github.com/mattn/go-isatty v0.0.24 h1:tGZZoVgT/KiqK1c8ocVLeDS8BSWMRd47J3Lbz7vsReI= +github.com/mattn/go-isatty v0.0.24/go.mod h1:nMCL3Zebbrt45jsMDgnfIwz6ydEQApk5oEI3HqDio6A= +github.com/oklog/run v1.1.0 h1:GEenZ1cK0+q0+wsJew9qUg/DyD8k3JzYsZAi5gYi2mA= +github.com/oklog/run v1.1.0/go.mod h1:sVPdnTZT1zYwAJeCMu2Th4T21pA3FPOQRfWjQlk7DVU= +github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= +github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= +github.com/stretchr/testify v1.7.2 h1:4jaiDzPyXQvSd7D0EjG45355tLlV3VOECpq10pLC+8s= +github.com/stretchr/testify v1.7.2/go.mod h1:R6va5+xMeoiuVRoj+gSkQ7d3FALtqAAGI1FQKckRals= +go.opentelemetry.io/auto/sdk v1.2.1 h1:jXsnJ4Lmnqd11kwkBV2LgLoFMZKizbCi5fNZ/ipaZ64= +go.opentelemetry.io/auto/sdk v1.2.1/go.mod h1:KRTj+aOaElaLi+wW1kO/DZRXwkF4C5xPbEe3ZiIhN7Y= +go.opentelemetry.io/otel v1.44.0 h1:JjwHmHpA4iZ3wBxluu2fbbE7j4kqlE8jXyAyPXH7HqU= +go.opentelemetry.io/otel v1.44.0/go.mod h1:BMgjTHL9WPRlRjL2oZCBTL4whCGtXch2H4BhOPIAyYc= +go.opentelemetry.io/otel/metric v1.44.0 h1:1w0gILTcHdr3YI+ixLyjemwrVnsMURbTZFrSYCdDdmc= +go.opentelemetry.io/otel/metric v1.44.0/go.mod h1:8O7hanEPBNgEMmybD3s2VBKcgWOCsA6tzHBPODAiquo= +go.opentelemetry.io/otel/sdk v1.44.0 h1:nHYwb9lK+fJPU/dnT6s7W7Z8itMWyqrnVfbheVYrZ58= +go.opentelemetry.io/otel/sdk v1.44.0/go.mod h1:Osuydd3Se74nqjAKxid74N5eC+jfEqfTegHRnq58oK0= +go.opentelemetry.io/otel/sdk/metric v1.44.0 h1:3LlKgI+VjbVsjNRFZJZAJ30WjXC5VkNRks6si09iEfI= +go.opentelemetry.io/otel/sdk/metric v1.44.0/go.mod h1:5B5pMARnXxKhltooO4xUuCBorl65a4EpnTalObqOigA= +go.opentelemetry.io/otel/trace v1.44.0 h1:jxF5CsGYCe74MCRx2X4g7WsY/VBKRqqpNvXlX/6gtIk= +go.opentelemetry.io/otel/trace v1.44.0/go.mod h1:oLl1jrMQAVo6v3GAggN+1VH9VIz9iUSvW53sW1Q8PIE= +golang.org/x/net v0.58.0 h1:ynWG7rqYi4ccpTEuPZ2QGWHktVEM9DMCj9yzDE0Q7To= +golang.org/x/net v0.58.0/go.mod h1:YwCddHnFlT7eLQqVprV19OnhLGtc5xOKgE0RyqgfWAU= +golang.org/x/sys v0.0.0-20200116001909-b77594299b42/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= +golang.org/x/sys v0.0.0-20200223170610-d5e6a3e2c0ae/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= +golang.org/x/sys v0.0.0-20210630005230-0f9fa26af87c/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.0.0-20210927094055-39ccf1dd6fa6/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.0.0-20220503163025-988cb79eb6c6/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs= +golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= +golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8= +golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M= +gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4= +gonum.org/v1/gonum v0.17.0/go.mod h1:El3tOrEuMpv2UdMrbNlKEh9vd86bmQ6vqIcDwxEOc1E= +google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa h1:mZHHdPZl0dbGHCflZgAq/Q468DWVFcU2whhB2KAo8fk= +google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8= +google.golang.org/grpc v1.83.2 h1:EManeRomTObA0BU7I8vXgg/78uE5MJ9M8B39EX2WscU= +google.golang.org/grpc v1.83.2/go.mod h1:YPI1hK3kDked6iHvgX3tR0y+nX/qpMFKhPgFsokw1S8= +google.golang.org/protobuf v1.36.12 h1:pJOKDDOyeXErUroCihFAd5LQuwXBSpVnKGrj5o/fwxc= +google.golang.org/protobuf v1.36.12/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/plugin/example/main.go b/plugin/example/main.go new file mode 100644 index 0000000..a41d9b2 --- /dev/null +++ b/plugin/example/main.go @@ -0,0 +1,77 @@ +// example 是给作者拷贝的模板。实现 sdk.Plugin,再按需实现各口的 Handler,main 里 sdk.Serve 即可。 +// 只 import codedock/pkg/plugin。复制本目录,改 go.mod 模块名、Manifest.Name 和下面的策略。 +// +// 本模板不订阅、不改正文、不换向、不登记方法。编进 PLUGIN_DIR 也不会改对话。 +// +// 六个口的参数和回包就是 SDK 里的结构体,跳进类型看字段: +// +// OnAgentInput sdk.AgentInput / sdk.AgentInputResult +// OnAgentPreStep sdk.AgentPreStep / sdk.AgentPreStepResult +// OnAgentRequest sdk.AgentRequest / sdk.AgentRequestResult +// OnLLMStream sdk.LLMStream / sdk.LLMStreamResult +// OnToolPreExecute sdk.ToolPreExecute / sdk.ToolPreExecuteResult +// OnToolPostExecute sdk.ToolPostExecute / sdk.ToolPostExecuteResult +// OnLedgerNotify sdk.LedgerNotify +// +// 没实现的口原样通过。Reply 继续,Handle / Block / Deny / AskApproval 换向。 +// 跨口、跨插件传参数用 in.Context.Set("example.xxx", v),不要塞 Hidden。 +package main + +import ( + "context" + + sdk "codedock/pkg/plugin" +) + +// plugin 是空模板:各口原样 Reply。 +type plugin struct{} + +// Bootstrap 不订阅、不登记方法。 +func (plugin) Bootstrap(context.Context, sdk.Host) (sdk.Manifest, error) { + return sdk.Manifest{Name: "example"}, nil +} + +// OnAgentInput 原样通过。复制后可改 Content,或 Handle() 不建 Run。 +func (plugin) OnAgentInput(_ context.Context, in sdk.AgentInput) (sdk.AgentInputResult, error) { + return in.Reply(), nil +} + +// OnAgentPreStep 原样通过。复制后可改 Hidden,或 Block() 取消本轮。 +func (plugin) OnAgentPreStep(_ context.Context, in sdk.AgentPreStep) (sdk.AgentPreStepResult, error) { + return in.Reply(), nil +} + +// OnAgentRequest 原样通过。只能 Reply。 +func (plugin) OnAgentRequest(_ context.Context, in sdk.AgentRequest) (sdk.AgentRequestResult, error) { + return in.Reply(), nil +} + +// OnLLMStream 原样通过。只能 Reply;fake 模型不插这个口。 +func (plugin) OnLLMStream(_ context.Context, in sdk.LLMStream) (sdk.LLMStreamResult, error) { + return in.Reply(), nil +} + +// OnToolPreExecute 原样通过。复制后可 Deny() 或 AskApproval()。 +func (plugin) OnToolPreExecute(_ context.Context, in sdk.ToolPreExecute) (sdk.ToolPreExecuteResult, error) { + return in.Reply(), nil +} + +// OnToolPostExecute 原样通过。只能 Reply。 +func (plugin) OnToolPostExecute(_ context.Context, in sdk.ToolPostExecute) (sdk.ToolPostExecuteResult, error) { + return in.Reply(), nil +} + +// OnLedgerNotify 忽略账本通知。 +func (plugin) OnLedgerNotify(context.Context, sdk.LedgerNotify) error { + return nil +} + +// ExecuteMethod 本模板不登记方法。 +func (plugin) ExecuteMethod(context.Context, sdk.MethodInput) (sdk.MethodResult, error) { + return sdk.MethodResult{Success: false, Error: "no methods"}, nil +} + +// main 启动 example 插件进程。 +func main() { + sdk.Serve(plugin{}) +} diff --git a/plugin/redact/go.mod b/plugin/redact/go.mod new file mode 100644 index 0000000..7c5cc76 --- /dev/null +++ b/plugin/redact/go.mod @@ -0,0 +1,25 @@ +module codedock-example-redact + +go 1.26.5 + +require codedock v0.0.0 + +require ( + github.com/fatih/color v1.13.0 // indirect + github.com/golang/protobuf v1.5.4 // indirect + github.com/google/uuid v1.6.0 // indirect + github.com/hashicorp/go-hclog v1.6.3 // indirect + github.com/hashicorp/go-plugin v1.8.0 // indirect + github.com/hashicorp/yamux v0.1.2 // indirect + github.com/mattn/go-colorable v0.1.12 // indirect + github.com/mattn/go-isatty v0.0.24 // indirect + github.com/oklog/run v1.1.0 // indirect + golang.org/x/net v0.58.0 // indirect + golang.org/x/sys v0.47.0 // indirect + golang.org/x/text v0.41.0 // indirect + google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa // indirect + google.golang.org/grpc v1.83.2 // indirect + google.golang.org/protobuf v1.36.12 // indirect +) + +replace codedock => ../../server diff --git a/plugin/redact/go.sum b/plugin/redact/go.sum new file mode 100644 index 0000000..724238c --- /dev/null +++ b/plugin/redact/go.sum @@ -0,0 +1,75 @@ +github.com/bufbuild/protocompile v0.14.1 h1:iA73zAf/fyljNjQKwYzUHD6AD4R8KMasmwa/FBatYVw= +github.com/bufbuild/protocompile v0.14.1/go.mod h1:ppVdAIhbr2H8asPk6k4pY7t9zB1OU5DoEw9xY/FUi1c= +github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs= +github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= +github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= +github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/fatih/color v1.13.0 h1:8LOYc1KYPPmyKMuN8QV2DNRWNbLo6LZ0iLs8+mlH53w= +github.com/fatih/color v1.13.0/go.mod h1:kLAiJbzzSOZDVNGyDpeOxJ47H46qBXwg5ILebYFFOfk= +github.com/go-logr/logr v1.4.3 h1:CjnDlHq8ikf6E492q6eKboGOC0T8CDaOvkHCIg8idEI= +github.com/go-logr/logr v1.4.3/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY= +github.com/go-logr/stdr v1.2.2 h1:hSWxHoqTgW2S2qGc0LTAI563KZ5YKYRhT3MFKZMbjag= +github.com/go-logr/stdr v1.2.2/go.mod h1:mMo/vtBO5dYbehREoey6XUKy/eSumjCCveDpRre4VKE= +github.com/golang/protobuf v1.5.4 h1:i7eJL8qZTpSEXOPTxNKhASYpMn+8e5Q6AdndVa1dWek= +github.com/golang/protobuf v1.5.4/go.mod h1:lnTiLA8Wa4RWRcIUkrtSVa5nRhsEGBg48fD6rSs7xps= +github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= +github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= +github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= +github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +github.com/hashicorp/go-hclog v1.6.3 h1:Qr2kF+eVWjTiYmU7Y31tYlP1h0q/X3Nl3tPGdaB11/k= +github.com/hashicorp/go-hclog v1.6.3/go.mod h1:W4Qnvbt70Wk/zYJryRzDRU/4r0kIg0PVHBcfoyhpF5M= +github.com/hashicorp/go-plugin v1.8.0 h1:ie8S6RRY8RvB2usYZv+AAZ/wBvx2AU5p5QeP5j/FORs= +github.com/hashicorp/go-plugin v1.8.0/go.mod h1:BExt6KEaIYx804z8k4gRzRLEvxKVb+kn0NMcihqOqb8= +github.com/hashicorp/yamux v0.1.2 h1:XtB8kyFOyHXYVFnwT5C3+Bdo8gArse7j2AQ0DA0Uey8= +github.com/hashicorp/yamux v0.1.2/go.mod h1:C+zze2n6e/7wshOZep2A70/aQU6QBRWJO/G6FT1wIns= +github.com/jhump/protoreflect v1.17.0 h1:qOEr613fac2lOuTgWN4tPAtLL7fUSbuJL5X5XumQh94= +github.com/jhump/protoreflect v1.17.0/go.mod h1:h9+vUUL38jiBzck8ck+6G/aeMX8Z4QUY/NiJPwPNi+8= +github.com/mattn/go-colorable v0.1.9/go.mod h1:u6P/XSegPjTcexA+o6vUJrdnUu04hMope9wVRipJSqc= +github.com/mattn/go-colorable v0.1.12 h1:jF+Du6AlPIjs2BiUiQlKOX0rt3SujHxPnksPKZbaA40= +github.com/mattn/go-colorable v0.1.12/go.mod h1:u5H1YNBxpqRaxsYJYSkiCWKzEfiAb1Gb520KVy5xxl4= +github.com/mattn/go-isatty v0.0.12/go.mod h1:cbi8OIDigv2wuxKPP5vlRcQ1OAZbq2CE4Kysco4FUpU= +github.com/mattn/go-isatty v0.0.14/go.mod h1:7GGIvUiUoEMVVmxf/4nioHXj79iQHKdU27kJ6hsGG94= +github.com/mattn/go-isatty v0.0.24 h1:tGZZoVgT/KiqK1c8ocVLeDS8BSWMRd47J3Lbz7vsReI= +github.com/mattn/go-isatty v0.0.24/go.mod h1:nMCL3Zebbrt45jsMDgnfIwz6ydEQApk5oEI3HqDio6A= +github.com/oklog/run v1.1.0 h1:GEenZ1cK0+q0+wsJew9qUg/DyD8k3JzYsZAi5gYi2mA= +github.com/oklog/run v1.1.0/go.mod h1:sVPdnTZT1zYwAJeCMu2Th4T21pA3FPOQRfWjQlk7DVU= +github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= +github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= +github.com/stretchr/testify v1.7.2 h1:4jaiDzPyXQvSd7D0EjG45355tLlV3VOECpq10pLC+8s= +github.com/stretchr/testify v1.7.2/go.mod h1:R6va5+xMeoiuVRoj+gSkQ7d3FALtqAAGI1FQKckRals= +go.opentelemetry.io/auto/sdk v1.2.1 h1:jXsnJ4Lmnqd11kwkBV2LgLoFMZKizbCi5fNZ/ipaZ64= +go.opentelemetry.io/auto/sdk v1.2.1/go.mod h1:KRTj+aOaElaLi+wW1kO/DZRXwkF4C5xPbEe3ZiIhN7Y= +go.opentelemetry.io/otel v1.44.0 h1:JjwHmHpA4iZ3wBxluu2fbbE7j4kqlE8jXyAyPXH7HqU= +go.opentelemetry.io/otel v1.44.0/go.mod h1:BMgjTHL9WPRlRjL2oZCBTL4whCGtXch2H4BhOPIAyYc= +go.opentelemetry.io/otel/metric v1.44.0 h1:1w0gILTcHdr3YI+ixLyjemwrVnsMURbTZFrSYCdDdmc= +go.opentelemetry.io/otel/metric v1.44.0/go.mod h1:8O7hanEPBNgEMmybD3s2VBKcgWOCsA6tzHBPODAiquo= +go.opentelemetry.io/otel/sdk v1.44.0 h1:nHYwb9lK+fJPU/dnT6s7W7Z8itMWyqrnVfbheVYrZ58= +go.opentelemetry.io/otel/sdk v1.44.0/go.mod h1:Osuydd3Se74nqjAKxid74N5eC+jfEqfTegHRnq58oK0= +go.opentelemetry.io/otel/sdk/metric v1.44.0 h1:3LlKgI+VjbVsjNRFZJZAJ30WjXC5VkNRks6si09iEfI= +go.opentelemetry.io/otel/sdk/metric v1.44.0/go.mod h1:5B5pMARnXxKhltooO4xUuCBorl65a4EpnTalObqOigA= +go.opentelemetry.io/otel/trace v1.44.0 h1:jxF5CsGYCe74MCRx2X4g7WsY/VBKRqqpNvXlX/6gtIk= +go.opentelemetry.io/otel/trace v1.44.0/go.mod h1:oLl1jrMQAVo6v3GAggN+1VH9VIz9iUSvW53sW1Q8PIE= +golang.org/x/net v0.58.0 h1:ynWG7rqYi4ccpTEuPZ2QGWHktVEM9DMCj9yzDE0Q7To= +golang.org/x/net v0.58.0/go.mod h1:YwCddHnFlT7eLQqVprV19OnhLGtc5xOKgE0RyqgfWAU= +golang.org/x/sys v0.0.0-20200116001909-b77594299b42/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= +golang.org/x/sys v0.0.0-20200223170610-d5e6a3e2c0ae/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= +golang.org/x/sys v0.0.0-20210630005230-0f9fa26af87c/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.0.0-20210927094055-39ccf1dd6fa6/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.0.0-20220503163025-988cb79eb6c6/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs= +golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= +golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8= +golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M= +gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4= +gonum.org/v1/gonum v0.17.0/go.mod h1:El3tOrEuMpv2UdMrbNlKEh9vd86bmQ6vqIcDwxEOc1E= +google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa h1:mZHHdPZl0dbGHCflZgAq/Q468DWVFcU2whhB2KAo8fk= +google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8= +google.golang.org/grpc v1.83.2 h1:EManeRomTObA0BU7I8vXgg/78uE5MJ9M8B39EX2WscU= +google.golang.org/grpc v1.83.2/go.mod h1:YPI1hK3kDked6iHvgX3tR0y+nX/qpMFKhPgFsokw1S8= +google.golang.org/protobuf v1.36.12 h1:pJOKDDOyeXErUroCihFAd5LQuwXBSpVnKGrj5o/fwxc= +google.golang.org/protobuf v1.36.12/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/plugin/redact/main.go b/plugin/redact/main.go new file mode 100644 index 0000000..dd09cf4 --- /dev/null +++ b/plugin/redact/main.go @@ -0,0 +1,62 @@ +// redact 在三个口把秘密换成占位符,不拦工具、不建方法。 +// 只 import codedock/pkg/plugin。plugin/example 是作者拷贝模板。 +// +// go build -o redact . +// PLUGIN_DIR=<仓根>/plugin 启动 API +package main + +import ( + "context" + + sdk "codedock/pkg/plugin" +) + +// plugin 订阅 input / request / post-execute,只改数据。 +type plugin struct{} + +// Bootstrap 声明听用户正文、发给模型前、工具回包后。 +func (p *plugin) Bootstrap(_ context.Context, _ sdk.Host) (sdk.Manifest, error) { + return sdk.Manifest{ + Name: "redact", + Subscriptions: []string{ + sdk.TypeInput, + sdk.TypeRequest, + sdk.TypePostExecute, + }, + }, nil +} + +// OnAgentInput 刮用户正文后再落库。 +func (p *plugin) OnAgentInput(_ context.Context, in sdk.AgentInput) (sdk.AgentInputResult, error) { + in.Content, _ = maskText(in.Content) + return in.Reply(), nil +} + +// OnAgentRequest 刮即将发给模型的提示、消息和工具参数。 +func (p *plugin) OnAgentRequest(_ context.Context, in sdk.AgentRequest) (sdk.AgentRequestResult, error) { + in.SystemPrompt, _ = maskText(in.SystemPrompt) + for i := range in.Messages { + in.Messages[i].Content, _ = maskJSON(in.Messages[i].Content) + for j := range in.Messages[i].ToolCalls { + in.Messages[i].ToolCalls[j].Arguments, _ = maskJSON(in.Messages[i].ToolCalls[j].Arguments) + } + } + return in.Reply(), nil +} + +// OnToolPostExecute 刮回包和失败信息,不改 Success。 +func (p *plugin) OnToolPostExecute(_ context.Context, in sdk.ToolPostExecute) (sdk.ToolPostExecuteResult, error) { + in.Result.Output = rewriteOutput(in.Result.Output) + in.Result.Error, _ = maskText(in.Result.Error) + return in.Reply(), nil +} + +// ExecuteMethod 本插件不登记方法。 +func (p *plugin) ExecuteMethod(_ context.Context, _ sdk.MethodInput) (sdk.MethodResult, error) { + return sdk.MethodResult{Success: false, Error: "redact has no methods"}, nil +} + +// main 启动 redact 插件进程。 +func main() { + sdk.Serve(&plugin{}) +} diff --git a/plugin/redact/mask.go b/plugin/redact/mask.go new file mode 100644 index 0000000..2240ef9 --- /dev/null +++ b/plugin/redact/mask.go @@ -0,0 +1,119 @@ +package main + +import ( + "crypto/sha256" + "fmt" + "regexp" + "sort" + "strings" +) + +type hit struct { + start int + end int + kind string + value string +} + +var ( + reAssignment = regexp.MustCompile(`(?i)(?:^|[^A-Za-z0-9_])((?:AWS_SECRET|PRIVATE_KEY|DATABASE_URL|API_KEY|SECRET|TOKEN|PASSWORD)[A-Za-z0-9_]*)\s*[=:]\s*(\S+)`) + reAWS = regexp.MustCompile(`AKIA[0-9A-Z]{16}`) + reGitHub = regexp.MustCompile(`(?:ghp_|github_pat_)[^\s]{20,}`) + reOpenAI = regexp.MustCompile(`sk-[A-Za-z0-9_-]{20,}`) + reSlack = regexp.MustCompile(`xox[baprs]-[^\s]{10,}`) + reConn = regexp.MustCompile(`(?:postgres|mysql|mongodb\+srv|redis)://([^/\s:@]+):([^@\s]+)@`) + rePEM = regexp.MustCompile(`-----BEGIN [A-Z0-9 ]*PRIVATE KEY-----[\s\S]*?-----END [A-Z0-9 ]*PRIVATE KEY-----`) +) + +// token 把原文收成稳定占位符;同值同 kind 永远同一记号。 +func token(kind, value string) string { + sum := sha256.Sum256([]byte(value)) + return fmt.Sprintf("«REDACTED:%s:%x»", kind, sum[:2]) +} + +// maskText 按规则替换一段文本里的秘密,返回新文本和命中数。 +func maskText(s string) (string, int) { + if s == "" { + return s, 0 + } + hits := collectHits(s) + if len(hits) == 0 { + return s, 0 + } + sort.SliceStable(hits, func(i, j int) bool { + if hits[i].start != hits[j].start { + return hits[i].start < hits[j].start + } + return (hits[i].end - hits[i].start) > (hits[j].end - hits[j].start) + }) + adopted := make([]hit, 0, len(hits)) + for _, h := range hits { + overlap := false + for _, a := range adopted { + if h.start < a.end && h.end > a.start { + overlap = true + break + } + } + if !overlap { + adopted = append(adopted, h) + } + } + sort.Slice(adopted, func(i, j int) bool { return adopted[i].start > adopted[j].start }) + out := s + for _, h := range adopted { + out = out[:h.start] + token(h.kind, h.value) + out[h.end:] + } + return out, len(adopted) +} + +// collectHits 收集全部正则命中,赋值行只记值的区间。 +func collectHits(s string) []hit { + var hits []hit + for _, loc := range reAssignment.FindAllStringSubmatchIndex(s, -1) { + if len(loc) < 6 { + continue + } + valStart, valEnd := loc[4], loc[5] + value := s[valStart:valEnd] + if q := quoted(value); q != "" { + value = q + } + hits = append(hits, hit{start: valStart, end: valEnd, kind: "assignment", value: value}) + } + addWhole := func(re *regexp.Regexp, kind string) { + for _, loc := range re.FindAllStringIndex(s, -1) { + hits = append(hits, hit{start: loc[0], end: loc[1], kind: kind, value: s[loc[0]:loc[1]]}) + } + } + addWhole(reAWS, "aws_ak") + addWhole(reGitHub, "github") + addWhole(reOpenAI, "openai") + addWhole(reSlack, "slack") + addWhole(rePEM, "pem") + for _, loc := range reConn.FindAllStringSubmatchIndex(s, -1) { + if len(loc) < 6 { + continue + } + user, pass := s[loc[2]:loc[3]], s[loc[4]:loc[5]] + start := loc[2] + end := loc[5] + 1 + if end > len(s) || s[loc[5]:end] != "@" { + end = loc[1] + } + hits = append(hits, hit{start: start, end: end, kind: "conn", value: user + ":" + pass + "@"}) + } + return hits +} + +// quoted 若整段被成对引号包住则返回去掉引号的正文,否则空串。 +func quoted(value string) string { + if len(value) < 2 { + return "" + } + if (strings.HasPrefix(value, `"`) && strings.HasSuffix(value, `"`)) || + (strings.HasPrefix(value, `'`) && strings.HasSuffix(value, `'`)) { + return value[1 : len(value)-1] + } + return "" +} diff --git a/plugin/redact/mask_test.go b/plugin/redact/mask_test.go new file mode 100644 index 0000000..e9cb9c5 --- /dev/null +++ b/plugin/redact/mask_test.go @@ -0,0 +1,188 @@ +package main + +import ( + "context" + "encoding/json" + "strings" + "testing" + + pkgagent "codedock/pkg/agent" + "codedock/pkg/agent/tool" + sdk "codedock/pkg/plugin" +) + +func TestMaskTextAssignmentKeepsKey(t *testing.T) { + out, n := maskText("API_KEY=sk-1234") + if n != 1 || !strings.HasPrefix(out, "API_KEY=«REDACTED:assignment:") || strings.Contains(out, "sk-1234") { + t.Fatalf("got %q n=%d", out, n) + } +} + +func TestMaskTextAssignmentVariants(t *testing.T) { + cases := []string{ + "export AWS_SECRET=wxyz", + "password: hunter2", + `TOKEN = "ghp_xxx"`, + } + for _, in := range cases { + out, n := maskText(in) + if n != 1 || strings.Contains(out, "wxyz") && strings.Contains(in, "wxyz") { + t.Fatalf("%q -> %q n=%d", in, out, n) + } + if strings.Contains(out, "hunter2") || strings.Contains(out, "ghp_xxx") { + t.Fatalf("value leaked: %q -> %q", in, out) + } + if !strings.Contains(out, "«REDACTED:assignment:") { + t.Fatalf("want assignment placeholder: %q", out) + } + } +} + +func TestMaskTextIgnoresPlainKeys(t *testing.T) { + for _, in := range []string{"timeout=30", "user_name=alice"} { + out, n := maskText(in) + if n != 0 || out != in { + t.Fatalf("%q -> %q n=%d", in, out, n) + } + } +} + +func TestMaskTextShapes(t *testing.T) { + aws := "AKIAIOSFODNN7EXAMPLE" + out, n := maskText("id=" + aws) + if n != 1 || strings.Contains(out, aws) || !strings.Contains(out, "«REDACTED:aws_ak:") { + t.Fatalf("aws %q n=%d", out, n) + } + again, _ := maskText(aws) + if again != token("aws_ak", aws) { + t.Fatalf("stable token %q", again) + } + + gh := "ghp_01234567890123456789" + out, n = maskText(gh) + if n != 1 || !strings.Contains(out, "«REDACTED:github:") { + t.Fatalf("github %q n=%d", out, n) + } + + sk := "sk-abcdefghijklmnopqrstuvwxyz" + out, n = maskText(sk) + if n != 1 || !strings.Contains(out, "«REDACTED:openai:") { + t.Fatalf("openai %q n=%d", out, n) + } + short, n := maskText("see sk-xxx and sk-test") + if n != 0 || short != "see sk-xxx and sk-test" { + t.Fatalf("short sk %q n=%d", short, n) + } + + slack := "xoxb-1234567890" + out, n = maskText(slack) + if n != 1 || !strings.Contains(out, "«REDACTED:slack:") { + t.Fatalf("slack %q n=%d", out, n) + } + + conn := "postgres://me:secret@host/db" + out, n = maskText(conn) + if n != 1 || strings.Contains(out, "me:secret") || !strings.HasPrefix(out, "postgres://«REDACTED:conn:") || !strings.Contains(out, "host/db") { + t.Fatalf("conn %q n=%d", out, n) + } +} + +func TestMaskTextAssignmentWinsOverShape(t *testing.T) { + in := "AWS_SECRET=" + "AKIAIOSFODNN7EXAMPLE" + out, n := maskText(in) + if n != 1 || !strings.Contains(out, "AWS_SECRET=«REDACTED:assignment:") || strings.Contains(out, "aws_ak") { + t.Fatalf("overlap %q n=%d", out, n) + } +} + +func TestMaskTextPEM(t *testing.T) { + pem := "-----BEGIN RSA PRIVATE KEY-----\nMIIE\n-----END RSA PRIVATE KEY-----" + out, n := maskText("before\n" + pem + "\nafter") + if n != 1 || strings.Contains(out, "MIIE") || !strings.Contains(out, "«REDACTED:pem:") || !strings.Contains(out, "before") { + t.Fatalf("pem %q n=%d", out, n) + } + half := "-----BEGIN RSA PRIVATE KEY-----\nMIIE" + out, n = maskText(half) + if n != 0 || out != half { + t.Fatalf("half pem %q n=%d", out, n) + } +} + +func TestMaskJSONSkipsImageData(t *testing.T) { + raw := json.RawMessage(`{"type":"image","data":"AKIAIOSFODNN7EXAMPLE","mimeType":"image/png"}`) + out, n := maskJSON(raw) + if n != 0 || !strings.Contains(string(out), "AKIAIOSFODNN7EXAMPLE") { + t.Fatalf("image %s n=%d", out, n) + } +} + +func TestMaskJSONWalksTextAndInvalid(t *testing.T) { + raw := json.RawMessage(`{"content":[{"type":"text","text":"API_KEY=sk-1234"}]}`) + out, n := maskJSON(raw) + if n != 1 || strings.Contains(string(out), "sk-1234") || !strings.Contains(string(out), "API_KEY=") { + t.Fatalf("json %s n=%d", out, n) + } + plain, n := maskJSON(json.RawMessage(`API_KEY=sk-1234 not-json`)) + if n != 1 || strings.Contains(string(plain), "sk-1234") { + t.Fatalf("invalid %s n=%d", plain, n) + } +} + +func TestRewriteOutputDropsPathAndFooter(t *testing.T) { + raw := json.RawMessage(`{"content":[{"type":"text","text":"API_KEY=sk-1234"}],"details":{"fullOutputPath":"/tmp/out"}}`) + out := rewriteOutput(raw) + if strings.Contains(string(out), "sk-1234") || strings.Contains(string(out), "fullOutputPath") || strings.Contains(string(out), "/tmp/out") { + t.Fatalf("rewrite %s", out) + } + if !strings.Contains(string(out), "[redact] masked 1 secrets") { + t.Fatalf("missing footer %s", out) + } + clean := rewriteOutput(json.RawMessage(`{"content":[{"type":"text","text":"ok"}],"details":{"fullOutputPath":"/tmp/out"}}`)) + if strings.Contains(string(clean), "fullOutputPath") || strings.Contains(string(clean), "[redact]") { + t.Fatalf("path-only %s", clean) + } +} + +func TestPluginHooks(t *testing.T) { + p := &plugin{} + in, err := p.OnAgentInput(context.Background(), sdk.AgentInput{Content: "use AKIAIOSFODNN7EXAMPLE"}) + if err != nil || strings.Contains(in.Content, "AKIA") || !strings.Contains(in.Content, "«REDACTED:aws_ak:") { + t.Fatalf("input %+v err=%v", in, err) + } + + req, err := p.OnAgentRequest(context.Background(), sdk.AgentRequest{ + SystemPrompt: "key AKIAIOSFODNN7EXAMPLE", + Messages: []pkgagent.Message{{ + Content: json.RawMessage(`{"text":"API_KEY=sk-1234"}`), + ToolCalls: []tool.Call{{ + Arguments: json.RawMessage(`{"cmd":"echo AKIAIOSFODNN7EXAMPLE"}`), + }}, + }}, + }) + if err != nil { + t.Fatal(err) + } + if strings.Contains(req.SystemPrompt, "AKIA") { + t.Fatalf("prompt %q", req.SystemPrompt) + } + if strings.Contains(string(req.Messages[0].Content), "sk-1234") { + t.Fatalf("msg %s", req.Messages[0].Content) + } + if strings.Contains(string(req.Messages[0].ToolCalls[0].Arguments), "AKIAIOSFODNN7EXAMPLE") { + t.Fatalf("args %s", req.Messages[0].ToolCalls[0].Arguments) + } + + post, err := p.OnToolPostExecute(context.Background(), sdk.ToolPostExecute{ + Result: tool.Result{ + Success: true, + Output: json.RawMessage(`{"content":[{"type":"text","text":"API_KEY=sk-1234"}]}`), + Error: "boom AKIAIOSFODNN7EXAMPLE", + }, + }) + if err != nil || !post.Result.Success { + t.Fatalf("post %+v err=%v", post, err) + } + if strings.Contains(string(post.Result.Output), "sk-1234") || strings.Contains(post.Result.Error, "AKIA") { + t.Fatalf("post leaked %+v", post.Result) + } +} diff --git a/plugin/redact/walk.go b/plugin/redact/walk.go new file mode 100644 index 0000000..d51de4d --- /dev/null +++ b/plugin/redact/walk.go @@ -0,0 +1,151 @@ +package main + +import ( + "bytes" + "encoding/json" + "fmt" +) + +// maskJSON 走访 JSON 里的字符串;图片 data 跳过;解析失败当整段文本刮。 +func maskJSON(raw json.RawMessage) (json.RawMessage, int) { + if len(bytes.TrimSpace(raw)) == 0 { + return raw, 0 + } + var v any + if err := json.Unmarshal(raw, &v); err != nil { + out, n := maskText(string(raw)) + return json.RawMessage(out), n + } + nv, n := walk(v) + if n == 0 { + return raw, 0 + } + body, err := json.Marshal(nv) + if err != nil { + return raw, 0 + } + return body, n +} + +// walk 递归处理对象和数组;type=image 的 data 不扫。 +func walk(v any) (any, int) { + switch x := v.(type) { + case map[string]any: + n := 0 + image := false + if typ, ok := x["type"].(string); ok && typ == "image" { + image = true + } + for key, val := range x { + if image && key == "data" { + continue + } + nv, c := walk(val) + x[key] = nv + n += c + } + return x, n + case []any: + n := 0 + for i, val := range x { + nv, c := walk(val) + x[i] = nv + n += c + } + return x, n + case string: + return maskText(x) + default: + return v, 0 + } +} + +// rewriteOutput 刮工具回包:删 fullOutputPath,命中则在最后一个 text 块补脚注。 +func rewriteOutput(raw json.RawMessage) json.RawMessage { + out, n := maskJSON(raw) + out = dropFullOutputPath(out) + if n == 0 { + return out + } + return appendFooter(out, n) +} + +// dropFullOutputPath 从任意嵌套对象里拿掉 fullOutputPath。 +func dropFullOutputPath(raw json.RawMessage) json.RawMessage { + if len(bytes.TrimSpace(raw)) == 0 { + return raw + } + var v any + if err := json.Unmarshal(raw, &v); err != nil { + return raw + } + nv, dropped := dropPath(v) + if !dropped { + return raw + } + body, err := json.Marshal(nv) + if err != nil { + return raw + } + return body +} + +// dropPath 删除对象里的 fullOutputPath,并往下找。 +func dropPath(v any) (any, bool) { + switch x := v.(type) { + case map[string]any: + dropped := false + if _, ok := x["fullOutputPath"]; ok { + delete(x, "fullOutputPath") + dropped = true + } + for key, val := range x { + nv, d := dropPath(val) + x[key] = nv + dropped = dropped || d + } + return x, dropped + case []any: + dropped := false + for i, val := range x { + nv, d := dropPath(val) + x[i] = nv + dropped = dropped || d + } + return x, dropped + default: + return v, false + } +} + +// appendFooter 在最后一个 type=text 的块末尾写 masked 条数。 +func appendFooter(raw json.RawMessage, n int) json.RawMessage { + var root map[string]any + if err := json.Unmarshal(raw, &root); err != nil { + return raw + } + content, ok := root["content"].([]any) + if !ok { + return raw + } + note := fmt.Sprintf("\n[redact] masked %d secrets", n) + for i := len(content) - 1; i >= 0; i-- { + block, ok := content[i].(map[string]any) + if !ok { + continue + } + if typ, _ := block["type"].(string); typ != "text" { + continue + } + text, _ := block["text"].(string) + block["text"] = text + note + content[i] = block + root["content"] = content + body, err := json.Marshal(root) + if err != nil { + return raw + } + return body + } + return raw +} diff --git a/server/cmd/server/main.go b/server/cmd/server/main.go index 953e8dc..250fc09 100644 --- a/server/cmd/server/main.go +++ b/server/cmd/server/main.go @@ -19,6 +19,7 @@ import ( "codedock/internal/handler" codexhttp "codedock/internal/handler/codex" "codedock/internal/logger" + "codedock/internal/pluginhost" pkgagent "codedock/pkg/agent" "codedock/pkg/db" ) @@ -63,6 +64,32 @@ func main() { runtime.SetModel(model) runtime.SetConcurrency(cfg.LLMConcurrency, cfg.ToolConcurrency) log.Info("concurrency", "llm", cfg.LLMConcurrency, "tool", cfg.ToolConcurrency) + + var pluginHost *pluginhost.Host + if cfg.PluginDir != "" { + host, err := pluginhost.Load(ctx, pluginhost.Options{ + Dir: cfg.PluginDir, + Timeout: cfg.PluginRPCTimeout, + Registry: runtime.Tools(), + Queries: queries, + Model: model, + Log: logger.NewLogger("plugin"), + }) + if err != nil { + log.Error("load plugins", "error", err) + os.Exit(1) + } + pluginHost = host + if pluginHost != nil { + runtime.SetDispatcher(pluginHost) + pluginHost.Attach(bus) + log.Info("plugins loaded", "dir", cfg.PluginDir) + } + } + if pluginHost != nil { + defer func() { _ = pluginHost.Close() }() + } + runtime.Start(ctx) defaults := pkgagent.DefaultRunConfig(pkgagent.WorkAgent, model) diff --git a/server/go.mod b/server/go.mod index d1fbe13..98bc777 100644 --- a/server/go.mod +++ b/server/go.mod @@ -6,18 +6,29 @@ require ( github.com/go-chi/chi/v5 v5.3.0 github.com/google/jsonschema-go v0.4.3 github.com/google/uuid v1.6.0 + github.com/hashicorp/go-hclog v1.6.3 + github.com/hashicorp/go-plugin v1.8.0 github.com/lmittmann/tint v1.2.0 + golang.org/x/image v0.30.0 + golang.org/x/text v0.41.0 + google.golang.org/grpc v1.83.2 + google.golang.org/protobuf v1.36.12 modernc.org/sqlite v1.57.0 ) require ( github.com/dustin/go-humanize v1.0.1 // indirect + github.com/fatih/color v1.13.0 // indirect + github.com/golang/protobuf v1.5.4 // indirect + github.com/hashicorp/yamux v0.1.2 // indirect + github.com/mattn/go-colorable v0.1.12 // indirect github.com/mattn/go-isatty v0.0.24 // indirect github.com/ncruces/go-strftime v1.0.0 // indirect + github.com/oklog/run v1.1.0 // indirect github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect - golang.org/x/image v0.30.0 // indirect + golang.org/x/net v0.58.0 // indirect golang.org/x/sys v0.47.0 // indirect - golang.org/x/text v0.30.0 // indirect + google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa // indirect modernc.org/libc v1.74.4 // indirect modernc.org/mathutil v1.7.1 // indirect modernc.org/memory v1.11.0 // indirect diff --git a/server/go.sum b/server/go.sum index 0cbb4f0..7a3fb73 100644 --- a/server/go.sum +++ b/server/go.sum @@ -1,7 +1,22 @@ +github.com/bufbuild/protocompile v0.14.1 h1:iA73zAf/fyljNjQKwYzUHD6AD4R8KMasmwa/FBatYVw= +github.com/bufbuild/protocompile v0.14.1/go.mod h1:ppVdAIhbr2H8asPk6k4pY7t9zB1OU5DoEw9xY/FUi1c= +github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs= +github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= +github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= +github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY= github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto= +github.com/fatih/color v1.13.0 h1:8LOYc1KYPPmyKMuN8QV2DNRWNbLo6LZ0iLs8+mlH53w= +github.com/fatih/color v1.13.0/go.mod h1:kLAiJbzzSOZDVNGyDpeOxJ47H46qBXwg5ILebYFFOfk= github.com/go-chi/chi/v5 v5.3.0 h1:halUjDxhshgXHMrao5bB8eNBXo/rnzwr8m5m36glehM= github.com/go-chi/chi/v5 v5.3.0/go.mod h1:R+tYY2hNuVUUjxoPtqUdgBqevM9s9njzkTLutVsOCto= +github.com/go-logr/logr v1.4.3 h1:CjnDlHq8ikf6E492q6eKboGOC0T8CDaOvkHCIg8idEI= +github.com/go-logr/logr v1.4.3/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY= +github.com/go-logr/stdr v1.2.2 h1:hSWxHoqTgW2S2qGc0LTAI563KZ5YKYRhT3MFKZMbjag= +github.com/go-logr/stdr v1.2.2/go.mod h1:mMo/vtBO5dYbehREoey6XUKy/eSumjCCveDpRre4VKE= +github.com/golang/protobuf v1.5.4 h1:i7eJL8qZTpSEXOPTxNKhASYpMn+8e5Q6AdndVa1dWek= +github.com/golang/protobuf v1.5.4/go.mod h1:lnTiLA8Wa4RWRcIUkrtSVa5nRhsEGBg48fD6rSs7xps= github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= github.com/google/jsonschema-go v0.4.3 h1:/DBOLZTfDow7pe2GmaJNhltueGTtDKICi8V8p+DQPd0= @@ -10,28 +25,78 @@ github.com/google/pprof v0.0.0-20260802141513-ef3492d7dac3 h1:LMLX+LgTNWpfvCBdFe github.com/google/pprof v0.0.0-20260802141513-ef3492d7dac3/go.mod h1:jl5iWTm0/hd5PjEYEOuwAJ57L/CibdZfrqZ5XA5GrCk= github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +github.com/hashicorp/go-hclog v1.6.3 h1:Qr2kF+eVWjTiYmU7Y31tYlP1h0q/X3Nl3tPGdaB11/k= +github.com/hashicorp/go-hclog v1.6.3/go.mod h1:W4Qnvbt70Wk/zYJryRzDRU/4r0kIg0PVHBcfoyhpF5M= +github.com/hashicorp/go-plugin v1.8.0 h1:ie8S6RRY8RvB2usYZv+AAZ/wBvx2AU5p5QeP5j/FORs= +github.com/hashicorp/go-plugin v1.8.0/go.mod h1:BExt6KEaIYx804z8k4gRzRLEvxKVb+kn0NMcihqOqb8= github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k= github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM= +github.com/hashicorp/yamux v0.1.2 h1:XtB8kyFOyHXYVFnwT5C3+Bdo8gArse7j2AQ0DA0Uey8= +github.com/hashicorp/yamux v0.1.2/go.mod h1:C+zze2n6e/7wshOZep2A70/aQU6QBRWJO/G6FT1wIns= +github.com/jhump/protoreflect v1.17.0 h1:qOEr613fac2lOuTgWN4tPAtLL7fUSbuJL5X5XumQh94= +github.com/jhump/protoreflect v1.17.0/go.mod h1:h9+vUUL38jiBzck8ck+6G/aeMX8Z4QUY/NiJPwPNi+8= github.com/lmittmann/tint v1.2.0 h1:AogHRHy8HUJUnNJBHJlYa+fR4YY8mko2cnCp67xn9JY= github.com/lmittmann/tint v1.2.0/go.mod h1:HIS3gSy7qNwGCj+5oRjAutErFBl4BzdQP6cJZ0NfMwE= +github.com/mattn/go-colorable v0.1.9/go.mod h1:u6P/XSegPjTcexA+o6vUJrdnUu04hMope9wVRipJSqc= +github.com/mattn/go-colorable v0.1.12 h1:jF+Du6AlPIjs2BiUiQlKOX0rt3SujHxPnksPKZbaA40= +github.com/mattn/go-colorable v0.1.12/go.mod h1:u5H1YNBxpqRaxsYJYSkiCWKzEfiAb1Gb520KVy5xxl4= +github.com/mattn/go-isatty v0.0.12/go.mod h1:cbi8OIDigv2wuxKPP5vlRcQ1OAZbq2CE4Kysco4FUpU= +github.com/mattn/go-isatty v0.0.14/go.mod h1:7GGIvUiUoEMVVmxf/4nioHXj79iQHKdU27kJ6hsGG94= github.com/mattn/go-isatty v0.0.24 h1:tGZZoVgT/KiqK1c8ocVLeDS8BSWMRd47J3Lbz7vsReI= github.com/mattn/go-isatty v0.0.24/go.mod h1:nMCL3Zebbrt45jsMDgnfIwz6ydEQApk5oEI3HqDio6A= github.com/ncruces/go-strftime v1.0.0 h1:HMFp8mLCTPp341M/ZnA4qaf7ZlsbTc+miZjCLOFAw7w= github.com/ncruces/go-strftime v1.0.0/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls= +github.com/oklog/run v1.1.0 h1:GEenZ1cK0+q0+wsJew9qUg/DyD8k3JzYsZAi5gYi2mA= +github.com/oklog/run v1.1.0/go.mod h1:sVPdnTZT1zYwAJeCMu2Th4T21pA3FPOQRfWjQlk7DVU= +github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= +github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE= github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo= +github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= +github.com/stretchr/testify v1.7.2 h1:4jaiDzPyXQvSd7D0EjG45355tLlV3VOECpq10pLC+8s= +github.com/stretchr/testify v1.7.2/go.mod h1:R6va5+xMeoiuVRoj+gSkQ7d3FALtqAAGI1FQKckRals= +go.opentelemetry.io/auto/sdk v1.2.1 h1:jXsnJ4Lmnqd11kwkBV2LgLoFMZKizbCi5fNZ/ipaZ64= +go.opentelemetry.io/auto/sdk v1.2.1/go.mod h1:KRTj+aOaElaLi+wW1kO/DZRXwkF4C5xPbEe3ZiIhN7Y= +go.opentelemetry.io/otel v1.44.0 h1:JjwHmHpA4iZ3wBxluu2fbbE7j4kqlE8jXyAyPXH7HqU= +go.opentelemetry.io/otel v1.44.0/go.mod h1:BMgjTHL9WPRlRjL2oZCBTL4whCGtXch2H4BhOPIAyYc= +go.opentelemetry.io/otel/metric v1.44.0 h1:1w0gILTcHdr3YI+ixLyjemwrVnsMURbTZFrSYCdDdmc= +go.opentelemetry.io/otel/metric v1.44.0/go.mod h1:8O7hanEPBNgEMmybD3s2VBKcgWOCsA6tzHBPODAiquo= +go.opentelemetry.io/otel/sdk v1.44.0 h1:nHYwb9lK+fJPU/dnT6s7W7Z8itMWyqrnVfbheVYrZ58= +go.opentelemetry.io/otel/sdk v1.44.0/go.mod h1:Osuydd3Se74nqjAKxid74N5eC+jfEqfTegHRnq58oK0= +go.opentelemetry.io/otel/sdk/metric v1.44.0 h1:3LlKgI+VjbVsjNRFZJZAJ30WjXC5VkNRks6si09iEfI= +go.opentelemetry.io/otel/sdk/metric v1.44.0/go.mod h1:5B5pMARnXxKhltooO4xUuCBorl65a4EpnTalObqOigA= +go.opentelemetry.io/otel/trace v1.44.0 h1:jxF5CsGYCe74MCRx2X4g7WsY/VBKRqqpNvXlX/6gtIk= +go.opentelemetry.io/otel/trace v1.44.0/go.mod h1:oLl1jrMQAVo6v3GAggN+1VH9VIz9iUSvW53sW1Q8PIE= golang.org/x/image v0.30.0 h1:jD5RhkmVAnjqaCUXfbGBrn3lpxbknfN9w2UhHHU+5B4= golang.org/x/image v0.30.0/go.mod h1:SAEUTxCCMWSrJcCy/4HwavEsfZZJlYxeHLc6tTiAe/c= -golang.org/x/mod v0.37.0 h1:vF1DjpVEshcIqoEaauuHebaLk1O1forxjxBaVn884JQ= -golang.org/x/mod v0.37.0/go.mod h1:m8S8VeM9r4dzDwjrKO0a1sZP3YjeMamRRlD+fmR2Q/0= -golang.org/x/sync v0.21.0 h1:HLII4xRRTtCRkxYp4HNFF0Js/Og6q2i++KXbg0gHCwM= -golang.org/x/sync v0.21.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= +golang.org/x/mod v0.38.0 h1:MECBjubtXD7yj4HrhIUcywNaGeNVUdfVnxmPajOk4yk= +golang.org/x/mod v0.38.0/go.mod h1:V6Xz0pq8TQ3dGqVQ1FVHuelZpAL0uNhSkk9ogYP3c40= +golang.org/x/net v0.58.0 h1:ynWG7rqYi4ccpTEuPZ2QGWHktVEM9DMCj9yzDE0Q7To= +golang.org/x/net v0.58.0/go.mod h1:YwCddHnFlT7eLQqVprV19OnhLGtc5xOKgE0RyqgfWAU= +golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek= +golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= +golang.org/x/sys v0.0.0-20200116001909-b77594299b42/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= +golang.org/x/sys v0.0.0-20200223170610-d5e6a3e2c0ae/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= +golang.org/x/sys v0.0.0-20210630005230-0f9fa26af87c/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.0.0-20210927094055-39ccf1dd6fa6/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.0.0-20220503163025-988cb79eb6c6/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs= golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= -golang.org/x/text v0.30.0 h1:yznKA/E9zq54KzlzBEAWn1NXSQ8DIp/NYMy88xJjl4k= -golang.org/x/text v0.30.0/go.mod h1:yDdHFIX9t+tORqspjENWgzaCVXgk0yYnYuSZ8UzzBVM= -golang.org/x/tools v0.47.0 h1:7Kn5x/d1svx/PzryTsqeoZN4TZwqeH5pGWjefhLi/1Q= -golang.org/x/tools v0.47.0/go.mod h1:dFHnyTvFWY212G+h7ZY4Vsp/K3U4/7W9TyVaAul8uCA= +golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8= +golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M= +golang.org/x/tools v0.48.0 h1:3+hClM1aLL5mjMKm5ovokw9epgRXPuu2tILgismM6RE= +golang.org/x/tools v0.48.0/go.mod h1:08xX0orndb/F7jJxGDicx061tyd5pcMto75YMAXr6lk= +gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4= +gonum.org/v1/gonum v0.17.0/go.mod h1:El3tOrEuMpv2UdMrbNlKEh9vd86bmQ6vqIcDwxEOc1E= +google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa h1:mZHHdPZl0dbGHCflZgAq/Q468DWVFcU2whhB2KAo8fk= +google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8= +google.golang.org/grpc v1.83.2 h1:EManeRomTObA0BU7I8vXgg/78uE5MJ9M8B39EX2WscU= +google.golang.org/grpc v1.83.2/go.mod h1:YPI1hK3kDked6iHvgX3tR0y+nX/qpMFKhPgFsokw1S8= +google.golang.org/protobuf v1.36.12 h1:pJOKDDOyeXErUroCihFAd5LQuwXBSpVnKGrj5o/fwxc= +google.golang.org/protobuf v1.36.12/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= modernc.org/cc/v4 v4.29.1 h1:MKgdCV3WykTSPqpVrnxdEDS0HEd2FHpKZDzxzU5LyeI= modernc.org/cc/v4 v4.29.1/go.mod h1:OnovgIhbbMXMu1aISnJ0wvVD1KnW+cAUJkIrAWh+kVI= modernc.org/ccgo/v4 v4.34.6 h1:sBgfIwyN0TQ9C5hwIeuqyeAKyMWnbvj2fvpF4L11uzU= diff --git a/server/internal/agent/coordinator.go b/server/internal/agent/coordinator.go index e22a608..fcd4319 100644 --- a/server/internal/agent/coordinator.go +++ b/server/internal/agent/coordinator.go @@ -287,8 +287,22 @@ func (r *Runtime) LoadAgentState(ctx context.Context, runID string) (pkgagent.Ag return pkgagent.AgentState{}, pkgagent.History{}, err } + names := unionNames(run.Config.Profile.Tools.Names, r.extraMethodNames()) + run.Config.Profile.Tools.Names = names + state.Config.Profile.Tools.Names = names tools := tool.Definitions(r.tools) prompt := pkgagent.DefaultSystemPrompt + var hidden []pkgagent.Message + if overlay, err := q.GetRunOverlay(ctx, runID); err == nil { + if overlay.SystemPrompt != "" { + prompt = overlay.SystemPrompt + } + if overlay.Hidden != "" && overlay.Hidden != "[]" { + _ = json.Unmarshal([]byte(overlay.Hidden), &hidden) + } + } else if !errors.Is(err, sql.ErrNoRows) { + return pkgagent.AgentState{}, pkgagent.History{}, err + } hist := pkgagent.History{ Run: run, @@ -297,6 +311,7 @@ func (r *Runtime) LoadAgentState(ctx context.Context, runID string) (pkgagent.Ag Messages: messages, Tools: tools, Prompt: prompt, + Hidden: hidden, MemoryIndexes: r.loadMemoryIndexes(ctx, sess.UserID, sess.WorkspaceID), } return state, hist, nil @@ -1047,3 +1062,44 @@ func containsString(items []string, want string) bool { } return false } + +// unionNames 按出现顺序合并两份工具名,去掉空串和重复。 +func unionNames(left, right []string) []string { + seen := make(map[string]struct{}, len(left)+len(right)) + out := make([]string, 0, len(left)+len(right)) + for _, name := range append(append([]string{}, left...), right...) { + if name == "" { + continue + } + if _, ok := seen[name]; ok { + continue + } + seen[name] = struct{}{} + out = append(out, name) + } + return out +} + +// hasOverlay 表示本轮已经跑过 pre-step 并落过 overlay,恢复时不要再喊插件。 +func (r *Runtime) hasOverlay(ctx context.Context, runID string) bool { + if r == nil || r.q(ctx) == nil || runID == "" { + return false + } + _, err := r.q(ctx).GetRunOverlay(ctx, runID) + return err == nil +} + +// saveOverlay 把 pre-step 改过的提示和隐藏消息写入本轮 overlay。 +func (r *Runtime) saveOverlay(ctx context.Context, runID string, payload pkgagent.PreStepPayload) error { + if r == nil || r.q(ctx) == nil || runID == "" { + return nil + } + hidden := marshalJSON(payload.Hidden) + _, err := r.q(ctx).UpsertRunOverlay(ctx, sqlite.UpsertRunOverlayParams{ + RunID: runID, + SystemPrompt: payload.SystemPrompt, + Hidden: hidden, + UpdatedAt: util.FormatTime(util.Now()), + }) + return err +} diff --git a/server/internal/agent/coordinator_test.go b/server/internal/agent/coordinator_test.go index 7ab4a67..4ea0b81 100644 --- a/server/internal/agent/coordinator_test.go +++ b/server/internal/agent/coordinator_test.go @@ -7,6 +7,7 @@ import ( "codedock/internal/events" "codedock/internal/util" pkgagent "codedock/pkg/agent" + "codedock/pkg/agent/seam" "codedock/pkg/agent/tool" "codedock/pkg/db" "codedock/pkg/db/sqlite" @@ -454,6 +455,221 @@ func TestWorkerSubmitAndCancel(t *testing.T) { t.Fatal("run did not finish") } +func waitRunStatus(t *testing.T, q *sqlite.Queries, ctx context.Context, runID string, want ...pkgagent.RunStatus) pkgagent.RunStatus { + t.Helper() + deadline := time.Now().Add(3 * time.Second) + var last pkgagent.RunStatus + for time.Now().Before(deadline) { + row, err := q.GetRun(ctx, runID) + if err == nil { + last = pkgagent.RunStatus(row.Status) + if len(want) == 0 && pkgagent.IsTerminal(last) { + return last + } + for _, status := range want { + if last == status { + return last + } + } + } + time.Sleep(20 * time.Millisecond) + } + t.Fatalf("run %s did not reach %v last=%s", runID, want, last) + return last +} + +func TestPreStepWritesOverlay(t *testing.T) { + rt, q, ctx := testRuntime(t, true) + sessionID := insertSession(t, q, ctx) + cfg := pkgagent.DefaultYoloConfig(pkgagent.ModelConfig{ + Provider: "fake", + Model: "fake", + Options: mustJSON(pkgagent.FakeOptions{Turns: []pkgagent.FakeTurn{{Text: "ok"}}}), + }) + rt.SetDispatcher(seam.Func(func(_ context.Context, ev seam.Envelope) (seam.Envelope, error) { + if ev.Type == seam.TypePreStep { + ev.Payload = pkgagent.MarshalPayload(pkgagent.PreStepPayload{ + SystemPrompt: "overlay-prompt", + Hidden: []pkgagent.Message{{Role: pkgagent.RoleSystem, Content: pkgagent.EncodeText("hidden-note")}}, + }) + } + return ev, nil + })) + runID, err := rt.CreateAgentState(ctx, sessionID, "go", cfg.Mode, cfg) + if err != nil { + t.Fatal(err) + } + if _, err := rt.ClaimSession(ctx, sessionID, runID); err != nil { + t.Fatal(err) + } + if err := rt.Enqueue(ctx, pkgagent.StepJob{RunID: runID, StepIndex: 1, Phase: pkgagent.PhaseUserInput}); err != nil { + t.Fatal(err) + } + waitRunStatus(t, q, ctx, runID, pkgagent.RunCompleted) + _, hist, err := rt.LoadAgentState(ctx, runID) + if err != nil { + t.Fatal(err) + } + if hist.Prompt != "overlay-prompt" { + t.Fatalf("prompt=%q", hist.Prompt) + } + if len(hist.Hidden) != 1 || pkgagent.DecodeText(hist.Hidden[0].Content) != "hidden-note" { + t.Fatalf("hidden=%+v", hist.Hidden) + } +} + +func TestPreStepDoesNotDuplicateHiddenOnReplay(t *testing.T) { + rt, q, ctx := testRuntime(t, true) + sessionID := insertSession(t, q, ctx) + cfg := pkgagent.DefaultYoloConfig(pkgagent.ModelConfig{ + Provider: "fake", + Model: "fake", + Options: mustJSON(pkgagent.FakeOptions{Turns: []pkgagent.FakeTurn{{Text: "ok"}}}), + }) + rt.SetDispatcher(seam.Func(func(_ context.Context, ev seam.Envelope) (seam.Envelope, error) { + if ev.Type != seam.TypePreStep { + return ev, nil + } + var payload pkgagent.PreStepPayload + _ = json.Unmarshal(ev.Payload, &payload) + payload.Hidden = append(payload.Hidden, pkgagent.Message{ + Role: pkgagent.RoleSystem, + Content: pkgagent.EncodeText("hidden-note"), + }) + ev.Payload = pkgagent.MarshalPayload(payload) + return ev, nil + })) + runID, err := rt.CreateAgentState(ctx, sessionID, "go", cfg.Mode, cfg) + if err != nil { + t.Fatal(err) + } + if _, err := rt.ClaimSession(ctx, sessionID, runID); err != nil { + t.Fatal(err) + } + if err := rt.Enqueue(ctx, pkgagent.StepJob{RunID: runID, StepIndex: 1, Phase: pkgagent.PhaseUserInput}); err != nil { + t.Fatal(err) + } + waitRunStatus(t, q, ctx, runID, pkgagent.RunCompleted) + state, hist, err := rt.LoadAgentState(ctx, runID) + if err != nil { + t.Fatal(err) + } + if len(hist.Hidden) != 1 { + t.Fatalf("first hidden=%+v", hist.Hidden) + } + blocked, replayed, err := rt.worker.applyPreStep(ctx, pkgagent.StepJob{ + RunID: runID, + StepIndex: 1, + Phase: pkgagent.PhaseUserInput, + }, state, hist) + if err != nil || blocked { + t.Fatalf("replay blocked=%v err=%v", blocked, err) + } + if len(replayed.Hidden) != 1 || pkgagent.DecodeText(replayed.Hidden[0].Content) != "hidden-note" { + t.Fatalf("replay hidden=%+v", replayed.Hidden) + } +} + +func TestPreStepBlockedCancelsRun(t *testing.T) { + rt, q, ctx := testRuntime(t, true) + sessionID := insertSession(t, q, ctx) + cfg := pkgagent.DefaultYoloConfig(pkgagent.ModelConfig{ + Provider: "fake", + Model: "fake", + Options: mustJSON(pkgagent.FakeOptions{Turns: []pkgagent.FakeTurn{{Text: "ok"}}}), + }) + rt.SetDispatcher(seam.Func(func(_ context.Context, ev seam.Envelope) (seam.Envelope, error) { + if ev.Type == seam.TypePreStep { + ev.Type = seam.TypeRunBlocked + } + return ev, nil + })) + runID, err := rt.CreateAgentState(ctx, sessionID, "go", cfg.Mode, cfg) + if err != nil { + t.Fatal(err) + } + if _, err := rt.ClaimSession(ctx, sessionID, runID); err != nil { + t.Fatal(err) + } + if err := rt.Enqueue(ctx, pkgagent.StepJob{RunID: runID, StepIndex: 1, Phase: pkgagent.PhaseUserInput}); err != nil { + t.Fatal(err) + } + if got := waitRunStatus(t, q, ctx, runID, pkgagent.RunCancelled); got != pkgagent.RunCancelled { + t.Fatalf("status=%s", got) + } +} + +type extraNamer struct { + seam.Func + names []string +} + +// MethodNames 返回要并进本轮可执行绑定的插件方法名。 +func (e extraNamer) MethodNames() []string { return e.names } + +type extraEcho struct{} + +// Definition 登记 echo 方法,默认 allow,便于测绑定并入。 +func (extraEcho) Definition() tool.Definition { + return tool.Definition{Name: "echo", Prompt: "echo", Permission: tool.Permission{Effect: tool.EffectAllow}} +} + +// Execute 占位成功,不读参数。 +func (extraEcho) Execute(_ context.Context, input tool.Input) (tool.Result, error) { + return tool.Result{CallID: input.Call.ID, Name: "echo", Success: true}, nil +} + +// TestLoadAgentStateMergesExtraMethods 确认插件方法在注册表里模型可见,并入 Names 后才可执行。 +func TestLoadAgentStateMergesExtraMethods(t *testing.T) { + rt, q, ctx := testRuntime(t, false) + if err := rt.Tools().Register(extraEcho{}); err != nil { + t.Fatal(err) + } + sessionID := insertSession(t, q, ctx) + cfg := pkgagent.DefaultYoloConfig(pkgagent.ModelConfig{Provider: "fake", Model: "fake"}) + runID, err := rt.CreateAgentState(ctx, sessionID, "hi", cfg.Mode, cfg) + if err != nil { + t.Fatal(err) + } + _, hist, err := rt.LoadAgentState(ctx, runID) + if err != nil { + t.Fatal(err) + } + found := false + for _, def := range hist.Tools { + if def.Name == "echo" { + found = true + } + } + if !found { + t.Fatal("registered methods stay in the full tool table") + } + if containsString(hist.Run.Config.Profile.Tools.Names, "echo") { + t.Fatal("echo should not be bound without extra names") + } + rt.SetDispatcher(extraNamer{names: []string{"echo"}}) + _, hist, err = rt.LoadAgentState(ctx, runID) + if err != nil { + t.Fatal(err) + } + if !containsString(hist.Run.Config.Profile.Tools.Names, "echo") { + t.Fatalf("names=%+v", hist.Run.Config.Profile.Tools.Names) + } + + askCfg := pkgagent.DefaultRunConfig(pkgagent.WorkAsk, pkgagent.ModelConfig{Provider: "fake", Model: "fake"}) + askID, err := rt.CreateAgentState(ctx, sessionID, "ask", askCfg.Mode, askCfg) + if err != nil { + t.Fatal(err) + } + _, askHist, err := rt.LoadAgentState(ctx, askID) + if err != nil { + t.Fatal(err) + } + if !containsString(askHist.Run.Config.Profile.Tools.Names, "echo") { + t.Fatalf("ask names should still include plugin methods: %+v", askHist.Run.Config.Profile.Tools.Names) + } +} + // TestNilRuntimeGuards 覆盖空 Runtime 上主要入口的防护。 func TestNilRuntimeGuards(t *testing.T) { var rt *Runtime diff --git a/server/internal/agent/runner.go b/server/internal/agent/runner.go index a1cdd7c..2c61341 100644 --- a/server/internal/agent/runner.go +++ b/server/internal/agent/runner.go @@ -8,6 +8,7 @@ import ( agenttools "codedock/internal/agent/tools" "codedock/internal/events" pkgagent "codedock/pkg/agent" + "codedock/pkg/agent/seam" "codedock/pkg/agent/tool" "codedock/pkg/db" "codedock/pkg/db/sqlite" @@ -27,6 +28,13 @@ type Runtime struct { compactWG sync.WaitGroup claimMu sync.Mutex claimedSteps map[string]struct{} + dispatcher seam.Dispatcher + extraMethods methodNamer +} + +// methodNamer 由插件宿主实现,用来把插件方法并进本轮可执行绑定。 +type methodNamer interface { + MethodNames() []string } // New 创建 Runtime 及其 Worker。工具定义在 tools 包注册;ports 只注入工具 Execute 所需的外部实现。 @@ -98,6 +106,37 @@ func (r *Runtime) SetConcurrency(llm, tools int) { r.engine.SetGates(pkgagent.NewSlotLimiter(llm), pkgagent.NewSlotLimiter(tools)) } +// SetDispatcher 注入喊话器,并同步给 Engine。若实现 MethodNames,则记为额外可执行方法。 +func (r *Runtime) SetDispatcher(d seam.Dispatcher) { + if r == nil { + return + } + r.dispatcher = d + if r.engine != nil { + r.engine.SetDispatcher(d) + } + r.extraMethods = nil + if namer, ok := d.(methodNamer); ok { + r.extraMethods = namer + } +} + +// Dispatcher 返回当前喊话器。 +func (r *Runtime) Dispatcher() seam.Dispatcher { + if r == nil { + return nil + } + return r.dispatcher +} + +// extraMethodNames 返回插件挂上的方法名,Load 时并进本轮可执行绑定。 +func (r *Runtime) extraMethodNames() []string { + if r == nil || r.extraMethods == nil { + return nil + } + return r.extraMethods.MethodNames() +} + // Start 启动 Worker。不自动恢复库里未完成的 Job,需用户显式 RecoverRun。 func (r *Runtime) Start(ctx context.Context) { if r.worker != nil { diff --git a/server/internal/agent/worker.go b/server/internal/agent/worker.go index a1fbf75..5aae77e 100644 --- a/server/internal/agent/worker.go +++ b/server/internal/agent/worker.go @@ -2,6 +2,7 @@ package agent import ( "context" + "encoding/json" "fmt" "strings" "sync" @@ -9,6 +10,7 @@ import ( cderr "codedock/internal/errors" pkgagent "codedock/pkg/agent" + "codedock/pkg/agent/seam" ) // stepJobKey 返回 StepJob 的唯一去重键:run_id + step_index。 @@ -183,7 +185,15 @@ func (w *Worker) execute(parent context.Context, job pkgagent.StepJob) { w.mu.Lock() _, skipped := w.skipped[runID] w.mu.Unlock() - if skipped || state.CancelRequested { + blocked := false + if !skipped && !state.CancelRequested && job.Phase == pkgagent.PhaseUserInput { + var err error + blocked, history, err = w.applyPreStep(ctx, job, state, history) + if err != nil { + blocked = true + } + } + if skipped || state.CancelRequested || blocked { if !pkgagent.IsTerminal(state.Status) { result, ferr := w.runtime.engine.Step(ctx, pkgagent.StepInput{ State: cancelState(state), @@ -248,3 +258,44 @@ func failState(state pkgagent.AgentState, err error) pkgagent.AgentState { func timeNow() time.Time { return time.Now().UTC() } + +// applyPreStep 在首拍把系统提示和隐藏消息递给插件;换向或出错则取消本轮。 +// 本轮已有 overlay 时不再喊插件,避免恢复或重入第一步时把 Hidden 再追加一遍。 +func (w *Worker) applyPreStep(ctx context.Context, job pkgagent.StepJob, state pkgagent.AgentState, history pkgagent.History) (bool, pkgagent.History, error) { + if w == nil || w.runtime == nil || w.runtime.Dispatcher() == nil { + return false, history, nil + } + if w.runtime.hasOverlay(ctx, job.RunID) { + return false, history, nil + } + ev, err := seam.Dispatch(ctx, w.runtime.Dispatcher(), seam.Envelope{ + Type: seam.TypePreStep, + SessionID: state.SessionID, + RunID: job.RunID, + Payload: pkgagent.MarshalPayload(pkgagent.PreStepPayload{ + SystemPrompt: history.Prompt, + Hidden: history.Hidden, + }), + }) + if err != nil { + w.runtime.logger().Info("pre-step blocked", "run_id", job.RunID, "error", err) + return true, history, err + } + if ev.Type == seam.TypeRunBlocked { + w.runtime.logger().Info("pre-step blocked", "run_id", job.RunID, "type", ev.Type) + return true, history, nil + } + if ev.Type != seam.TypePreStep || len(ev.Payload) == 0 { + return false, history, nil + } + var payload pkgagent.PreStepPayload + if err := json.Unmarshal(ev.Payload, &payload); err != nil { + return false, history, nil + } + if err := w.runtime.saveOverlay(ctx, job.RunID, payload); err != nil { + w.runtime.logger().Error("save overlay failed", "run_id", job.RunID, "error", err) + } + history.Prompt = payload.SystemPrompt + history.Hidden = payload.Hidden + return false, history, nil +} diff --git a/server/internal/config/config.go b/server/internal/config/config.go index 73c2efe..475ff1f 100644 --- a/server/internal/config/config.go +++ b/server/internal/config/config.go @@ -5,39 +5,44 @@ import ( "path/filepath" "strconv" "strings" + "time" ) // Config 是进程启动时一次性读取的环境配置。 type Config struct { - HTTPAddr string // HTTP 监听地址 - LogLevel string - DBEngine string - DBDSN string - LLMProvider string - LLMModel string - LLMAPIKey string - LLMBaseURL string - GitRepo string // 默认仓库根;会话未指定工作目录时回落到这里,再否则 cwd - LLMConcurrency int // 进程内同时进行的模型调用上限;0 表示不限制 - ToolConcurrency int // 进程内同时执行的工具调用上限;0 表示不限制 - CodexBin string // 本机 Codex CLI;未安装时主服务仍可启动 + HTTPAddr string // HTTP 监听地址 + LogLevel string + DBEngine string + DBDSN string + LLMProvider string + LLMModel string + LLMAPIKey string + LLMBaseURL string + GitRepo string // 默认仓库根;会话未指定工作目录时回落到这里,再否则 cwd + LLMConcurrency int // 进程内同时进行的模型调用上限;0 表示不限制 + ToolConcurrency int // 进程内同时执行的工具调用上限;0 表示不限制 + PluginDir string // 插件目录;空则不拉进程 + PluginRPCTimeout time.Duration // 单次插件 RPC 超时 + CodexBin string // 本机 Codex CLI;未安装时主服务仍可启动 } // Load 从环境变量读取配置,未设置时使用默认值。 func Load() Config { return Config{ - HTTPAddr: env("HTTP_ADDR", ":8080"), - LogLevel: env("LOG_LEVEL", "debug"), - DBEngine: env("DB_ENGINE", "sqlite"), - DBDSN: env("DB_DSN", defaultSQLiteDSN()), - LLMProvider: env("LLM_PROVIDER", "fake"), - LLMModel: env("LLM_MODEL", "fake"), - LLMAPIKey: env("LLM_API_KEY", ""), - LLMBaseURL: env("LLM_BASE_URL", ""), - GitRepo: env("GIT_REPO", ""), - LLMConcurrency: envInt("LLM_CONCURRENCY", 4), - ToolConcurrency: envInt("TOOL_CONCURRENCY", 8), - CodexBin: env("CODEX_BIN", "codex"), + HTTPAddr: env("HTTP_ADDR", ":8080"), + LogLevel: env("LOG_LEVEL", "debug"), + DBEngine: env("DB_ENGINE", "sqlite"), + DBDSN: env("DB_DSN", defaultSQLiteDSN()), + LLMProvider: env("LLM_PROVIDER", "fake"), + LLMModel: env("LLM_MODEL", "fake"), + LLMAPIKey: env("LLM_API_KEY", ""), + LLMBaseURL: env("LLM_BASE_URL", ""), + GitRepo: env("GIT_REPO", ""), + LLMConcurrency: envInt("LLM_CONCURRENCY", 4), + ToolConcurrency: envInt("TOOL_CONCURRENCY", 8), + PluginDir: env("PLUGIN_DIR", ""), + PluginRPCTimeout: envDuration("PLUGIN_RPC_TIMEOUT", 10*time.Second), + CodexBin: env("CODEX_BIN", "codex"), } } @@ -51,6 +56,7 @@ func DataDir() string { return filepath.Join(repoRoot(), "data") } +// repoRoot 从 cwd 向上找仓根(有 pnpm-workspace.yaml 或 AGENTS.md+server/go.mod)。 func repoRoot() string { cwd, err := os.Getwd() if err != nil { @@ -70,6 +76,7 @@ func repoRoot() string { return cwd } +// isRepoRoot 判断 dir 是否是本仓根。 func isRepoRoot(dir string) bool { if _, err := os.Stat(filepath.Join(dir, "pnpm-workspace.yaml")); err == nil { return true @@ -106,6 +113,7 @@ func env(key, fallback string) string { return value } +// envInt 读整数环境变量,未设或解析失败时用 fallback。 func envInt(key string, fallback int) int { value := os.Getenv(key) if value == "" { @@ -117,3 +125,16 @@ func envInt(key string, fallback int) int { } return n } + +// envDuration 读 duration 环境变量,未设或解析失败时用 fallback。 +func envDuration(key string, fallback time.Duration) time.Duration { + value := os.Getenv(key) + if value == "" { + return fallback + } + d, err := time.ParseDuration(value) + if err != nil { + return fallback + } + return d +} diff --git a/server/internal/config/config_test.go b/server/internal/config/config_test.go index e8e52f5..8ceaafe 100644 --- a/server/internal/config/config_test.go +++ b/server/internal/config/config_test.go @@ -5,6 +5,7 @@ import ( "path/filepath" "strings" "testing" + "time" ) // TestLoadDefaults 校验未设置环境变量时的默认配置。 @@ -18,6 +19,8 @@ func TestLoadDefaults(t *testing.T) { t.Setenv("LLM_API_KEY", "") t.Setenv("LLM_BASE_URL", "") t.Setenv("GIT_REPO", "") + t.Setenv("PLUGIN_DIR", "") + t.Setenv("PLUGIN_RPC_TIMEOUT", "") t.Setenv("CODEX_BIN", "") cfg := Load() @@ -49,6 +52,12 @@ func TestLoadDefaults(t *testing.T) { if cfg.ToolConcurrency != 8 { t.Fatalf("ToolConcurrency = %d, want 8", cfg.ToolConcurrency) } + if cfg.PluginDir != "" { + t.Fatalf("PluginDir = %q, want empty", cfg.PluginDir) + } + if cfg.PluginRPCTimeout != 10*time.Second { + t.Fatalf("PluginRPCTimeout = %s, want 10s", cfg.PluginRPCTimeout) + } if cfg.CodexBin != "codex" { t.Fatalf("CodexBin = %q, want codex", cfg.CodexBin) } @@ -65,6 +74,8 @@ func TestLoadFromEnv(t *testing.T) { t.Setenv("LLM_API_KEY", "sk-test") t.Setenv("LLM_BASE_URL", "https://api.example.com/v1") t.Setenv("GIT_REPO", "/tmp/repo") + t.Setenv("PLUGIN_DIR", "/tmp/plugins") + t.Setenv("PLUGIN_RPC_TIMEOUT", "2s") t.Setenv("CODEX_BIN", "/usr/local/bin/codex") cfg := Load() @@ -77,6 +88,9 @@ func TestLoadFromEnv(t *testing.T) { if cfg.GitRepo != "/tmp/repo" { t.Fatalf("GitRepo = %q, want /tmp/repo", cfg.GitRepo) } + if cfg.PluginDir != "/tmp/plugins" || cfg.PluginRPCTimeout != 2*time.Second { + t.Fatalf("plugin cfg = %+v", cfg) + } if cfg.CodexBin != "/usr/local/bin/codex" { t.Fatalf("CodexBin = %q, want /usr/local/bin/codex", cfg.CodexBin) } diff --git a/server/internal/handler/hello_plugin_test.go b/server/internal/handler/hello_plugin_test.go new file mode 100644 index 0000000..dd1bef2 --- /dev/null +++ b/server/internal/handler/hello_plugin_test.go @@ -0,0 +1,221 @@ +package handler_test + +import ( + "context" + "encoding/json" + "net/http" + "os" + "os/exec" + "path/filepath" + "strings" + "testing" + "time" + + "codedock/internal/handler" + "codedock/internal/pluginhost" + pkgagent "codedock/pkg/agent" +) + +// TestLoopHelloPlugin 对照示例:/skip 不建 Run,world 改正文并注入 Hidden,hello 方法可执行。 +func TestLoopHelloPlugin(t *testing.T) { + if testing.Short() { + t.Skip("hello plugin") + } + dir := t.TempDir() + buildHelloPlugin(t, filepath.Join(dir, "hello", "hello")) + + f := newFixture(t) + host, err := pluginhost.Load(context.Background(), pluginhost.Options{ + Dir: dir, + Timeout: 3 * time.Second, + Registry: f.runtime.Tools(), + Queries: f.queries, + Model: pkgagent.ModelConfig{Provider: "fake", Model: "fake"}, + }) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = host.Close() }) + f.runtime.SetDispatcher(host) + + sessionID := f.createSession(t) + rec := f.do(t, http.MethodPost, "/sessions/"+sessionID+"/runs", handler.StartRunRequest{ + Content: "/skip this", + Mode: pkgagent.WorkAgent, + }) + if rec.Code != http.StatusOK { + t.Fatalf("skip %d %s", rec.Code, rec.Body.String()) + } + var skipped handler.StartRunResponse + if err := json.Unmarshal(rec.Body.Bytes(), &skipped); err != nil { + t.Fatal(err) + } + if !skipped.Handled || skipped.RunID != "" { + t.Fatalf("skip resp=%+v", skipped) + } + + runID := f.start(t, sessionID, handler.StartRunRequest{ + Content: "world", + Mode: pkgagent.WorkAgent, + Config: withFake(pkgagent.DefaultYoloConfig(pkgagent.ModelConfig{}), pkgagent.FakeOptions{ + Turns: []pkgagent.FakeTurn{{Text: "ok"}}, + }), + }) + f.waitRun(t, runID, pkgagent.RunCompleted) + msgs := listMessages(t, f, sessionID, "") + var userText string + for _, msg := range msgs.Messages { + if msg.Role == pkgagent.RoleUser { + userText = pkgagent.DecodeText(msg.Content) + break + } + } + if userText != "[hello] world" { + t.Fatalf("user message=%q all=%+v", userText, msgs.Messages) + } + _, hist, err := f.runtime.LoadAgentState(context.Background(), runID) + if err != nil { + t.Fatal(err) + } + if len(hist.Hidden) != 1 || pkgagent.DecodeText(hist.Hidden[0].Content) == "" { + t.Fatalf("overlay hidden=%+v", hist.Hidden) + } + found := false + for _, def := range hist.Tools { + if def.Name == "hello" { + found = true + } + } + if !found { + t.Fatalf("hello method not visible: %+v", hist.Tools) + } + bound := false + for _, name := range hist.Run.Config.Profile.Tools.Names { + if name == "hello" { + bound = true + } + } + if !bound { + t.Fatalf("hello should be executable: names=%v", hist.Run.Config.Profile.Tools.Names) + } +} + +// TestLoopHelloPluginMethodsAndDeny 确认 hello 能被模型调用,forbidden 参数会被否决。 +func TestLoopHelloPluginMethodsAndDeny(t *testing.T) { + if testing.Short() { + t.Skip("hello plugin") + } + dir := t.TempDir() + buildHelloPlugin(t, filepath.Join(dir, "hello", "hello")) + + f := newFixture(t) + host, err := pluginhost.Load(context.Background(), pluginhost.Options{ + Dir: dir, + Timeout: 3 * time.Second, + Registry: f.runtime.Tools(), + Queries: f.queries, + Model: pkgagent.ModelConfig{Provider: "fake", Model: "fake"}, + }) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = host.Close() }) + f.runtime.SetDispatcher(host) + + sessionID := f.createSession(t) + skip := f.do(t, http.MethodPost, "/sessions/"+sessionID+"/runs", handler.StartRunRequest{ + Content: "/skip notice", + Mode: pkgagent.WorkAgent, + }) + if skip.Code != http.StatusOK { + t.Fatalf("skip %d %s", skip.Code, skip.Body.String()) + } + notices := listMessages(t, f, sessionID, "") + sawNotice := false + for _, msg := range notices.Messages { + if msg.Role == pkgagent.RoleSystem && pkgagent.DecodeText(msg.Content) == "hello 已跳过本次对话。" { + sawNotice = true + } + } + if !sawNotice { + t.Fatalf("missing skip notice: %+v", notices.Messages) + } + + callID := f.start(t, sessionID, handler.StartRunRequest{ + Content: "call hello", + Mode: pkgagent.WorkAgent, + Config: withFake(pkgagent.DefaultYoloConfig(pkgagent.ModelConfig{}), pkgagent.FakeOptions{ + Turns: []pkgagent.FakeTurn{ + {ToolCalls: []pkgagent.FakeToolCall{{ + Name: "hello", + Arguments: json.RawMessage(`{"text":"hi"}`), + }}}, + {Text: "hello-done"}, + }, + }), + }) + f.waitRun(t, callID, pkgagent.RunCompleted) + called := listMessages(t, f, sessionID, "") + var sawHello bool + for _, msg := range called.Messages { + if msg.Role == pkgagent.RoleTool && strings.Contains(string(msg.Content), `"text":"hi"`) { + sawHello = true + } + } + if !sawHello { + t.Fatalf("hello method should execute: %+v", called.Messages) + } + + denyID := f.start(t, sessionID, handler.StartRunRequest{ + Content: "deny ping", + Mode: pkgagent.WorkAgent, + Config: withFake(pkgagent.DefaultYoloConfig(pkgagent.ModelConfig{}), pkgagent.FakeOptions{ + Turns: []pkgagent.FakeTurn{ + {ToolCalls: []pkgagent.FakeToolCall{{ + Name: "ping", + Arguments: json.RawMessage(`{"x":"forbidden"}`), + }}}, + {Text: "after-deny"}, + }, + }), + }) + f.waitRun(t, denyID, pkgagent.RunCompleted) + denied := listMessages(t, f, sessionID, "") + var sawDenied bool + for _, msg := range denied.Messages { + if msg.Role == pkgagent.RoleTool && (strings.Contains(pkgagent.DecodeText(msg.Content), "denied by plugin") || strings.Contains(string(msg.Content), "denied by plugin")) { + sawDenied = true + } + } + if !sawDenied { + t.Fatalf("forbidden ping should be denied: %+v", denied.Messages) + } +} + +func buildHelloPlugin(t *testing.T, dest string) { + t.Helper() + if err := os.MkdirAll(filepath.Dir(dest), 0o755); err != nil { + t.Fatal(err) + } + dir, err := os.Getwd() + if err != nil { + t.Fatal(err) + } + root := dir + for i := 0; i < 6; i++ { + if _, err := os.Stat(filepath.Join(root, "go.mod")); err == nil { + break + } + parent := filepath.Dir(root) + if parent == root { + t.Fatal("go.mod not found") + } + root = parent + } + cmd := exec.Command("go", "build", "-o", dest, ".") + cmd.Dir = filepath.Join(root, "internal", "pluginhost", "testdata", "hello") + cmd.Env = append(os.Environ(), "CGO_ENABLED=0") + if out, err := cmd.CombinedOutput(); err != nil { + t.Fatalf("build hello: %v\n%s", err, out) + } +} diff --git a/server/internal/handler/loop_test.go b/server/internal/handler/loop_test.go index c8edc5c..ce1f080 100644 --- a/server/internal/handler/loop_test.go +++ b/server/internal/handler/loop_test.go @@ -3,12 +3,14 @@ package handler_test import ( "context" "encoding/json" + "fmt" "net/http" "strings" "testing" "codedock/internal/handler" pkgagent "codedock/pkg/agent" + "codedock/pkg/agent/seam" ) func TestLoopPlainText(t *testing.T) { @@ -238,6 +240,83 @@ func TestLoopModeSmoke(t *testing.T) { } } +func TestLoopStartInputHandled(t *testing.T) { + f := newFixture(t) + sessionID := f.createSession(t) + f.runtime.SetDispatcher(seam.Func(func(_ context.Context, ev seam.Envelope) (seam.Envelope, error) { + if ev.Type == seam.TypeInput { + ev.Type = seam.TypeInputHandled + } + return ev, nil + })) + rec := f.do(t, http.MethodPost, "/sessions/"+sessionID+"/runs", handler.StartRunRequest{ + Content: "skip me", + Mode: pkgagent.WorkAgent, + }) + if rec.Code != http.StatusOK { + t.Fatalf("handled start %d %s", rec.Code, rec.Body.String()) + } + var resp handler.StartRunResponse + if err := json.Unmarshal(rec.Body.Bytes(), &resp); err != nil { + t.Fatal(err) + } + if !resp.Handled || resp.RunID != "" || resp.SessionID != sessionID { + t.Fatalf("resp=%+v", resp) + } + f.runtime.SetDispatcher(nil) + second := f.start(t, sessionID, handler.StartRunRequest{ + Content: "now start", + Mode: pkgagent.WorkAgent, + Config: withFake(pkgagent.DefaultYoloConfig(pkgagent.ModelConfig{}), pkgagent.FakeOptions{ + Turns: []pkgagent.FakeTurn{{Text: "ok"}}, + }), + }) + f.waitRun(t, second, pkgagent.RunCompleted) +} + +func TestLoopStartInputRewritesContent(t *testing.T) { + f := newFixture(t) + sessionID := f.createSession(t) + f.runtime.SetDispatcher(seam.Func(func(_ context.Context, ev seam.Envelope) (seam.Envelope, error) { + if ev.Type != seam.TypeInput { + return ev, nil + } + var payload pkgagent.InputPayload + _ = json.Unmarshal(ev.Payload, &payload) + payload.Content = "rewritten" + ev.Payload = pkgagent.MarshalPayload(payload) + return ev, nil + })) + runID := f.start(t, sessionID, handler.StartRunRequest{ + Content: "original", + Mode: pkgagent.WorkAgent, + Config: withFake(pkgagent.DefaultYoloConfig(pkgagent.ModelConfig{}), pkgagent.FakeOptions{ + Turns: []pkgagent.FakeTurn{{Text: "ok"}}, + }), + }) + f.waitRun(t, runID, pkgagent.RunCompleted) + msgs := listMessages(t, f, sessionID, "") + if len(msgs.Messages) == 0 || pkgagent.DecodeText(msgs.Messages[0].Content) != "rewritten" { + t.Fatalf("messages=%+v", msgs.Messages) + } +} + +func TestLoopStartInputDispatchError(t *testing.T) { + f := newFixture(t) + sessionID := f.createSession(t) + f.runtime.SetDispatcher(seam.Func(func(context.Context, seam.Envelope) (seam.Envelope, error) { + return seam.Envelope{}, fmt.Errorf("plugin down") + })) + rec := f.do(t, http.MethodPost, "/sessions/"+sessionID+"/runs", handler.StartRunRequest{ + Content: "x", + Mode: pkgagent.WorkAgent, + }) + if rec.Code != http.StatusServiceUnavailable { + t.Fatalf("expected 503, got %d %s", rec.Code, rec.Body.String()) + } +} + +// TestLoopStartWhileActiveConflicts 确认已有活跃 Run 时再开一轮 409,但插件 handled 不算冲突。 func TestLoopStartWhileActiveConflicts(t *testing.T) { f := newFixture(t) sessionID := f.createSession(t) @@ -257,6 +336,27 @@ func TestLoopStartWhileActiveConflicts(t *testing.T) { if rec.Code != http.StatusConflict { t.Fatalf("expected 409, got %d %s", rec.Code, rec.Body.String()) } + + f.runtime.SetDispatcher(seam.Func(func(_ context.Context, ev seam.Envelope) (seam.Envelope, error) { + if ev.Type == seam.TypeInput { + ev.Type = seam.TypeInputHandled + } + return ev, nil + })) + handled := f.do(t, http.MethodPost, "/sessions/"+sessionID+"/runs", handler.StartRunRequest{ + Content: "plugin takes this", + Mode: pkgagent.WorkAgent, + }) + if handled.Code != http.StatusOK { + t.Fatalf("handled during active %d %s", handled.Code, handled.Body.String()) + } + var resp handler.StartRunResponse + if err := json.Unmarshal(handled.Body.Bytes(), &resp); err != nil { + t.Fatal(err) + } + if !resp.Handled || resp.RunID != "" { + t.Fatalf("handled during active resp=%+v", resp) + } } func startPingApproval(t *testing.T, f *fixture, sessionID string) string { diff --git a/server/internal/handler/redact_plugin_test.go b/server/internal/handler/redact_plugin_test.go new file mode 100644 index 0000000..83f539e --- /dev/null +++ b/server/internal/handler/redact_plugin_test.go @@ -0,0 +1,141 @@ +package handler_test + +import ( + "context" + "encoding/json" + "net/http" + "os" + "os/exec" + "path/filepath" + "strings" + "testing" + "time" + + "codedock/internal/handler" + "codedock/internal/pluginhost" + pkgagent "codedock/pkg/agent" +) + +const redactSecret = "super-secret-value-not-a-shape" +const redactAWS = "AKIAIOSFODNN7EXAMPLE" + +// TestLoopRedactPlugin 对照脱敏示例:read .env 和用户粘贴的 key 落库都没有原文。 +func TestLoopRedactPlugin(t *testing.T) { + if testing.Short() { + t.Skip("redact plugin") + } + pluginDir := t.TempDir() + buildRedactPlugin(t, filepath.Join(pluginDir, "redact", "redact")) + + ws := t.TempDir() + if err := os.WriteFile(filepath.Join(ws, ".env"), []byte("API_KEY="+redactSecret+"\n"), 0o600); err != nil { + t.Fatal(err) + } + + f := newFixture(t) + host, err := pluginhost.Load(context.Background(), pluginhost.Options{ + Dir: pluginDir, + Timeout: 3 * time.Second, + Registry: f.runtime.Tools(), + Queries: f.queries, + Model: pkgagent.ModelConfig{Provider: "fake", Model: "fake"}, + }) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = host.Close() }) + f.runtime.SetDispatcher(host) + + rec := f.do(t, http.MethodPost, "/sessions", handler.CreateSessionRequest{ + UserID: "u1", TenantID: "t1", WorkspaceID: ws, + }) + if rec.Code != http.StatusOK { + t.Fatalf("create session %d %s", rec.Code, rec.Body.String()) + } + var created handler.SessionResponse + if err := json.Unmarshal(rec.Body.Bytes(), &created); err != nil { + t.Fatal(err) + } + sessionID := created.Session.ID + + readID := f.start(t, sessionID, handler.StartRunRequest{ + Content: "read env", + Mode: pkgagent.WorkAgent, + Config: withFake(pkgagent.DefaultYoloConfig(pkgagent.ModelConfig{}), pkgagent.FakeOptions{ + Turns: []pkgagent.FakeTurn{ + {ToolCalls: []pkgagent.FakeToolCall{{ + Name: "read", + Arguments: json.RawMessage(`{"path":".env"}`), + }}}, + {Text: "done"}, + }, + }), + }) + f.waitRun(t, readID, pkgagent.RunCompleted) + msgs := listMessages(t, f, sessionID, "") + var toolBody string + for _, msg := range msgs.Messages { + if msg.Role == pkgagent.RoleTool { + toolBody = string(msg.Content) + } + } + if toolBody == "" { + t.Fatalf("missing tool message: %+v", msgs.Messages) + } + if strings.Contains(toolBody, redactSecret) { + t.Fatalf("secret leaked in tool message: %s", toolBody) + } + if !strings.Contains(toolBody, "API_KEY=") || !strings.Contains(toolBody, "«REDACTED:assignment:") { + t.Fatalf("want keyed placeholder: %s", toolBody) + } + + pasteID := f.start(t, sessionID, handler.StartRunRequest{ + Content: "use " + redactAWS, + Mode: pkgagent.WorkAgent, + Config: withFake(pkgagent.DefaultYoloConfig(pkgagent.ModelConfig{}), pkgagent.FakeOptions{ + Turns: []pkgagent.FakeTurn{{Text: "ok"}}, + }), + }) + f.waitRun(t, pasteID, pkgagent.RunCompleted) + pasted := listMessages(t, f, sessionID, "") + var userText string + for _, msg := range pasted.Messages { + if msg.Role == pkgagent.RoleUser { + text := pkgagent.DecodeText(msg.Content) + if strings.Contains(text, "use ") || strings.Contains(text, "REDACTED") { + userText = text + } + } + } + if userText == "" || strings.Contains(userText, redactAWS) || !strings.Contains(userText, "«REDACTED:aws_ak:") { + t.Fatalf("user paste %q all=%+v", userText, pasted.Messages) + } +} + +func buildRedactPlugin(t *testing.T, dest string) { + t.Helper() + if err := os.MkdirAll(filepath.Dir(dest), 0o755); err != nil { + t.Fatal(err) + } + dir, err := os.Getwd() + if err != nil { + t.Fatal(err) + } + root := dir + for i := 0; i < 6; i++ { + if _, err := os.Stat(filepath.Join(root, "go.mod")); err == nil { + break + } + parent := filepath.Dir(root) + if parent == root { + t.Fatal("go.mod not found") + } + root = parent + } + cmd := exec.Command("go", "build", "-o", dest, ".") + cmd.Dir = filepath.Join(filepath.Dir(root), "plugin", "redact") + cmd.Env = append(os.Environ(), "CGO_ENABLED=0") + if out, err := cmd.CombinedOutput(); err != nil { + t.Fatalf("build redact: %v\n%s", err, out) + } +} diff --git a/server/internal/handler/run.go b/server/internal/handler/run.go index 5abb49b..1691d66 100644 --- a/server/internal/handler/run.go +++ b/server/internal/handler/run.go @@ -2,6 +2,7 @@ package handler import ( "context" + "encoding/json" "net/http" "github.com/go-chi/chi/v5" @@ -9,6 +10,7 @@ import ( cderr "codedock/internal/errors" pkgagent "codedock/pkg/agent" "codedock/pkg/agent/profile" + "codedock/pkg/agent/seam" ) type StartRunRequest struct { @@ -20,7 +22,8 @@ type StartRunRequest struct { type StartRunResponse struct { SessionID string `json:"session_id"` - RunID string `json:"run_id"` + RunID string `json:"run_id,omitempty"` + Handled bool `json:"handled,omitempty"` } type RunResponse struct { @@ -133,6 +136,34 @@ func (a *API) start(ctx context.Context, sessionID string, req StartRunRequest) } config.Profile = profile.For(string(config.Mode)) + input, err := seam.Dispatch(ctx, a.runtime.Dispatcher(), seam.Envelope{ + Type: seam.TypeInput, + SessionID: sessionID, + Payload: pkgagent.MarshalPayload(pkgagent.InputPayload{ + Content: req.Content, + Mode: req.Mode, + }), + }) + if err != nil { + return StartRunResponse{}, cderr.Unavailable("%s", err.Error()) + } + if input.Type == seam.TypeInputHandled { + return StartRunResponse{SessionID: sessionID, Handled: true}, nil + } + if input.Type == seam.TypeInput && len(input.Payload) > 0 { + var payload pkgagent.InputPayload + if err := json.Unmarshal(input.Payload, &payload); err == nil { + if payload.Content != "" { + req.Content = payload.Content + } + if payload.Mode != "" { + req.Mode = payload.Mode + config.Mode = payload.Mode + config.Profile = profile.For(string(payload.Mode)) + } + } + } + if session.ActiveRunID != nil && *session.ActiveRunID != "" { return StartRunResponse{}, cderr.Conflict("session already has an active run") } diff --git a/server/internal/pluginhost/context.go b/server/internal/pluginhost/context.go new file mode 100644 index 0000000..987e570 --- /dev/null +++ b/server/internal/pluginhost/context.go @@ -0,0 +1,159 @@ +package pluginhost + +import ( + "encoding/json" + + pkgagent "codedock/pkg/agent" + "codedock/pkg/agent/seam" + sdk "codedock/pkg/plugin" +) + +// attachPluginContext 把已存的袋子填进信封;信封上已有合法键则盖在已存之上。 +func (h *Host) attachPluginContext(ev seam.Envelope) seam.Envelope { + if h == nil { + return ev + } + stored := h.lookupPluginContext(ev.SessionID, ev.RunID) + if len(ev.Context) > 0 && sdk.AcceptableContext(ev.Context) { + bag := sdk.DecodePluginContext(stored) + incoming := sdk.DecodePluginContext(ev.Context) + for key, value := range incoming.Values { + bag.Values[key] = value + } + ev.Context = bag.Raw() + return ev + } + if len(stored) > 0 { + ev.Context = append(json.RawMessage(nil), stored...) + return ev + } + if len(ev.Context) == 0 { + ev.Context = json.RawMessage(`{}`) + } + return ev +} + +// persistPluginContext 按 Run(没有则按会话)保存袋子;超限或非法则丢掉这次写入。 +func (h *Host) persistPluginContext(ev seam.Envelope) { + if h == nil { + return + } + raw := ev.Context + if len(raw) == 0 { + raw = json.RawMessage(`{}`) + } + if !sdk.AcceptableContext(raw) { + if h.log != nil { + h.log.Warn("plugin context not persisted", "bytes", len(raw), "session_id", ev.SessionID, "run_id", ev.RunID) + } + return + } + key := pluginContextKey(ev.SessionID, ev.RunID) + if key == "" { + return + } + h.ctxMu.Lock() + defer h.ctxMu.Unlock() + if h.contexts == nil { + h.contexts = map[string]json.RawMessage{} + } + h.contexts[key] = append(json.RawMessage(nil), raw...) + if ev.RunID != "" && ev.SessionID != "" { + delete(h.contexts, sessionPluginContextKey(ev.SessionID)) + } +} + +// lookupPluginContext 取出已存袋子;有 Run 但还没有对应条目时,把会话袋迁过去。 +func (h *Host) lookupPluginContext(sessionID, runID string) json.RawMessage { + if h == nil { + return nil + } + h.ctxMu.Lock() + defer h.ctxMu.Unlock() + if h.contexts == nil { + return nil + } + if runID != "" { + if raw, ok := h.contexts[runPluginContextKey(runID)]; ok { + return append(json.RawMessage(nil), raw...) + } + if sessionID != "" { + if raw, ok := h.contexts[sessionPluginContextKey(sessionID)]; ok { + copied := append(json.RawMessage(nil), raw...) + h.contexts[runPluginContextKey(runID)] = copied + delete(h.contexts, sessionPluginContextKey(sessionID)) + return append(json.RawMessage(nil), copied...) + } + } + return nil + } + if sessionID == "" { + return nil + } + raw, ok := h.contexts[sessionPluginContextKey(sessionID)] + if !ok { + return nil + } + return append(json.RawMessage(nil), raw...) +} + +// dropSessionPluginContext 丢掉还没挂到 Run 的会话袋。 +func (h *Host) dropSessionPluginContext(sessionID string) { + if h == nil || sessionID == "" { + return + } + h.ctxMu.Lock() + defer h.ctxMu.Unlock() + if h.contexts != nil { + delete(h.contexts, sessionPluginContextKey(sessionID)) + } +} + +// dropPluginContext 丢掉该 Run 和该会话上的袋子。 +func (h *Host) dropPluginContext(runID, sessionID string) { + if h == nil { + return + } + h.ctxMu.Lock() + defer h.ctxMu.Unlock() + if h.contexts == nil { + return + } + if runID != "" { + delete(h.contexts, runPluginContextKey(runID)) + } + if sessionID != "" { + delete(h.contexts, sessionPluginContextKey(sessionID)) + } +} + +// pluginContextKey 优先用 Run,否则用会话。 +func pluginContextKey(sessionID, runID string) string { + if runID != "" { + return runPluginContextKey(runID) + } + if sessionID != "" { + return sessionPluginContextKey(sessionID) + } + return "" +} + +// sessionPluginContextKey 是建 Run 之前的袋子键。 +func sessionPluginContextKey(sessionID string) string { + return "session:" + sessionID +} + +// runPluginContextKey 是本轮 Run 的袋子键。 +func runPluginContextKey(runID string) string { + return "run:" + runID +} + +// isTerminalLedger 判断账本事件是否表示 Run 已结束。 +func isTerminalLedger(typ string) bool { + switch typ { + case string(pkgagent.EventRunCompleted), string(pkgagent.EventRunFailed), string(pkgagent.EventRunCancelled): + return true + default: + return false + } +} diff --git a/server/internal/pluginhost/context_test.go b/server/internal/pluginhost/context_test.go new file mode 100644 index 0000000..f262d18 --- /dev/null +++ b/server/internal/pluginhost/context_test.go @@ -0,0 +1,86 @@ +package pluginhost + +import ( + "context" + "encoding/json" + "strings" + "testing" + + pkgagent "codedock/pkg/agent" + "codedock/pkg/agent/seam" + sdk "codedock/pkg/plugin" +) + +// TestHostPluginContextPersistsAcrossSeams 确认 input 写下的键在 pre-step(已有 Run)还能读到。 +func TestHostPluginContextPersistsAcrossSeams(t *testing.T) { + t.Parallel() + h := &Host{contexts: map[string]json.RawMessage{}} + first, err := h.Dispatch(context.Background(), seam.Envelope{ + Type: seam.TypeInput, + SessionID: "s1", + Context: json.RawMessage(`{"hello.marked":true}`), + }) + if err != nil { + t.Fatal(err) + } + if string(sdk.DecodePluginContext(first.Context).Get("hello.marked")) != "true" { + t.Fatalf("input context=%s", first.Context) + } + + second, err := h.Dispatch(context.Background(), seam.Envelope{ + Type: seam.TypePreStep, + SessionID: "s1", + RunID: "r1", + }) + if err != nil { + t.Fatal(err) + } + if string(sdk.DecodePluginContext(second.Context).Get("hello.marked")) != "true" { + t.Fatalf("pre-step context=%s", second.Context) + } +} + +// TestHostPluginContextRejectedKeepsPrevious 确认超限袋子不会盖掉上一份。 +func TestHostPluginContextRejectedKeepsPrevious(t *testing.T) { + t.Parallel() + h := &Host{contexts: map[string]json.RawMessage{}} + if _, err := h.Dispatch(context.Background(), seam.Envelope{ + Type: seam.TypeInput, + SessionID: "s1", + Context: json.RawMessage(`{"hello.marked":true}`), + }); err != nil { + t.Fatal(err) + } + tooBig := json.RawMessage(`{"k":"` + strings.Repeat("x", sdk.MaxPluginContextBytes) + `"}`) + got := h.attachPluginContext(seam.Envelope{ + Type: seam.TypePreStep, + SessionID: "s1", + Context: tooBig, + }) + if string(sdk.DecodePluginContext(got.Context).Get("hello.marked")) != "true" { + t.Fatalf("kept=%s", got.Context) + } +} + +// TestHostPluginContextDroppedOnHandledAndTerminal 确认 handled 与终态会清掉袋子。 +func TestHostPluginContextDroppedOnHandledAndTerminal(t *testing.T) { + t.Parallel() + h := &Host{contexts: map[string]json.RawMessage{}} + if _, err := h.Dispatch(context.Background(), seam.Envelope{ + Type: seam.TypeInput, + SessionID: "s1", + Context: json.RawMessage(`{"hello.marked":true}`), + }); err != nil { + t.Fatal(err) + } + h.dropSessionPluginContext("s1") + if raw := h.lookupPluginContext("s1", ""); raw != nil { + t.Fatalf("handled should drop session bag: %s", raw) + } + + h.persistPluginContext(seam.Envelope{SessionID: "s2", RunID: "r2", Context: json.RawMessage(`{"a":1}`)}) + h.notify(context.Background(), seam.Envelope{Type: string(pkgagent.EventRunCompleted), SessionID: "s2", RunID: "r2"}) + if raw := h.lookupPluginContext("s2", "r2"); raw != nil { + t.Fatalf("terminal should drop: %s", raw) + } +} diff --git a/server/internal/pluginhost/host.go b/server/internal/pluginhost/host.go new file mode 100644 index 0000000..e25ff32 --- /dev/null +++ b/server/internal/pluginhost/host.go @@ -0,0 +1,550 @@ +package pluginhost + +import ( + "context" + "database/sql" + "encoding/json" + "fmt" + "log/slog" + "os" + "os/exec" + "path/filepath" + "runtime" + "sort" + "strings" + "sync" + "time" + + "github.com/hashicorp/go-hclog" + goplugin "github.com/hashicorp/go-plugin" + + "codedock/internal/agent/memory" + "codedock/internal/events" + "codedock/internal/util" + pkgagent "codedock/pkg/agent" + "codedock/pkg/agent/seam" + "codedock/pkg/agent/tool" + "codedock/pkg/db/sqlite" + sdk "codedock/pkg/plugin" +) + +var reservedMethods = map[string]struct{}{ + "ping": {}, + "memory_read": {}, + "memory_write": {}, + "memory_search": {}, +} + +// Options 是加载插件宿主的入参。 +type Options struct { + Dir string // PLUGIN_DIR,每个子目录一个插件 + Timeout time.Duration // 单次 RPC 超时;零则用 10s + Registry tool.Registry // 用来挂插件方法;空则新建 + Queries *sqlite.Queries // 记忆和 AppendNotice 用;可空 + Model pkgagent.ModelConfig // Complete 自己打模型时用 + Log *slog.Logger +} + +// Host 按子目录拉起插件进程,并实现 seam.Dispatcher。 +type Host struct { + plugins []*instance + registry tool.Registry + queries *sqlite.Queries + model pkgagent.ModelConfig + timeout time.Duration + log *slog.Logger + unsub func() // Attach 订阅账本后的取消函数 + ctxMu sync.Mutex // 保护 contexts + contexts map[string]json.RawMessage // session: / run: 下的插件共享袋子 +} + +// instance 是一个已拉起的插件进程。 +type instance struct { + name string // PLUGIN_DIR 子目录名,也是 Seen 里的名字 + client *goplugin.Client + plugin sdk.Handler + subs map[string]struct{} + methods map[string]struct{} + mu sync.Mutex // 同一插件串行 RPC +} + +// hostRPC 把单个插件的 Host 回调转到本 Host。 +type hostRPC struct { + host *Host + inst *instance +} + +// Load 扫描 PLUGIN_DIR 的每个子目录,按名字排序拉起进程并 Bootstrap。 +func Load(ctx context.Context, opts Options) (*Host, error) { + dir := strings.TrimSpace(opts.Dir) + if dir == "" { + return nil, nil + } + if opts.Timeout <= 0 { + opts.Timeout = 10 * time.Second + } + if opts.Log == nil { + opts.Log = slog.Default() + } + if opts.Registry == nil { + opts.Registry = tool.NewRegistry() + } + entries, err := os.ReadDir(dir) + if err != nil { + if os.IsNotExist(err) { + return nil, nil + } + return nil, err + } + names := make([]string, 0, len(entries)) + for _, entry := range entries { + if entry.IsDir() && !strings.HasPrefix(entry.Name(), ".") { + names = append(names, entry.Name()) + } + } + sort.Strings(names) + + h := &Host{ + registry: opts.Registry, + queries: opts.Queries, + model: opts.Model, + timeout: opts.Timeout, + log: opts.Log, + contexts: map[string]json.RawMessage{}, + } + for _, name := range names { + bin := filepath.Join(dir, name, name) + if runtime.GOOS == "windows" { + bin += ".exe" + } + if _, err := os.Stat(bin); err != nil { + h.log.Warn("skip plugin without binary", "name", name, "path", bin) + continue + } + if err := h.start(ctx, name, bin); err != nil { + _ = h.Close() + return nil, fmt.Errorf("load plugin %s: %w", name, err) + } + } + return h, nil +} + +// start 拉起一个插件二进制并完成 Bootstrap。 +func (h *Host) start(ctx context.Context, name, bin string) error { + inst := &instance{name: name, subs: map[string]struct{}{}, methods: map[string]struct{}{}} + cmd := exec.Command(bin) + cmd.Dir = filepath.Dir(bin) + client := goplugin.NewClient(&goplugin.ClientConfig{ + HandshakeConfig: sdk.Handshake, + Plugins: map[string]goplugin.Plugin{ + sdk.PluginName: &sdk.GRPCPlugin{Host: &hostRPC{host: h, inst: inst}}, + }, + Cmd: cmd, + AllowedProtocols: []goplugin.Protocol{goplugin.ProtocolGRPC}, + Logger: hclog.New(&hclog.LoggerOptions{Name: "plugin-" + name, Level: hclog.Error, Output: os.Stderr}), + }) + inst.client = client + rpcClient, err := client.Client() + if err != nil { + client.Kill() + return err + } + raw, err := rpcClient.Dispense(sdk.PluginName) + if err != nil { + client.Kill() + return err + } + p, ok := raw.(sdk.Handler) + if !ok { + client.Kill() + return fmt.Errorf("dispensed plugin has unexpected type %T", raw) + } + inst.plugin = p + man, err := p.Bootstrap(ctx, nil) + if err != nil { + client.Kill() + return err + } + if man.Name == "" { + man.Name = name + } + for _, sub := range man.Subscriptions { + if sub != "" { + inst.subs[sub] = struct{}{} + } + } + h.plugins = append(h.plugins, inst) + h.log.Info("plugin loaded", "name", man.Name, "subscriptions", man.Subscriptions) + return nil +} + +// Dispatch 按 Seen 逐个问订阅了该口的插件;类型变了立刻返回。 +func (h *Host) Dispatch(ctx context.Context, ev seam.Envelope) (seam.Envelope, error) { + if h == nil { + return ev, nil + } + if ev.ChainID == "" { + ev.ChainID = util.NewID() + } + if seam.IsSeam(ev.Type) { + ev = h.attachPluginContext(ev) + } + seen := append([]string{}, ev.Seen...) + var err error + defer func() { + h.persistPluginContext(ev) + if ev.Type == seam.TypeInputHandled { + h.dropSessionPluginContext(ev.SessionID) + } + }() + for _, inst := range h.plugins { + if !inst.subscribed(ev.Type) || contains(seen, inst.name) { + continue + } + sentType := ev.Type + var out seam.Envelope + out, err = inst.callOnEvent(ctx, h.timeout, ev) + if err != nil { + return ev, fmt.Errorf("plugin %s: %w", inst.name, err) + } + out.ChainID = ev.ChainID + out.SessionID = ev.SessionID + out.RunID = ev.RunID + out.TurnID = ev.TurnID + if !sdk.AcceptableContext(out.Context) { + if h.log != nil { + h.log.Warn("plugin context rejected", "plugin", inst.name, "bytes", len(out.Context)) + } + out.Context = ev.Context + } + seen = append(seen, inst.name) + out.Seen = append([]string{}, seen...) + if out.Type == "" { + out.Type = sentType + } + ev = out + if ev.Type != sentType { + return ev, nil + } + } + ev.Seen = seen + return ev, nil +} + +// Emit 另发一条与当前口无关的事件。六个口的同名事件会被拒绝。 +func (h *Host) Emit(ctx context.Context, ev seam.Envelope) error { + if h == nil { + return fmt.Errorf("plugin host is nil") + } + if seam.IsSeam(ev.Type) { + return fmt.Errorf("cannot emit seam event %q", ev.Type) + } + ev.ChainID = util.NewID() + ev.Seen = nil + go h.notify(context.WithoutCancel(ctx), ev) + return nil +} + +// MethodNames 返回已挂进 Registry 的插件方法名。 +func (h *Host) MethodNames() []string { + if h == nil { + return nil + } + var out []string + for _, inst := range h.plugins { + for name := range inst.methods { + out = append(out, name) + } + } + sort.Strings(out) + return out +} + +// Attach 把账本事实异步转给订阅了该类型的插件;跳过 assistant.delta。 +func (h *Host) Attach(bus *events.Bus) { + if h == nil || bus == nil { + return + } + h.unsub = bus.SubscribeAll(func(e events.Event) { + if e.Type == string(pkgagent.EventAssistantDelta) { + return + } + ev := seam.Envelope{Type: e.Type, ChainID: util.NewID()} + if ae, ok := e.Payload.(pkgagent.AgentEvent); ok { + ev.SessionID = ae.SessionID + ev.RunID = ae.RunID + if ae.TurnID != nil { + ev.TurnID = *ae.TurnID + } + ev.Payload = ae.Payload + } + go h.notify(context.Background(), ev) + }) +} + +// Close 停掉全部插件进程。 +func (h *Host) Close() error { + if h == nil { + return nil + } + if h.unsub != nil { + h.unsub() + h.unsub = nil + } + for _, inst := range h.plugins { + if inst.client != nil { + inst.client.Kill() + } + } + h.plugins = nil + h.ctxMu.Lock() + h.contexts = nil + h.ctxMu.Unlock() + return nil +} + +// notify 把非换向事件异步发给订阅了该类型的插件。 +func (h *Host) notify(ctx context.Context, ev seam.Envelope) { + ev = h.attachPluginContext(ev) + for _, inst := range h.plugins { + if !inst.subscribed(ev.Type) { + continue + } + if _, err := inst.callOnEvent(ctx, h.timeout, ev); err != nil { + h.log.Warn("plugin notify failed", "plugin", inst.name, "type", ev.Type, "error", err) + } + } + if isTerminalLedger(ev.Type) { + h.dropPluginContext(ev.RunID, ev.SessionID) + } +} + +// registerMethod 把插件方法挂进工具 Registry,并拒绝保留名和重名。 +func (h *Host) registerMethod(inst *instance, m sdk.Method) error { + if m.Name == "" { + return fmt.Errorf("method name is required") + } + if _, ok := reservedMethods[m.Name]; ok { + return fmt.Errorf("method %q is reserved", m.Name) + } + if _, err := h.registry.Get(tool.Reference{Name: m.Name}); err == nil { + return fmt.Errorf("method %q already registered", m.Name) + } + def := sdk.MethodToDefinition(m) + if err := h.registry.Register(&methodTool{inst: inst, def: def, timeout: h.timeout}); err != nil { + return err + } + inst.methods[m.Name] = struct{}{} + return nil +} + +// subscribed 判断该插件是否订阅了这个事件类型。 +func (i *instance) subscribed(typ string) bool { + if i == nil { + return false + } + _, ok := i.subs[typ] + return ok +} + +// callOnEvent 在超时内串行调用插件的 OnEvent。 +func (i *instance) callOnEvent(ctx context.Context, timeout time.Duration, ev seam.Envelope) (seam.Envelope, error) { + if i == nil || i.plugin == nil { + return ev, fmt.Errorf("plugin is not started") + } + if i.client != nil && i.client.Exited() { + return ev, fmt.Errorf("plugin %s exited", i.name) + } + i.mu.Lock() + defer i.mu.Unlock() + ctx, cancel := context.WithTimeout(ctx, timeout) + defer cancel() + return i.plugin.OnEvent(ctx, ev) +} + +// callExecute 在超时内串行调用插件的 ExecuteMethod。 +func (i *instance) callExecute(ctx context.Context, timeout time.Duration, input tool.Input) (tool.Result, error) { + if i == nil || i.plugin == nil { + return tool.Result{}, fmt.Errorf("plugin is not started") + } + if i.client != nil && i.client.Exited() { + return tool.Result{}, fmt.Errorf("plugin %s exited", i.name) + } + i.mu.Lock() + defer i.mu.Unlock() + ctx, cancel := context.WithTimeout(ctx, timeout) + defer cancel() + out, err := i.plugin.ExecuteMethod(ctx, sdk.MethodInput{ + SessionID: input.SessionID, + RunID: input.RunID, + TurnID: input.TurnID, + CallID: input.Call.ID, + Name: input.Call.Name, + Arguments: input.Call.Arguments, + }) + if err != nil { + return tool.Result{}, err + } + return tool.Result{ + CallID: input.Call.ID, + Name: input.Call.Name, + Success: out.Success, + Output: out.Output, + Error: out.Error, + }, nil +} + +// Emit 转发插件的另发事件请求。 +func (r *hostRPC) Emit(ctx context.Context, ev seam.Envelope) error { + return r.host.Emit(ctx, ev) +} + +// RegisterMethod 把该方法记到发起回调的那个插件上。 +func (r *hostRPC) RegisterMethod(_ context.Context, method sdk.Method) error { + return r.host.registerMethod(r.inst, method) +} + +// MemoryGet 按会话读一篇专题记忆。 +func (r *hostRPC) MemoryGet(ctx context.Context, key sdk.MemoryKey) (string, error) { + if r.host.queries == nil { + return "", fmt.Errorf("memory is not available") + } + if key.SessionID == "" || key.Name == "" { + return "", fmt.Errorf("session_id and name are required") + } + sess, err := r.host.queries.GetSession(ctx, key.SessionID) + if err != nil { + return "", err + } + scope, scopeID := memoryScope(key.Scope, sess) + item, err := memory.Get(ctx, r.host.queries, memory.TextMemoryKey{ + Scope: scope, + ScopeID: scopeID, + Kind: memory.KindTopic, + Name: key.Name, + }) + if err != nil { + return "", err + } + return item.Content, nil +} + +// MemoryUpsert 按会话写一篇专题记忆。 +func (r *hostRPC) MemoryUpsert(ctx context.Context, key sdk.MemoryKey, text string) error { + if r.host.queries == nil { + return fmt.Errorf("memory is not available") + } + if key.SessionID == "" || key.Name == "" { + return fmt.Errorf("session_id and name are required") + } + sess, err := r.host.queries.GetSession(ctx, key.SessionID) + if err != nil { + return err + } + scope, scopeID := memoryScope(key.Scope, sess) + _, err = memory.Upsert(ctx, r.host.queries, memory.TextMemory{ + Scope: scope, + ScopeID: scopeID, + Kind: memory.KindTopic, + Name: key.Name, + Content: text, + }) + return err +} + +// Complete 用宿主模型配置单独打一次模型,不进当前助手流。 +func (r *hostRPC) Complete(ctx context.Context, req sdk.CompleteRequest) (sdk.CompleteResult, error) { + chat := pkgagent.Chat{ + SessionID: req.SessionID, + RunID: req.RunID, + Model: r.host.model, + SystemPrompt: req.Prompt, + Messages: []pkgagent.Message{{Role: pkgagent.RoleUser, Content: pkgagent.EncodeText(req.Text)}}, + } + stream, err := pkgagent.Stream(ctx, chat) + if err != nil { + return sdk.CompleteResult{}, err + } + defer stream.Close() + result, err := stream.Result(ctx) + if err != nil { + return sdk.CompleteResult{}, err + } + return sdk.CompleteResult{Text: pkgagent.DecodeText(result.Message.Content)}, nil +} + +// AppendNotice 写一条用户看得见的 system 消息(本期只落库)。 +func (r *hostRPC) AppendNotice(ctx context.Context, sessionID, runID, text string) error { + if r.host.queries == nil { + return fmt.Errorf("queries are not available") + } + if sessionID == "" || strings.TrimSpace(text) == "" { + return fmt.Errorf("session_id and text are required") + } + now := util.FormatTime(util.Now()) + seq, err := r.host.queries.IncrementEventSeq(ctx, sqlite.IncrementEventSeqParams{UpdatedAt: now, ID: sessionID}) + if err != nil { + return err + } + _, err = r.host.queries.InsertMessage(ctx, sqlite.InsertMessageParams{ + ID: util.NewID(), + SessionID: sessionID, + RunID: nullString(runID), + Role: string(pkgagent.RoleSystem), + Content: string(pkgagent.EncodeText(text)), + EventSeq: seq, + CreatedAt: now, + }) + return err +} + +// memoryScope 把插件传来的 scope 落到用户或工作区。 +func memoryScope(scope string, sess sqlite.Session) (memory.TextMemoryScope, string) { + if scope == string(memory.ScopeUser) { + return memory.ScopeUser, sess.UserID + } + return memory.ScopeWorkspace, sess.WorkspaceID +} + +// nullString 把空串收成 SQL NULL。 +func nullString(value string) sql.NullString { + if value == "" { + return sql.NullString{} + } + return sql.NullString{String: value, Valid: true} +} + +// contains 判断字符串是否已在切片里。 +func contains(items []string, want string) bool { + for _, item := range items { + if item == want { + return true + } + } + return false +} + +// methodTool 把插件方法暴露成 Registry 里的 Tool。 +type methodTool struct { + inst *instance + def tool.Definition + timeout time.Duration +} + +// Definition 返回挂进 Registry 的工具描述。 +func (m *methodTool) Definition() tool.Definition { return m.def } + +// Execute 把模型的工具调用转到插件进程。 +func (m *methodTool) Execute(ctx context.Context, input tool.Input) (tool.Result, error) { + result, err := m.inst.callExecute(ctx, m.timeout, input) + if err != nil { + return tool.Result{CallID: input.Call.ID, Name: m.def.Name, Success: false, Error: err.Error()}, nil + } + if result.CallID == "" { + result.CallID = input.Call.ID + } + if result.Name == "" { + result.Name = m.def.Name + } + return result, nil +} diff --git a/server/internal/pluginhost/host_test.go b/server/internal/pluginhost/host_test.go new file mode 100644 index 0000000..e90285a --- /dev/null +++ b/server/internal/pluginhost/host_test.go @@ -0,0 +1,379 @@ +package pluginhost + +import ( + "context" + "encoding/json" + "os" + "os/exec" + "path/filepath" + "testing" + "time" + + pkgagent "codedock/pkg/agent" + "codedock/pkg/agent/seam" + "codedock/pkg/agent/tool" + sdk "codedock/pkg/plugin" +) + +// jsonGet 读信封袋子里的一个键。 +func jsonGet(raw json.RawMessage, key string) json.RawMessage { + return sdk.DecodePluginContext(raw).Get(key) +} + +// TestHostDispatchMethodsAndEmit 覆盖链式改写、换向、方法执行和 Emit 拒绝同名口。 +func TestHostDispatchMethodsAndEmit(t *testing.T) { + if testing.Short() { + t.Skip("plugin host integration") + } + dir := t.TempDir() + buildPlugin(t, "codedock/internal/pluginhost/testdata/hello", filepath.Join(dir, "hello", "hello")) + buildPlugin(t, "codedock/internal/pluginhost/testdata/probe", filepath.Join(dir, "probe", "probe")) + + reg := tool.NewRegistry() + host, err := Load(context.Background(), Options{ + Dir: dir, + Timeout: 3 * time.Second, + Registry: reg, + Model: pkgagent.ModelConfig{Provider: "fake", Model: "fake"}, + Log: nil, + }) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = host.Close() }) + + ev, err := host.Dispatch(context.Background(), seam.Envelope{ + Type: seam.TypeInput, + Payload: pkgagent.MarshalPayload(pkgagent.InputPayload{Content: "hello"}), + }) + if err != nil { + t.Fatal(err) + } + var payload pkgagent.InputPayload + if err := json.Unmarshal(ev.Payload, &payload); err != nil { + t.Fatal(err) + } + if ev.Type != seam.TypeInput || payload.Content != "[hello] hello +probe" { + t.Fatalf("ev=%+v payload=%+v", ev, payload) + } + if len(ev.Seen) != 2 || ev.Seen[0] != "hello" || ev.Seen[1] != "probe" { + t.Fatalf("seen=%v", ev.Seen) + } + + handled, err := host.Dispatch(context.Background(), seam.Envelope{ + Type: seam.TypeInput, + Payload: pkgagent.MarshalPayload(pkgagent.InputPayload{Content: "/skip later"}), + }) + if err != nil || handled.Type != seam.TypeInputHandled { + t.Fatalf("handled=%+v err=%v", handled, err) + } + + denied, err := host.Dispatch(context.Background(), seam.Envelope{ + Type: seam.TypePreExecute, + Payload: pkgagent.MarshalPayload(tool.PreExecutePayload{ + Call: tool.Call{ID: "c1", Name: "ping", Arguments: json.RawMessage(`{"x":"forbidden"}`)}, + }), + }) + if err != nil || denied.Type != seam.TypeToolsDenied { + t.Fatalf("denied=%+v err=%v", denied, err) + } + + names := host.MethodNames() + if len(names) != 1 || names[0] != "hello" { + t.Fatalf("methods=%v", names) + } + item, err := reg.Get(tool.Reference{Name: "hello"}) + if err != nil { + t.Fatal(err) + } + result, err := item.Execute(context.Background(), tool.Input{ + Call: tool.Call{ID: "c1", Name: "hello", Arguments: json.RawMessage(`{"text":"hi"}`)}, + }) + if err != nil || !result.Success || string(result.Output) != `{"text":"hi"}` { + t.Fatalf("hello result=%+v err=%v", result, err) + } + + if err := host.Emit(context.Background(), seam.Envelope{Type: seam.TypeInput}); err == nil { + t.Fatal("expected emit of seam type to fail") + } +} + +// TestHostProbeHandled 确认 probe 能把 input 换成 handled。 +func TestHostProbeHandled(t *testing.T) { + if testing.Short() { + t.Skip("plugin host integration") + } + dir := t.TempDir() + buildPlugin(t, "codedock/internal/pluginhost/testdata/probe", filepath.Join(dir, "probe", "probe")) + host, err := Load(context.Background(), Options{ + Dir: dir, + Timeout: 3 * time.Second, + Model: pkgagent.ModelConfig{Provider: "fake", Model: "fake"}, + }) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = host.Close() }) + handled, err := host.Dispatch(context.Background(), seam.Envelope{ + Type: seam.TypeInput, + Payload: pkgagent.MarshalPayload(pkgagent.InputPayload{Content: "handle-me"}), + }) + if err != nil || handled.Type != seam.TypeInputHandled { + t.Fatalf("handled=%+v err=%v", handled, err) + } +} + +// TestHostTimeout 确认插件拖延超过 RPC 超时会失败。 +func TestHostTimeout(t *testing.T) { + if testing.Short() { + t.Skip("plugin host integration") + } + dir := t.TempDir() + buildPlugin(t, "codedock/internal/pluginhost/testdata/probe", filepath.Join(dir, "probe", "probe")) + host, err := Load(context.Background(), Options{ + Dir: dir, + Timeout: 200 * time.Millisecond, + Model: pkgagent.ModelConfig{Provider: "fake", Model: "fake"}, + }) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = host.Close() }) + _, err = host.Dispatch(context.Background(), seam.Envelope{ + Type: seam.TypeInput, + Payload: pkgagent.MarshalPayload(pkgagent.InputPayload{Content: "stall:now"}), + }) + if err == nil { + t.Fatal("expected timeout") + } +} + +// TestHelloFixture 跑一遍测试夹具的改正文、跳过、隐藏提示、否决和方法。 +func TestHelloFixture(t *testing.T) { + if testing.Short() { + t.Skip("plugin host integration") + } + dir := t.TempDir() + buildPlugin(t, "codedock/internal/pluginhost/testdata/hello", filepath.Join(dir, "hello", "hello")) + reg := tool.NewRegistry() + host, err := Load(context.Background(), Options{ + Dir: dir, + Timeout: 3 * time.Second, + Registry: reg, + Model: pkgagent.ModelConfig{Provider: "fake", Model: "fake"}, + }) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = host.Close() }) + + rewritten, err := host.Dispatch(context.Background(), seam.Envelope{ + Type: seam.TypeInput, + SessionID: "sess-hello", + Payload: pkgagent.MarshalPayload(pkgagent.InputPayload{Content: "world"}), + }) + if err != nil { + t.Fatal(err) + } + var input pkgagent.InputPayload + if err := json.Unmarshal(rewritten.Payload, &input); err != nil { + t.Fatal(err) + } + if rewritten.Type != seam.TypeInput || input.Content != "[hello] world" || len(rewritten.Seen) != 1 || rewritten.Seen[0] != "hello" { + t.Fatalf("rewrite=%+v payload=%+v", rewritten, input) + } + if string(jsonGet(rewritten.Context, "hello.marked")) != "true" { + t.Fatalf("input context=%s", rewritten.Context) + } + + skipped, err := host.Dispatch(context.Background(), seam.Envelope{ + Type: seam.TypeInput, + Payload: pkgagent.MarshalPayload(pkgagent.InputPayload{Content: "/skip later"}), + }) + if err != nil || skipped.Type != seam.TypeInputHandled { + t.Fatalf("skip=%+v err=%v", skipped, err) + } + + pre, err := host.Dispatch(context.Background(), seam.Envelope{ + Type: seam.TypePreStep, + SessionID: "sess-hello", + RunID: "run-hello", + Payload: pkgagent.MarshalPayload(pkgagent.PreStepPayload{SystemPrompt: "base"}), + }) + if err != nil { + t.Fatal(err) + } + var step pkgagent.PreStepPayload + if err := json.Unmarshal(pre.Payload, &step); err != nil { + t.Fatal(err) + } + if len(step.Hidden) != 1 || pkgagent.DecodeText(step.Hidden[0].Content) == "" { + t.Fatalf("pre-step=%+v", step) + } + if string(jsonGet(pre.Context, "hello.marked")) != "true" || string(jsonGet(pre.Context, "hello.seen_at_pre_step")) != "true" { + t.Fatalf("pre-step context=%s", pre.Context) + } + + denied, err := host.Dispatch(context.Background(), seam.Envelope{ + Type: seam.TypePreExecute, + Payload: pkgagent.MarshalPayload(tool.PreExecutePayload{ + Call: tool.Call{ID: "c1", Name: "ping", Arguments: json.RawMessage(`{"x":"forbidden"}`)}, + }), + }) + if err != nil || denied.Type != seam.TypeToolsDenied { + t.Fatalf("denied=%+v err=%v", denied, err) + } + + item, err := reg.Get(tool.Reference{Name: "hello"}) + if err != nil { + t.Fatal(err) + } + result, err := item.Execute(context.Background(), tool.Input{ + Call: tool.Call{ID: "c1", Name: "hello", Arguments: json.RawMessage(`{"text":"hi"}`)}, + }) + if err != nil || !result.Success || string(result.Output) != `{"text":"hi"}` { + t.Fatalf("hello method=%+v err=%v", result, err) + } +} + +// TestLoadEmptyDir 确认未设插件目录时不拉进程。 +func TestLoadEmptyDir(t *testing.T) { + host, err := Load(context.Background(), Options{Dir: ""}) + if err != nil || host != nil { + t.Fatalf("empty dir host=%v err=%v", host, err) + } +} + +// TestExampleTemplate 确认 plugin/example 不改正文、不换向、不登记方法。 +func TestExampleTemplate(t *testing.T) { + if testing.Short() { + t.Skip("plugin host integration") + } + dir := t.TempDir() + buildRepoPlugin(t, filepath.Join("plugin", "example"), filepath.Join(dir, "example", "example")) + reg := tool.NewRegistry() + host, err := Load(context.Background(), Options{ + Dir: dir, + Timeout: 3 * time.Second, + Registry: reg, + Model: pkgagent.ModelConfig{Provider: "fake", Model: "fake"}, + }) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = host.Close() }) + + rewritten, err := host.Dispatch(context.Background(), seam.Envelope{ + Type: seam.TypeInput, + Payload: pkgagent.MarshalPayload(pkgagent.InputPayload{Content: "world"}), + }) + if err != nil { + t.Fatal(err) + } + var input pkgagent.InputPayload + if err := json.Unmarshal(rewritten.Payload, &input); err != nil { + t.Fatal(err) + } + if rewritten.Type != seam.TypeInput || input.Content != "world" || len(rewritten.Seen) != 0 { + t.Fatalf("example mutated input: ev=%+v payload=%+v", rewritten, input) + } + + skipped, err := host.Dispatch(context.Background(), seam.Envelope{ + Type: seam.TypeInput, + Payload: pkgagent.MarshalPayload(pkgagent.InputPayload{Content: "/skip later"}), + }) + if err != nil || skipped.Type != seam.TypeInput { + t.Fatalf("example handled skip: ev=%+v err=%v", skipped, err) + } + + denied, err := host.Dispatch(context.Background(), seam.Envelope{ + Type: seam.TypePreExecute, + Payload: pkgagent.MarshalPayload(tool.PreExecutePayload{ + Call: tool.Call{ID: "c1", Name: "ping", Arguments: json.RawMessage(`{"x":"forbidden"}`)}, + }), + }) + if err != nil || denied.Type != seam.TypePreExecute { + t.Fatalf("example denied tool: ev=%+v err=%v", denied, err) + } + if names := host.MethodNames(); len(names) != 0 { + t.Fatalf("example registered methods: %v", names) + } +} + +// TestRegisterReservedMethods 确认 ping / memory_* 不能被插件覆盖。 +func TestRegisterReservedMethods(t *testing.T) { + h := &Host{registry: tool.NewRegistry()} + inst := &instance{methods: map[string]struct{}{}} + for _, name := range []string{"ping", "memory_read", "memory_write", "memory_search"} { + if err := h.registerMethod(inst, sdk.Method{Name: name, Prompt: "no"}); err == nil { + t.Fatalf("reserved %s should be rejected", name) + } + } + if err := h.registerMethod(inst, sdk.Method{Name: "echo", Prompt: "echo"}); err != nil { + t.Fatal(err) + } + if err := h.registerMethod(inst, sdk.Method{Name: "echo", Prompt: "dup"}); err == nil { + t.Fatal("duplicate method should be rejected") + } +} + +// buildRepoPlugin 编译仓根相对路径下的插件到 dest。 +func buildRepoPlugin(t *testing.T, rel, dest string) { + t.Helper() + if err := os.MkdirAll(filepath.Dir(dest), 0o755); err != nil { + t.Fatal(err) + } + cmd := exec.Command("go", "build", "-o", dest, ".") + cmd.Dir = filepath.Join(repoRoot(t), rel) + cmd.Env = append(os.Environ(), "CGO_ENABLED=0") + out, err := cmd.CombinedOutput() + if err != nil { + t.Fatalf("build %s: %v\n%s", rel, err, out) + } +} + +// buildPlugin 用 server 模块路径编译一个测试插件到 dest。 +func buildPlugin(t *testing.T, pkg, dest string) { + t.Helper() + if err := os.MkdirAll(filepath.Dir(dest), 0o755); err != nil { + t.Fatal(err) + } + cmd := exec.Command("go", "build", "-o", dest, pkg) + cmd.Dir = moduleRoot(t) + cmd.Env = append(os.Environ(), "CGO_ENABLED=0") + out, err := cmd.CombinedOutput() + if err != nil { + t.Fatalf("build %s: %v\n%s", pkg, err, out) + } +} + +// repoRoot 返回仓根(server 的上一级)。 +func repoRoot(t *testing.T) string { + t.Helper() + root := filepath.Dir(moduleRoot(t)) + if _, err := os.Stat(filepath.Join(root, "plugin", "example")); err != nil { + t.Fatalf("plugin/example not found at %s", root) + } + return root +} + +// moduleRoot 从当前工作目录向上找到 server/go.mod。 +func moduleRoot(t *testing.T) string { + t.Helper() + dir, err := os.Getwd() + if err != nil { + t.Fatal(err) + } + for i := 0; i < 6; i++ { + if _, err := os.Stat(filepath.Join(dir, "go.mod")); err == nil { + return dir + } + parent := filepath.Dir(dir) + if parent == dir { + break + } + dir = parent + } + t.Fatal("go.mod not found") + return "" +} diff --git a/server/internal/pluginhost/testdata/hello/main.go b/server/internal/pluginhost/testdata/hello/main.go new file mode 100644 index 0000000..628eda6 --- /dev/null +++ b/server/internal/pluginhost/testdata/hello/main.go @@ -0,0 +1,112 @@ +// hello 只给宿主和回路测试用:改正文、跳过、否决、登记方法。 +// 不是作者模板,不要编进仓根 plugin/。 +package main + +import ( + "bytes" + "context" + "encoding/json" + "log/slog" + "strings" + + sdk "codedock/pkg/plugin" +) + +// hello 演示四个口和方法登记。 +type hello struct { + host sdk.Host +} + +// helloInput 是 hello 方法的入参。 +type helloInput struct { + Text string `json:"text"` +} + +// helloOutput 是 hello 方法的出参。 +type helloOutput struct { + Text string `json:"text"` +} + +// Bootstrap 登记 hello 方法,并订阅 input / pre-step / pre-execute / run.completed。 +func (p *hello) Bootstrap(ctx context.Context, host sdk.Host) (sdk.Manifest, error) { + p.host = host + err := host.RegisterMethod(ctx, sdk.Method{ + Name: "hello", + Prompt: "连通性检查:把 text 原样返回。", + ParametersSchema: json.RawMessage(`{ + "type": "object", + "properties": {"text": {"type": "string", "description": "要回显的文本"}}, + "required": ["text"] + }`), + }) + return sdk.Manifest{ + Name: "hello", + Subscriptions: []string{ + sdk.TypeInput, + sdk.TypePreStep, + sdk.TypePreExecute, + "run.completed", + }, + }, err +} + +// OnAgentInput 给正文加 [hello] 前缀;以 /skip 开头则不建 Run。 +func (p *hello) OnAgentInput(ctx context.Context, in sdk.AgentInput) (sdk.AgentInputResult, error) { + content := strings.TrimSpace(in.Content) + if strings.HasPrefix(content, "/skip") { + if p.host != nil && in.SessionID != "" { + _ = p.host.AppendNotice(ctx, in.SessionID, in.RunID, "hello 已跳过本次对话。") + } + return in.Handle(), nil + } + if content != "" && !strings.HasPrefix(content, "[hello] ") { + in.Content = "[hello] " + content + } + in.Context.Set("hello.marked", true) + return in.Reply(), nil +} + +// OnAgentPreStep 注入一条告诉模型可以调用 hello 的隐藏提示。 +func (p *hello) OnAgentPreStep(_ context.Context, in sdk.AgentPreStep) (sdk.AgentPreStepResult, error) { + in.Hidden = append(in.Hidden, sdk.HiddenText("hello 插件已加载。需要回显时调用 hello。")) + if string(in.Context.Get("hello.marked")) == "true" { + in.Context.Set("hello.seen_at_pre_step", true) + } + return in.Reply(), nil +} + +// OnToolPreExecute 在工具参数含 forbidden 时否决该次调用。 +func (p *hello) OnToolPreExecute(_ context.Context, in sdk.ToolPreExecute) (sdk.ToolPreExecuteResult, error) { + if bytes.Contains(in.Call.Arguments, []byte("forbidden")) { + return in.Deny(), nil + } + return in.Reply(), nil +} + +// OnLedgerNotify 在 run 结束时打一条日志。 +func (p *hello) OnLedgerNotify(_ context.Context, in sdk.LedgerNotify) error { + if in.Type == "run.completed" { + slog.Info("hello: run completed", "session_id", in.SessionID, "run_id", in.RunID) + } + return nil +} + +// ExecuteMethod 把 hello 方法的 text 原样返回。 +func (p *hello) ExecuteMethod(_ context.Context, in sdk.MethodInput) (sdk.MethodResult, error) { + var args helloInput + if len(in.Arguments) > 0 { + if err := json.Unmarshal(in.Arguments, &args); err != nil { + return sdk.MethodResult{Success: false, Error: err.Error()}, nil + } + } + out, err := json.Marshal(helloOutput{Text: args.Text}) + if err != nil { + return sdk.MethodResult{Success: false, Error: err.Error()}, nil + } + return sdk.MethodResult{Success: true, Output: out}, nil +} + +// main 启动 hello 测试插件进程。 +func main() { + sdk.Serve(&hello{}) +} diff --git a/server/internal/pluginhost/testdata/probe/main.go b/server/internal/pluginhost/testdata/probe/main.go new file mode 100644 index 0000000..11aee5e --- /dev/null +++ b/server/internal/pluginhost/testdata/probe/main.go @@ -0,0 +1,43 @@ +package main + +import ( + "context" + "strings" + "time" + + sdk "codedock/pkg/plugin" +) + +// probe 给宿主测试用:改正文、换向、故意拖延。 +type probe struct{} + +// Bootstrap 只订阅 agent/input。 +func (probe) Bootstrap(context.Context, sdk.Host) (sdk.Manifest, error) { + return sdk.Manifest{ + Name: "probe", + Subscriptions: []string{sdk.TypeInput}, + }, nil +} + +// OnAgentInput 测拖延、input/handled,以及给正文加 +probe。 +func (probe) OnAgentInput(_ context.Context, in sdk.AgentInput) (sdk.AgentInputResult, error) { + if strings.HasPrefix(in.Content, "stall:") { + time.Sleep(2 * time.Second) + return in.Reply(), nil + } + if in.Content == "handle-me" { + return in.Handle(), nil + } + in.Content += " +probe" + return in.Reply(), nil +} + +// ExecuteMethod 声明本插件没有方法。 +func (probe) ExecuteMethod(context.Context, sdk.MethodInput) (sdk.MethodResult, error) { + return sdk.MethodResult{Success: false, Error: "no methods"}, nil +} + +// main 启动 probe 插件进程。 +func main() { + sdk.Serve(probe{}) +} diff --git a/server/migrations/0009_run_overlays.sql b/server/migrations/0009_run_overlays.sql new file mode 100644 index 0000000..9d6c6f1 --- /dev/null +++ b/server/migrations/0009_run_overlays.sql @@ -0,0 +1,6 @@ +CREATE TABLE IF NOT EXISTS run_overlays ( + run_id TEXT PRIMARY KEY, + system_prompt TEXT NOT NULL DEFAULT '', + hidden TEXT NOT NULL DEFAULT '[]', + updated_at TEXT NOT NULL +); diff --git a/server/pkg/agent/context.go b/server/pkg/agent/context.go index 3062a8e..237b4dc 100644 --- a/server/pkg/agent/context.go +++ b/server/pkg/agent/context.go @@ -15,6 +15,7 @@ type History struct { Messages []Message Tools []tool.Definition Prompt string + Hidden []Message MemoryIndexes []string } @@ -25,6 +26,7 @@ func Load(_ context.Context, hist History) (ContextSnapshot, error) { Messages: hist.Messages, Tools: hist.Tools, SystemPrompt: hist.Prompt, + Hidden: hist.Hidden, MemoryIndexes: hist.MemoryIndexes, } if hist.Checkpoint != nil { @@ -85,6 +87,9 @@ func EstimateTokens(snapshot ContextSnapshot) int64 { for _, index := range snapshot.MemoryIndexes { total += CountTokens(index) } + for _, msg := range snapshot.Hidden { + total += CountTokens(DecodeText(msg.Content)) + } if snapshot.Summary != nil { total += CountTokens(snapshot.Summary.Content) } diff --git a/server/pkg/agent/engine.go b/server/pkg/agent/engine.go index 65b7f46..4b68a24 100644 --- a/server/pkg/agent/engine.go +++ b/server/pkg/agent/engine.go @@ -3,21 +3,24 @@ package agent import ( "context" "encoding/json" + "log/slog" "strings" "time" "github.com/google/uuid" + "codedock/pkg/agent/seam" "codedock/pkg/agent/tool" ) // Engine 执行一步:按 Brain 的指令调用对应执行器,自身不直接写库、不发事件、不调度下一步。 type Engine struct { - brain *Brain - facts FactWriter - tools tool.Registry - llmGate tool.Gate - toolGate tool.Gate + brain *Brain + facts FactWriter + tools tool.Registry + llmGate tool.Gate + toolGate tool.Gate + dispatcher seam.Dispatcher } // NewEngine 创建执行引擎。brain 为空时自动构造一个空 Brain。 @@ -37,6 +40,14 @@ func (e *Engine) SetGates(llm, tools tool.Gate) { e.toolGate = tools } +// SetDispatcher 设置步骤内各口使用的喊话器。nil 表示各口原样通过。 +func (e *Engine) SetDispatcher(d seam.Dispatcher) { + if e == nil { + return + } + e.dispatcher = d +} + // Step 执行一步:先让 Brain 决策,再按指令类型分发到对应执行器。 func (e *Engine) Step(ctx context.Context, in StepInput) (StepResult, error) { if e == nil { @@ -122,6 +133,8 @@ func (e *Engine) callLLM(ctx context.Context, in StepInput, _ Instruction) (Step chat.SessionID = state.SessionID chat.RunID = state.RunID chat.TurnID = turnID + chat.Dispatcher = e.dispatcher + chat = applyRequestSeam(ctx, e.dispatcher, chat) stream, err := Stream(ctx, chat) if err != nil { @@ -262,6 +275,7 @@ func (e *Engine) callToolsBatch(ctx context.Context, in StepInput, inst Instruct DeniedCallIDs: state.Checkpoint.Denied, OnEvent: liveHook, Gate: e.toolGate, + Dispatcher: e.dispatcher, } if state.Config.Approval == ApprovalAuto { inv.OnEvent = nil @@ -508,6 +522,34 @@ func (e *Engine) appendFact(ctx context.Context, runID string, fact Fact) error return e.facts.Append(ctx, runID, fact) } +// applyRequestSeam 在发模型前递系统提示和消息;出错或换向都保留原数据。 +func applyRequestSeam(ctx context.Context, d seam.Dispatcher, chat Chat) Chat { + ev, err := seam.Dispatch(ctx, d, seam.Envelope{ + Type: seam.TypeRequest, + SessionID: chat.SessionID, + RunID: chat.RunID, + TurnID: chat.TurnID, + Payload: MarshalPayload(RequestPayload{ + SystemPrompt: chat.SystemPrompt, + Messages: chat.Messages, + }), + }) + if err != nil { + slog.Warn("agent/request dispatch failed", "run_id", chat.RunID, "error", err) + return chat + } + if ev.Type != seam.TypeRequest || len(ev.Payload) == 0 { + return chat + } + var payload RequestPayload + if err := json.Unmarshal(ev.Payload, &payload); err != nil { + return chat + } + chat.SystemPrompt = payload.SystemPrompt + chat.Messages = payload.Messages + return chat +} + // newEntityID 生成去掉连字符的 UUID,用作消息/Turn 等实体 id。 func newEntityID() string { return strings.ReplaceAll(uuid.NewString(), "-", "") diff --git a/server/pkg/agent/engine_test.go b/server/pkg/agent/engine_test.go index 0580d5c..5047d98 100644 --- a/server/pkg/agent/engine_test.go +++ b/server/pkg/agent/engine_test.go @@ -6,11 +6,16 @@ import ( "encoding/json" "errors" "fmt" + "io" + "net/http" + "net/http/httptest" "strings" "sync" "sync/atomic" "testing" "time" + + "codedock/pkg/agent/seam" ) type memFacts struct { @@ -71,6 +76,65 @@ func mustRaw(v any) json.RawMessage { return body } +func TestEngineRequestSeamRewritesChat(t *testing.T) { + var gotSystem string + var gotHeader string + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + gotHeader = r.Header.Get("X-Plugin") + body, _ := io.ReadAll(r.Body) + var req openaiChatRequest + _ = json.Unmarshal(body, &req) + if len(req.Messages) > 0 && req.Messages[0].Role == "system" { + gotSystem = req.Messages[0].Content + } + w.Header().Set("Content-Type", "text/event-stream") + fmt.Fprint(w, "data: {\"choices\":[{\"delta\":{\"content\":\"ok\"}}]}\n\n") + fmt.Fprint(w, "data: [DONE]\n\n") + })) + t.Cleanup(srv.Close) + + engine, _, _ := testEngine(t) + engine.SetDispatcher(seam.Func(func(_ context.Context, ev seam.Envelope) (seam.Envelope, error) { + switch ev.Type { + case seam.TypeRequest: + var payload RequestPayload + _ = json.Unmarshal(ev.Payload, &payload) + payload.SystemPrompt = "from-plugin" + ev.Payload = MarshalPayload(payload) + case seam.TypeStream: + var payload StreamPayload + _ = json.Unmarshal(ev.Payload, &payload) + if payload.Headers == nil { + payload.Headers = map[string]string{} + } + payload.Headers["X-Plugin"] = "1" + ev.Payload = MarshalPayload(payload) + } + return ev, nil + })) + opts, _ := json.Marshal(map[string]string{"api_key": "sk-test", "base_url": srv.URL}) + cfg := DefaultRunConfig(WorkAgent, ModelConfig{Provider: "openai", Model: "gpt-test", Options: opts}) + hist := fakeHistory("run-1", FakeOptions{}) + hist.Run.Config = cfg + got, err := engine.Step(context.Background(), StepInput{ + State: AgentState{SessionID: "sess-1", RunID: "run-1", Config: cfg}, + Job: StepJob{RunID: "run-1", StepIndex: 1, Phase: PhaseUserInput}, + History: hist, + }) + if err != nil { + t.Fatal(err) + } + if got.State.Status != RunRunningLLM { + t.Fatalf("status=%s", got.State.Status) + } + if gotSystem != "from-plugin" { + t.Fatalf("system=%q", gotSystem) + } + if gotHeader != "1" { + t.Fatalf("header=%q", gotHeader) + } +} + func TestEngineStepCallsDecide(t *testing.T) { engine := NewEngine(&Brain{}, nil, nil) got, err := engine.Step(context.Background(), StepInput{ diff --git a/server/pkg/agent/model.go b/server/pkg/agent/model.go index d33ec7a..8478c8d 100644 --- a/server/pkg/agent/model.go +++ b/server/pkg/agent/model.go @@ -7,6 +7,7 @@ import ( "strings" "time" + "codedock/pkg/agent/seam" "codedock/pkg/agent/tool" ) @@ -33,6 +34,7 @@ type Chat struct { MaxInputTokens int64 MaxOutputTokens int64 Attempt int + Dispatcher seam.Dispatcher } // ModelStreamEvent 是一条模型流事件。 diff --git a/server/pkg/agent/openai.go b/server/pkg/agent/openai.go index 18681ed..6a41e62 100644 --- a/server/pkg/agent/openai.go +++ b/server/pkg/agent/openai.go @@ -11,6 +11,9 @@ import ( "strings" "time" + "log/slog" + + "codedock/pkg/agent/seam" "codedock/pkg/agent/tool" ) @@ -101,13 +104,19 @@ func streamOpenAI(ctx context.Context, chat Chat) (ModelStream, error) { if err != nil { return nil, err } + headers := map[string]string{ + "Authorization": "Bearer " + opts.APIKey, + "Content-Type": "application/json", + "Accept": "text/event-stream", + } + body, headers = applyStreamSeam(ctx, chat, body, headers) req, err := http.NewRequestWithContext(ctx, http.MethodPost, base+"/chat/completions", bytes.NewReader(body)) if err != nil { return nil, err } - req.Header.Set("Authorization", "Bearer "+opts.APIKey) - req.Header.Set("Content-Type", "application/json") - req.Header.Set("Accept", "text/event-stream") + for key, value := range headers { + req.Header.Set(key, value) + } resp, err := http.DefaultClient.Do(req) if err != nil { @@ -132,6 +141,37 @@ func streamOpenAI(ctx context.Context, chat Chat) (ModelStream, error) { return stream, nil } +// applyStreamSeam 在真正发出 HTTP 前递请求头和请求体;出错保留原数据。 +func applyStreamSeam(ctx context.Context, chat Chat, body json.RawMessage, headers map[string]string) (json.RawMessage, map[string]string) { + ev, err := seam.Dispatch(ctx, chat.Dispatcher, seam.Envelope{ + Type: seam.TypeStream, + SessionID: chat.SessionID, + RunID: chat.RunID, + TurnID: chat.TurnID, + Payload: MarshalPayload(StreamPayload{Headers: headers, Body: body}), + }) + if err != nil { + slog.Warn("llm/stream dispatch failed", "run_id", chat.RunID, "error", err) + return body, headers + } + if ev.Type != seam.TypeStream || len(ev.Payload) == 0 { + return body, headers + } + var payload StreamPayload + if err := json.Unmarshal(ev.Payload, &payload); err != nil { + return body, headers + } + if payload.Headers != nil { + for key, value := range payload.Headers { + headers[key] = value + } + } + if len(payload.Body) > 0 { + body = payload.Body + } + return body, headers +} + func applyOutputLimit(req *openaiChatRequest, model string, n int64) { if req == nil || n <= 0 { return diff --git a/server/pkg/agent/payload.go b/server/pkg/agent/payload.go new file mode 100644 index 0000000..fac581d --- /dev/null +++ b/server/pkg/agent/payload.go @@ -0,0 +1,27 @@ +package agent + +import "encoding/json" + +// InputPayload 是 agent/input 口能改的数据。 +type InputPayload struct { + Content string `json:"content"` + Mode WorkMode `json:"mode,omitempty"` // ask / plan / agent;空表示不改 +} + +// PreStepPayload 是 agent/pre-step 口能改的数据。 +type PreStepPayload struct { + SystemPrompt string `json:"system_prompt"` + Hidden []Message `json:"hidden,omitempty"` +} + +// RequestPayload 是 agent/request 口能改的数据。 +type RequestPayload struct { + SystemPrompt string `json:"system_prompt"` + Messages []Message `json:"messages,omitempty"` +} + +// StreamPayload 是 llm/stream 口能改的数据。 +type StreamPayload struct { + Headers map[string]string `json:"headers,omitempty"` + Body json.RawMessage `json:"body,omitempty"` +} diff --git a/server/pkg/agent/prompt.go b/server/pkg/agent/prompt.go index 842d39e..684f1a3 100644 --- a/server/pkg/agent/prompt.go +++ b/server/pkg/agent/prompt.go @@ -140,6 +140,12 @@ func Build(_ context.Context, req Prompt) (Chat, error) { } prefix = append(prefix, Message{Role: RoleSystem, Content: EncodeText(index)}) } + for _, msg := range req.Context.Hidden { + if len(msg.Content) == 0 { + continue + } + prefix = append(prefix, Message{Role: RoleSystem, Content: msg.Content}) + } if req.Context.Summary != nil && req.Context.Summary.Content != "" { prefix = append(prefix, Message{ Role: RoleSystem, diff --git a/server/pkg/agent/prompt_test.go b/server/pkg/agent/prompt_test.go index cae20bc..b99d932 100644 --- a/server/pkg/agent/prompt_test.go +++ b/server/pkg/agent/prompt_test.go @@ -164,3 +164,27 @@ func lastDeveloper(messages []Message) string { } return "" } + +func TestBuildInsertsHiddenAfterMemory(t *testing.T) { + t.Parallel() + chat, err := Build(context.Background(), Prompt{ + Context: ContextSnapshot{ + SystemPrompt: "base", + MemoryIndexes: []string{"index-note"}, + Hidden: []Message{{Role: RoleSystem, Content: EncodeText("hidden-note")}}, + Messages: []Message{{Role: RoleUser, Content: EncodeText("hi")}}, + }, + }) + if err != nil { + t.Fatal(err) + } + if len(chat.Messages) < 3 { + t.Fatalf("messages=%d", len(chat.Messages)) + } + if DecodeText(chat.Messages[0].Content) != "index-note" || chat.Messages[0].Role != RoleSystem { + t.Fatalf("memory first: %+v", chat.Messages[0]) + } + if DecodeText(chat.Messages[1].Content) != "hidden-note" || chat.Messages[1].Role != RoleSystem { + t.Fatalf("hidden second: %+v", chat.Messages[1]) + } +} diff --git a/server/pkg/agent/seam/seam.go b/server/pkg/agent/seam/seam.go new file mode 100644 index 0000000..1442298 --- /dev/null +++ b/server/pkg/agent/seam/seam.go @@ -0,0 +1,67 @@ +package seam + +import ( + "context" + "encoding/json" +) + +const ( + TypeInput = "agent/input" + TypePreStep = "agent/pre-step" + TypeRequest = "agent/request" + TypeStream = "llm/stream" + TypePreExecute = "tools/pre-execute" + TypePostExecute = "tools/post-execute" + + TypeInputHandled = "input/handled" + TypeRunBlocked = "run/blocked" + TypeToolsDenied = "tools/denied" + TypeToolsAsk = "tools/ask" +) + +// Envelope 是一次喊话:当前口、已问过谁、以及还能改的数据。 +type Envelope struct { + Type string `json:"type"` // 口或事件名 + ChainID string `json:"chain_id"` // 这一次 Dispatch 的链,宿主生成 + Seen []string `json:"seen,omitempty"` // 这条链上已经问过的插件(目录名) + SessionID string `json:"session_id,omitempty"` + RunID string `json:"run_id,omitempty"` + TurnID string `json:"turn_id,omitempty"` + Payload json.RawMessage `json:"payload,omitempty"` // 该口的 JSON 载荷 + Context json.RawMessage `json:"context,omitempty"` // 插件共享袋子,不进模型、不换向 +} + +// Dispatcher 把信封交给订阅者。没有实现时由 Dispatch 原样返回。 +type Dispatcher interface { + // Dispatch 按订阅依次处理信封;类型变了应立刻返回。 + Dispatch(ctx context.Context, ev Envelope) (Envelope, error) +} + +// Dispatch 是所有口统一调用的入口:d 为 nil 时原信封原样返回。 +func Dispatch(ctx context.Context, d Dispatcher, ev Envelope) (Envelope, error) { + if d == nil { + return ev, nil + } + return d.Dispatch(ctx, ev) +} + +// Func 让测试用闭包充当 Dispatcher。 +type Func func(ctx context.Context, ev Envelope) (Envelope, error) + +// Dispatch 调用闭包;闭包为 nil 时原信封原样返回。 +func (f Func) Dispatch(ctx context.Context, ev Envelope) (Envelope, error) { + if f == nil { + return ev, nil + } + return f(ctx, ev) +} + +// IsSeam 判断 typ 是否为六个拦截口之一。 +func IsSeam(typ string) bool { + switch typ { + case TypeInput, TypePreStep, TypeRequest, TypeStream, TypePreExecute, TypePostExecute: + return true + default: + return false + } +} diff --git a/server/pkg/agent/seam/seam_test.go b/server/pkg/agent/seam/seam_test.go new file mode 100644 index 0000000..e5da21d --- /dev/null +++ b/server/pkg/agent/seam/seam_test.go @@ -0,0 +1,57 @@ +package seam + +import ( + "context" + "errors" + "testing" +) + +func TestDispatchNilReturnsOriginal(t *testing.T) { + t.Parallel() + in := Envelope{Type: TypeInput, Payload: []byte(`{"content":"hi"}`)} + got, err := Dispatch(context.Background(), nil, in) + if err != nil { + t.Fatal(err) + } + if got.Type != in.Type || string(got.Payload) != string(in.Payload) { + t.Fatalf("got=%+v", got) + } +} + +func TestFuncRewrites(t *testing.T) { + t.Parallel() + d := Func(func(_ context.Context, ev Envelope) (Envelope, error) { + ev.Type = TypeInputHandled + return ev, nil + }) + got, err := Dispatch(context.Background(), d, Envelope{Type: TypeInput}) + if err != nil || got.Type != TypeInputHandled { + t.Fatalf("got=%+v err=%v", got, err) + } +} + +func TestFuncNilAndError(t *testing.T) { + t.Parallel() + var empty Func + got, err := Dispatch(context.Background(), empty, Envelope{Type: TypeRequest}) + if err != nil || got.Type != TypeRequest { + t.Fatalf("nil func: %+v %v", got, err) + } + want := errors.New("boom") + _, err = Dispatch(context.Background(), Func(func(context.Context, Envelope) (Envelope, error) { + return Envelope{}, want + }), Envelope{Type: TypeInput}) + if !errors.Is(err, want) { + t.Fatalf("err=%v", err) + } +} + +func TestIsSeam(t *testing.T) { + t.Parallel() + if !IsSeam(TypeInput) || !IsSeam(TypePostExecute) { + t.Fatal("seams should match") + } + if IsSeam(TypeInputHandled) || IsSeam("run.completed") || IsSeam("") { + t.Fatal("non-seams should not match") + } +} diff --git a/server/pkg/agent/tool/dispatch.go b/server/pkg/agent/tool/dispatch.go index 8261ce5..003f4c4 100644 --- a/server/pkg/agent/tool/dispatch.go +++ b/server/pkg/agent/tool/dispatch.go @@ -5,8 +5,11 @@ import ( "encoding/json" "errors" "fmt" + "log/slog" "sync" "time" + + "codedock/pkg/agent/seam" ) // Dispatch 按权限与审批策略调度工具调用。 @@ -48,9 +51,10 @@ func Dispatch(ctx context.Context, inv Invocation) (DispatchResult, error) { }) continue } - item, wait := prepareCall(ctx, inv, call, approved) + item, wait, updated := prepareCall(ctx, inv, call, approved) + inv.Calls[i] = updated if wait { - approvalCalls = append(approvalCalls, call) + approvalCalls = append(approvalCalls, updated) continue } prepared = append(prepared, item) @@ -83,12 +87,13 @@ type preparedCall struct { } // prepareCall 查找工具并做参数、权限、审批校验;wait=true 表示需先审批。 -// 查不到、参数错、权限不足都写成失败 Result,不返回 error。 -func prepareCall(ctx context.Context, inv Invocation, call Call, approved map[string]struct{}) (preparedCall, bool) { +// 查不到、参数错、权限不足、插件否决都写成失败 Result,不返回 error。 +// 未批准的调用在流水线之后再走 tools/pre-execute,可改参、否决或抬到审批。 +func prepareCall(ctx context.Context, inv Invocation, call Call, approved map[string]struct{}) (preparedCall, bool, Call) { emit(inv, "call_started", call, max(1, call.Attempt), nil) item, err := inv.Registry.Get(Reference{Name: call.Name}) if err != nil { - return preparedCall{call: call, result: failResult(call, err.Error()), skip: true}, false + return preparedCall{call: call, result: failResult(call, err.Error()), skip: true}, false, call } def := item.Definition() bound := Bound(inv.BoundNames, call.Name) @@ -136,12 +141,97 @@ func prepareCall(ctx context.Context, inv Invocation, call Call, approved map[st if inspectErr != nil { msg = inspectErr.Error() } - return preparedCall{call: call, result: failResult(call, msg), skip: true}, false + return preparedCall{call: call, result: failResult(call, msg), skip: true}, false, call + } + if !already { + updated, denied, ask := applyPreExecute(ctx, inv, call) + call = updated + if denied { + return preparedCall{call: call, result: failResult(call, "denied by plugin"), skip: true}, false, call + } + if ask { + return preparedCall{}, true, call + } } if effect == EffectAsk { - return preparedCall{}, true + return preparedCall{}, true, call + } + return preparedCall{call: call, tool: item}, false, call +} + +// applyPreExecute 在审批判断前递工具入参。已批准的调用不会走到这里。 +func applyPreExecute(ctx context.Context, inv Invocation, call Call) (Call, bool, bool) { + ev, err := seam.Dispatch(ctx, inv.Dispatcher, seam.Envelope{ + Type: seam.TypePreExecute, + SessionID: inv.SessionID, + RunID: inv.RunID, + TurnID: inv.TurnID, + Payload: mustJSON(PreExecutePayload{Call: call}), + }) + if err != nil { + slog.Warn("tools/pre-execute dispatch failed", "run_id", inv.RunID, "call_id", call.ID, "error", err) + return call, true, false + } + updated := call + if ev.Type == seam.TypePreExecute || ev.Type == seam.TypeToolsAsk { + if len(ev.Payload) > 0 { + var payload PreExecutePayload + if err := json.Unmarshal(ev.Payload, &payload); err == nil { + if payload.Call.ID == "" { + payload.Call.ID = call.ID + } + if payload.Call.Name == "" { + payload.Call.Name = call.Name + } + updated = payload.Call + } + } + } + switch ev.Type { + case seam.TypeToolsDenied: + return call, true, false + case seam.TypeToolsAsk: + return updated, false, true + default: + return updated, false, false + } +} + +func mustJSON(v any) json.RawMessage { + body, err := json.Marshal(v) + if err != nil { + return json.RawMessage("{}") + } + return body +} + +// applyPostExecute 在写结果前递工具结果;出错保留原结果。 +func applyPostExecute(ctx context.Context, inv Invocation, call Call, result Result) Result { + ev, err := seam.Dispatch(ctx, inv.Dispatcher, seam.Envelope{ + Type: seam.TypePostExecute, + SessionID: inv.SessionID, + RunID: inv.RunID, + TurnID: inv.TurnID, + Payload: mustJSON(PostExecutePayload{Call: call, Result: result}), + }) + if err != nil { + slog.Warn("tools/post-execute dispatch failed", "run_id", inv.RunID, "call_id", call.ID, "error", err) + return result + } + if ev.Type != seam.TypePostExecute || len(ev.Payload) == 0 { + return result + } + var payload PostExecutePayload + if err := json.Unmarshal(ev.Payload, &payload); err != nil { + return result + } + if payload.Result.CallID == "" { + payload.Result.CallID = result.CallID + } + if payload.Result.Name == "" { + payload.Result.Name = result.Name } - return preparedCall{call: call, tool: item}, false + return payload.Result } // runSerial 按调用顺序执行工具;fail_fast 遇失败 Result 后不再跑后续调用。 @@ -248,6 +338,7 @@ func executeOne(ctx context.Context, inv Invocation, item preparedCall) (Result, } last = result if last.Success { + last = applyPostExecute(ctx, inv, item.call, last) emit(inv, "execution_result", item.call, attempt, &last) return last, nil } @@ -256,6 +347,7 @@ func executeOne(ctx context.Context, inv Invocation, item preparedCall) (Result, retryErr = fmt.Errorf("tool failed") } if !def.SupportsRetry || !retryableTool(retryErr) || attempt >= maxAttempts(inv) { + last = applyPostExecute(ctx, inv, item.call, last) emit(inv, "execution_result", item.call, attempt, &last) return last, nil } diff --git a/server/pkg/agent/tool/dispatch_test.go b/server/pkg/agent/tool/dispatch_test.go index 99b489c..2634fc6 100644 --- a/server/pkg/agent/tool/dispatch_test.go +++ b/server/pkg/agent/tool/dispatch_test.go @@ -7,6 +7,8 @@ import ( "sync/atomic" "testing" "time" + + "codedock/pkg/agent/seam" ) type stubTool struct { @@ -163,10 +165,12 @@ func TestDispatchFailFastStopsOnFirstError(t *testing.T) { type outsideTool struct{ stubTool } +// Inspect 假装路径落在工作区外。 func (outsideTool) Inspect(context.Context, Input) error { return ErrOutsideWorkspace } type allowAskTool struct{ stubTool } +// ResolveEffect 把默认 ask 抬成 allow。 func (t allowAskTool) ResolveEffect(context.Context, Input) Effect { return EffectAllow } @@ -201,6 +205,126 @@ func TestDispatchOutsideWorkspaceNeedsApproval(t *testing.T) { } } +func TestDispatchPreExecuteRewritesArguments(t *testing.T) { + reg := NewRegistry() + _ = reg.Register(stubTool{def: Definition{Name: "ping", Permission: Permission{Effect: EffectAllow}}}) + inv := dispatchInv(reg, []Call{{ID: "c1", Name: "ping", Arguments: json.RawMessage(`{"x":1}`)}}, ApprovalYolo) + inv.Dispatcher = seam.Func(func(_ context.Context, ev seam.Envelope) (seam.Envelope, error) { + if ev.Type != seam.TypePreExecute { + return ev, nil + } + var payload PreExecutePayload + if err := json.Unmarshal(ev.Payload, &payload); err != nil { + t.Fatal(err) + } + payload.Call.Arguments = json.RawMessage(`{"x":2}`) + body, _ := json.Marshal(payload) + ev.Payload = body + return ev, nil + }) + out, err := Dispatch(context.Background(), inv) + if err != nil || len(out.Results) != 1 || !out.Results[0].Success { + t.Fatalf("err=%v out=%+v", err, out) + } +} + +func TestDispatchPreExecuteDenied(t *testing.T) { + reg := NewRegistry() + _ = reg.Register(stubTool{def: Definition{Name: "ping", Permission: Permission{Effect: EffectAllow}}}) + inv := dispatchInv(reg, []Call{{ID: "c1", Name: "ping", Arguments: json.RawMessage(`{}`)}}, ApprovalYolo) + inv.Dispatcher = seam.Func(func(_ context.Context, ev seam.Envelope) (seam.Envelope, error) { + ev.Type = seam.TypeToolsDenied + return ev, nil + }) + out, err := Dispatch(context.Background(), inv) + if err != nil || len(out.Results) != 1 || out.Results[0].Success || out.Results[0].Error != "denied by plugin" { + t.Fatalf("err=%v out=%+v", err, out) + } +} + +func TestDispatchPreExecuteAsk(t *testing.T) { + reg := NewRegistry() + _ = reg.Register(stubTool{def: Definition{Name: "ping", Permission: Permission{Effect: EffectAllow}}}) + inv := dispatchInv(reg, []Call{{ID: "c1", Name: "ping", Arguments: json.RawMessage(`{"x":1}`)}}, ApprovalYolo) + inv.Dispatcher = seam.Func(func(_ context.Context, ev seam.Envelope) (seam.Envelope, error) { + var payload PreExecutePayload + _ = json.Unmarshal(ev.Payload, &payload) + payload.Call.Arguments = json.RawMessage(`{"x":9}`) + body, _ := json.Marshal(payload) + ev.Payload = body + ev.Type = seam.TypeToolsAsk + return ev, nil + }) + out, err := Dispatch(context.Background(), inv) + if err != nil || !out.WaitingApproval || len(out.PendingCalls) != 1 || string(out.PendingCalls[0].Arguments) != `{"x":9}` { + t.Fatalf("err=%v out=%+v", err, out) + } +} + +func TestDispatchPreExecuteSkipsApproved(t *testing.T) { + reg := NewRegistry() + _ = reg.Register(stubTool{def: Definition{Name: "ping", Permission: Permission{Effect: EffectAsk}}}) + called := 0 + inv := dispatchInv(reg, []Call{{ID: "c1", Name: "ping", Arguments: json.RawMessage(`{}`)}}, ApprovalManual) + inv.ApprovedCallIDs = []string{"c1"} + inv.Dispatcher = seam.Func(func(_ context.Context, ev seam.Envelope) (seam.Envelope, error) { + if ev.Type == seam.TypePreExecute { + called++ + } + return ev, nil + }) + out, err := Dispatch(context.Background(), inv) + if err != nil || called != 0 || len(out.Results) != 1 || !out.Results[0].Success { + t.Fatalf("called=%d err=%v out=%+v", called, err, out) + } +} + +func TestDispatchPostExecuteRewritesResult(t *testing.T) { + reg := NewRegistry() + _ = reg.Register(stubTool{def: Definition{Name: "ping", Permission: Permission{Effect: EffectAllow}}}) + inv := dispatchInv(reg, []Call{{ID: "c1", Name: "ping", Arguments: json.RawMessage(`{}`)}}, ApprovalYolo) + inv.Dispatcher = seam.Func(func(_ context.Context, ev seam.Envelope) (seam.Envelope, error) { + if ev.Type != seam.TypePostExecute { + return ev, nil + } + var payload PostExecutePayload + _ = json.Unmarshal(ev.Payload, &payload) + payload.Result.Output = json.RawMessage(`{"ok":false}`) + body, _ := json.Marshal(payload) + ev.Payload = body + return ev, nil + }) + out, err := Dispatch(context.Background(), inv) + if err != nil || len(out.Results) != 1 || string(out.Results[0].Output) != `{"ok":false}` { + t.Fatalf("err=%v out=%+v", err, out) + } +} + +func TestDispatchFailFastStopsOnPluginDenied(t *testing.T) { + reg := NewRegistry() + _ = reg.Register(stubTool{def: Definition{Name: "ping", Permission: Permission{Effect: EffectAllow}}}) + inv := dispatchInv(reg, []Call{ + {ID: "c1", Name: "ping", Arguments: json.RawMessage(`{}`)}, + {ID: "c2", Name: "ping", Arguments: json.RawMessage(`{}`)}, + }, ApprovalYolo) + inv.FailurePolicy = FailureFast + inv.Dispatcher = seam.Func(func(_ context.Context, ev seam.Envelope) (seam.Envelope, error) { + var payload PreExecutePayload + _ = json.Unmarshal(ev.Payload, &payload) + if payload.Call.ID == "c1" { + ev.Type = seam.TypeToolsDenied + } + return ev, nil + }) + out, err := Dispatch(context.Background(), inv) + if err != nil { + t.Fatal(err) + } + if len(out.Results) != 2 || out.Results[0].Success || out.Results[1].Success || out.Results[1].Error != "tool did not execute" { + t.Fatalf("out=%+v", out) + } +} + func TestDispatchUnknownToolIsFailedResult(t *testing.T) { reg := NewRegistry() out, err := Dispatch(context.Background(), Invocation{ diff --git a/server/pkg/agent/tool/payload.go b/server/pkg/agent/tool/payload.go new file mode 100644 index 0000000..94373d2 --- /dev/null +++ b/server/pkg/agent/tool/payload.go @@ -0,0 +1,12 @@ +package tool + +// PreExecutePayload 是 tools/pre-execute 口能改的数据。 +type PreExecutePayload struct { + Call Call `json:"call"` +} + +// PostExecutePayload 是 tools/post-execute 口能改的数据。 +type PostExecutePayload struct { + Call Call `json:"call"` + Result Result `json:"result"` +} diff --git a/server/pkg/agent/tool/tool.go b/server/pkg/agent/tool/tool.go index 753546b..baa8d17 100644 --- a/server/pkg/agent/tool/tool.go +++ b/server/pkg/agent/tool/tool.go @@ -4,6 +4,8 @@ import ( "context" "encoding/json" "errors" + + "codedock/pkg/agent/seam" ) // ErrOutsideWorkspace 表示路径落在会话工作区外。第 1 层校验不算失败,但必须走审批。 @@ -148,14 +150,15 @@ type Invocation struct { Mode ExecutionMode FailurePolicy FailurePolicy MaxParallel int - BoundNames []string - Effects map[string]Effect - Approval ApprovalMode + BoundNames []string // 本 Agent 可执行名;未列入则 deny + Effects map[string]Effect // 第 2 层覆盖表 + Approval ApprovalMode // 第 3 层 Registry Registry ApprovedCallIDs []string DeniedCallIDs []string OnEvent DispatchHook Gate Gate + Dispatcher seam.Dispatcher // 插件 tools/pre-execute 与 post-execute;空则原样通过 } // DispatchResult 按模型调用顺序保存结果,并标识是否因审批暂停。 diff --git a/server/pkg/agent/types.go b/server/pkg/agent/types.go index 9c5670a..908925c 100644 --- a/server/pkg/agent/types.go +++ b/server/pkg/agent/types.go @@ -263,6 +263,7 @@ type ContextSnapshot struct { Tools []tool.Definition `json:"tools"` // 本轮可见工具定义 SystemPrompt string `json:"system_prompt"` // 注入的系统提示(身份段;Compose 再拼工具与 Guidelines) WorkspaceRoot string `json:"workspace_root,omitempty"` // 会话冻结的工作目录,写入 Current working directory + Hidden []Message `json:"hidden,omitempty"` // 插件注入的隐藏消息,不入库 MemoryIndexes []string `json:"memory_indexes,omitempty"` // 冻结记忆目录 EstimatedTokens int64 `json:"estimated_tokens"` // 估算 token 数 Version int64 `json:"version"` // 快照版本 diff --git a/server/pkg/db/queries/run_overlays.sql b/server/pkg/db/queries/run_overlays.sql new file mode 100644 index 0000000..65b1e1b --- /dev/null +++ b/server/pkg/db/queries/run_overlays.sql @@ -0,0 +1,15 @@ +-- name: GetRunOverlay :one +SELECT * FROM run_overlays +WHERE run_id = ?; + +-- name: UpsertRunOverlay :one +INSERT INTO run_overlays ( + run_id, system_prompt, hidden, updated_at +) VALUES ( + ?, ?, ?, ? +) +ON CONFLICT(run_id) DO UPDATE SET + system_prompt = excluded.system_prompt, + hidden = excluded.hidden, + updated_at = excluded.updated_at +RETURNING *; diff --git a/server/pkg/db/sqlite/models.go b/server/pkg/db/sqlite/models.go index e1b6e58..0d9f209 100644 --- a/server/pkg/db/sqlite/models.go +++ b/server/pkg/db/sqlite/models.go @@ -82,6 +82,13 @@ type Run struct { FinishedAt sql.NullString } +type RunOverlay struct { + RunID string + SystemPrompt string + Hidden string + UpdatedAt string +} + type RunToolCheckpoint struct { RunID string TurnID string diff --git a/server/pkg/db/sqlite/run_overlays.sql.go b/server/pkg/db/sqlite/run_overlays.sql.go new file mode 100644 index 0000000..a914255 --- /dev/null +++ b/server/pkg/db/sqlite/run_overlays.sql.go @@ -0,0 +1,64 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.31.1 +// source: run_overlays.sql + +package sqlite + +import ( + "context" +) + +const getRunOverlay = `-- name: GetRunOverlay :one +SELECT run_id, system_prompt, hidden, updated_at FROM run_overlays +WHERE run_id = ? +` + +func (q *Queries) GetRunOverlay(ctx context.Context, runID string) (RunOverlay, error) { + row := q.db.QueryRowContext(ctx, getRunOverlay, runID) + var i RunOverlay + err := row.Scan( + &i.RunID, + &i.SystemPrompt, + &i.Hidden, + &i.UpdatedAt, + ) + return i, err +} + +const upsertRunOverlay = `-- name: UpsertRunOverlay :one +INSERT INTO run_overlays ( + run_id, system_prompt, hidden, updated_at +) VALUES ( + ?, ?, ?, ? +) +ON CONFLICT(run_id) DO UPDATE SET + system_prompt = excluded.system_prompt, + hidden = excluded.hidden, + updated_at = excluded.updated_at +RETURNING run_id, system_prompt, hidden, updated_at +` + +type UpsertRunOverlayParams struct { + RunID string + SystemPrompt string + Hidden string + UpdatedAt string +} + +func (q *Queries) UpsertRunOverlay(ctx context.Context, arg UpsertRunOverlayParams) (RunOverlay, error) { + row := q.db.QueryRowContext(ctx, upsertRunOverlay, + arg.RunID, + arg.SystemPrompt, + arg.Hidden, + arg.UpdatedAt, + ) + var i RunOverlay + err := row.Scan( + &i.RunID, + &i.SystemPrompt, + &i.Hidden, + &i.UpdatedAt, + ) + return i, err +} diff --git a/server/pkg/plugin/adapt.go b/server/pkg/plugin/adapt.go new file mode 100644 index 0000000..f4d0eac --- /dev/null +++ b/server/pkg/plugin/adapt.go @@ -0,0 +1,216 @@ +package plugin + +import ( + "context" + "encoding/json" + + pkgagent "codedock/pkg/agent" + "codedock/pkg/agent/seam" + "codedock/pkg/agent/tool" +) + +// runtime 把作者的 Plugin 接到宿主用的信封 OnEvent。 +type runtime struct { + Plugin +} + +// OnEvent 按信封类型拆成各口结构体,再调作者实现。 +func (r *runtime) OnEvent(ctx context.Context, ev seam.Envelope) (seam.Envelope, error) { + return applyHooks(ctx, r.Plugin, ev) +} + +// applyHooks 按 Type 分发给对应口;未实现该口则原样返回。 +func applyHooks(ctx context.Context, p Plugin, ev seam.Envelope) (seam.Envelope, error) { + switch ev.Type { + case TypeInput: + return applyInput(ctx, p, ev) + case TypePreStep: + return applyPreStep(ctx, p, ev) + case TypeRequest: + return applyRequest(ctx, p, ev) + case TypeStream: + return applyStream(ctx, p, ev) + case TypePreExecute: + return applyPreExecute(ctx, p, ev) + case TypePostExecute: + return applyPostExecute(ctx, p, ev) + default: + if h, ok := p.(LedgerNotifyHandler); ok { + return ev, h.OnLedgerNotify(ctx, LedgerNotify{ + Type: ev.Type, + SessionID: ev.SessionID, + RunID: ev.RunID, + TurnID: ev.TurnID, + Payload: ev.Payload, + Context: DecodePluginContext(ev.Context), + }) + } + return ev, nil + } +} + +// applyInput 把信封转成 AgentInput / AgentInputResult。 +func applyInput(ctx context.Context, p Plugin, ev seam.Envelope) (seam.Envelope, error) { + h, ok := p.(AgentInputHandler) + if !ok { + return ev, nil + } + var payload pkgagent.InputPayload + _ = json.Unmarshal(ev.Payload, &payload) + out, err := h.OnAgentInput(ctx, AgentInput{ + SessionID: ev.SessionID, + RunID: ev.RunID, + TurnID: ev.TurnID, + Content: payload.Content, + Mode: payload.Mode, + Context: DecodePluginContext(ev.Context), + }) + if err != nil { + return ev, err + } + ev.Payload = mustPayload(pkgagent.InputPayload{Content: out.Content, Mode: out.Mode}) + ev.Context = out.Context.Raw() + if out.Handled { + ev.Type = TypeInputHandled + } + return ev, nil +} + +// applyPreStep 把信封转成 AgentPreStep / AgentPreStepResult。 +func applyPreStep(ctx context.Context, p Plugin, ev seam.Envelope) (seam.Envelope, error) { + h, ok := p.(AgentPreStepHandler) + if !ok { + return ev, nil + } + var payload pkgagent.PreStepPayload + _ = json.Unmarshal(ev.Payload, &payload) + out, err := h.OnAgentPreStep(ctx, AgentPreStep{ + SessionID: ev.SessionID, + RunID: ev.RunID, + TurnID: ev.TurnID, + SystemPrompt: payload.SystemPrompt, + Hidden: payload.Hidden, + Context: DecodePluginContext(ev.Context), + }) + if err != nil { + return ev, err + } + ev.Payload = mustPayload(pkgagent.PreStepPayload{SystemPrompt: out.SystemPrompt, Hidden: out.Hidden}) + ev.Context = out.Context.Raw() + if out.Blocked { + ev.Type = TypeRunBlocked + } + return ev, nil +} + +// applyRequest 把信封转成 AgentRequest / AgentRequestResult。 +func applyRequest(ctx context.Context, p Plugin, ev seam.Envelope) (seam.Envelope, error) { + h, ok := p.(AgentRequestHandler) + if !ok { + return ev, nil + } + var payload pkgagent.RequestPayload + _ = json.Unmarshal(ev.Payload, &payload) + out, err := h.OnAgentRequest(ctx, AgentRequest{ + SessionID: ev.SessionID, + RunID: ev.RunID, + TurnID: ev.TurnID, + SystemPrompt: payload.SystemPrompt, + Messages: payload.Messages, + Context: DecodePluginContext(ev.Context), + }) + if err != nil { + return ev, err + } + ev.Payload = mustPayload(pkgagent.RequestPayload{SystemPrompt: out.SystemPrompt, Messages: out.Messages}) + ev.Context = out.Context.Raw() + return ev, nil +} + +// applyStream 把信封转成 LLMStream / LLMStreamResult。 +func applyStream(ctx context.Context, p Plugin, ev seam.Envelope) (seam.Envelope, error) { + h, ok := p.(LLMStreamHandler) + if !ok { + return ev, nil + } + var payload pkgagent.StreamPayload + _ = json.Unmarshal(ev.Payload, &payload) + out, err := h.OnLLMStream(ctx, LLMStream{ + SessionID: ev.SessionID, + RunID: ev.RunID, + TurnID: ev.TurnID, + Headers: payload.Headers, + Body: payload.Body, + Context: DecodePluginContext(ev.Context), + }) + if err != nil { + return ev, err + } + ev.Payload = mustPayload(pkgagent.StreamPayload{Headers: out.Headers, Body: out.Body}) + ev.Context = out.Context.Raw() + return ev, nil +} + +// applyPreExecute 把信封转成 ToolPreExecute / ToolPreExecuteResult。 +func applyPreExecute(ctx context.Context, p Plugin, ev seam.Envelope) (seam.Envelope, error) { + h, ok := p.(ToolPreExecuteHandler) + if !ok { + return ev, nil + } + var payload tool.PreExecutePayload + _ = json.Unmarshal(ev.Payload, &payload) + out, err := h.OnToolPreExecute(ctx, ToolPreExecute{ + SessionID: ev.SessionID, + RunID: ev.RunID, + TurnID: ev.TurnID, + Call: payload.Call, + Context: DecodePluginContext(ev.Context), + }) + if err != nil { + return ev, err + } + ev.Payload = mustPayload(tool.PreExecutePayload{Call: out.Call}) + ev.Context = out.Context.Raw() + if out.Denied { + ev.Type = TypeToolsDenied + } else if out.Ask { + ev.Type = TypeToolsAsk + } + return ev, nil +} + +// applyPostExecute 把信封转成 ToolPostExecute / ToolPostExecuteResult。 +func applyPostExecute(ctx context.Context, p Plugin, ev seam.Envelope) (seam.Envelope, error) { + h, ok := p.(ToolPostExecuteHandler) + if !ok { + return ev, nil + } + var payload tool.PostExecutePayload + _ = json.Unmarshal(ev.Payload, &payload) + out, err := h.OnToolPostExecute(ctx, ToolPostExecute{ + SessionID: ev.SessionID, + RunID: ev.RunID, + TurnID: ev.TurnID, + Call: payload.Call, + Result: payload.Result, + Context: DecodePluginContext(ev.Context), + }) + if err != nil { + return ev, err + } + ev.Payload = mustPayload(tool.PostExecutePayload{Call: out.Call, Result: out.Result}) + ev.Context = out.Context.Raw() + return ev, nil +} + +// mustPayload 把结构体编成信封载荷;失败时退回空对象。 +func mustPayload(v any) json.RawMessage { + if raw, ok := v.(json.RawMessage); ok { + return raw + } + body, err := json.Marshal(v) + if err != nil { + return json.RawMessage("{}") + } + return body +} diff --git a/server/pkg/plugin/adapt_test.go b/server/pkg/plugin/adapt_test.go new file mode 100644 index 0000000..a46dfd0 --- /dev/null +++ b/server/pkg/plugin/adapt_test.go @@ -0,0 +1,217 @@ +package plugin + +import ( + "context" + "encoding/json" + "testing" + + pkgagent "codedock/pkg/agent" + "codedock/pkg/agent/seam" + "codedock/pkg/agent/tool" +) + +type stubPlugin struct{} + +// Bootstrap 返回空订阅的占位清单。 +func (stubPlugin) Bootstrap(context.Context, Host) (Manifest, error) { + return Manifest{Name: "stub"}, nil +} + +// ExecuteMethod 声明占位插件没有方法。 +func (stubPlugin) ExecuteMethod(context.Context, MethodInput) (MethodResult, error) { + return MethodResult{Success: false, Error: "no methods"}, nil +} + +type rewritePlugin struct{ stubPlugin } + +// OnAgentInput 给正文加方括号;内容为 skip 则不建 Run。 +func (rewritePlugin) OnAgentInput(_ context.Context, in AgentInput) (AgentInputResult, error) { + if in.Content == "skip" { + return in.Handle(), nil + } + in.Content = "[" + in.Content + "]" + in.Context.Set("rewrite.seen", true) + return in.Reply(), nil +} + +// OnAgentPreStep 把系统提示改成 new;stop 则取消本轮。 +func (rewritePlugin) OnAgentPreStep(_ context.Context, in AgentPreStep) (AgentPreStepResult, error) { + if in.SystemPrompt == "stop" { + return in.Block(), nil + } + in.SystemPrompt = "new" + return in.Reply(), nil +} + +// OnAgentRequest 把系统提示改成 req。 +func (rewritePlugin) OnAgentRequest(_ context.Context, in AgentRequest) (AgentRequestResult, error) { + in.SystemPrompt = "req" + return in.Reply(), nil +} + +// OnLLMStream 加上 X=1 请求头。 +func (rewritePlugin) OnLLMStream(_ context.Context, in LLMStream) (LLMStreamResult, error) { + if in.Headers == nil { + in.Headers = map[string]string{} + } + in.Headers["X"] = "1" + return in.Reply(), nil +} + +// OnToolPreExecute 按工具名否决、送审或改参。 +func (rewritePlugin) OnToolPreExecute(_ context.Context, in ToolPreExecute) (ToolPreExecuteResult, error) { + switch in.Call.Name { + case "deny": + return in.Deny(), nil + case "ask": + return in.AskApproval(), nil + } + in.Call.Arguments = json.RawMessage(`{"x":2}`) + return in.Reply(), nil +} + +// OnToolPostExecute 把工具输出改成 ok。 +func (rewritePlugin) OnToolPostExecute(_ context.Context, in ToolPostExecute) (ToolPostExecuteResult, error) { + in.Result.Output = json.RawMessage(`{"ok":true}`) + return in.Reply(), nil +} + +// TestApplyHooksIdentityWithoutHandler 确认未实现的口原样通过。 +func TestApplyHooksIdentityWithoutHandler(t *testing.T) { + t.Parallel() + ev := seam.Envelope{Type: TypeInput, Payload: json.RawMessage(`{"content":"hi"}`)} + got, err := applyHooks(context.Background(), stubPlugin{}, ev) + if err != nil || got.Type != TypeInput || string(got.Payload) != `{"content":"hi"}` { + t.Fatalf("got=%+v err=%v", got, err) + } +} + +// TestApplyHooksTypedResults 覆盖六个口的改写和换向。 +func TestApplyHooksTypedResults(t *testing.T) { + t.Parallel() + p := rewritePlugin{} + + input, err := applyHooks(context.Background(), p, seam.Envelope{ + Type: TypeInput, + SessionID: "s1", + Payload: pkgagent.MarshalPayload(pkgagent.InputPayload{Content: "hi"}), + }) + if err != nil { + t.Fatal(err) + } + var in pkgagent.InputPayload + _ = json.Unmarshal(input.Payload, &in) + if input.Type != TypeInput || in.Content != "[hi]" { + t.Fatalf("input=%+v payload=%+v", input, in) + } + if string(DecodePluginContext(input.Context).Get("rewrite.seen")) != "true" { + t.Fatalf("context=%s", input.Context) + } + + handled, err := applyHooks(context.Background(), p, seam.Envelope{ + Type: TypeInput, + Payload: pkgagent.MarshalPayload(pkgagent.InputPayload{Content: "skip"}), + }) + if err != nil || handled.Type != TypeInputHandled { + t.Fatalf("handled=%+v err=%v", handled, err) + } + + blocked, err := applyHooks(context.Background(), p, seam.Envelope{ + Type: TypePreStep, + Payload: pkgagent.MarshalPayload(pkgagent.PreStepPayload{SystemPrompt: "stop"}), + }) + if err != nil || blocked.Type != TypeRunBlocked { + t.Fatalf("blocked=%+v err=%v", blocked, err) + } + + req, err := applyHooks(context.Background(), p, seam.Envelope{ + Type: TypeRequest, + Payload: pkgagent.MarshalPayload(pkgagent.RequestPayload{SystemPrompt: "old"}), + }) + if err != nil || req.Type != TypeRequest { + t.Fatalf("request=%+v err=%v", req, err) + } + var reqPayload pkgagent.RequestPayload + _ = json.Unmarshal(req.Payload, &reqPayload) + if reqPayload.SystemPrompt != "req" { + t.Fatalf("request payload=%+v", reqPayload) + } + + stream, err := applyHooks(context.Background(), p, seam.Envelope{ + Type: TypeStream, + Payload: pkgagent.MarshalPayload(pkgagent.StreamPayload{Headers: map[string]string{}}), + }) + if err != nil { + t.Fatal(err) + } + var streamPayload pkgagent.StreamPayload + _ = json.Unmarshal(stream.Payload, &streamPayload) + if streamPayload.Headers["X"] != "1" { + t.Fatalf("stream=%+v", streamPayload) + } + + rewritten, err := applyHooks(context.Background(), p, seam.Envelope{ + Type: TypePreExecute, + Payload: pkgagent.MarshalPayload(tool.PreExecutePayload{Call: tool.Call{Name: "ping"}}), + }) + if err != nil || rewritten.Type != TypePreExecute { + t.Fatalf("pre=%+v err=%v", rewritten, err) + } + var pre tool.PreExecutePayload + _ = json.Unmarshal(rewritten.Payload, &pre) + if string(pre.Call.Arguments) != `{"x":2}` { + t.Fatalf("pre payload=%+v", pre) + } + + denied, err := applyHooks(context.Background(), p, seam.Envelope{ + Type: TypePreExecute, + Payload: pkgagent.MarshalPayload(tool.PreExecutePayload{Call: tool.Call{Name: "deny"}}), + }) + if err != nil || denied.Type != TypeToolsDenied { + t.Fatalf("denied=%+v err=%v", denied, err) + } + + ask, err := applyHooks(context.Background(), p, seam.Envelope{ + Type: TypePreExecute, + Payload: pkgagent.MarshalPayload(tool.PreExecutePayload{Call: tool.Call{Name: "ask"}}), + }) + if err != nil || ask.Type != TypeToolsAsk { + t.Fatalf("ask=%+v err=%v", ask, err) + } + + post, err := applyHooks(context.Background(), p, seam.Envelope{ + Type: TypePostExecute, + Payload: pkgagent.MarshalPayload(tool.PostExecutePayload{ + Result: tool.Result{Output: json.RawMessage(`{}`)}, + }), + }) + if err != nil || post.Type != TypePostExecute { + t.Fatalf("post=%+v err=%v", post, err) + } + var postPayload tool.PostExecutePayload + _ = json.Unmarshal(post.Payload, &postPayload) + if string(postPayload.Result.Output) != `{"ok":true}` { + t.Fatalf("post payload=%+v", postPayload) + } +} + +type notifyPlugin struct { + stubPlugin + got string +} + +// OnLedgerNotify 记下事件类型供断言。 +func (p *notifyPlugin) OnLedgerNotify(_ context.Context, in LedgerNotify) error { + p.got = in.Type + return nil +} + +// TestApplyHooksNotify 确认非口事件走进 OnLedgerNotify。 +func TestApplyHooksNotify(t *testing.T) { + t.Parallel() + p := ¬ifyPlugin{} + ev, err := applyHooks(context.Background(), p, seam.Envelope{Type: "run.completed"}) + if err != nil || ev.Type != "run.completed" || p.got != "run.completed" { + t.Fatalf("ev=%+v got=%s err=%v", ev, p.got, err) + } +} diff --git a/server/pkg/plugin/context.go b/server/pkg/plugin/context.go new file mode 100644 index 0000000..2914941 --- /dev/null +++ b/server/pkg/plugin/context.go @@ -0,0 +1,80 @@ +package plugin + +import "encoding/json" + +// MaxPluginContextBytes 是宿主接受的袋子上限;超限则丢掉这次改动、保留上一份。 +const MaxPluginContextBytes = 8 << 10 + +// PluginContext 是插件之间共享的参数袋。不进模型、不进消息表、不换向。 +// 键建议写成 plugin.field,避免互相覆盖。宿主按会话暂存,有 Run 后挂到该 Run;进程重启即丢。 +type PluginContext struct { + Values map[string]json.RawMessage // 插件自约定的键值,值是 JSON +} + +// DecodePluginContext 把信封上的 JSON 对象解成袋子;空或非法时得到空 Values。 +func DecodePluginContext(raw json.RawMessage) PluginContext { + values := map[string]json.RawMessage{} + if len(raw) > 0 { + _ = json.Unmarshal(raw, &values) + } + if values == nil { + values = map[string]json.RawMessage{} + } + return PluginContext{Values: values} +} + +// Raw 把袋子编成 JSON 对象;空袋是 {}。 +func (c PluginContext) Raw() json.RawMessage { + if len(c.Values) == 0 { + return json.RawMessage(`{}`) + } + body, err := json.Marshal(c.Values) + if err != nil { + return json.RawMessage(`{}`) + } + return body +} + +// Clone 复制一份袋子,后续 Set 不会改到原来那份。 +func (c PluginContext) Clone() PluginContext { + out := PluginContext{Values: make(map[string]json.RawMessage, len(c.Values))} + for key, value := range c.Values { + out.Values[key] = append(json.RawMessage(nil), value...) + } + return out +} + +// Set 写入一个键;value 按 JSON 编码。空键或编码失败则忽略。 +func (c *PluginContext) Set(key string, value any) { + if c == nil || key == "" { + return + } + if c.Values == nil { + c.Values = map[string]json.RawMessage{} + } + body, err := json.Marshal(value) + if err != nil { + return + } + c.Values[key] = body +} + +// Get 读取一个键的 JSON;没有则 nil。 +func (c PluginContext) Get(key string) json.RawMessage { + if c.Values == nil { + return nil + } + return c.Values[key] +} + +// AcceptableContext 判断 raw 是否可作为袋子保存:空、{} 或未超限的 JSON 对象。 +func AcceptableContext(raw json.RawMessage) bool { + if len(raw) == 0 { + return true + } + if len(raw) > MaxPluginContextBytes { + return false + } + var values map[string]json.RawMessage + return json.Unmarshal(raw, &values) == nil +} diff --git a/server/pkg/plugin/context_test.go b/server/pkg/plugin/context_test.go new file mode 100644 index 0000000..ee14e75 --- /dev/null +++ b/server/pkg/plugin/context_test.go @@ -0,0 +1,42 @@ +package plugin + +import ( + "encoding/json" + "strings" + "testing" +) + +// TestPluginContextSetGet 确认 Set / Get / Clone / Raw 往返。 +func TestPluginContextSetGet(t *testing.T) { + t.Parallel() + var bag PluginContext + bag.Set("hello.marked", true) + bag.Set("", "ignore") + if string(bag.Get("hello.marked")) != "true" || bag.Get("missing") != nil { + t.Fatalf("bag=%+v", bag) + } + clone := bag.Clone() + clone.Set("hello.marked", false) + if string(bag.Get("hello.marked")) != "true" { + t.Fatalf("clone mutated original: %s", bag.Get("hello.marked")) + } + decoded := DecodePluginContext(bag.Raw()) + if string(decoded.Get("hello.marked")) != "true" { + t.Fatalf("raw=%s", bag.Raw()) + } +} + +// TestAcceptableContext 确认空对象可通过、超限和数组被拒。 +func TestAcceptableContext(t *testing.T) { + t.Parallel() + if !AcceptableContext(nil) || !AcceptableContext(json.RawMessage(`{}`)) { + t.Fatal("empty should pass") + } + if AcceptableContext(json.RawMessage(`[]`)) { + t.Fatal("array should fail") + } + tooBig := json.RawMessage(`{"k":"` + strings.Repeat("x", MaxPluginContextBytes) + `"}`) + if AcceptableContext(tooBig) { + t.Fatal("oversize should fail") + } +} diff --git a/server/pkg/plugin/convert.go b/server/pkg/plugin/convert.go new file mode 100644 index 0000000..f5665d4 --- /dev/null +++ b/server/pkg/plugin/convert.go @@ -0,0 +1,127 @@ +package plugin + +import ( + "encoding/json" + + "codedock/pkg/agent/seam" + "codedock/pkg/agent/tool" + pluginpb "codedock/pkg/plugin/proto" +) + +// envelopeFromPB 把 proto 信封转成 seam.Envelope。 +func envelopeFromPB(in *pluginpb.Envelope) seam.Envelope { + if in == nil { + return seam.Envelope{} + } + return seam.Envelope{ + Type: in.GetType(), + ChainID: in.GetChainId(), + Seen: append([]string(nil), in.GetSeen()...), + SessionID: in.GetSessionId(), + RunID: in.GetRunId(), + TurnID: in.GetTurnId(), + Payload: json.RawMessage(in.GetPayload()), + Context: json.RawMessage(in.GetContext()), + } +} + +// envelopeToPB 把 seam.Envelope 转成 proto 信封。 +func envelopeToPB(in seam.Envelope) *pluginpb.Envelope { + return &pluginpb.Envelope{ + Type: in.Type, + ChainId: in.ChainID, + Seen: append([]string(nil), in.Seen...), + SessionId: in.SessionID, + RunId: in.RunID, + TurnId: in.TurnID, + Payload: []byte(in.Payload), + Context: []byte(in.Context), + } +} + +// methodFromPB 把 proto 方法描述转成 Method。 +func methodFromPB(in *pluginpb.Method) Method { + if in == nil { + return Method{} + } + return Method{ + Name: in.GetName(), + Prompt: in.GetPrompt(), + ParametersSchema: json.RawMessage(in.GetParametersSchema()), + Capabilities: append([]string(nil), in.GetCapabilities()...), + RequiresApproval: in.GetRequiresApproval(), + } +} + +// methodToPB 把 Method 转成 proto 方法描述。 +func methodToPB(in Method) *pluginpb.Method { + return &pluginpb.Method{ + Name: in.Name, + Prompt: in.Prompt, + ParametersSchema: []byte(in.ParametersSchema), + Capabilities: append([]string(nil), in.Capabilities...), + RequiresApproval: in.RequiresApproval, + } +} + +// methodInputFromPB 把 proto 方法入参转成 MethodInput。 +func methodInputFromPB(in *pluginpb.MethodInput) MethodInput { + if in == nil { + return MethodInput{} + } + return MethodInput{ + SessionID: in.GetSessionId(), + RunID: in.GetRunId(), + TurnID: in.GetTurnId(), + CallID: in.GetCallId(), + Name: in.GetName(), + Arguments: json.RawMessage(in.GetArguments()), + } +} + +// methodInputToPB 把 MethodInput 转成 proto 方法入参。 +func methodInputToPB(in MethodInput) *pluginpb.MethodInput { + return &pluginpb.MethodInput{ + SessionId: in.SessionID, + RunId: in.RunID, + TurnId: in.TurnID, + CallId: in.CallID, + Name: in.Name, + Arguments: []byte(in.Arguments), + } +} + +// methodResultFromPB 把 proto 方法结果转成 MethodResult。 +func methodResultFromPB(in *pluginpb.MethodResult) MethodResult { + if in == nil { + return MethodResult{} + } + return MethodResult{ + Success: in.GetSuccess(), + Output: json.RawMessage(in.GetOutput()), + Error: in.GetError(), + } +} + +// methodResultToPB 把 MethodResult 转成 proto 方法结果。 +func methodResultToPB(in MethodResult) *pluginpb.MethodResult { + return &pluginpb.MethodResult{ + Success: in.Success, + Output: []byte(in.Output), + Error: in.Error, + } +} + +// MethodToDefinition 把插件方法映射成工具定义。RequiresApproval 映射为 EffectAsk,否则 allow。 +func MethodToDefinition(m Method) tool.Definition { + effect := tool.EffectAllow + if m.RequiresApproval { + effect = tool.EffectAsk + } + return tool.Definition{ + Name: m.Name, + Prompt: m.Prompt, + ParametersSchema: m.ParametersSchema, + Permission: tool.Permission{Effect: effect}, + } +} diff --git a/server/pkg/plugin/convert_test.go b/server/pkg/plugin/convert_test.go new file mode 100644 index 0000000..3c9bd88 --- /dev/null +++ b/server/pkg/plugin/convert_test.go @@ -0,0 +1,77 @@ +package plugin + +import ( + "encoding/json" + "testing" + + pkgagent "codedock/pkg/agent" + "codedock/pkg/agent/seam" + "codedock/pkg/agent/tool" +) + +// TestEnvelopeRoundTrip 确认信封与 proto 往返不丢字段。 +func TestEnvelopeRoundTrip(t *testing.T) { + t.Parallel() + in := seam.Envelope{ + Type: TypeInput, + ChainID: "c1", + Seen: []string{"echo"}, + SessionID: "s1", + RunID: "r1", + TurnID: "t1", + Payload: json.RawMessage(`{"content":"hi"}`), + Context: json.RawMessage(`{"hello.marked":true}`), + } + got := envelopeFromPB(envelopeToPB(in)) + if got.Type != in.Type || got.ChainID != in.ChainID || got.SessionID != in.SessionID || got.RunID != in.RunID || got.TurnID != in.TurnID { + t.Fatalf("got=%+v", got) + } + if string(got.Payload) != string(in.Payload) || len(got.Seen) != 1 || got.Seen[0] != "echo" { + t.Fatalf("got=%+v", got) + } + if string(got.Context) != string(in.Context) { + t.Fatalf("context=%s", got.Context) + } +} + +// TestMethodRoundTripAndDefinition 确认方法描述往返并能映射成工具定义。 +func TestMethodRoundTripAndDefinition(t *testing.T) { + t.Parallel() + in := Method{ + Name: "echo", + Prompt: "echo text", + ParametersSchema: json.RawMessage(`{"type":"object"}`), + Capabilities: []string{"read"}, + RequiresApproval: true, + } + got := methodFromPB(methodToPB(in)) + if got.Name != in.Name || got.Prompt != in.Prompt || got.RequiresApproval != in.RequiresApproval || string(got.ParametersSchema) != string(in.ParametersSchema) { + t.Fatalf("got=%+v", got) + } + def := MethodToDefinition(got) + if def.Name != "echo" || def.Permission.Effect != tool.EffectAsk { + t.Fatalf("def=%+v", def) + } +} + +// TestHiddenText 确认隐藏提示是一条 system 消息。 +func TestHiddenText(t *testing.T) { + t.Parallel() + msg := HiddenText("note") + if msg.Role != pkgagent.RoleSystem || pkgagent.DecodeText(msg.Content) != "note" { + t.Fatalf("hidden=%+v", msg) + } +} + +// TestInputReplyAndHandle 确认 Reply 继续、Handle 标记不建 Run。 +func TestInputReplyAndHandle(t *testing.T) { + t.Parallel() + in := AgentInput{Content: "x", Mode: "yolo"} + in.Context.Set("hello.marked", true) + if got := in.Reply(); got.Handled || got.Content != "x" || got.Mode != "yolo" || string(got.Context.Get("hello.marked")) != "true" { + t.Fatalf("reply=%+v", got) + } + if got := in.Handle(); !got.Handled || got.Content != "x" { + t.Fatalf("handle=%+v", got) + } +} diff --git a/server/pkg/plugin/grpc.go b/server/pkg/plugin/grpc.go new file mode 100644 index 0000000..10bb85e --- /dev/null +++ b/server/pkg/plugin/grpc.go @@ -0,0 +1,271 @@ +package plugin + +import ( + "context" + "fmt" + + goplugin "github.com/hashicorp/go-plugin" + "google.golang.org/grpc" + + "codedock/pkg/agent/seam" + pluginpb "codedock/pkg/plugin/proto" +) + +// GRPCPlugin 同时给宿主当客户端、给插件进程当服务端。 +type GRPCPlugin struct { + goplugin.NetRPCUnsupportedPlugin + Impl Handler // 插件进程里的实现;宿主侧由 Serve 包成 runtime + Host Host // 宿主白名单,经 broker 回传给插件 +} + +// GRPCServer 在插件进程里挂 Plugin 服务。 +func (p *GRPCPlugin) GRPCServer(broker *goplugin.GRPCBroker, s *grpc.Server) error { + pluginpb.RegisterPluginServer(s, &pluginServer{impl: p.Impl, broker: broker}) + return nil +} + +// GRPCClient 在宿主进程里拿到插件客户端,并拉起 Host 回调服务。 +func (p *GRPCPlugin) GRPCClient(_ context.Context, broker *goplugin.GRPCBroker, conn *grpc.ClientConn) (interface{}, error) { + return &pluginClient{ + client: pluginpb.NewPluginClient(conn), + broker: broker, + host: p.Host, + }, nil +} + +var _ goplugin.GRPCPlugin = (*GRPCPlugin)(nil) + +// pluginClient 是宿主调用插件进程的 gRPC 客户端。 +type pluginClient struct { + client pluginpb.PluginClient + broker *goplugin.GRPCBroker + host Host +} + +// Bootstrap 把 Host 服务端交给插件,并取回 Manifest。 +func (c *pluginClient) Bootstrap(ctx context.Context, _ Host) (Manifest, error) { + id := c.broker.NextId() + go c.broker.AcceptAndServe(id, func(opts []grpc.ServerOption) *grpc.Server { + s := grpc.NewServer(opts...) + pluginpb.RegisterHostServer(s, &hostServer{host: c.host}) + return s + }) + resp, err := c.client.Bootstrap(ctx, &pluginpb.BootstrapRequest{HostServerId: id}) + if err != nil { + return Manifest{}, err + } + return Manifest{Name: resp.GetName(), Subscriptions: append([]string(nil), resp.GetSubscriptions()...)}, nil +} + +// OnEvent 把信封发给插件进程。 +func (c *pluginClient) OnEvent(ctx context.Context, ev seam.Envelope) (seam.Envelope, error) { + resp, err := c.client.OnEvent(ctx, envelopeToPB(ev)) + if err != nil { + return ev, err + } + return envelopeFromPB(resp), nil +} + +// ExecuteMethod 在插件进程里执行已登记的方法。 +func (c *pluginClient) ExecuteMethod(ctx context.Context, in MethodInput) (MethodResult, error) { + resp, err := c.client.ExecuteMethod(ctx, methodInputToPB(in)) + if err != nil { + return MethodResult{}, err + } + return methodResultFromPB(resp), nil +} + +// pluginServer 是插件进程里的 Plugin 服务。 +type pluginServer struct { + pluginpb.UnimplementedPluginServer + impl Handler + broker *goplugin.GRPCBroker +} + +// Bootstrap 拨通宿主回调,再调作者的 Bootstrap。 +func (s *pluginServer) Bootstrap(ctx context.Context, req *pluginpb.BootstrapRequest) (*pluginpb.Manifest, error) { + if s.impl == nil { + return nil, fmt.Errorf("plugin implementation is nil") + } + conn, err := s.broker.Dial(req.GetHostServerId()) + if err != nil { + return nil, err + } + host := &hostClient{client: pluginpb.NewHostClient(conn)} + man, err := s.impl.Bootstrap(ctx, host) + if err != nil { + return nil, err + } + return &pluginpb.Manifest{Name: man.Name, Subscriptions: man.Subscriptions}, nil +} + +// OnEvent 把 proto 信封交给 Handler。 +func (s *pluginServer) OnEvent(ctx context.Context, req *pluginpb.Envelope) (*pluginpb.Envelope, error) { + if s.impl == nil { + return nil, fmt.Errorf("plugin implementation is nil") + } + out, err := s.impl.OnEvent(ctx, envelopeFromPB(req)) + if err != nil { + return nil, err + } + return envelopeToPB(out), nil +} + +// ExecuteMethod 把 proto 方法入参交给 Handler。 +func (s *pluginServer) ExecuteMethod(ctx context.Context, req *pluginpb.MethodInput) (*pluginpb.MethodResult, error) { + if s.impl == nil { + return nil, fmt.Errorf("plugin implementation is nil") + } + out, err := s.impl.ExecuteMethod(ctx, methodInputFromPB(req)) + if err != nil { + return nil, err + } + return methodResultToPB(out), nil +} + +// hostClient 是插件进程回调宿主的客户端。 +type hostClient struct { + client pluginpb.HostClient +} + +// Emit 请宿主另发一条事件。 +func (c *hostClient) Emit(ctx context.Context, ev seam.Envelope) error { + _, err := c.client.Emit(ctx, envelopeToPB(ev)) + return err +} + +// RegisterMethod 请宿主给模型加方法。 +func (c *hostClient) RegisterMethod(ctx context.Context, method Method) error { + _, err := c.client.RegisterMethod(ctx, methodToPB(method)) + return err +} + +// MemoryGet 请宿主读一篇专题记忆。 +func (c *hostClient) MemoryGet(ctx context.Context, key MemoryKey) (string, error) { + resp, err := c.client.MemoryGet(ctx, &pluginpb.MemoryKey{ + SessionId: key.SessionID, + Scope: key.Scope, + Name: key.Name, + }) + if err != nil { + return "", err + } + return resp.GetText(), nil +} + +// MemoryUpsert 请宿主写一篇专题记忆。 +func (c *hostClient) MemoryUpsert(ctx context.Context, key MemoryKey, text string) error { + _, err := c.client.MemoryUpsert(ctx, &pluginpb.MemoryEntry{ + Key: &pluginpb.MemoryKey{ + SessionId: key.SessionID, + Scope: key.Scope, + Name: key.Name, + }, + Text: text, + }) + return err +} + +// Complete 请宿主单独打一次模型。 +func (c *hostClient) Complete(ctx context.Context, req CompleteRequest) (CompleteResult, error) { + resp, err := c.client.Complete(ctx, &pluginpb.CompleteRequest{ + SessionId: req.SessionID, + RunId: req.RunID, + Prompt: req.Prompt, + Text: req.Text, + }) + if err != nil { + return CompleteResult{}, err + } + return CompleteResult{Text: resp.GetText()}, nil +} + +// AppendNotice 请宿主写一条用户可见的 system 消息。 +func (c *hostClient) AppendNotice(ctx context.Context, sessionID, runID, text string) error { + _, err := c.client.AppendNotice(ctx, &pluginpb.Notice{ + SessionId: sessionID, + RunId: runID, + Text: text, + }) + return err +} + +// hostServer 是宿主进程里给插件回调的 Host 服务。 +type hostServer struct { + pluginpb.UnimplementedHostServer + host Host +} + +// Emit 转发插件的另发事件请求。 +func (s *hostServer) Emit(ctx context.Context, req *pluginpb.Envelope) (*pluginpb.Empty, error) { + if s.host == nil { + return nil, fmt.Errorf("host is nil") + } + if err := s.host.Emit(ctx, envelopeFromPB(req)); err != nil { + return nil, err + } + return &pluginpb.Empty{}, nil +} + +// RegisterMethod 转发插件的方法登记。 +func (s *hostServer) RegisterMethod(ctx context.Context, req *pluginpb.Method) (*pluginpb.Empty, error) { + if s.host == nil { + return nil, fmt.Errorf("host is nil") + } + if err := s.host.RegisterMethod(ctx, methodFromPB(req)); err != nil { + return nil, err + } + return &pluginpb.Empty{}, nil +} + +// MemoryGet 转发插件的记忆读取。 +func (s *hostServer) MemoryGet(ctx context.Context, req *pluginpb.MemoryKey) (*pluginpb.MemoryValue, error) { + if s.host == nil { + return nil, fmt.Errorf("host is nil") + } + text, err := s.host.MemoryGet(ctx, MemoryKey{SessionID: req.GetSessionId(), Scope: req.GetScope(), Name: req.GetName()}) + if err != nil { + return nil, err + } + return &pluginpb.MemoryValue{Text: text}, nil +} + +// MemoryUpsert 转发插件的记忆写入。 +func (s *hostServer) MemoryUpsert(ctx context.Context, req *pluginpb.MemoryEntry) (*pluginpb.Empty, error) { + if s.host == nil { + return nil, fmt.Errorf("host is nil") + } + key := req.GetKey() + if err := s.host.MemoryUpsert(ctx, MemoryKey{SessionID: key.GetSessionId(), Scope: key.GetScope(), Name: key.GetName()}, req.GetText()); err != nil { + return nil, err + } + return &pluginpb.Empty{}, nil +} + +// Complete 转发插件的单独模型调用。 +func (s *hostServer) Complete(ctx context.Context, req *pluginpb.CompleteRequest) (*pluginpb.CompleteResult, error) { + if s.host == nil { + return nil, fmt.Errorf("host is nil") + } + out, err := s.host.Complete(ctx, CompleteRequest{ + SessionID: req.GetSessionId(), + RunID: req.GetRunId(), + Prompt: req.GetPrompt(), + Text: req.GetText(), + }) + if err != nil { + return nil, err + } + return &pluginpb.CompleteResult{Text: out.Text}, nil +} + +// AppendNotice 转发插件的 system 通知写入。 +func (s *hostServer) AppendNotice(ctx context.Context, req *pluginpb.Notice) (*pluginpb.Empty, error) { + if s.host == nil { + return nil, fmt.Errorf("host is nil") + } + if err := s.host.AppendNotice(ctx, req.GetSessionId(), req.GetRunId(), req.GetText()); err != nil { + return nil, err + } + return &pluginpb.Empty{}, nil +} diff --git a/server/pkg/plugin/plugin.go b/server/pkg/plugin/plugin.go new file mode 100644 index 0000000..ebde824 --- /dev/null +++ b/server/pkg/plugin/plugin.go @@ -0,0 +1,179 @@ +package plugin + +import ( + "context" + "encoding/json" + + goplugin "github.com/hashicorp/go-plugin" + + pkgagent "codedock/pkg/agent" + "codedock/pkg/agent/seam" + "codedock/pkg/agent/tool" +) + +// PluginName 是 go-plugin 握手时登记的插件键。 +const PluginName = "codedock" + +// Handshake 是宿主与插件进程的约定口令。 +var Handshake = goplugin.HandshakeConfig{ + ProtocolVersion: 1, + MagicCookieKey: "CODEDOCK_PLUGIN", + MagicCookieValue: "codedock-plugin-v1", +} + +const ( + TypeInput = seam.TypeInput + TypePreStep = seam.TypePreStep + TypeRequest = seam.TypeRequest + TypeStream = seam.TypeStream + TypePreExecute = seam.TypePreExecute + TypePostExecute = seam.TypePostExecute + TypeInputHandled = seam.TypeInputHandled + TypeRunBlocked = seam.TypeRunBlocked + TypeToolsDenied = seam.TypeToolsDenied + TypeToolsAsk = seam.TypeToolsAsk +) + +type ( + Envelope = seam.Envelope + Message = pkgagent.Message + // Call 是一次工具调用:ID、Name、Arguments(模型填的 JSON)。 + Call = tool.Call + // Result 是一次工具结果:CallID、Name、Output、Success、Error。 + Result = tool.Result + // WorkMode 是本轮内置 Agent:ask / plan / agent。 + WorkMode = pkgagent.WorkMode +) + +// Plugin 是作者要实现的接口。六个口用可选的 *Handler 接口,入参和回包都是结构体。 +type Plugin interface { + // Bootstrap 进程起来时调用一次,登记方法和订阅。 + Bootstrap(ctx context.Context, host Host) (Manifest, error) + // ExecuteMethod 执行本插件登记给模型的方法。 + ExecuteMethod(ctx context.Context, in MethodInput) (MethodResult, error) +} + +// Handler 是宿主侧进程:信封进出。作者实现 Plugin,Serve 会转成 Handler。 +type Handler interface { + // Bootstrap 拉起插件进程后立刻调用。 + Bootstrap(ctx context.Context, host Host) (Manifest, error) + // OnEvent 把信封交给插件;作者侧已拆成各口结构体。 + OnEvent(ctx context.Context, ev seam.Envelope) (seam.Envelope, error) + // ExecuteMethod 执行本插件登记给模型的方法。 + ExecuteMethod(ctx context.Context, in MethodInput) (MethodResult, error) +} + +// AgentInputHandler 拦 agent/input。 +type AgentInputHandler interface { + // OnAgentInput 用户刚提交、Run 还没建。 + OnAgentInput(ctx context.Context, in AgentInput) (AgentInputResult, error) +} + +// AgentPreStepHandler 拦 agent/pre-step。 +type AgentPreStepHandler interface { + // OnAgentPreStep 本轮第一拍,可改系统提示和隐藏消息。 + OnAgentPreStep(ctx context.Context, in AgentPreStep) (AgentPreStepResult, error) +} + +// AgentRequestHandler 拦 agent/request。只能改数据。 +type AgentRequestHandler interface { + // OnAgentRequest 即将发给模型的提示和消息。 + OnAgentRequest(ctx context.Context, in AgentRequest) (AgentRequestResult, error) +} + +// LLMStreamHandler 拦 llm/stream。只能改数据;fake 模型不经过这里。 +type LLMStreamHandler interface { + // OnLLMStream 即将发出的模型 HTTP 请求。 + OnLLMStream(ctx context.Context, in LLMStream) (LLMStreamResult, error) +} + +// ToolPreExecuteHandler 拦 tools/pre-execute。已批准的调用不会再进。 +type ToolPreExecuteHandler interface { + // OnToolPreExecute 某个工具马上要跑。 + OnToolPreExecute(ctx context.Context, in ToolPreExecute) (ToolPreExecuteResult, error) +} + +// ToolPostExecuteHandler 拦 tools/post-execute。只能改数据。 +type ToolPostExecuteHandler interface { + // OnToolPostExecute 工具已经跑完。 + OnToolPostExecute(ctx context.Context, in ToolPostExecute) (ToolPostExecuteResult, error) +} + +// LedgerNotifyHandler 收账本通知(如 run.completed)。改回包没有换向效果。 +type LedgerNotifyHandler interface { + // OnLedgerNotify 处理非六个口的账本事件。 + OnLedgerNotify(ctx context.Context, in LedgerNotify) error +} + +// Host 是插件回调宿主的白名单。 +type Host interface { + // Emit 另发一条与当前口无关的事件;不能发六个口的同名类型。 + Emit(ctx context.Context, ev seam.Envelope) error + // RegisterMethod 给模型加方法;不能覆盖 ping / memory_*。 + RegisterMethod(ctx context.Context, method Method) error + // MemoryGet 按会话读一篇专题记忆。 + MemoryGet(ctx context.Context, key MemoryKey) (string, error) + // MemoryUpsert 按会话写一篇专题记忆。 + MemoryUpsert(ctx context.Context, key MemoryKey, text string) error + // Complete 自己打一次模型,不进当前助手流。 + Complete(ctx context.Context, req CompleteRequest) (CompleteResult, error) + // AppendNotice 写一条用户看得见的 system 消息(本期只落库)。 + AppendNotice(ctx context.Context, sessionID, runID, text string) error +} + +// Manifest 是插件启动时声明的名字与订阅。 +type Manifest struct { + Name string // 日志用;空则用 PLUGIN_DIR 子目录名。Seen 永远是子目录名 + Subscriptions []string // 要听的口(TypeInput 等)或账本事件(如 run.completed) +} + +// Method 是插件登记给模型的方法。 +type Method struct { + Name string // 模型看到的工具名;不能覆盖 ping / memory_* + Prompt string // 什么时候该调这个方法 + ParametersSchema json.RawMessage // 模型填参用的 JSON Schema;ExecuteMethod 按同一份解 Arguments + Capabilities []string // proto 保留字段,不映射到工具 Permission + RequiresApproval bool // true 时默认 EffectAsk,需过审批流水线 +} + +// MethodInput 是一次方法执行的入参。 +type MethodInput struct { + SessionID string // 当前会话 + RunID string // 本轮 Run + TurnID string // 当前 Turn + CallID string // 这次调用号 + Name string // 被点到的方法名(一个插件可登记多个) + Arguments json.RawMessage // 模型按 ParametersSchema 填的 JSON +} + +// MethodResult 是一次方法执行的结果。业务失败用 Success:false,error 只表示取消或超时。 +type MethodResult struct { + Success bool // 业务是否成功 + Output json.RawMessage // 成功时的 JSON + Error string // 失败原因 +} + +// MemoryKey 定位一篇专题记忆。 +type MemoryKey struct { + SessionID string // 当前会话 + Scope string // 记忆范围 + Name string // 专题名 +} + +// CompleteRequest 是插件自己打一次模型的入参。 +type CompleteRequest struct { + SessionID string // 当前会话 + RunID string // 本轮 Run + Prompt string // 系统提示 + Text string // 用户侧文本 +} + +// CompleteResult 是插件自己打一次模型的出参。 +type CompleteResult struct { + Text string // 模型回的正文 +} + +// HiddenText 构造一条不入库的隐藏 system 消息,供 agent/pre-step 注入。 +func HiddenText(text string) pkgagent.Message { + return pkgagent.Message{Role: pkgagent.RoleSystem, Content: pkgagent.EncodeText(text)} +} diff --git a/server/pkg/plugin/points.go b/server/pkg/plugin/points.go new file mode 100644 index 0000000..55f54c8 --- /dev/null +++ b/server/pkg/plugin/points.go @@ -0,0 +1,179 @@ +package plugin + +import ( + "encoding/json" + + pkgagent "codedock/pkg/agent" +) + +// AgentInput 是 agent/input 的入参。此时 Run 还没建。 +type AgentInput struct { + SessionID string // 当前会话;Host 回调要用 + RunID string // 这个口上为空 + TurnID string // 这个口上为空 + Content string // 用户正文,可改 + Mode pkgagent.WorkMode // 本轮内置 Agent:ask / plan / agent,可改 + Context PluginContext // 插件共享袋子;Reply 会带回 +} + +// AgentInputResult 是 agent/input 的回包。 +type AgentInputResult struct { + Content string // 写进用户消息的正文 + Mode pkgagent.WorkMode // 本轮模式;空表示不改 + Handled bool // true:不建 Run。用 Handle() 设置 + Context PluginContext // 盖写后的共享袋子 +} + +// Reply 继续本口,带回改过的正文、模式和袋子。 +func (in AgentInput) Reply() AgentInputResult { + return AgentInputResult{Content: in.Content, Mode: in.Mode, Context: in.Context.Clone()} +} + +// Handle 换到 input/handled,不建 Run。 +func (in AgentInput) Handle() AgentInputResult { + out := in.Reply() + out.Handled = true + return out +} + +// AgentPreStep 是 agent/pre-step 的入参。 +type AgentPreStep struct { + SessionID string // 当前会话 + RunID string // 本轮 Run + TurnID string // 当前 Turn,可能为空 + SystemPrompt string // 本轮系统提示,可改 + Hidden []Message // 不入库、只进模型上下文的 system 消息,可追加 + Context PluginContext // 插件共享袋子;Reply 会带回 +} + +// AgentPreStepResult 是 agent/pre-step 的回包。 +type AgentPreStepResult struct { + SystemPrompt string // 盖写后的系统提示 + Hidden []Message // 盖写后的隐藏消息 + Blocked bool // true:取消本轮。用 Block() 设置 + Context PluginContext // 盖写后的共享袋子 +} + +// Reply 继续本口,带回改过的系统提示、隐藏消息和袋子。 +func (in AgentPreStep) Reply() AgentPreStepResult { + return AgentPreStepResult{SystemPrompt: in.SystemPrompt, Hidden: in.Hidden, Context: in.Context.Clone()} +} + +// Block 换到 run/blocked,取消本轮。 +func (in AgentPreStep) Block() AgentPreStepResult { + out := in.Reply() + out.Blocked = true + return out +} + +// AgentRequest 是 agent/request 的入参。只能改数据。 +type AgentRequest struct { + SessionID string // 当前会话 + RunID string // 本轮 Run + TurnID string // 当前 Turn + SystemPrompt string // 即将发给模型的系统提示,可改 + Messages []Message // 即将发给模型的消息,可改 + Context PluginContext // 插件共享袋子;Reply 会带回 +} + +// AgentRequestResult 是 agent/request 的回包。 +type AgentRequestResult struct { + SystemPrompt string // 盖写后的系统提示 + Messages []Message // 盖写后的消息 + Context PluginContext // 盖写后的共享袋子 +} + +// Reply 继续本口,带回改过的提示、消息和袋子。 +func (in AgentRequest) Reply() AgentRequestResult { + return AgentRequestResult{SystemPrompt: in.SystemPrompt, Messages: in.Messages, Context: in.Context.Clone()} +} + +// LLMStream 是 llm/stream 的入参。只能改数据;fake 模型不经过这里。 +type LLMStream struct { + SessionID string // 当前会话 + RunID string // 本轮 Run + TurnID string // 当前 Turn + Headers map[string]string // 即将发出的 HTTP 头,可改 + Body json.RawMessage // 即将发出的 HTTP 体,可改 + Context PluginContext // 插件共享袋子;Reply 会带回 +} + +// LLMStreamResult 是 llm/stream 的回包。 +type LLMStreamResult struct { + Headers map[string]string // 盖写后的请求头;nil 表示不改头 + Body json.RawMessage // 盖写后的请求体;空表示不改体 + Context PluginContext // 盖写后的共享袋子 +} + +// Reply 继续本口,带回改过的请求头、请求体和袋子。 +func (in LLMStream) Reply() LLMStreamResult { + return LLMStreamResult{Headers: in.Headers, Body: in.Body, Context: in.Context.Clone()} +} + +// ToolPreExecute 是 tools/pre-execute 的入参。已批准的调用不会再进。 +type ToolPreExecute struct { + SessionID string // 当前会话 + RunID string // 本轮 Run + TurnID string // 当前 Turn + Call Call // 马上要跑的工具调用,可改 Arguments + Context PluginContext // 插件共享袋子;Reply 会带回 +} + +// ToolPreExecuteResult 是 tools/pre-execute 的回包。Denied 与 Ask 同时为 true 时按 Denied。 +type ToolPreExecuteResult struct { + Call Call // 盖写后的调用(改参后仍走这个 Call) + Denied bool // true:当失败。用 Deny() 设置 + Ask bool // true:进审批。用 AskApproval() 设置 + Context PluginContext // 盖写后的共享袋子 +} + +// Reply 继续本口,带回改过的工具调用和袋子。 +func (in ToolPreExecute) Reply() ToolPreExecuteResult { + return ToolPreExecuteResult{Call: in.Call, Context: in.Context.Clone()} +} + +// Deny 换到 tools/denied,这次调用当失败。 +func (in ToolPreExecute) Deny() ToolPreExecuteResult { + out := in.Reply() + out.Denied = true + return out +} + +// AskApproval 换到 tools/ask,进审批。 +func (in ToolPreExecute) AskApproval() ToolPreExecuteResult { + out := in.Reply() + out.Ask = true + return out +} + +// ToolPostExecute 是 tools/post-execute 的入参。只能改数据。 +type ToolPostExecute struct { + SessionID string // 当前会话 + RunID string // 本轮 Run + TurnID string // 当前 Turn + Call Call // 刚跑完的调用 + Result Result // 工具结果,可改 Output / Success / Error + Context PluginContext // 插件共享袋子;Reply 会带回 +} + +// ToolPostExecuteResult 是 tools/post-execute 的回包。 +type ToolPostExecuteResult struct { + Call Call // 原样带回即可 + Result Result // 盖写后的结果 + Context PluginContext // 盖写后的共享袋子 +} + +// Reply 继续本口,带回改过的工具结果和袋子。 +func (in ToolPostExecute) Reply() ToolPostExecuteResult { + return ToolPostExecuteResult{Call: in.Call, Result: in.Result, Context: in.Context.Clone()} +} + +// LedgerNotify 是账本通知,不是六个口。改回包没有换向效果。 +type LedgerNotify struct { + Type string // 事件名,例如 run.completed + SessionID string // 当前会话 + RunID string // 本轮 Run + TurnID string // 当前 Turn,可能为空 + Payload json.RawMessage // 事件载荷,只读 + Context PluginContext // 当时的共享袋子,只读 +} diff --git a/server/pkg/plugin/proto/gen.go b/server/pkg/plugin/proto/gen.go new file mode 100644 index 0000000..f67b8d7 --- /dev/null +++ b/server/pkg/plugin/proto/gen.go @@ -0,0 +1,3 @@ +package pluginpb + +//go:generate protoc --go_out=. --go_opt=paths=source_relative --go-grpc_out=. --go-grpc_opt=paths=source_relative plugin.proto diff --git a/server/pkg/plugin/proto/plugin.pb.go b/server/pkg/plugin/proto/plugin.pb.go new file mode 100644 index 0000000..8d8d19e --- /dev/null +++ b/server/pkg/plugin/proto/plugin.pb.go @@ -0,0 +1,955 @@ +// Code generated by protoc-gen-go. DO NOT EDIT. +// versions: +// protoc-gen-go v1.36.12 +// protoc v7.36.1 +// source: plugin.proto + +package pluginpb + +import ( + protoreflect "google.golang.org/protobuf/reflect/protoreflect" + protoimpl "google.golang.org/protobuf/runtime/protoimpl" + reflect "reflect" + sync "sync" + unsafe "unsafe" +) + +const ( + // Verify that this generated code is sufficiently up-to-date. + _ = protoimpl.EnforceVersion(20 - protoimpl.MinVersion) + // Verify that runtime/protoimpl is sufficiently up-to-date. + _ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20) +) + +type Empty struct { + state protoimpl.MessageState `protogen:"open.v1"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *Empty) Reset() { + *x = Empty{} + mi := &file_plugin_proto_msgTypes[0] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *Empty) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*Empty) ProtoMessage() {} + +func (x *Empty) ProtoReflect() protoreflect.Message { + mi := &file_plugin_proto_msgTypes[0] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use Empty.ProtoReflect.Descriptor instead. +func (*Empty) Descriptor() ([]byte, []int) { + return file_plugin_proto_rawDescGZIP(), []int{0} +} + +type BootstrapRequest struct { + state protoimpl.MessageState `protogen:"open.v1"` + HostServerId uint32 `protobuf:"varint,1,opt,name=host_server_id,json=hostServerId,proto3" json:"host_server_id,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *BootstrapRequest) Reset() { + *x = BootstrapRequest{} + mi := &file_plugin_proto_msgTypes[1] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *BootstrapRequest) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*BootstrapRequest) ProtoMessage() {} + +func (x *BootstrapRequest) ProtoReflect() protoreflect.Message { + mi := &file_plugin_proto_msgTypes[1] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use BootstrapRequest.ProtoReflect.Descriptor instead. +func (*BootstrapRequest) Descriptor() ([]byte, []int) { + return file_plugin_proto_rawDescGZIP(), []int{1} +} + +func (x *BootstrapRequest) GetHostServerId() uint32 { + if x != nil { + return x.HostServerId + } + return 0 +} + +type Manifest struct { + state protoimpl.MessageState `protogen:"open.v1"` + Name string `protobuf:"bytes,1,opt,name=name,proto3" json:"name,omitempty"` + Subscriptions []string `protobuf:"bytes,2,rep,name=subscriptions,proto3" json:"subscriptions,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *Manifest) Reset() { + *x = Manifest{} + mi := &file_plugin_proto_msgTypes[2] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *Manifest) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*Manifest) ProtoMessage() {} + +func (x *Manifest) ProtoReflect() protoreflect.Message { + mi := &file_plugin_proto_msgTypes[2] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use Manifest.ProtoReflect.Descriptor instead. +func (*Manifest) Descriptor() ([]byte, []int) { + return file_plugin_proto_rawDescGZIP(), []int{2} +} + +func (x *Manifest) GetName() string { + if x != nil { + return x.Name + } + return "" +} + +func (x *Manifest) GetSubscriptions() []string { + if x != nil { + return x.Subscriptions + } + return nil +} + +type Envelope struct { + state protoimpl.MessageState `protogen:"open.v1"` + Type string `protobuf:"bytes,1,opt,name=type,proto3" json:"type,omitempty"` + ChainId string `protobuf:"bytes,2,opt,name=chain_id,json=chainId,proto3" json:"chain_id,omitempty"` + Seen []string `protobuf:"bytes,3,rep,name=seen,proto3" json:"seen,omitempty"` + SessionId string `protobuf:"bytes,4,opt,name=session_id,json=sessionId,proto3" json:"session_id,omitempty"` + RunId string `protobuf:"bytes,5,opt,name=run_id,json=runId,proto3" json:"run_id,omitempty"` + TurnId string `protobuf:"bytes,6,opt,name=turn_id,json=turnId,proto3" json:"turn_id,omitempty"` + Payload []byte `protobuf:"bytes,7,opt,name=payload,proto3" json:"payload,omitempty"` + Context []byte `protobuf:"bytes,8,opt,name=context,proto3" json:"context,omitempty"` // 插件共享袋子,JSON object;不进模型 + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *Envelope) Reset() { + *x = Envelope{} + mi := &file_plugin_proto_msgTypes[3] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *Envelope) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*Envelope) ProtoMessage() {} + +func (x *Envelope) ProtoReflect() protoreflect.Message { + mi := &file_plugin_proto_msgTypes[3] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use Envelope.ProtoReflect.Descriptor instead. +func (*Envelope) Descriptor() ([]byte, []int) { + return file_plugin_proto_rawDescGZIP(), []int{3} +} + +func (x *Envelope) GetType() string { + if x != nil { + return x.Type + } + return "" +} + +func (x *Envelope) GetChainId() string { + if x != nil { + return x.ChainId + } + return "" +} + +func (x *Envelope) GetSeen() []string { + if x != nil { + return x.Seen + } + return nil +} + +func (x *Envelope) GetSessionId() string { + if x != nil { + return x.SessionId + } + return "" +} + +func (x *Envelope) GetRunId() string { + if x != nil { + return x.RunId + } + return "" +} + +func (x *Envelope) GetTurnId() string { + if x != nil { + return x.TurnId + } + return "" +} + +func (x *Envelope) GetPayload() []byte { + if x != nil { + return x.Payload + } + return nil +} + +func (x *Envelope) GetContext() []byte { + if x != nil { + return x.Context + } + return nil +} + +type Method struct { + state protoimpl.MessageState `protogen:"open.v1"` + Name string `protobuf:"bytes,1,opt,name=name,proto3" json:"name,omitempty"` + Prompt string `protobuf:"bytes,2,opt,name=prompt,proto3" json:"prompt,omitempty"` + ParametersSchema []byte `protobuf:"bytes,3,opt,name=parameters_schema,json=parametersSchema,proto3" json:"parameters_schema,omitempty"` + Capabilities []string `protobuf:"bytes,4,rep,name=capabilities,proto3" json:"capabilities,omitempty"` + RequiresApproval bool `protobuf:"varint,5,opt,name=requires_approval,json=requiresApproval,proto3" json:"requires_approval,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *Method) Reset() { + *x = Method{} + mi := &file_plugin_proto_msgTypes[4] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *Method) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*Method) ProtoMessage() {} + +func (x *Method) ProtoReflect() protoreflect.Message { + mi := &file_plugin_proto_msgTypes[4] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use Method.ProtoReflect.Descriptor instead. +func (*Method) Descriptor() ([]byte, []int) { + return file_plugin_proto_rawDescGZIP(), []int{4} +} + +func (x *Method) GetName() string { + if x != nil { + return x.Name + } + return "" +} + +func (x *Method) GetPrompt() string { + if x != nil { + return x.Prompt + } + return "" +} + +func (x *Method) GetParametersSchema() []byte { + if x != nil { + return x.ParametersSchema + } + return nil +} + +func (x *Method) GetCapabilities() []string { + if x != nil { + return x.Capabilities + } + return nil +} + +func (x *Method) GetRequiresApproval() bool { + if x != nil { + return x.RequiresApproval + } + return false +} + +type MethodInput struct { + state protoimpl.MessageState `protogen:"open.v1"` + SessionId string `protobuf:"bytes,1,opt,name=session_id,json=sessionId,proto3" json:"session_id,omitempty"` + RunId string `protobuf:"bytes,2,opt,name=run_id,json=runId,proto3" json:"run_id,omitempty"` + TurnId string `protobuf:"bytes,3,opt,name=turn_id,json=turnId,proto3" json:"turn_id,omitempty"` + CallId string `protobuf:"bytes,4,opt,name=call_id,json=callId,proto3" json:"call_id,omitempty"` + Name string `protobuf:"bytes,5,opt,name=name,proto3" json:"name,omitempty"` + Arguments []byte `protobuf:"bytes,6,opt,name=arguments,proto3" json:"arguments,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *MethodInput) Reset() { + *x = MethodInput{} + mi := &file_plugin_proto_msgTypes[5] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *MethodInput) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*MethodInput) ProtoMessage() {} + +func (x *MethodInput) ProtoReflect() protoreflect.Message { + mi := &file_plugin_proto_msgTypes[5] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use MethodInput.ProtoReflect.Descriptor instead. +func (*MethodInput) Descriptor() ([]byte, []int) { + return file_plugin_proto_rawDescGZIP(), []int{5} +} + +func (x *MethodInput) GetSessionId() string { + if x != nil { + return x.SessionId + } + return "" +} + +func (x *MethodInput) GetRunId() string { + if x != nil { + return x.RunId + } + return "" +} + +func (x *MethodInput) GetTurnId() string { + if x != nil { + return x.TurnId + } + return "" +} + +func (x *MethodInput) GetCallId() string { + if x != nil { + return x.CallId + } + return "" +} + +func (x *MethodInput) GetName() string { + if x != nil { + return x.Name + } + return "" +} + +func (x *MethodInput) GetArguments() []byte { + if x != nil { + return x.Arguments + } + return nil +} + +type MethodResult struct { + state protoimpl.MessageState `protogen:"open.v1"` + Success bool `protobuf:"varint,1,opt,name=success,proto3" json:"success,omitempty"` + Output []byte `protobuf:"bytes,2,opt,name=output,proto3" json:"output,omitempty"` + Error string `protobuf:"bytes,3,opt,name=error,proto3" json:"error,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *MethodResult) Reset() { + *x = MethodResult{} + mi := &file_plugin_proto_msgTypes[6] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *MethodResult) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*MethodResult) ProtoMessage() {} + +func (x *MethodResult) ProtoReflect() protoreflect.Message { + mi := &file_plugin_proto_msgTypes[6] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use MethodResult.ProtoReflect.Descriptor instead. +func (*MethodResult) Descriptor() ([]byte, []int) { + return file_plugin_proto_rawDescGZIP(), []int{6} +} + +func (x *MethodResult) GetSuccess() bool { + if x != nil { + return x.Success + } + return false +} + +func (x *MethodResult) GetOutput() []byte { + if x != nil { + return x.Output + } + return nil +} + +func (x *MethodResult) GetError() string { + if x != nil { + return x.Error + } + return "" +} + +type MemoryKey struct { + state protoimpl.MessageState `protogen:"open.v1"` + SessionId string `protobuf:"bytes,1,opt,name=session_id,json=sessionId,proto3" json:"session_id,omitempty"` + Scope string `protobuf:"bytes,2,opt,name=scope,proto3" json:"scope,omitempty"` + Name string `protobuf:"bytes,3,opt,name=name,proto3" json:"name,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *MemoryKey) Reset() { + *x = MemoryKey{} + mi := &file_plugin_proto_msgTypes[7] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *MemoryKey) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*MemoryKey) ProtoMessage() {} + +func (x *MemoryKey) ProtoReflect() protoreflect.Message { + mi := &file_plugin_proto_msgTypes[7] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use MemoryKey.ProtoReflect.Descriptor instead. +func (*MemoryKey) Descriptor() ([]byte, []int) { + return file_plugin_proto_rawDescGZIP(), []int{7} +} + +func (x *MemoryKey) GetSessionId() string { + if x != nil { + return x.SessionId + } + return "" +} + +func (x *MemoryKey) GetScope() string { + if x != nil { + return x.Scope + } + return "" +} + +func (x *MemoryKey) GetName() string { + if x != nil { + return x.Name + } + return "" +} + +type MemoryValue struct { + state protoimpl.MessageState `protogen:"open.v1"` + Text string `protobuf:"bytes,1,opt,name=text,proto3" json:"text,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *MemoryValue) Reset() { + *x = MemoryValue{} + mi := &file_plugin_proto_msgTypes[8] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *MemoryValue) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*MemoryValue) ProtoMessage() {} + +func (x *MemoryValue) ProtoReflect() protoreflect.Message { + mi := &file_plugin_proto_msgTypes[8] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use MemoryValue.ProtoReflect.Descriptor instead. +func (*MemoryValue) Descriptor() ([]byte, []int) { + return file_plugin_proto_rawDescGZIP(), []int{8} +} + +func (x *MemoryValue) GetText() string { + if x != nil { + return x.Text + } + return "" +} + +type MemoryEntry struct { + state protoimpl.MessageState `protogen:"open.v1"` + Key *MemoryKey `protobuf:"bytes,1,opt,name=key,proto3" json:"key,omitempty"` + Text string `protobuf:"bytes,2,opt,name=text,proto3" json:"text,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *MemoryEntry) Reset() { + *x = MemoryEntry{} + mi := &file_plugin_proto_msgTypes[9] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *MemoryEntry) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*MemoryEntry) ProtoMessage() {} + +func (x *MemoryEntry) ProtoReflect() protoreflect.Message { + mi := &file_plugin_proto_msgTypes[9] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use MemoryEntry.ProtoReflect.Descriptor instead. +func (*MemoryEntry) Descriptor() ([]byte, []int) { + return file_plugin_proto_rawDescGZIP(), []int{9} +} + +func (x *MemoryEntry) GetKey() *MemoryKey { + if x != nil { + return x.Key + } + return nil +} + +func (x *MemoryEntry) GetText() string { + if x != nil { + return x.Text + } + return "" +} + +type CompleteRequest struct { + state protoimpl.MessageState `protogen:"open.v1"` + SessionId string `protobuf:"bytes,1,opt,name=session_id,json=sessionId,proto3" json:"session_id,omitempty"` + RunId string `protobuf:"bytes,2,opt,name=run_id,json=runId,proto3" json:"run_id,omitempty"` + Prompt string `protobuf:"bytes,3,opt,name=prompt,proto3" json:"prompt,omitempty"` + Text string `protobuf:"bytes,4,opt,name=text,proto3" json:"text,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *CompleteRequest) Reset() { + *x = CompleteRequest{} + mi := &file_plugin_proto_msgTypes[10] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *CompleteRequest) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*CompleteRequest) ProtoMessage() {} + +func (x *CompleteRequest) ProtoReflect() protoreflect.Message { + mi := &file_plugin_proto_msgTypes[10] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use CompleteRequest.ProtoReflect.Descriptor instead. +func (*CompleteRequest) Descriptor() ([]byte, []int) { + return file_plugin_proto_rawDescGZIP(), []int{10} +} + +func (x *CompleteRequest) GetSessionId() string { + if x != nil { + return x.SessionId + } + return "" +} + +func (x *CompleteRequest) GetRunId() string { + if x != nil { + return x.RunId + } + return "" +} + +func (x *CompleteRequest) GetPrompt() string { + if x != nil { + return x.Prompt + } + return "" +} + +func (x *CompleteRequest) GetText() string { + if x != nil { + return x.Text + } + return "" +} + +type CompleteResult struct { + state protoimpl.MessageState `protogen:"open.v1"` + Text string `protobuf:"bytes,1,opt,name=text,proto3" json:"text,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *CompleteResult) Reset() { + *x = CompleteResult{} + mi := &file_plugin_proto_msgTypes[11] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *CompleteResult) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*CompleteResult) ProtoMessage() {} + +func (x *CompleteResult) ProtoReflect() protoreflect.Message { + mi := &file_plugin_proto_msgTypes[11] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use CompleteResult.ProtoReflect.Descriptor instead. +func (*CompleteResult) Descriptor() ([]byte, []int) { + return file_plugin_proto_rawDescGZIP(), []int{11} +} + +func (x *CompleteResult) GetText() string { + if x != nil { + return x.Text + } + return "" +} + +type Notice struct { + state protoimpl.MessageState `protogen:"open.v1"` + SessionId string `protobuf:"bytes,1,opt,name=session_id,json=sessionId,proto3" json:"session_id,omitempty"` + RunId string `protobuf:"bytes,2,opt,name=run_id,json=runId,proto3" json:"run_id,omitempty"` + Text string `protobuf:"bytes,3,opt,name=text,proto3" json:"text,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *Notice) Reset() { + *x = Notice{} + mi := &file_plugin_proto_msgTypes[12] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *Notice) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*Notice) ProtoMessage() {} + +func (x *Notice) ProtoReflect() protoreflect.Message { + mi := &file_plugin_proto_msgTypes[12] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use Notice.ProtoReflect.Descriptor instead. +func (*Notice) Descriptor() ([]byte, []int) { + return file_plugin_proto_rawDescGZIP(), []int{12} +} + +func (x *Notice) GetSessionId() string { + if x != nil { + return x.SessionId + } + return "" +} + +func (x *Notice) GetRunId() string { + if x != nil { + return x.RunId + } + return "" +} + +func (x *Notice) GetText() string { + if x != nil { + return x.Text + } + return "" +} + +var File_plugin_proto protoreflect.FileDescriptor + +const file_plugin_proto_rawDesc = "" + + "\n" + + "\fplugin.proto\x12\x12codedock.plugin.v1\"\a\n" + + "\x05Empty\"8\n" + + "\x10BootstrapRequest\x12$\n" + + "\x0ehost_server_id\x18\x01 \x01(\rR\fhostServerId\"D\n" + + "\bManifest\x12\x12\n" + + "\x04name\x18\x01 \x01(\tR\x04name\x12$\n" + + "\rsubscriptions\x18\x02 \x03(\tR\rsubscriptions\"\xd0\x01\n" + + "\bEnvelope\x12\x12\n" + + "\x04type\x18\x01 \x01(\tR\x04type\x12\x19\n" + + "\bchain_id\x18\x02 \x01(\tR\achainId\x12\x12\n" + + "\x04seen\x18\x03 \x03(\tR\x04seen\x12\x1d\n" + + "\n" + + "session_id\x18\x04 \x01(\tR\tsessionId\x12\x15\n" + + "\x06run_id\x18\x05 \x01(\tR\x05runId\x12\x17\n" + + "\aturn_id\x18\x06 \x01(\tR\x06turnId\x12\x18\n" + + "\apayload\x18\a \x01(\fR\apayload\x12\x18\n" + + "\acontext\x18\b \x01(\fR\acontext\"\xb2\x01\n" + + "\x06Method\x12\x12\n" + + "\x04name\x18\x01 \x01(\tR\x04name\x12\x16\n" + + "\x06prompt\x18\x02 \x01(\tR\x06prompt\x12+\n" + + "\x11parameters_schema\x18\x03 \x01(\fR\x10parametersSchema\x12\"\n" + + "\fcapabilities\x18\x04 \x03(\tR\fcapabilities\x12+\n" + + "\x11requires_approval\x18\x05 \x01(\bR\x10requiresApproval\"\xa7\x01\n" + + "\vMethodInput\x12\x1d\n" + + "\n" + + "session_id\x18\x01 \x01(\tR\tsessionId\x12\x15\n" + + "\x06run_id\x18\x02 \x01(\tR\x05runId\x12\x17\n" + + "\aturn_id\x18\x03 \x01(\tR\x06turnId\x12\x17\n" + + "\acall_id\x18\x04 \x01(\tR\x06callId\x12\x12\n" + + "\x04name\x18\x05 \x01(\tR\x04name\x12\x1c\n" + + "\targuments\x18\x06 \x01(\fR\targuments\"V\n" + + "\fMethodResult\x12\x18\n" + + "\asuccess\x18\x01 \x01(\bR\asuccess\x12\x16\n" + + "\x06output\x18\x02 \x01(\fR\x06output\x12\x14\n" + + "\x05error\x18\x03 \x01(\tR\x05error\"T\n" + + "\tMemoryKey\x12\x1d\n" + + "\n" + + "session_id\x18\x01 \x01(\tR\tsessionId\x12\x14\n" + + "\x05scope\x18\x02 \x01(\tR\x05scope\x12\x12\n" + + "\x04name\x18\x03 \x01(\tR\x04name\"!\n" + + "\vMemoryValue\x12\x12\n" + + "\x04text\x18\x01 \x01(\tR\x04text\"R\n" + + "\vMemoryEntry\x12/\n" + + "\x03key\x18\x01 \x01(\v2\x1d.codedock.plugin.v1.MemoryKeyR\x03key\x12\x12\n" + + "\x04text\x18\x02 \x01(\tR\x04text\"s\n" + + "\x0fCompleteRequest\x12\x1d\n" + + "\n" + + "session_id\x18\x01 \x01(\tR\tsessionId\x12\x15\n" + + "\x06run_id\x18\x02 \x01(\tR\x05runId\x12\x16\n" + + "\x06prompt\x18\x03 \x01(\tR\x06prompt\x12\x12\n" + + "\x04text\x18\x04 \x01(\tR\x04text\"$\n" + + "\x0eCompleteResult\x12\x12\n" + + "\x04text\x18\x01 \x01(\tR\x04text\"R\n" + + "\x06Notice\x12\x1d\n" + + "\n" + + "session_id\x18\x01 \x01(\tR\tsessionId\x12\x15\n" + + "\x06run_id\x18\x02 \x01(\tR\x05runId\x12\x12\n" + + "\x04text\x18\x03 \x01(\tR\x04text2\xf4\x01\n" + + "\x06Plugin\x12O\n" + + "\tBootstrap\x12$.codedock.plugin.v1.BootstrapRequest\x1a\x1c.codedock.plugin.v1.Manifest\x12E\n" + + "\aOnEvent\x12\x1c.codedock.plugin.v1.Envelope\x1a\x1c.codedock.plugin.v1.Envelope\x12R\n" + + "\rExecuteMethod\x12\x1f.codedock.plugin.v1.MethodInput\x1a .codedock.plugin.v1.MethodResult2\xc5\x03\n" + + "\x04Host\x12?\n" + + "\x04Emit\x12\x1c.codedock.plugin.v1.Envelope\x1a\x19.codedock.plugin.v1.Empty\x12G\n" + + "\x0eRegisterMethod\x12\x1a.codedock.plugin.v1.Method\x1a\x19.codedock.plugin.v1.Empty\x12K\n" + + "\tMemoryGet\x12\x1d.codedock.plugin.v1.MemoryKey\x1a\x1f.codedock.plugin.v1.MemoryValue\x12J\n" + + "\fMemoryUpsert\x12\x1f.codedock.plugin.v1.MemoryEntry\x1a\x19.codedock.plugin.v1.Empty\x12S\n" + + "\bComplete\x12#.codedock.plugin.v1.CompleteRequest\x1a\".codedock.plugin.v1.CompleteResult\x12E\n" + + "\fAppendNotice\x12\x1a.codedock.plugin.v1.Notice\x1a\x19.codedock.plugin.v1.EmptyB$Z\"codedock/pkg/plugin/proto;pluginpbb\x06proto3" + +var ( + file_plugin_proto_rawDescOnce sync.Once + file_plugin_proto_rawDescData []byte +) + +func file_plugin_proto_rawDescGZIP() []byte { + file_plugin_proto_rawDescOnce.Do(func() { + file_plugin_proto_rawDescData = protoimpl.X.CompressGZIP(unsafe.Slice(unsafe.StringData(file_plugin_proto_rawDesc), len(file_plugin_proto_rawDesc))) + }) + return file_plugin_proto_rawDescData +} + +var file_plugin_proto_msgTypes = make([]protoimpl.MessageInfo, 13) +var file_plugin_proto_goTypes = []any{ + (*Empty)(nil), // 0: codedock.plugin.v1.Empty + (*BootstrapRequest)(nil), // 1: codedock.plugin.v1.BootstrapRequest + (*Manifest)(nil), // 2: codedock.plugin.v1.Manifest + (*Envelope)(nil), // 3: codedock.plugin.v1.Envelope + (*Method)(nil), // 4: codedock.plugin.v1.Method + (*MethodInput)(nil), // 5: codedock.plugin.v1.MethodInput + (*MethodResult)(nil), // 6: codedock.plugin.v1.MethodResult + (*MemoryKey)(nil), // 7: codedock.plugin.v1.MemoryKey + (*MemoryValue)(nil), // 8: codedock.plugin.v1.MemoryValue + (*MemoryEntry)(nil), // 9: codedock.plugin.v1.MemoryEntry + (*CompleteRequest)(nil), // 10: codedock.plugin.v1.CompleteRequest + (*CompleteResult)(nil), // 11: codedock.plugin.v1.CompleteResult + (*Notice)(nil), // 12: codedock.plugin.v1.Notice +} +var file_plugin_proto_depIdxs = []int32{ + 7, // 0: codedock.plugin.v1.MemoryEntry.key:type_name -> codedock.plugin.v1.MemoryKey + 1, // 1: codedock.plugin.v1.Plugin.Bootstrap:input_type -> codedock.plugin.v1.BootstrapRequest + 3, // 2: codedock.plugin.v1.Plugin.OnEvent:input_type -> codedock.plugin.v1.Envelope + 5, // 3: codedock.plugin.v1.Plugin.ExecuteMethod:input_type -> codedock.plugin.v1.MethodInput + 3, // 4: codedock.plugin.v1.Host.Emit:input_type -> codedock.plugin.v1.Envelope + 4, // 5: codedock.plugin.v1.Host.RegisterMethod:input_type -> codedock.plugin.v1.Method + 7, // 6: codedock.plugin.v1.Host.MemoryGet:input_type -> codedock.plugin.v1.MemoryKey + 9, // 7: codedock.plugin.v1.Host.MemoryUpsert:input_type -> codedock.plugin.v1.MemoryEntry + 10, // 8: codedock.plugin.v1.Host.Complete:input_type -> codedock.plugin.v1.CompleteRequest + 12, // 9: codedock.plugin.v1.Host.AppendNotice:input_type -> codedock.plugin.v1.Notice + 2, // 10: codedock.plugin.v1.Plugin.Bootstrap:output_type -> codedock.plugin.v1.Manifest + 3, // 11: codedock.plugin.v1.Plugin.OnEvent:output_type -> codedock.plugin.v1.Envelope + 6, // 12: codedock.plugin.v1.Plugin.ExecuteMethod:output_type -> codedock.plugin.v1.MethodResult + 0, // 13: codedock.plugin.v1.Host.Emit:output_type -> codedock.plugin.v1.Empty + 0, // 14: codedock.plugin.v1.Host.RegisterMethod:output_type -> codedock.plugin.v1.Empty + 8, // 15: codedock.plugin.v1.Host.MemoryGet:output_type -> codedock.plugin.v1.MemoryValue + 0, // 16: codedock.plugin.v1.Host.MemoryUpsert:output_type -> codedock.plugin.v1.Empty + 11, // 17: codedock.plugin.v1.Host.Complete:output_type -> codedock.plugin.v1.CompleteResult + 0, // 18: codedock.plugin.v1.Host.AppendNotice:output_type -> codedock.plugin.v1.Empty + 10, // [10:19] is the sub-list for method output_type + 1, // [1:10] is the sub-list for method input_type + 1, // [1:1] is the sub-list for extension type_name + 1, // [1:1] is the sub-list for extension extendee + 0, // [0:1] is the sub-list for field type_name +} + +func init() { file_plugin_proto_init() } +func file_plugin_proto_init() { + if File_plugin_proto != nil { + return + } + type x struct{} + out := protoimpl.TypeBuilder{ + File: protoimpl.DescBuilder{ + GoPackagePath: reflect.TypeOf(x{}).PkgPath(), + RawDescriptor: unsafe.Slice(unsafe.StringData(file_plugin_proto_rawDesc), len(file_plugin_proto_rawDesc)), + NumEnums: 0, + NumMessages: 13, + NumExtensions: 0, + NumServices: 2, + }, + GoTypes: file_plugin_proto_goTypes, + DependencyIndexes: file_plugin_proto_depIdxs, + MessageInfos: file_plugin_proto_msgTypes, + }.Build() + File_plugin_proto = out.File + file_plugin_proto_goTypes = nil + file_plugin_proto_depIdxs = nil +} diff --git a/server/pkg/plugin/proto/plugin.proto b/server/pkg/plugin/proto/plugin.proto new file mode 100644 index 0000000..73744f2 --- /dev/null +++ b/server/pkg/plugin/proto/plugin.proto @@ -0,0 +1,97 @@ +syntax = "proto3"; + +package codedock.plugin.v1; + +option go_package = "codedock/pkg/plugin/proto;pluginpb"; + +service Plugin { + rpc Bootstrap(BootstrapRequest) returns (Manifest); + rpc OnEvent(Envelope) returns (Envelope); + rpc ExecuteMethod(MethodInput) returns (MethodResult); +} + +service Host { + rpc Emit(Envelope) returns (Empty); + rpc RegisterMethod(Method) returns (Empty); + rpc MemoryGet(MemoryKey) returns (MemoryValue); + rpc MemoryUpsert(MemoryEntry) returns (Empty); + rpc Complete(CompleteRequest) returns (CompleteResult); + rpc AppendNotice(Notice) returns (Empty); +} + +message Empty {} + +message BootstrapRequest { + uint32 host_server_id = 1; +} + +message Manifest { + string name = 1; + repeated string subscriptions = 2; +} + +message Envelope { + string type = 1; + string chain_id = 2; + repeated string seen = 3; + string session_id = 4; + string run_id = 5; + string turn_id = 6; + bytes payload = 7; + bytes context = 8; // 插件共享袋子,JSON object;不进模型 +} + +message Method { + string name = 1; + string prompt = 2; + bytes parameters_schema = 3; + repeated string capabilities = 4; + bool requires_approval = 5; +} + +message MethodInput { + string session_id = 1; + string run_id = 2; + string turn_id = 3; + string call_id = 4; + string name = 5; + bytes arguments = 6; +} + +message MethodResult { + bool success = 1; + bytes output = 2; + string error = 3; +} + +message MemoryKey { + string session_id = 1; + string scope = 2; + string name = 3; +} + +message MemoryValue { + string text = 1; +} + +message MemoryEntry { + MemoryKey key = 1; + string text = 2; +} + +message CompleteRequest { + string session_id = 1; + string run_id = 2; + string prompt = 3; + string text = 4; +} + +message CompleteResult { + string text = 1; +} + +message Notice { + string session_id = 1; + string run_id = 2; + string text = 3; +} diff --git a/server/pkg/plugin/proto/plugin_grpc.pb.go b/server/pkg/plugin/proto/plugin_grpc.pb.go new file mode 100644 index 0000000..b302727 --- /dev/null +++ b/server/pkg/plugin/proto/plugin_grpc.pb.go @@ -0,0 +1,489 @@ +// Code generated by protoc-gen-go-grpc. DO NOT EDIT. +// versions: +// - protoc-gen-go-grpc v1.6.2 +// - protoc v7.36.1 +// source: plugin.proto + +package pluginpb + +import ( + context "context" + grpc "google.golang.org/grpc" + codes "google.golang.org/grpc/codes" + status "google.golang.org/grpc/status" +) + +// This is a compile-time assertion to ensure that this generated file +// is compatible with the grpc package it is being compiled against. +// Requires gRPC-Go v1.64.0 or later. +const _ = grpc.SupportPackageIsVersion9 + +const ( + Plugin_Bootstrap_FullMethodName = "/codedock.plugin.v1.Plugin/Bootstrap" + Plugin_OnEvent_FullMethodName = "/codedock.plugin.v1.Plugin/OnEvent" + Plugin_ExecuteMethod_FullMethodName = "/codedock.plugin.v1.Plugin/ExecuteMethod" +) + +// PluginClient is the client API for Plugin service. +// +// For semantics around ctx use and closing/ending streaming RPCs, please refer to https://pkg.go.dev/google.golang.org/grpc/?tab=doc#ClientConn.NewStream. +type PluginClient interface { + Bootstrap(ctx context.Context, in *BootstrapRequest, opts ...grpc.CallOption) (*Manifest, error) + OnEvent(ctx context.Context, in *Envelope, opts ...grpc.CallOption) (*Envelope, error) + ExecuteMethod(ctx context.Context, in *MethodInput, opts ...grpc.CallOption) (*MethodResult, error) +} + +type pluginClient struct { + cc grpc.ClientConnInterface +} + +func NewPluginClient(cc grpc.ClientConnInterface) PluginClient { + return &pluginClient{cc} +} + +func (c *pluginClient) Bootstrap(ctx context.Context, in *BootstrapRequest, opts ...grpc.CallOption) (*Manifest, error) { + cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) + out := new(Manifest) + err := c.cc.Invoke(ctx, Plugin_Bootstrap_FullMethodName, in, out, cOpts...) + if err != nil { + return nil, err + } + return out, nil +} + +func (c *pluginClient) OnEvent(ctx context.Context, in *Envelope, opts ...grpc.CallOption) (*Envelope, error) { + cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) + out := new(Envelope) + err := c.cc.Invoke(ctx, Plugin_OnEvent_FullMethodName, in, out, cOpts...) + if err != nil { + return nil, err + } + return out, nil +} + +func (c *pluginClient) ExecuteMethod(ctx context.Context, in *MethodInput, opts ...grpc.CallOption) (*MethodResult, error) { + cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) + out := new(MethodResult) + err := c.cc.Invoke(ctx, Plugin_ExecuteMethod_FullMethodName, in, out, cOpts...) + if err != nil { + return nil, err + } + return out, nil +} + +// PluginServer is the server API for Plugin service. +// All implementations must embed UnimplementedPluginServer +// for forward compatibility. +type PluginServer interface { + Bootstrap(context.Context, *BootstrapRequest) (*Manifest, error) + OnEvent(context.Context, *Envelope) (*Envelope, error) + ExecuteMethod(context.Context, *MethodInput) (*MethodResult, error) + mustEmbedUnimplementedPluginServer() +} + +// UnimplementedPluginServer must be embedded to have +// forward compatible implementations. +// +// NOTE: this should be embedded by value instead of pointer to avoid a nil +// pointer dereference when methods are called. +type UnimplementedPluginServer struct{} + +func (UnimplementedPluginServer) Bootstrap(context.Context, *BootstrapRequest) (*Manifest, error) { + return nil, status.Error(codes.Unimplemented, "method Bootstrap not implemented") +} +func (UnimplementedPluginServer) OnEvent(context.Context, *Envelope) (*Envelope, error) { + return nil, status.Error(codes.Unimplemented, "method OnEvent not implemented") +} +func (UnimplementedPluginServer) ExecuteMethod(context.Context, *MethodInput) (*MethodResult, error) { + return nil, status.Error(codes.Unimplemented, "method ExecuteMethod not implemented") +} +func (UnimplementedPluginServer) mustEmbedUnimplementedPluginServer() {} +func (UnimplementedPluginServer) testEmbeddedByValue() {} + +// UnsafePluginServer may be embedded to opt out of forward compatibility for this service. +// Use of this interface is not recommended, as added methods to PluginServer will +// result in compilation errors. +type UnsafePluginServer interface { + mustEmbedUnimplementedPluginServer() +} + +func RegisterPluginServer(s grpc.ServiceRegistrar, srv PluginServer) { + // If the following call panics, it indicates UnimplementedPluginServer was + // embedded by pointer and is nil. This will cause panics if an + // unimplemented method is ever invoked, so we test this at initialization + // time to prevent it from happening at runtime later due to I/O. + if t, ok := srv.(interface{ testEmbeddedByValue() }); ok { + t.testEmbeddedByValue() + } + s.RegisterService(&Plugin_ServiceDesc, srv) +} + +func _Plugin_Bootstrap_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { + in := new(BootstrapRequest) + if err := dec(in); err != nil { + return nil, err + } + if interceptor == nil { + return srv.(PluginServer).Bootstrap(ctx, in) + } + info := &grpc.UnaryServerInfo{ + Server: srv, + FullMethod: Plugin_Bootstrap_FullMethodName, + } + handler := func(ctx context.Context, req interface{}) (interface{}, error) { + return srv.(PluginServer).Bootstrap(ctx, req.(*BootstrapRequest)) + } + return interceptor(ctx, in, info, handler) +} + +func _Plugin_OnEvent_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { + in := new(Envelope) + if err := dec(in); err != nil { + return nil, err + } + if interceptor == nil { + return srv.(PluginServer).OnEvent(ctx, in) + } + info := &grpc.UnaryServerInfo{ + Server: srv, + FullMethod: Plugin_OnEvent_FullMethodName, + } + handler := func(ctx context.Context, req interface{}) (interface{}, error) { + return srv.(PluginServer).OnEvent(ctx, req.(*Envelope)) + } + return interceptor(ctx, in, info, handler) +} + +func _Plugin_ExecuteMethod_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { + in := new(MethodInput) + if err := dec(in); err != nil { + return nil, err + } + if interceptor == nil { + return srv.(PluginServer).ExecuteMethod(ctx, in) + } + info := &grpc.UnaryServerInfo{ + Server: srv, + FullMethod: Plugin_ExecuteMethod_FullMethodName, + } + handler := func(ctx context.Context, req interface{}) (interface{}, error) { + return srv.(PluginServer).ExecuteMethod(ctx, req.(*MethodInput)) + } + return interceptor(ctx, in, info, handler) +} + +// Plugin_ServiceDesc is the grpc.ServiceDesc for Plugin service. +// It's only intended for direct use with grpc.RegisterService, +// and not to be introspected or modified (even as a copy) +var Plugin_ServiceDesc = grpc.ServiceDesc{ + ServiceName: "codedock.plugin.v1.Plugin", + HandlerType: (*PluginServer)(nil), + Methods: []grpc.MethodDesc{ + { + MethodName: "Bootstrap", + Handler: _Plugin_Bootstrap_Handler, + }, + { + MethodName: "OnEvent", + Handler: _Plugin_OnEvent_Handler, + }, + { + MethodName: "ExecuteMethod", + Handler: _Plugin_ExecuteMethod_Handler, + }, + }, + Streams: []grpc.StreamDesc{}, + Metadata: "plugin.proto", +} + +const ( + Host_Emit_FullMethodName = "/codedock.plugin.v1.Host/Emit" + Host_RegisterMethod_FullMethodName = "/codedock.plugin.v1.Host/RegisterMethod" + Host_MemoryGet_FullMethodName = "/codedock.plugin.v1.Host/MemoryGet" + Host_MemoryUpsert_FullMethodName = "/codedock.plugin.v1.Host/MemoryUpsert" + Host_Complete_FullMethodName = "/codedock.plugin.v1.Host/Complete" + Host_AppendNotice_FullMethodName = "/codedock.plugin.v1.Host/AppendNotice" +) + +// HostClient is the client API for Host service. +// +// For semantics around ctx use and closing/ending streaming RPCs, please refer to https://pkg.go.dev/google.golang.org/grpc/?tab=doc#ClientConn.NewStream. +type HostClient interface { + Emit(ctx context.Context, in *Envelope, opts ...grpc.CallOption) (*Empty, error) + RegisterMethod(ctx context.Context, in *Method, opts ...grpc.CallOption) (*Empty, error) + MemoryGet(ctx context.Context, in *MemoryKey, opts ...grpc.CallOption) (*MemoryValue, error) + MemoryUpsert(ctx context.Context, in *MemoryEntry, opts ...grpc.CallOption) (*Empty, error) + Complete(ctx context.Context, in *CompleteRequest, opts ...grpc.CallOption) (*CompleteResult, error) + AppendNotice(ctx context.Context, in *Notice, opts ...grpc.CallOption) (*Empty, error) +} + +type hostClient struct { + cc grpc.ClientConnInterface +} + +func NewHostClient(cc grpc.ClientConnInterface) HostClient { + return &hostClient{cc} +} + +func (c *hostClient) Emit(ctx context.Context, in *Envelope, opts ...grpc.CallOption) (*Empty, error) { + cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) + out := new(Empty) + err := c.cc.Invoke(ctx, Host_Emit_FullMethodName, in, out, cOpts...) + if err != nil { + return nil, err + } + return out, nil +} + +func (c *hostClient) RegisterMethod(ctx context.Context, in *Method, opts ...grpc.CallOption) (*Empty, error) { + cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) + out := new(Empty) + err := c.cc.Invoke(ctx, Host_RegisterMethod_FullMethodName, in, out, cOpts...) + if err != nil { + return nil, err + } + return out, nil +} + +func (c *hostClient) MemoryGet(ctx context.Context, in *MemoryKey, opts ...grpc.CallOption) (*MemoryValue, error) { + cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) + out := new(MemoryValue) + err := c.cc.Invoke(ctx, Host_MemoryGet_FullMethodName, in, out, cOpts...) + if err != nil { + return nil, err + } + return out, nil +} + +func (c *hostClient) MemoryUpsert(ctx context.Context, in *MemoryEntry, opts ...grpc.CallOption) (*Empty, error) { + cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) + out := new(Empty) + err := c.cc.Invoke(ctx, Host_MemoryUpsert_FullMethodName, in, out, cOpts...) + if err != nil { + return nil, err + } + return out, nil +} + +func (c *hostClient) Complete(ctx context.Context, in *CompleteRequest, opts ...grpc.CallOption) (*CompleteResult, error) { + cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) + out := new(CompleteResult) + err := c.cc.Invoke(ctx, Host_Complete_FullMethodName, in, out, cOpts...) + if err != nil { + return nil, err + } + return out, nil +} + +func (c *hostClient) AppendNotice(ctx context.Context, in *Notice, opts ...grpc.CallOption) (*Empty, error) { + cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) + out := new(Empty) + err := c.cc.Invoke(ctx, Host_AppendNotice_FullMethodName, in, out, cOpts...) + if err != nil { + return nil, err + } + return out, nil +} + +// HostServer is the server API for Host service. +// All implementations must embed UnimplementedHostServer +// for forward compatibility. +type HostServer interface { + Emit(context.Context, *Envelope) (*Empty, error) + RegisterMethod(context.Context, *Method) (*Empty, error) + MemoryGet(context.Context, *MemoryKey) (*MemoryValue, error) + MemoryUpsert(context.Context, *MemoryEntry) (*Empty, error) + Complete(context.Context, *CompleteRequest) (*CompleteResult, error) + AppendNotice(context.Context, *Notice) (*Empty, error) + mustEmbedUnimplementedHostServer() +} + +// UnimplementedHostServer must be embedded to have +// forward compatible implementations. +// +// NOTE: this should be embedded by value instead of pointer to avoid a nil +// pointer dereference when methods are called. +type UnimplementedHostServer struct{} + +func (UnimplementedHostServer) Emit(context.Context, *Envelope) (*Empty, error) { + return nil, status.Error(codes.Unimplemented, "method Emit not implemented") +} +func (UnimplementedHostServer) RegisterMethod(context.Context, *Method) (*Empty, error) { + return nil, status.Error(codes.Unimplemented, "method RegisterMethod not implemented") +} +func (UnimplementedHostServer) MemoryGet(context.Context, *MemoryKey) (*MemoryValue, error) { + return nil, status.Error(codes.Unimplemented, "method MemoryGet not implemented") +} +func (UnimplementedHostServer) MemoryUpsert(context.Context, *MemoryEntry) (*Empty, error) { + return nil, status.Error(codes.Unimplemented, "method MemoryUpsert not implemented") +} +func (UnimplementedHostServer) Complete(context.Context, *CompleteRequest) (*CompleteResult, error) { + return nil, status.Error(codes.Unimplemented, "method Complete not implemented") +} +func (UnimplementedHostServer) AppendNotice(context.Context, *Notice) (*Empty, error) { + return nil, status.Error(codes.Unimplemented, "method AppendNotice not implemented") +} +func (UnimplementedHostServer) mustEmbedUnimplementedHostServer() {} +func (UnimplementedHostServer) testEmbeddedByValue() {} + +// UnsafeHostServer may be embedded to opt out of forward compatibility for this service. +// Use of this interface is not recommended, as added methods to HostServer will +// result in compilation errors. +type UnsafeHostServer interface { + mustEmbedUnimplementedHostServer() +} + +func RegisterHostServer(s grpc.ServiceRegistrar, srv HostServer) { + // If the following call panics, it indicates UnimplementedHostServer was + // embedded by pointer and is nil. This will cause panics if an + // unimplemented method is ever invoked, so we test this at initialization + // time to prevent it from happening at runtime later due to I/O. + if t, ok := srv.(interface{ testEmbeddedByValue() }); ok { + t.testEmbeddedByValue() + } + s.RegisterService(&Host_ServiceDesc, srv) +} + +func _Host_Emit_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { + in := new(Envelope) + if err := dec(in); err != nil { + return nil, err + } + if interceptor == nil { + return srv.(HostServer).Emit(ctx, in) + } + info := &grpc.UnaryServerInfo{ + Server: srv, + FullMethod: Host_Emit_FullMethodName, + } + handler := func(ctx context.Context, req interface{}) (interface{}, error) { + return srv.(HostServer).Emit(ctx, req.(*Envelope)) + } + return interceptor(ctx, in, info, handler) +} + +func _Host_RegisterMethod_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { + in := new(Method) + if err := dec(in); err != nil { + return nil, err + } + if interceptor == nil { + return srv.(HostServer).RegisterMethod(ctx, in) + } + info := &grpc.UnaryServerInfo{ + Server: srv, + FullMethod: Host_RegisterMethod_FullMethodName, + } + handler := func(ctx context.Context, req interface{}) (interface{}, error) { + return srv.(HostServer).RegisterMethod(ctx, req.(*Method)) + } + return interceptor(ctx, in, info, handler) +} + +func _Host_MemoryGet_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { + in := new(MemoryKey) + if err := dec(in); err != nil { + return nil, err + } + if interceptor == nil { + return srv.(HostServer).MemoryGet(ctx, in) + } + info := &grpc.UnaryServerInfo{ + Server: srv, + FullMethod: Host_MemoryGet_FullMethodName, + } + handler := func(ctx context.Context, req interface{}) (interface{}, error) { + return srv.(HostServer).MemoryGet(ctx, req.(*MemoryKey)) + } + return interceptor(ctx, in, info, handler) +} + +func _Host_MemoryUpsert_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { + in := new(MemoryEntry) + if err := dec(in); err != nil { + return nil, err + } + if interceptor == nil { + return srv.(HostServer).MemoryUpsert(ctx, in) + } + info := &grpc.UnaryServerInfo{ + Server: srv, + FullMethod: Host_MemoryUpsert_FullMethodName, + } + handler := func(ctx context.Context, req interface{}) (interface{}, error) { + return srv.(HostServer).MemoryUpsert(ctx, req.(*MemoryEntry)) + } + return interceptor(ctx, in, info, handler) +} + +func _Host_Complete_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { + in := new(CompleteRequest) + if err := dec(in); err != nil { + return nil, err + } + if interceptor == nil { + return srv.(HostServer).Complete(ctx, in) + } + info := &grpc.UnaryServerInfo{ + Server: srv, + FullMethod: Host_Complete_FullMethodName, + } + handler := func(ctx context.Context, req interface{}) (interface{}, error) { + return srv.(HostServer).Complete(ctx, req.(*CompleteRequest)) + } + return interceptor(ctx, in, info, handler) +} + +func _Host_AppendNotice_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { + in := new(Notice) + if err := dec(in); err != nil { + return nil, err + } + if interceptor == nil { + return srv.(HostServer).AppendNotice(ctx, in) + } + info := &grpc.UnaryServerInfo{ + Server: srv, + FullMethod: Host_AppendNotice_FullMethodName, + } + handler := func(ctx context.Context, req interface{}) (interface{}, error) { + return srv.(HostServer).AppendNotice(ctx, req.(*Notice)) + } + return interceptor(ctx, in, info, handler) +} + +// Host_ServiceDesc is the grpc.ServiceDesc for Host service. +// It's only intended for direct use with grpc.RegisterService, +// and not to be introspected or modified (even as a copy) +var Host_ServiceDesc = grpc.ServiceDesc{ + ServiceName: "codedock.plugin.v1.Host", + HandlerType: (*HostServer)(nil), + Methods: []grpc.MethodDesc{ + { + MethodName: "Emit", + Handler: _Host_Emit_Handler, + }, + { + MethodName: "RegisterMethod", + Handler: _Host_RegisterMethod_Handler, + }, + { + MethodName: "MemoryGet", + Handler: _Host_MemoryGet_Handler, + }, + { + MethodName: "MemoryUpsert", + Handler: _Host_MemoryUpsert_Handler, + }, + { + MethodName: "Complete", + Handler: _Host_Complete_Handler, + }, + { + MethodName: "AppendNotice", + Handler: _Host_AppendNotice_Handler, + }, + }, + Streams: []grpc.StreamDesc{}, + Metadata: "plugin.proto", +} diff --git a/server/pkg/plugin/serve.go b/server/pkg/plugin/serve.go new file mode 100644 index 0000000..3aacb67 --- /dev/null +++ b/server/pkg/plugin/serve.go @@ -0,0 +1,14 @@ +package plugin + +import goplugin "github.com/hashicorp/go-plugin" + +// Serve 在插件进程里启动 gRPC 服务,供宿主连接。 +func Serve(p Plugin) { + goplugin.Serve(&goplugin.ServeConfig{ + HandshakeConfig: Handshake, + Plugins: map[string]goplugin.Plugin{ + PluginName: &GRPCPlugin{Impl: &runtime{Plugin: p}}, + }, + GRPCServer: goplugin.DefaultGRPCServer, + }) +}