From 46cb78f4b72da3c77152ce23121bee83796bcbe3 Mon Sep 17 00:00:00 2001 From: fylorn <249551762+fylorn@users.noreply.github.com> Date: Fri, 2 Oct 2026 19:36:39 +0800 Subject: [PATCH] Run plugins in the tw-plugin sandbox The gateway's default plugin engine is now the tw-plugin sandbox instead of the stand-in that refused every plugin. `plugin::sandbox` is only a type adapter: manifests, hook outcomes, errors and log lines are converted one to one, and limits stay tw-plugin's defaults. - The runtime is created once per process, the first time a plugin is compiled, so processes and tests without plugins never start it. If it cannot start, every plugin fails to load with the reason, and the requests it covers follow `on_error`, the same as before. - Plugins are compiled on a thread with an 8 MiB stack. Compiling runs the module's top level in the sandbox, and it happens on the configuration path, which can be the main thread (1 MiB on Windows). Loading from a 128 KiB thread overflowed before this change and now works. - What a plugin returns is compared with `tw_plugin::js_equal`: a value that went through JavaScript comes back with `2.0` as `2` and large integers rounded. Untouched tool inputs, tool schemas and read-only parts were seen as edited (rewritten, or refused as read-only); now they keep their original bytes, so prompt caching of tool definitions survives. - An empty array from onToolCall drops the call, the same as null. Tests: the adapter against real JavaScript plugins, and three end-to-end runs through the configuration (file, approved hash, settings): a plugin rewrites the request the upstream receives and the answer the client receives, sees a placeholder instead of the user's key, drops a Bash call and logs on itself; a plugin that throws stops the request under `on_error: reject`; numbers a plugin did not touch reach the upstream as they were. Co-Authored-By: Claude Opus 5.5 --- Cargo.lock | 1 + crates/tw-gateway/Cargo.toml | 3 + crates/tw-gateway/src/plugin/engine.rs | 7 +- crates/tw-gateway/src/plugin/mod.rs | 9 +- crates/tw-gateway/src/plugin/pool.rs | 5 +- crates/tw-gateway/src/plugin/reply/mod.rs | 13 +- .../tw-gateway/src/plugin/reply/tests/mod.rs | 52 ++- crates/tw-gateway/src/plugin/sandbox.rs | 219 ++++++++++ .../src/plugin/sandbox/tests/mod.rs | 233 +++++++++++ crates/tw-gateway/src/plugin/view/mod.rs | 18 +- .../tw-gateway/src/plugin/view/tests/mod.rs | 54 +++ crates/tw-gateway/tests/plugins_js.rs | 385 ++++++++++++++++++ 12 files changed, 965 insertions(+), 34 deletions(-) create mode 100644 crates/tw-gateway/src/plugin/sandbox.rs create mode 100644 crates/tw-gateway/src/plugin/sandbox/tests/mod.rs create mode 100644 crates/tw-gateway/tests/plugins_js.rs diff --git a/Cargo.lock b/Cargo.lock index fd9bf94d..991eae07 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -3308,6 +3308,7 @@ dependencies = [ "tw-engine", "tw-guard", "tw-observe", + "tw-plugin", "tw-pricing", "tw-secret", "tw-types", diff --git a/crates/tw-gateway/Cargo.toml b/crates/tw-gateway/Cargo.toml index 1e82713a..b2859387 100644 --- a/crates/tw-gateway/Cargo.toml +++ b/crates/tw-gateway/Cargo.toml @@ -45,6 +45,9 @@ rand = { workspace = true } uuid = { workspace = true } tw-guard = { workspace = true } tw-dialect = { workspace = true } +# 脚本插件的沙箱。**编它要一个能出 wasm 的 clang**:依赖它的只能是网关和 twcore +# (tw-plugin 的 tests/boundary.rs 守着) +tw-plugin = { workspace = true } thiserror = { workspace = true } tokio = { workspace = true } tracing = { workspace = true } diff --git a/crates/tw-gateway/src/plugin/engine.rs b/crates/tw-gateway/src/plugin/engine.rs index cf1304d3..6d4dc610 100644 --- a/crates/tw-gateway/src/plugin/engine.rs +++ b/crates/tw-gateway/src/plugin/engine.rs @@ -1,8 +1,9 @@ //! 网关要从插件运行时那里拿到的东西:把一份源码编成插件,读出它的 manifest。 //! -//! **这里是一道接缝**:沙箱在 `tw-plugin` 里,网关只认这个 trait。真的运行时接上之前, -//! 由 [`Unavailable`] 顶着 —— 每个插件都「加载不了」,有一个就拒一个请求(出错时 -//! 拒绝是出厂的做法),而不是悄悄放过。测试拿一个假的引擎接在这里。 +//! **这里是一道接缝**:沙箱在 `tw-plugin` 里(接上它的是 [`crate::plugin::sandbox`]), +//! 网关只认这个 trait。没有运行时可用时由 [`Unavailable`] 顶着 —— 每个插件都「加载 +//! 不了」,有一个就拒一个请求(出错时拒绝是出厂的做法),而不是悄悄放过。测试拿一个 +//! 假的引擎接在这里。 use std::sync::Arc; diff --git a/crates/tw-gateway/src/plugin/mod.rs b/crates/tw-gateway/src/plugin/mod.rs index 7f8d8fd4..72247d35 100644 --- a/crates/tw-gateway/src/plugin/mod.rs +++ b/crates/tw-gateway/src/plugin/mod.rs @@ -3,7 +3,7 @@ //! 插件是 JavaScript,**只在沙箱里跑**(`tw-plugin`,Wasmtime)。网关这一侧分三块: //! //! - [`engine`]:网关要从运行时那里拿到的东西 —— 把一份源码编成一个插件,读出它的 -//! manifest。运行时还没接上时由一个替身顶着,所有插件都是「加载不了」; +//! manifest。真的运行时在 [`sandbox`],测试可以换一个假的; //! - [`host`]:一个编好的插件能做什么(跑请求钩子、回答钩子),数据面调它; //! - [`set`]:跟着配置一起换的那一份 —— 每个配置了的插件此刻的样子(能跑、文件 //! 变了、加载出错)、范围、出错时怎么办,以及跨重载存活的计数和日志。 @@ -35,6 +35,7 @@ pub mod load; pub mod pool; pub mod reply; pub mod request; +pub mod sandbox; pub mod set; pub mod trial; pub mod view; @@ -94,10 +95,10 @@ impl crate::AppState { } } -/// 这个进程用的插件运行时。**沙箱还没接上**:在那之前每个插件都「加载不了」, -/// 管得着的请求照它的 `on_error` 处置。 +/// 这个进程用的插件运行时:`tw-plugin` 的沙箱(见 [`sandbox`])。**第一次编插件时 +/// 才真的起来**;起不来时每个插件都「加载不了」,管得着的请求照它的 `on_error` 处置。 pub fn default_engine() -> std::sync::Arc { - std::sync::Arc::new(Unavailable::default()) + std::sync::Arc::new(sandbox::Sandbox) } /// 出错却没说为什么。数据面总该给一句,这里只是不让通知空着 diff --git a/crates/tw-gateway/src/plugin/pool.rs b/crates/tw-gateway/src/plugin/pool.rs index 037090e4..7b241f82 100644 --- a/crates/tw-gateway/src/plugin/pool.rs +++ b/crates/tw-gateway/src/plugin/pool.rs @@ -35,8 +35,9 @@ pub struct Pool { tx: OnceLock>, String>>, } -/// 一根线程的栈。沙箱里的调用要比默认的 2 MiB 深一些 -const STACK: usize = 8 * 1024 * 1024; +/// 一根线程的栈。**跑沙箱的线程至少要 2 MiB**(wasm 自己最多用 1 MiB,外面还有宿主 +/// 那一侧的调用),这里明着给足,不靠平台的默认值(Windows 上只有 1 MiB) +pub(crate) const STACK: usize = 8 * 1024 * 1024; impl Pool { /// `threads` 根线程,最多 `queue` 个任务在跑或者排着。 diff --git a/crates/tw-gateway/src/plugin/reply/mod.rs b/crates/tw-gateway/src/plugin/reply/mod.rs index 66503490..e10406f5 100644 --- a/crates/tw-gateway/src/plugin/reply/mod.rs +++ b/crates/tw-gateway/src/plugin/reply/mod.rs @@ -531,6 +531,11 @@ impl Chain { changed = true; self.stages[i].counts.dropped += 1; } + // 交回一个空数组也是去掉 + Ok(ToolCallOutcome::Replace(vals)) if vals.is_empty() => { + changed = true; + self.stages[i].counts.dropped += 1; + } Ok(ToolCallOutcome::Replace(vals)) => { // 原样交回来的一个调用就是没改 if let [one] = vals.as_slice() @@ -653,7 +658,8 @@ impl Drop for Chain { } } -/// 交回来的调用和交出去的一样(id、名字、参数都没变) +/// 交回来的调用和交出去的一样(id、名字、参数都没变)。参数按 JavaScript 的眼光比 +/// ([`tw_plugin::js_equal`]):进出一趟 JS 的 `1.0` 回来是 `1`,那不算改 fn same_call(v: &Value, given: &Value) -> bool { let o = match v.as_object() { Some(o) => o, @@ -662,7 +668,10 @@ fn same_call(v: &Value, given: &Value) -> bool { o.keys() .all(|k| matches!(k.as_str(), "id" | "name" | "input")) && o.get("name") == given.get("name") - && o.get("input") == given.get("input") + && match (o.get("input"), given.get("input")) { + (Some(a), Some(b)) => tw_plugin::js_equal(a, b), + (a, b) => a == b, + } && o.get("id").is_none_or(|id| Some(id) == given.get("id")) } diff --git a/crates/tw-gateway/src/plugin/reply/tests/mod.rs b/crates/tw-gateway/src/plugin/reply/tests/mod.rs index cb94f935..94ba3749 100644 --- a/crates/tw-gateway/src/plugin/reply/tests/mod.rs +++ b/crates/tw-gateway/src/plugin/reply/tests/mod.rs @@ -328,25 +328,39 @@ async fn a_replaced_tool_call_is_written_whole_and_later_blocks_move_up() { #[tokio::test] async fn dropping_every_tool_call_ends_the_turn_instead_of_waiting_for_results() { let input = anthropic_stream(); - let mut s = Stream::new( - chain_of( - Dialect::Anthropic, - vec![tools(|_| ToolCallOutcome::Drop)], - &json!({}), - ) - .await, - Framing::Sse, - ); - let (out, _) = run(&mut s, &input, 4096).await; - let blocks = anthropic_blocks(&out); - let idx: Vec = blocks.iter().map(|b| b.0).collect(); - assert_eq!(idx, [0, 1, 2]); - assert_eq!(blocks[2], (2, "text".into(), "bye".into())); - let stop = frames(&out) - .into_iter() - .find(|(_, v)| v["type"] == "message_delta") - .unwrap(); - assert_eq!(stop.1["delta"]["stop_reason"], "end_turn"); + // 交回 null 和交回空数组是同一个意思 + for drop in [ + tools(|_| ToolCallOutcome::Drop), + tools(|_| ToolCallOutcome::Replace(Vec::new())), + ] { + let mut s = Stream::new( + chain_of(Dialect::Anthropic, vec![drop], &json!({})).await, + Framing::Sse, + ); + let (out, _) = run(&mut s, &input, 4096).await; + let blocks = anthropic_blocks(&out); + let idx: Vec = blocks.iter().map(|b| b.0).collect(); + assert_eq!(idx, [0, 1, 2]); + assert_eq!(blocks[2], (2, "text".into(), "bye".into())); + let stop = frames(&out) + .into_iter() + .find(|(_, v)| v["type"] == "message_delta") + .unwrap(); + assert_eq!(stop.1["delta"]["stop_reason"], "end_turn"); + } +} + +/// 进出一趟 JavaScript 的调用:`2.0` 回来是 `2`,超过 2^53 的整数丢了精度 —— 这还是 +/// 原来那个调用,不能算改过 +#[test] +fn a_call_that_went_through_javascript_unchanged_is_the_same_call() { + let given = json!({"id": "t1", "name": "Read", + "input": {"limit": 2.0, "seed": 12345678901234567890u64}}); + let back = json!({"id": "t1", "name": "Read", + "input": {"seed": 12345678901234567000u64, "limit": 2}}); + assert!(same_call(&back, &given)); + let other = json!({"id": "t1", "name": "Read", "input": {"limit": 3, "seed": 1}}); + assert!(!same_call(&other, &given)); } #[tokio::test] diff --git a/crates/tw-gateway/src/plugin/sandbox.rs b/crates/tw-gateway/src/plugin/sandbox.rs new file mode 100644 index 00000000..82cd7e22 --- /dev/null +++ b/crates/tw-gateway/src/plugin/sandbox.rs @@ -0,0 +1,219 @@ +//! 真的插件运行时:`tw-plugin` 的沙箱(QuickJS 跑在 Wasmtime 里)接到 [`Engine`] 这道 +//! 接缝上。 +//! +//! **只是一层把类型对上的适配**:manifest、钩子的结局、错误、日志,一样一样换成网关的 +//! 那一份,语义一点不改 —— 怎么跑、上限多少都是 `tw-plugin` 的事,视图、权限、写回 +//! 都是网关的事。 +//! +//! 运行时一个进程一份,**第一次编插件时才起**:起它要加载沙箱、起一个计时线程,没装 +//! 插件的进程(以及绝大多数测试)不该为它付这个钱。起不来(比如地址空间受限的机器) +//! 就和没有运行时一样:每个插件都加载不了,原因照实说,管得着的请求按 `on_error` 处置。 + +use std::sync::{Arc, OnceLock}; + +use serde_json::Value; + +use crate::plugin::engine::{Engine, Hooks, LoadError, Manifest, SettingSpec}; +use crate::plugin::host::{ + Invocation, PluginHost, ReplyHost, RequestOutcome, RunError, ToolCallOutcome, +}; +use crate::plugin::set::{LogLine, Scope}; + +/// 进程里那一份运行时。起不来时是起不来的原因 +fn runtime() -> Result<&'static tw_plugin::Runtime, LoadError> { + static RT: OnceLock> = OnceLock::new(); + RT.get_or_init(|| { + tw_plugin::Runtime::new(tw_plugin::Limits::default()).map_err(|e| e.to_string()) + }) + .as_ref() + .map_err(|e| LoadError::Engine(e.clone())) +} + +/// `tw-plugin` 的沙箱。生产上用的就是它(见 [`crate::plugin::default_engine`])。 +#[derive(Debug, Default, Clone, Copy)] +pub struct Sandbox; + +impl Engine for Sandbox { + fn load(&self, source: &[u8]) -> Result, LoadError> { + let rt = runtime()?; + // 编译要在沙箱里跑一遍模块顶层,**调用方的线程栈未必够**:插件是在换配置的那一路 + // 上编的,那可能是主线程(Windows 上只有 1 MiB)。在一根栈给足了的线程上编 —— + // 编插件只在换配置、装插件时发生,多起一根线程不算什么 + let plugin = std::thread::scope(|s| { + let compiling = std::thread::Builder::new() + .name("tw-plugin-load".into()) + .stack_size(crate::plugin::pool::STACK) + .spawn_scoped(s, || rt.load(source).map_err(load_error)) + .map_err(|e| { + LoadError::Engine(format!("cannot start a thread to compile the plugin: {e}")) + })?; + compiling + .join() + .unwrap_or_else(|_| Err(LoadError::Engine("compiling the plugin crashed".into()))) + })?; + let manifest = manifest(plugin.manifest()); + Ok(Arc::new(Host { plugin, manifest })) + } +} + +/// 编好的一个插件。克隆、跨线程共享都便宜(`tw_plugin::Plugin` 里是一个 `Arc`) +struct Host { + plugin: tw_plugin::Plugin, + /// 换成网关那一份的 manifest。编的时候换一次,之后每次问都是它 + manifest: Manifest, +} + +impl PluginHost for Host { + fn manifest(&self) -> &Manifest { + &self.manifest + } + + fn sha256(&self) -> [u8; 32] { + self.plugin.sha256() + } + + fn on_request(&self, view: Value, ctx: Value) -> Invocation { + invocation(self.plugin.on_request(view, ctx), |o| match o { + tw_plugin::RequestOutcome::Unchanged => RequestOutcome::Unchanged, + tw_plugin::RequestOutcome::Changed(v) => RequestOutcome::Changed(v), + tw_plugin::RequestOutcome::Rejected(r) => RequestOutcome::Rejected(r), + }) + } + + fn reply(&self, ctx: Value) -> Result, RunError> { + match self.plugin.reply(ctx) { + Ok(r) => Ok(Box::new(Reply(r))), + Err(e) => Err(run_error(e)), + } + } +} + +/// 一个回答的实例 +struct Reply(tw_plugin::Reply); + +impl ReplyHost for Reply { + fn on_text(&mut self, text: &str) -> Invocation> { + invocation(self.0.on_text(text), |t| t) + } + + fn on_text_end(&mut self) -> Invocation> { + invocation(self.0.on_text_end(), |t| t) + } + + fn on_tool_call(&mut self, call: Value) -> Invocation { + invocation(self.0.on_tool_call(call), |o| match o { + tw_plugin::ToolCallOutcome::Unchanged => ToolCallOutcome::Unchanged, + tw_plugin::ToolCallOutcome::Replace(calls) => ToolCallOutcome::Replace(calls), + tw_plugin::ToolCallOutcome::Drop => ToolCallOutcome::Drop, + }) + } +} + +// ── 换类型 ─────────────────────────────────────────────────────── + +fn invocation(inv: tw_plugin::Invocation, f: impl FnOnce(T) -> U) -> Invocation { + Invocation { + result: inv.result.map(f).map_err(run_error), + logs: inv.logs.into_iter().map(log_line).collect(), + cpu: inv.cpu, + } +} + +fn log_line(l: tw_plugin::LogLine) -> LogLine { + LogLine { + level: match l.level { + tw_plugin::LogLevel::Log => tw_api::PluginLogLevel::Log, + tw_plugin::LogLevel::Info => tw_api::PluginLogLevel::Info, + tw_plugin::LogLevel::Warn => tw_api::PluginLogLevel::Warn, + tw_plugin::LogLevel::Error => tw_api::PluginLogLevel::Error, + }, + text: l.text, + } +} + +fn run_error(e: tw_plugin::RunError) -> RunError { + match e { + tw_plugin::RunError::CpuLimit => RunError::CpuLimit, + tw_plugin::RunError::MemoryLimit => RunError::MemoryLimit, + tw_plugin::RunError::OutputLimit => RunError::OutputLimit, + tw_plugin::RunError::Threw { message, stack } => RunError::Threw { message, stack }, + tw_plugin::RunError::BadOutput(d) => RunError::BadOutput(d), + tw_plugin::RunError::Trap(d) => RunError::Trap(d), + } +} + +fn load_error(e: tw_plugin::LoadError) -> LoadError { + match e { + tw_plugin::LoadError::TooLarge => LoadError::TooLarge, + tw_plugin::LoadError::Syntax { + message, + line, + column, + } => LoadError::Syntax { + message, + line, + column, + }, + tw_plugin::LoadError::Manifest(d) => LoadError::Manifest(d), + tw_plugin::LoadError::UnsupportedApi(api) => LoadError::UnsupportedApi(api), + tw_plugin::LoadError::Engine(d) => LoadError::Engine(d), + } +} + +fn permission(p: tw_plugin::Permission) -> tw_api::Permission { + match p { + tw_plugin::Permission::System => tw_api::Permission::System, + tw_plugin::Permission::Messages => tw_api::Permission::Messages, + tw_plugin::Permission::Tools => tw_api::Permission::Tools, + tw_plugin::Permission::Params => tw_api::Permission::Params, + tw_plugin::Permission::ReplyText => tw_api::Permission::ReplyText, + tw_plugin::Permission::ReplyToolCalls => tw_api::Permission::ReplyToolCalls, + } +} + +fn manifest(m: &tw_plugin::Manifest) -> Manifest { + let granted: Vec = m.permissions.iter().copied().map(permission).collect(); + Manifest { + name: m.name.clone(), + api: m.api, + description: m.description.clone(), + // 网关这边按 `Permission::ALL` 的顺序排 + permissions: tw_api::Permission::ALL + .iter() + .copied() + .filter(|p| granted.contains(p)) + .collect(), + scope: Scope { + clients: m.scope.clients.clone(), + models: m.scope.models.clone(), + upstreams: m.scope.upstreams.clone(), + }, + reply_mode: match m.reply_mode { + tw_plugin::ReplyMode::Block => tw_api::ReplyMode::Block, + tw_plugin::ReplyMode::Stream => tw_api::ReplyMode::Stream, + }, + settings: m + .settings + .iter() + .map(|s| SettingSpec { + key: s.key.clone(), + kind: match s.kind { + tw_plugin::SettingKind::String => tw_api::SettingKind::String, + tw_plugin::SettingKind::Number => tw_api::SettingKind::Number, + tw_plugin::SettingKind::Boolean => tw_api::SettingKind::Boolean, + }, + label: s.label.clone(), + default: s.default.clone(), + }) + .collect(), + hooks: Hooks { + request: m.hooks.request, + reply_text: m.hooks.reply_text, + reply_text_end: m.hooks.reply_text_end, + tool_call: m.hooks.tool_call, + }, + } +} + +#[cfg(test)] +mod tests; diff --git a/crates/tw-gateway/src/plugin/sandbox/tests/mod.rs b/crates/tw-gateway/src/plugin/sandbox/tests/mod.rs new file mode 100644 index 00000000..aae02a52 --- /dev/null +++ b/crates/tw-gateway/src/plugin/sandbox/tests/mod.rs @@ -0,0 +1,233 @@ +//! 适配层:真的 JavaScript 插件编出来、跑起来,交回的都是网关那一份类型。 +//! +//! 沙箱本身的行为(上限、隔离、清单的每条规则)在 `tw-plugin` 的测试里;这里只看 +//! 换过来的东西对不对得上。 + +use serde_json::{Value, json}; +use sha2::{Digest, Sha256}; + +use super::*; + +fn load(src: &str) -> Arc { + match Sandbox.load(src.as_bytes()) { + Ok(h) => h, + Err(e) => panic!("load failed: {e:?}\n{src}"), + } +} + +fn ctx() -> Value { + json!({ "client": "claude-code", "model": "claude-sonnet-4-5", "format": "anthropic", + "upstream": null, "settings": { "note": "Friday" } }) +} + +fn view() -> Value { + json!({ + "format": "anthropic", + "model": "claude-sonnet-4-5", + "system": "Be brief.", + "messages": [ + { "key": "m0", "role": "user", "parts": [ { "key": "m0.p0", "type": "text", "text": "hi" } ] } + ] + }) +} + +const BOTH: &str = r#" +export const manifest = { + name: "Both", + api: 1, + description: "adds a note and shouts", + permissions: ["reply.text", "system", "reply.tool_calls"], + match: { clients: ["claude-*"], models: [], upstreams: ["anthropic"] }, + reply: "stream", + settings: { + note: { type: "string", label: "Note", default: "today" }, + loud: { type: "boolean", label: "Loud", default: true }, + }, +}; +export function onRequest(req, ctx) { + console.info("saw", req.system); + req.system = req.system + " Note: " + ctx.settings.note; + return req; +} +let held = ""; +export function onReplyText(text) { + held += text; + return ""; +} +export function onReplyTextEnd() { + const out = held.toUpperCase(); + held = ""; + return out; +} +export function onToolCall(call) { + if (call.name === "Bash") return null; + return { ...call, input: { ...call.input, checked: true } }; +} +"#; + +/// manifest 换成网关那一份:权限按 `Permission::ALL` 排,设置按作者写的先后 +#[test] +fn the_manifest_is_carried_over() { + let h = load(BOTH); + let m = h.manifest(); + assert_eq!(m.name, "Both"); + assert_eq!(m.api, 1); + assert_eq!(m.description.as_deref(), Some("adds a note and shouts")); + assert_eq!( + m.permissions, + [ + tw_api::Permission::System, + tw_api::Permission::ReplyText, + tw_api::Permission::ReplyToolCalls + ] + ); + assert_eq!(m.scope.clients, ["claude-*"]); + assert!(m.scope.models.is_empty()); + assert_eq!(m.scope.upstreams, ["anthropic"]); + assert_eq!(m.reply_mode, tw_api::ReplyMode::Stream); + let keys: Vec<(&str, tw_api::SettingKind)> = m + .settings + .iter() + .map(|s| (s.key.as_str(), s.kind)) + .collect(); + assert_eq!( + keys, + [ + ("note", tw_api::SettingKind::String), + ("loud", tw_api::SettingKind::Boolean) + ] + ); + assert_eq!(m.settings[0].label, "Note"); + assert_eq!(m.settings[0].default, json!("today")); + assert_eq!( + m.hooks, + Hooks { + request: true, + reply_text: true, + reply_text_end: true, + tool_call: true + } + ); + // 哈希的就是交进来的那些字节(不变式 I9) + let sha: [u8; 32] = Sha256::digest(BOTH.as_bytes()).into(); + assert_eq!(h.sha256(), sha); +} + +#[test] +fn a_request_hook_changes_the_view_and_its_log_comes_along() { + let h = load(BOTH); + let inv = h.on_request(view(), ctx()); + let Ok(RequestOutcome::Changed(v)) = &inv.result else { + panic!("{:?}", inv.result); + }; + assert_eq!(v["system"], "Be brief. Note: Friday"); + assert_eq!(inv.logs.len(), 1, "{:?}", inv.logs); + assert_eq!(inv.logs[0].level, tw_api::PluginLogLevel::Info); + assert_eq!(inv.logs[0].text, "saw Be brief."); + assert!(inv.cpu > std::time::Duration::ZERO); +} + +#[test] +fn a_rejection_and_a_throw_keep_their_meaning() { + let src = |body: &str| { + format!( + "export const manifest = {{ name: \"t\", api: 1, permissions: [\"messages\"] }};\n\ + export function onRequest(req, ctx) {{ {body} }}" + ) + }; + let no = load(&src("reject(\"not on Fridays\");")); + assert_eq!( + no.on_request(view(), ctx()).result, + Ok(RequestOutcome::Rejected("not on Fridays".into())) + ); + let same = load(&src("return req;")); + assert_eq!( + same.on_request(view(), ctx()).result, + Ok(RequestOutcome::Unchanged) + ); + let boom = load(&src( + "console.error(\"about to fail\"); throw new Error(\"boom\");", + )); + let inv = boom.on_request(view(), ctx()); + match &inv.result { + Err(RunError::Threw { message, .. }) => assert!(message.contains("boom"), "{message}"), + other => panic!("{other:?}"), + } + assert_eq!(inv.logs[0].level, tw_api::PluginLogLevel::Error); + // 记在运行上的那一句是网关的消息码 + assert_eq!(inv.result.unwrap_err().msg().code, "gw.plugin.threw"); +} + +#[test] +fn a_reply_instance_keeps_its_state_for_one_answer() { + let h = load(BOTH); + let mut r = h.reply(ctx()).unwrap(); + assert_eq!(r.on_text("hel").result, Ok(Some(String::new()))); + assert_eq!(r.on_text("lo").result, Ok(Some(String::new()))); + assert_eq!(r.on_text_end().result, Ok(Some("HELLO".into()))); + // 工具调用:丢掉一个、改一个 + assert_eq!( + r.on_tool_call(json!({"id": "t1", "name": "Bash", "input": {"command": "ls"}})) + .result, + Ok(ToolCallOutcome::Drop) + ); + let Ok(ToolCallOutcome::Replace(calls)) = r + .on_tool_call(json!({"id": "t2", "name": "Read", "input": {"path": "a"}})) + .result + else { + panic!("not replaced"); + }; + assert_eq!(calls.len(), 1); + assert_eq!(calls[0]["input"], json!({"path": "a", "checked": true})); + + // 另一个回答是另一个实例:上一个攒着的不会漏过来 + let mut fresh = h.reply(ctx()).unwrap(); + assert_eq!(fresh.on_text_end().result, Ok(Some(String::new()))); +} + +#[test] +fn load_errors_keep_their_line_and_their_code() { + let e = Sandbox + .load(b"export const manifest = { name: \"t\", api: 1, permissions: [\"system\"] };\nexport function onRequest( {\n") + .err() + .expect("a syntax error loaded"); + match &e { + LoadError::Syntax { line, .. } => assert!(line.is_some(), "{e:?}"), + other => panic!("{other:?}"), + } + assert!(e.msg().code.starts_with("gw.plugin.syntax"), "{e:?}"); + + let e = Sandbox + .load(b"export const manifest = { name: \"t\", api: 1, permissions: [\"system\"] };\n") + .err() + .expect("a manifest without its hook loaded"); + assert!(matches!(e, LoadError::Manifest(_)), "{e:?}"); + assert_eq!(e.msg().code, "gw.plugin.manifest"); + + let e = Sandbox + .load(b"export const manifest = { name: \"t\", api: 2, permissions: [\"system\"] };\nexport function onRequest(r) {}\n") + .err() + .expect("API 2 loaded"); + assert_eq!(e, LoadError::UnsupportedApi(2)); +} + +/// 编插件不看调用方的栈有多大:换配置那一路可能在一根小栈的线程上(Windows 的主线程 +/// 只有 1 MiB),而模块顶层是在沙箱里真跑的 +#[test] +fn a_plugin_compiles_from_a_thread_with_a_small_stack() { + let src = "export const manifest = { name: \"deep\", api: 1, permissions: [\"system\"] };\n\ + function depth(n) { return n === 0 ? 0 : 1 + depth(n - 1); }\n\ + const d = depth(400);\n\ + export function onRequest(req) { req.system = String(d); return req; }\n"; + let loaded = std::thread::Builder::new() + .stack_size(128 * 1024) + .spawn(move || { + Sandbox + .load(src.as_bytes()) + .map(|h| h.manifest().name.clone()) + }) + .unwrap() + .join() + .unwrap(); + assert_eq!(loaded.unwrap(), "deep"); +} diff --git a/crates/tw-gateway/src/plugin/view/mod.rs b/crates/tw-gateway/src/plugin/view/mod.rs index d86a98c5..91822705 100644 --- a/crates/tw-gateway/src/plugin/view/mod.rs +++ b/crates/tw-gateway/src/plugin/view/mod.rs @@ -337,6 +337,14 @@ impl ParamsEdit { // ───────────────────────────────────────────────────────── 核对 +/// 插件交回来的一项和它拿到的那一项是不是同一个值。**按 JavaScript 的眼光比** +/// ([`tw_plugin::js_equal`]):值进出一趟 JS,`1.0` 回来是 `1`,超过 2^53 的整数丢了 +/// 精度 —— 插件没碰的那一项不能因此算成改过(改过的要写回去,写回去就丢了原来的 +/// 写法,只读的那些更会被当成越权) +fn same(a: &Value, b: &Value) -> bool { + tw_plugin::js_equal(a, b) +} + /// 把插件交回来的视图对着它拿到的那一份核一遍,得出改了什么。 /// /// `input` 是插件拿到的那一份(裁过、占位符换过);`hidden_tools` 是视图里看不到的 @@ -354,7 +362,7 @@ pub fn check( for (k, v) in out { match k.as_str() { "format" | "model" => { - if input.get(k) != Some(v) { + if !input.get(k).is_some_and(|i| same(i, v)) { return Err(denied(if k == "model" { "`model` is read-only; change `params.model` instead".to_string() } else { @@ -595,7 +603,8 @@ fn check_parts( let Some(input) = o.get("input") else { return Err(bad(format!("tool call `{key}` has no `input`"))); }; - (Some(input) != before.get("input")).then(|| Change::Input(input.clone())) + (!before.get("input").is_some_and(|b| same(b, input))) + .then(|| Change::Input(input.clone())) } "tool_result" => { only_fields( @@ -618,7 +627,7 @@ fn check_parts( } // 推理、图片、别的:整个只读 _ => { - if p != before { + if !same(p, before) { return Err(denied(format!("part `{key}` ({kind}) is read-only"))); } None @@ -712,7 +721,8 @@ fn check_tools( let description = (before.get("description").and_then(Value::as_str) != Some(description)) .then(|| description.to_string()); - let schema = (before.get("input_schema") != Some(schema)).then(|| schema.clone()); + let schema = (!before.get("input_schema").is_some_and(|b| same(b, schema))) + .then(|| schema.clone()); changed |= description.is_some() || schema.is_some(); edits.push(ToolEdit::Keep { from: i, diff --git a/crates/tw-gateway/src/plugin/view/tests/mod.rs b/crates/tw-gateway/src/plugin/view/tests/mod.rs index dd5a2f21..0c1e8cda 100644 --- a/crates/tw-gateway/src/plugin/view/tests/mod.rs +++ b/crates/tw-gateway/src/plugin/view/tests/mod.rs @@ -147,6 +147,60 @@ fn returning_the_view_untouched_changes_nothing() { } } +/// 一个值进出一趟 JavaScript 之后的样子:数字都是双精度浮点,整数值写成不带小数点的 +fn through_js(v: &Value) -> Value { + match v { + Value::Number(n) => { + let f = n.as_f64().unwrap(); + if f.fract() == 0.0 && f.abs() < 1e21 { + serde_json::from_str(&format!("{f:.0}")).unwrap() + } else { + json!(f) + } + } + Value::Array(a) => Value::Array(a.iter().map(through_js).collect()), + Value::Object(o) => { + Value::Object(o.iter().map(|(k, v)| (k.clone(), through_js(v))).collect()) + } + other => other.clone(), + } +} + +/// 插件没碰的数字回来变了写法(`2.0` → `2`、大整数丢了精度)不算改:只读的不报越权, +/// 能改的也不写回 —— 写回就丢了原来的写法,工具定义还会让缓存失效 +#[test] +fn numbers_that_went_through_javascript_are_not_changes() { + let mut raw = anthropic(); + raw["temperature"] = json!(1.0); + raw["messages"][1]["content"][2]["input"] = + json!({ "file_path": "/a.txt", "limit": 2.0, "seed": 12345678901234567890u64 }); + raw["tools"][0]["input_schema"]["properties"]["limit"] = + json!({ "type": "number", "maximum": 2.0 }); + let built = build(Dialect::Anthropic, &raw, MESSAGES).unwrap(); + let input = trim(&built.view, &all()); + let back = through_js(&input); + assert_ne!( + back, input, + "the round trip changed nothing, so this proves nothing" + ); + let edits = check(&input, &back, &all(), built.src.hidden_tools()).unwrap(); + assert!(edits.is_empty(), "{edits:?}"); + + // 只改了系统提示:别的照原样写回,一个字节都不动 + let mut sys = back.clone(); + sys["system"] = json!(format!( + "{} Today is Friday.", + sys["system"].as_str().unwrap() + )); + let edits = check(&input, &sys, &all(), built.src.hidden_tools()).unwrap(); + let mut next = raw.clone(); + apply(&mut next, &built.src, &edits, MESSAGES).unwrap(); + assert_ne!(next["system"], raw["system"]); + for k in ["messages", "tools", "temperature"] { + assert_eq!(next[k].to_string(), raw[k].to_string(), "{k}"); + } +} + #[test] fn appending_to_the_anthropic_system_prompt_adds_a_block_after_the_cached_one() { let (raw, _) = edit(Dialect::Anthropic, &anthropic(), MESSAGES, |v| { diff --git a/crates/tw-gateway/tests/plugins_js.rs b/crates/tw-gateway/tests/plugins_js.rs new file mode 100644 index 00000000..e6a28047 --- /dev/null +++ b/crates/tw-gateway/tests/plugins_js.rs @@ -0,0 +1,385 @@ +//! 真的 JavaScript 插件,从配置一路跑到假上游再回到客户端。 +//! +//! 别的插件测试用替身证明接线;这里证明**真的运行时接在了那条线上**:插件文件按配置 +//! 读进来、哈希对上了才编(沙箱是 `tw-plugin`),请求钩子改的请求是上游收到的那一份, +//! 回答钩子改的是客户端收到的那一份,插件看到的密钥是占位符,日志和计数记在插件上。 + +use std::net::SocketAddr; +use std::sync::{Arc, Mutex}; +use std::time::Duration; + +use axum::Router; +use bytes::Bytes; +use serde_json::{Value, json}; +use sha2::{Digest, Sha256}; +use tw_config::{Client, Config, Listen, Plugin, PluginOnError, Protocol, Provider}; + +const USER_KEY: &str = "sk-ant-api03-USERSOWNKEYAAAAAAAAAAAAAA"; + +const FRIDAY: &str = r#" +export const manifest = { + name: "Friday", + api: 1, + permissions: ["system", "messages", "reply.text", "reply.tool_calls"], + settings: { day: { type: "string", label: "Day", default: "Friday" } }, +}; + +export function onRequest(req, ctx) { + const last = req.messages[req.messages.length - 1]; + console.log("the user said: " + last.parts.map((p) => p.text || "").join("")); + req.system = (req.system || "") + " Today is " + ctx.settings.day + "."; + return req; +} + +export function onReplyText(text) { + return text.toUpperCase(); +} + +export function onToolCall(call) { + if (call.name === "Bash") return null; +} +"#; + +/// 假上游:记下收到的请求体,回一条带一段文字、两个工具调用的流 +async fn upstream(seen: Arc>>) -> SocketAddr { + let app = Router::new().fallback(axum::routing::post(move |body: Bytes| { + let seen = seen.clone(); + async move { + seen.lock() + .unwrap() + .push(serde_json::from_slice(&body).unwrap_or(Value::Null)); + axum::response::Response::builder() + .header("content-type", "text/event-stream") + .body(axum::body::Body::from(answer())) + .unwrap() + } + })); + let l = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = l.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(l, app).await.unwrap() }); + addr +} + +fn ev(v: Value) -> String { + format!("event: {}\ndata: {v}\n\n", v["type"].as_str().unwrap()) +} + +fn tool(index: u32, id: &str, name: &str, input: &str) -> String { + [ + ev( + json!({"type":"content_block_start","index":index,"content_block":{"type":"tool_use","id":id,"name":name,"input":{}}}), + ), + ev( + json!({"type":"content_block_delta","index":index,"delta":{"type":"input_json_delta","partial_json":input}}), + ), + ev(json!({"type":"content_block_stop","index":index})), + ] + .concat() +} + +fn answer() -> String { + [ + ev( + json!({"type":"message_start","message":{"id":"msg_1","type":"message","role":"assistant","model":"claude-sonnet-4-5","content":[],"usage":{"input_tokens":3,"output_tokens":1}}}), + ), + ev( + json!({"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}), + ), + ev( + json!({"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"hello "}}), + ), + ev( + json!({"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"there"}}), + ), + ev(json!({"type":"content_block_stop","index":0})), + tool(1, "toolu_1", "Bash", r#"{"command":"ls"}"#), + tool(2, "toolu_2", "Read", r#"{"path":"a.txt"}"#), + ev( + json!({"type":"message_delta","delta":{"stop_reason":"tool_use"},"usage":{"output_tokens":5}}), + ), + ev(json!({"type":"message_stop"})), + ] + .concat() +} + +fn sha256_hex(b: &[u8]) -> String { + Sha256::digest(b) + .iter() + .map(|x| format!("{x:02x}")) + .collect() +} + +#[tokio::test] +async fn a_javascript_plugin_rewrites_the_request_and_the_answer() { + let seen = Arc::new(Mutex::new(Vec::new())); + let up = upstream(seen.clone()).await; + + // 插件文件放在配置旁边,配置里记着批准的那一份的哈希 + let tmp = tempfile::tempdir().unwrap(); + std::fs::create_dir_all(tmp.path().join("plugins")).unwrap(); + std::fs::write(tmp.path().join("plugins/friday.js"), FRIDAY).unwrap(); + let cfg = Config { + version: 1, + listen: Listen::default(), + clients: vec![Client { + name: "claude-code".into(), + key: "tw-testkey".into(), + ..Default::default() + }], + providers: vec![Provider { + name: "anthropic".into(), + base_url: format!("http://{up}"), + key: Some("sk-upstream".into()), + protocol: Some(Protocol::Anthropic), + ..Default::default() + }], + plugins: vec![Plugin { + id: "friday".into(), + file: "plugins/friday.js".into(), + sha256: sha256_hex(FRIDAY.as_bytes()), + enabled: true, + on_error: PluginOnError::Reject, + scope: Default::default(), + settings: [("day".to_string(), serde_yaml_ng::Value::from("Saturday"))] + .into_iter() + .collect(), + }], + ..Default::default() + }; + let state = tw_gateway::AppState::new(cfg).unwrap(); + // 控制面拿到配置文件的位置时就是这么告诉网关的:插件这时才读得到文件 + state.set_config_dir(tmp.path().to_path_buf()); + let active = state.runtime().plugins.get("friday").cloned().unwrap(); + assert!( + active.ready().is_some(), + "the plugin did not load: {:?}", + active.broken() + ); + assert_eq!(active.name, "Friday"); + + let addr = tw_gateway::serve(state.clone(), ([127, 0, 0, 1], 0).into()) + .await + .unwrap(); + tokio::time::sleep(Duration::from_millis(30)).await; + let r = reqwest::Client::new() + .post(format!("http://{addr}/v1/messages")) + .header("x-api-key", "tw-testkey") + .header("content-type", "application/json") + .body( + json!({ + "model": "claude-sonnet-4-5", "max_tokens": 64, "stream": true, + "system": "Be brief.", + "messages": [{ "role": "user", "content": format!("deploy with {USER_KEY}") }] + }) + .to_string(), + ) + .send() + .await + .unwrap(); + assert_eq!(r.status(), 200); + let body = r.text().await.unwrap(); + + // 请求钩子:上游收到的是改过的那一份,用的是配置里的设置 + let sent = seen.lock().unwrap().clone(); + assert_eq!(sent.len(), 1); + let system = sent[0]["system"].to_string(); + assert!(system.contains("Be brief. Today is Saturday."), "{system}"); + + // 回答钩子:文字是大写的,Bash 那个调用被丢掉,Read 那个留着、编号接上 + let frames: Vec = body + .lines() + .filter_map(|l| l.strip_prefix("data: ")) + .filter_map(|d| serde_json::from_str(d).ok()) + .collect(); + let text: String = frames + .iter() + .filter_map(|v| v["delta"]["text"].as_str()) + .collect(); + assert_eq!(text, "HELLO THERE", "{body}"); + let tools: Vec<(u64, &str)> = frames + .iter() + .filter(|v| v["type"] == "content_block_start" && v["content_block"]["type"] == "tool_use") + .map(|v| { + ( + v["index"].as_u64().unwrap(), + v["content_block"]["name"].as_str().unwrap(), + ) + }) + .collect(); + assert_eq!(tools, [(1, "Read")], "{body}"); + assert!(!body.contains("Bash"), "{body}"); + + // 插件看到的是占位符,不是密钥;日志记在这个插件上、挂着请求号 + let logs = active.logs.lines(); + assert_eq!(logs.len(), 1, "{logs:?}"); + assert!( + logs[0] + .text + .starts_with("the user said: deploy with <