From 29cf3e3d61da8126c1f9af7b022374a3d95369ec Mon Sep 17 00:00:00 2001 From: fylorn <249551762+fylorn@users.noreply.github.com> Date: Sat, 3 Oct 2026 01:39:14 +0800 Subject: [PATCH] Run request hooks on every body sent upstream; cap live reply instances Token-count and compaction requests reached the upstream with the client's original body: the request hook ran only for generating calls, so a plugin that removed sensitive content from requests did not remove it from /v1/messages/count_tokens, :countTokens or /responses/compact. Request hooks now run on every request body sent upstream, with the same placeholder bridge, the same scope (client, sent model, upstream) and the same on_error as for generating calls: - Token counting (Anthropic count_tokens, Gemini :countTokens, Responses input_tokens) and Responses compaction (/responses/compact and Codex's /backend-api/codex/responses/compact) carry a conversation. Plugins see the normal view; params other than the model are not written back, since these endpoints do not take sampling parameters. A Gemini count wrapped in generateContentRequest is edited inside; a bare one is wrapped when a plugin adds a system prompt or tools. The post-plugin body is recorded like any other. - Counts core answers locally (a different-format upstream, Bedrock's 501) send nothing upstream and run no plugin. - Bodies plugins cannot read (embeddings, legacy completions, paths core does not recognize) follow each in-scope plugin's on_error: reject refuses the request with gw.plugin.cannot_read, skip passes it through unchanged and records a skipped run. Empty bodies are left alone. WebSocket connections that are not the Responses WebSocket are judged the same way at the upgrade. - Trial runs read stored requests the same way. Each streamed reply with a reply plugin holds one sandbox instance (up to 64 MiB) for the whole answer, so memory grew with concurrent streams. At most MAX_LIVE_REPLIES (32) reply instances are now alive at once. A slot is taken before an instance starts and travels with it, so it comes back when the answer ends (Chain::finish now drops the instances), when a plugin fails and is removed, when the chain is dropped (client gone, request cancelled) and when a call panics on a plugin thread. When no slot is free the plugin's on_error decides: reject fails the request with gw.plugin.reply_busy (403), skip lets the answer through unchanged. Both are recorded as plugin errors. Co-Authored-By: Claude Opus 5.5 --- crates/tw-api/msg-codes.txt | 2 + crates/tw-gateway/src/client_api.rs | 36 ++ crates/tw-gateway/src/plugin/mod.rs | 3 + crates/tw-gateway/src/plugin/pool.rs | 73 ++- crates/tw-gateway/src/plugin/reply/mod.rs | 112 ++++- .../tw-gateway/src/plugin/reply/tests/mod.rs | 260 ++++++++++ crates/tw-gateway/src/plugin/request.rs | 208 +++++++- crates/tw-gateway/src/plugin/trial.rs | 47 +- .../tw-gateway/src/plugin/trial/tests/mod.rs | 66 +++ crates/tw-gateway/src/server/pipeline/plug.rs | 21 +- crates/tw-gateway/src/server/upgrade.rs | 60 ++- crates/tw-gateway/src/ws.rs | 4 + .../tw-gateway/tests/plugins_reply_slots.rs | 475 ++++++++++++++++++ crates/tw-gateway/tests/plugins_request.rs | 426 +++++++++++++++- crates/tw-gateway/tests/plugins_security.rs | 75 ++- crates/tw-gateway/tests/plugins_ws.rs | 104 ++++ 16 files changed, 1861 insertions(+), 111 deletions(-) create mode 100644 crates/tw-gateway/tests/plugins_reply_slots.rs diff --git a/crates/tw-api/msg-codes.txt b/crates/tw-api/msg-codes.txt index 7b6ec496..ba956260 100644 --- a/crates/tw-api/msg-codes.txt +++ b/crates/tw-api/msg-codes.txt @@ -282,6 +282,7 @@ gw.output_limit.withheld gw.plugin.answer_unreadable gw.plugin.api gw.plugin.bad_output +gw.plugin.cannot_read gw.plugin.changed gw.plugin.cpu_limit gw.plugin.engine @@ -296,6 +297,7 @@ gw.plugin.output_limit gw.plugin.permission_violation gw.plugin.reason passthrough gw.plugin.rejected +gw.plugin.reply_busy gw.plugin.reply_failed gw.plugin.request_failed gw.plugin.request_unreadable diff --git a/crates/tw-gateway/src/client_api.rs b/crates/tw-gateway/src/client_api.rs index 740c24a0..7c76fbd1 100644 --- a/crates/tw-gateway/src/client_api.rs +++ b/crates/tw-gateway/src/client_api.rs @@ -92,6 +92,21 @@ impl ClientApi { || (p.contains("/models/") && p.ends_with(":countTokens")) } + /// 请求体和生成回答**同一种形状**、回来的却不是一次回答的接口:数 token(见 + /// [`ClientApi::counts_tokens`]),Responses 的压缩和数 token(`/responses/compact`、 + /// `/responses/input_tokens`,Codex 后端的 `/backend-api/codex/responses/compact`)。 + /// + /// 插件的请求钩子照样看得懂它们(见 [`crate::plugin::request::Shape`]) + pub fn like_generation(path: &str) -> bool { + if Self::counts_tokens(path) { + return true; + } + let p = path.trim_end_matches('/'); + let tail = p.strip_prefix("/v1").unwrap_or(p); + matches!(tail, "/responses/compact" | "/responses/input_tokens") + || p == "/backend-api/codex/responses/compact" + } + /// 转换库里对应的格式 pub fn dialect(&self) -> Dialect { match self { @@ -252,6 +267,27 @@ mod tests { } } + #[test] + fn counting_and_compacting_take_a_generation_shaped_body() { + for (path, like) in [ + ("/v1/messages/count_tokens", true), + ("/v1beta/models/gemini-2.5-pro:countTokens", true), + ("/v1/responses/compact", true), + ("/responses/compact/", true), + ("/v1/responses/input_tokens", true), + ("/backend-api/codex/responses/compact", true), + // 生成回答本身不算:它是「生成」那一类 + ("/v1/messages", false), + ("/v1/responses", false), + ("/v1/embeddings", false), + ("/v1/completions", false), + ("/v1/responses/resp_1/cancel", false), + ("/v1beta/models/gemini-embedding-001:embedContent", false), + ] { + assert_eq!(ClientApi::like_generation(path), like, "{path}"); + } + } + #[test] fn a_path_we_do_not_know_is_not_guessed() { for path in ["/v1/models", "/v1/files", "/healthz", "/v1/messagesx", "/"] { diff --git a/crates/tw-gateway/src/plugin/mod.rs b/crates/tw-gateway/src/plugin/mod.rs index 7e685015..2007ab22 100644 --- a/crates/tw-gateway/src/plugin/mod.rs +++ b/crates/tw-gateway/src/plugin/mod.rs @@ -22,9 +22,12 @@ //! - [`request`]:请求钩子。排在路由之后,**每发往一个上游跑一次**(契约附录二的 I7、 //! I8):按这一次的客户端、发出去的模型和上游挑插件,从客户端的原话起改;换上游从 //! 原话重来,同一家重发不重跑。改过的请求再过一遍内容审查,然后才转换格式、脱敏。 +//! **发往上游的每个请求体都过**:数 token、Responses 的压缩也改,插件看不懂的接口按 +//! 插件的 `on_error` 处置(见 [`request::Shape`])。 //! - [`reply`]:回答钩子。排在格式转换之后、工具调用审查和输出长度之前(I7)—— //! 这两道防护看的就是插件改过的那一版。 //! - [`pool`]:插件调用都是阻塞的、吃 CPU 的,放在专用线程池上跑,不占 tokio 的线程。 +//! 同时活着的回答实例也在这里限数([`pool::MAX_LIVE_REPLIES`])。 //! - [`trial`]:对着存下来的请求和回答试跑一个插件。 //! //! [`defaults`] 是随 core 一起发的那几个插件(清单和源码);[`manifests`] 是编过的插件的 diff --git a/crates/tw-gateway/src/plugin/pool.rs b/crates/tw-gateway/src/plugin/pool.rs index 7b241f82..6dcec636 100644 --- a/crates/tw-gateway/src/plugin/pool.rs +++ b/crates/tw-gateway/src/plugin/pool.rs @@ -1,4 +1,4 @@ -//! 跑插件的专用线程池。 +//! 跑插件的专用线程池,和回答实例的名额。 //! //! 插件调用是阻塞的、吃 CPU 的(一次请求钩子最多跑两百毫秒)。放在 tokio 的工作线程 //! 上调,几个慢插件就能把整个数据面的线程占满 —— 那时连不走插件的请求也一起卡住。 @@ -6,10 +6,24 @@ //! 不占着 tokio 的线程。 //! //! 线程**第一次用到时才起**:绝大多数用户一个插件都没装,不该为它多几根闲着的线程。 +//! +//! # 回答实例的名额 +//! +//! 回答钩子的实例从回答开始活到回答结束(不变式 I3),一个最多占 +//! `tw_plugin::Limits::reply_memory`(64 MiB)。流式的回答一开就是几分钟,同时在流的 +//! 回答越多,活着的实例就越多 —— 不设上限的话,内存跟着并发的流一起涨。所以整个进程 +//! 同时活着的回答实例最多 [`MAX_LIVE_REPLIES`] 个:起实例之前先拿一个名额 +//! ([`Pool::reply_slot`])。**拿不到不等**:回答已经到了,等一个名额就是让客户端干等, +//! 满了就按那个插件的 `on_error` 处置 —— 拒绝是这个请求失败,跳过是这次回答绕过它。 +//! 名额跟着实例走([`Slot`]),实例扔掉的那一刻还回来。 use std::sync::mpsc; use std::sync::{Arc, Mutex, OnceLock}; +/// 同时活着的回答实例最多这么多个(见模块说明)。按一个实例 64 MiB 的上限算,最坏 +/// 2 GiB;实际的插件一个实例多半不到 1 MiB +pub const MAX_LIVE_REPLIES: usize = 32; + type Job = Box; /// 池子坏了:任务 panic 了,或者线程起不来。 @@ -32,24 +46,64 @@ pub struct Pool { threads: usize, /// 正在跑的加排着的,最多这么多 permits: Arc, + /// 回答实例的名额(见模块说明) + replies: Arc, + /// 名额一共几个 + reply_cap: usize, tx: OnceLock>, String>>, } +/// 一个回答实例占着的名额(见 [`Pool::reply_slot`])。**和实例放在一起**:实例在哪儿被 +/// 扔掉 —— 回答收尾、插件出错被拿掉、客户端走了、插件线程上 panic —— 名额就在哪儿还回去 +#[derive(Debug)] +pub struct Slot { + _permit: tokio::sync::OwnedSemaphorePermit, +} + /// 一根线程的栈。**跑沙箱的线程至少要 2 MiB**(wasm 自己最多用 1 MiB,外面还有宿主 /// 那一侧的调用),这里明着给足,不靠平台的默认值(Windows 上只有 1 MiB) pub(crate) const STACK: usize = 8 * 1024 * 1024; impl Pool { - /// `threads` 根线程,最多 `queue` 个任务在跑或者排着。 + /// `threads` 根线程,最多 `queue` 个任务在跑或者排着;回答实例的名额是 + /// [`MAX_LIVE_REPLIES`] 个。 pub fn new(threads: usize, queue: usize) -> Self { + Self::with_replies(threads, queue, MAX_LIVE_REPLIES) + } + + /// 同 [`Pool::new`],回答实例的名额是 `replies` 个。**测试用**:拿一个小数,不用开几十 + /// 条流就能把名额占满 + pub fn with_replies(threads: usize, queue: usize, replies: usize) -> Self { let threads = threads.max(1); Self { threads, permits: Arc::new(tokio::sync::Semaphore::new(queue.max(threads))), + replies: Arc::new(tokio::sync::Semaphore::new(replies)), + reply_cap: replies, tx: OnceLock::new(), } } + /// 给一个回答实例拿一个名额。**满了是 `None`,不等**(理由见模块说明) + pub fn reply_slot(&self) -> Option { + self.replies + .clone() + .try_acquire_owned() + .ok() + .map(|p| Slot { _permit: p }) + } + + /// 回答实例的名额一共几个 + pub fn reply_cap(&self) -> usize { + self.reply_cap + } + + /// 此刻活着的回答实例(占着的名额) + pub fn live_replies(&self) -> usize { + self.reply_cap + .saturating_sub(self.replies.available_permits()) + } + /// 按机器的核数定:至少两根,最多八根;排队的是线程数的四倍。 pub fn default_size() -> Self { let cores = std::thread::available_parallelism().map_or(2, |n| n.get()); @@ -142,6 +196,21 @@ mod tests { assert_eq!(pool.run(|| 7).await, Ok(7)); } + #[test] + fn reply_slots_run_out_without_waiting_and_come_back_when_dropped() { + let pool = Pool::with_replies(1, 1, 2); + let a = pool.reply_slot().expect("the first slot"); + let b = pool.reply_slot().expect("the second slot"); + assert_eq!(pool.live_replies(), 2); + assert!(pool.reply_slot().is_none(), "a third slot past the cap"); + drop(a); + assert_eq!(pool.live_replies(), 1); + let c = pool.reply_slot().expect("the slot that came back"); + drop((b, c)); + assert_eq!(pool.live_replies(), 0); + assert_eq!(Pool::default_size().reply_cap(), MAX_LIVE_REPLIES); + } + #[tokio::test] async fn more_calls_than_threads_wait_their_turn() { let pool = Arc::new(Pool::new(2, 2)); diff --git a/crates/tw-gateway/src/plugin/reply/mod.rs b/crates/tw-gateway/src/plugin/reply/mod.rs index f6abd962..a332ba74 100644 --- a/crates/tw-gateway/src/plugin/reply/mod.rs +++ b/crates/tw-gateway/src/plugin/reply/mod.rs @@ -13,6 +13,11 @@ //! 回答开始时给每个范围内的插件起一个实例([`Chain::start`]),这次回答的文字和工具 //! 调用都交给它,回答结束就扔掉(约定 I3)。 //! +//! 每个实例占一个名额,整个进程同时活着的回答实例有上限(见 [`super::pool`])。名额满了, +//! 这个插件这次回答不起实例,按它的 `on_error`:拒绝就是这个请求失败,跳过就是这次回答 +//! 绕过它。名额和实例放在一起,回答收尾([`Chain::finish`])、插件出错被拿掉、客户端走了 +//! (整条链被扔掉)时跟着实例一起还回去。 +//! //! - **文字**按块交:整块模式攒齐一块再交一次,交回来的才发给客户端;流式模式每段 //! 增量交一次,交回什么现在就发什么(空串是先扣着),块结束时调 `onReplyTextEnd` //! 把扣着的补上。几个插件串起来,前一个交出的是后一个收到的。 @@ -36,8 +41,8 @@ use tw_dialect::ir::Dialect; use tw_types::{Msg, msg}; use super::bridge::Bridge; -use super::host::{Invocation, ReplyHost, RunError, ToolCallOutcome}; -use super::pool::Pool; +use super::host::{Invocation, PluginHost, ReplyHost, RunError, ToolCallOutcome}; +use super::pool::{Pool, Slot}; use super::set::{Active, LogLine, PluginRun, PluginSet}; use crate::error::GatewayError; @@ -82,8 +87,8 @@ struct Stage { text: bool, text_end: bool, tools: bool, - /// 出错之后被拿掉了(跳过)的是 None - instance: Option>, + /// 出错之后被拿掉了(跳过)的、回答收尾了的是 None + instance: Option, lanes: HashMap, cpu: Duration, counts: Counts, @@ -92,6 +97,14 @@ struct Stage { logs: Vec, } +/// 一个插件在这次回答里的实例,连同它占着的名额:**一起扔掉,一起还回去**。调用时整个 +/// 交给插件线程、跑完再交回来,所以在插件线程上没了(panic、调用方已经走了)的实例, +/// 名额也跟着还 +struct Instance { + host: Box, + _slot: Slot, +} + #[derive(Default)] struct StageLane { /// 整块模式:攒着的这一块。流式模式:为了不把半截密钥交给插件先扣着的尾巴 @@ -132,11 +145,56 @@ fn failed(name: &str, detail: &Msg) -> GatewayError { )) } +/// 名额满了:这个插件这次回答没起实例(见 [`super::pool`])。记在这次运行上,拒绝时 +/// 也是报给客户端的那一句 +fn busy(plugin: &str, max: usize) -> Msg { + msg!( + "gw.plugin.reply_busy", plugin = plugin, max = max => + "Plugin `{plugin}` was not started for this answer: the limit of {max} plugins running \ + on answers at the same time was reached." + ) +} + +/// 一个插件这次回答没起来。 +enum NotStarted { + /// 名额满了。记的、拒绝时报给客户端的都是 [`busy`] 那一句 + Busy(Msg), + /// 起实例出错了。记的是这个错误,拒绝时报给客户端的是 [`failed`] 那一句 + Failed(Msg), +} + +impl NotStarted { + /// 记在这次运行上的那一句 + fn why(&self) -> &Msg { + match self { + NotStarted::Busy(m) | NotStarted::Failed(m) => m, + } + } +} + +/// 拿一个名额、在插件线程上起这个插件的实例。**名额交给插件线程上的那一步**:调用方 +/// 半路走了,实例照样起完,和名额一起扔掉 +async fn instantiate( + pool: &Pool, + host: Arc, + name: &str, + ctx: Value, +) -> Result { + let Some(slot) = pool.reply_slot() else { + return Err(NotStarted::Busy(busy(name, pool.reply_cap()))); + }; + pool.run(move || host.reply(ctx).map(|host| Instance { host, _slot: slot })) + .await + .map_err(|e| RunError::Trap(e.to_string())) + .and_then(|r| r) + .map_err(|e| NotStarted::Failed(e.msg())) +} + fn stage( active: Option>, m: &super::engine::Manifest, on_error: OnError, - instance: Box, + instance: Instance, ) -> Stage { Stage { name: active @@ -161,8 +219,8 @@ impl Chain { /// 给这次回答起插件实例。范围内一个回答钩子都没有时是 `None` —— 这次回答原样走, /// 不付任何代价。 /// - /// 起实例失败按 `on_error`:拒绝就是这个错误(这时一个字节都还没发给客户端), - /// 跳过就不要它。 + /// 起实例失败、名额满了(见 [`super::pool`])按 `on_error`:拒绝就是这个错误(这时 + /// 一个字节都还没发给客户端),跳过就不要它。两样都和别的插件错误一样记一笔。 pub async fn start( state: &crate::AppState, set: &PluginSet, @@ -184,28 +242,21 @@ impl Chain { ctx.upstream, &a.settings, ); - let made = state - .plugin_pool - .run(move || host.reply(c)) - .await - .map_err(|e| RunError::Trap(e.to_string())) - .and_then(|r| r); - match made { + match instantiate(&state.plugin_pool, host, &a.name, c).await { Ok(instance) => stages.push(stage(Some(a.clone()), &m, a.on_error, instance)), - Err(err) => { - let why = err.msg(); + Err(not) => { let run = PluginRun { plugin_id: a.id.clone(), plugin_name: a.name.clone(), hook: PluginHook::Reply, outcome: PluginOutcome::Error, - error: Some(why.clone()), + error: Some(not.why().clone()), cpu_us: 0, detail: Some(json!({ "attempt": ctx.attempt })), }; state.plugin_ran(ctx.request_id, &a, run, Vec::new()); if a.on_error == OnError::Reject { - // 已经起好的那几个也记一笔(一次都没调用过) + // 已经起好的那几个也记一笔(一次都没调用过),名额还回去 let mut started = Chain { pool: state.plugin_pool.clone(), state: Some(state.clone()), @@ -217,7 +268,10 @@ impl Chain { recorded: false, }; started.finish(); - return Err(failed(&a.name, &why)); + return Err(match not { + NotStarted::Busy(why) => GatewayError::denied(why), + NotStarted::Failed(why) => failed(&a.name, &why), + }); } } } @@ -238,10 +292,10 @@ impl Chain { } /// 一条试跑用的链:只有这一个插件,日志收下来交给调用方,不进统计、日志圈和记录。 - /// 起不来是那个错误 + /// 起不来(名额满了也算)是那个错误 pub(crate) async fn trial( pool: Arc, - host: Arc, + host: Arc, settings: &serde_json::Map, ctx: &ReplyCtx<'_>, ) -> Result, Msg> { @@ -257,11 +311,9 @@ impl Chain { ctx.upstream, settings, ); - let instance = pool - .run(move || host.reply(c)) + let instance = instantiate(&pool, host, &m.name, c) .await - .map_err(|e| RunError::Trap(e.to_string()).msg())? - .map_err(|e| e.msg())?; + .map_err(|not| not.why().clone())?; Ok(Some(Chain { pool, state: None, @@ -305,7 +357,7 @@ impl Chain { let ran = self .pool .run(move || { - let inv = f(inst.as_mut()); + let inv = f(inst.host.as_mut()); (inst, inv) }) .await; @@ -621,8 +673,14 @@ impl Chain { Ok(out) } - /// 回答结束了(或者断了):每个插件一条记录,改了几处写在 `detail` 里。只记一次 + /// 回答结束了(或者断了):每个插件一条记录,改了几处写在 `detail` 里。只记一次。 + /// + /// **实例这时就扔掉**,名额还回去:调用方还攥着这条链(流还要补一段收尾、整包还要过 + /// 一遍审查)的那一会儿,实例已经用不上了 pub fn finish(&mut self) { + for s in &mut self.stages { + s.instance = None; + } if self.recorded { return; } diff --git a/crates/tw-gateway/src/plugin/reply/tests/mod.rs b/crates/tw-gateway/src/plugin/reply/tests/mod.rs index b08eae95..2552e44c 100644 --- a/crates/tw-gateway/src/plugin/reply/tests/mod.rs +++ b/crates/tw-gateway/src/plugin/reply/tests/mod.rs @@ -930,3 +930,263 @@ async fn reply_runs_are_counted_on_the_plugins() { assert_eq!((v.calls, v.changed, v.errors), (1, 1, 0), "{}", a.id); } } + +// ───────────────────────────────────────────────────────── 名额 + +/// 回答实例的名额只有 `n` 个的网关 +fn state_with_slots(n: usize) -> crate::AppState { + let mut s = state(); + s.plugin_pool = Arc::new(Pool::with_replies(2, 8, n)); + s +} + +/// 一份插件,每个出错时怎么办各自给 +fn set_on_error(doubles: Vec<(Double, OnError)>) -> PluginSet { + PluginSet::new( + doubles + .into_iter() + .enumerate() + .map(|(i, (d, on_error))| { + let mut a = crate::plugin::host::double::active(&format!("p{i}"), d); + a.on_error = on_error; + Arc::new(a) + }) + .collect(), + ) +} + +async fn start(state: &crate::AppState, set: &PluginSet) -> Result, GatewayError> { + Chain::start( + state, + set, + Bridge::new(Arc::new(tw_guard::redact::rules::RuleSet::none())), + &ReplyCtx { + dialect: Dialect::Anthropic, + client: None, + model: "m", + requested_model: "m", + upstream: "u", + request_id: 7, + attempt: 0, + }, + ) + .await +} + +fn live(state: &crate::AppState) -> usize { + state.plugin_pool.live_replies() +} + +/// 一个实例的名额从回答开始占到回答结束;**收尾时就还**,不等调用方扔掉这条流 +#[tokio::test] +async fn a_slot_is_held_for_the_whole_answer_and_returned_when_it_ends() { + let state = state_with_slots(4); + let set = set_of(vec![upper(), hold_until_end()]); + let mut s = Stream::new( + start(&state, &set).await.unwrap().expect("in scope"), + Framing::Sse, + ); + assert_eq!(live(&state), 2, "one slot per instance"); + let input = anthropic_stream(); + let (half, rest) = input.split_at(input.len() / 2); + let (_, err) = s.feed(half.as_bytes()).await; + assert!(err.is_none()); + assert_eq!(live(&state), 2, "the answer is still streaming"); + let (_, err) = s.feed(rest.as_bytes()).await; + assert!(err.is_none()); + let (_, err) = s.finish(false).await; + assert!(err.is_none()); + assert_eq!( + live(&state), + 0, + "the answer ended and the slots stayed taken" + ); + drop(s); + + // 整包:改完那一份就还 + let mut chain = start(&state, &set).await.unwrap().unwrap(); + assert_eq!(live(&state), 2); + let body = json!({"id":"m","type":"message","role":"assistant", + "content":[{"type":"text","text":"hello"}],"stop_reason":"end_turn"}); + let out = whole(&mut chain, body.to_string().as_bytes()) + .await + .unwrap(); + assert!(String::from_utf8_lossy(&out).contains("HELLO")); + assert_eq!(live(&state), 0); +} + +/// 名额满了:**不等**,按这个插件的 `on_error` —— 拒绝是这个请求失败(码是 +/// `gw.plugin.reply_busy`),跳过是这次回答绕过它。两样都记成这个插件的一次出错 +#[tokio::test] +async fn a_full_house_turns_plugins_away_by_their_on_error() { + let state = state_with_slots(1); + let holder = set_of(vec![upper()]); + let first = start(&state, &holder).await.unwrap().expect("in scope"); + assert_eq!(live(&state), 1); + + let mut events = state.bus.subscribe(); + let rejecting = set_on_error(vec![(upper(), OnError::Reject)]); + let Err(err) = start(&state, &rejecting).await else { + panic!("started past the cap"); + }; + assert_eq!(err.detail.code, "gw.plugin.reply_busy"); + assert_eq!(err.source, crate::error::Source::Denied); + assert_eq!( + err.detail.text, + "Plugin `upper` was not started for this answer: the limit of 1 plugins running on \ + answers at the same time was reached." + ); + let skipping = set_on_error(vec![(upper(), OnError::Skip)]); + assert!( + start(&state, &skipping).await.unwrap().is_none(), + "the only plugin was skipped: the answer goes through as it is" + ); + for set in [&rejecting, &skipping] { + let a = &set.all()[0]; + let v = a.stats.view(); + assert_eq!((v.calls, v.errors), (1, 1)); + assert_eq!( + v.last_error.map(|e| e.message.code).as_deref(), + Some("gw.plugin.reply_busy") + ); + } + // 和别的插件错误一样发一条通知 + let mut failed = 0; + while let Ok(ev) = events.try_recv() { + if let tw_api::Event::PluginFailed { message, .. } = ev { + assert_eq!(message.code, "gw.plugin.reply_busy"); + failed += 1; + } + } + assert_eq!(failed, 2); + + // 名额还回来,下一个回答照常起 + drop(first); + assert_eq!(live(&state), 0); + let again = start(&state, &skipping).await.unwrap(); + assert!(again.is_some()); + drop(again); + + // 前一个插件拿到了名额、后一个没拿到而策略是拒绝:已经起好的那个也还回去 + let both = set_on_error(vec![(upper(), OnError::Skip), (upper(), OnError::Reject)]); + assert!(start(&state, &both).await.is_err()); + assert_eq!(live(&state), 0, "the first plugin's slot was not returned"); +} + +/// 回答半路被扔掉(客户端走了、请求被取消):名额跟着实例一起还 +#[tokio::test] +async fn an_answer_dropped_halfway_returns_its_slots() { + let state = state_with_slots(4); + let set = set_of(vec![upper(), hold_until_end()]); + let mut s = Stream::new(start(&state, &set).await.unwrap().unwrap(), Framing::Sse); + let input = anthropic_stream(); + let _ = s.feed(&input.as_bytes()[..input.len() / 2]).await; + assert_eq!(live(&state), 2); + drop(s); + assert_eq!(live(&state), 0); + + // 一次都没用过就被扔掉的也一样 + let chain = start(&state, &set).await.unwrap().unwrap(); + assert_eq!(live(&state), 2); + drop(chain); + assert_eq!(live(&state), 0); +} + +/// 插件出错:被拿掉的那一刻就还它的名额(跳过时回答还在接着流);拒绝时这条流收尾就全还 +#[tokio::test] +async fn a_plugin_that_errors_out_returns_its_slot() { + let boom = || { + Double::new("boom") + .permit(&[Permission::ReplyText]) + .on_reply(true, false, false, |_| { + Ok(Box::new(Closures { + text: Box::new(|_| Invocation::err(RunError::CpuLimit)), + end: Box::new(|| Invocation::ok(None)), + tool: Box::new(|_| Invocation::ok(ToolCallOutcome::Unchanged)), + })) + }) + }; + let state = state_with_slots(4); + let input = anthropic_stream(); + // 跳过:出错的那个插件被拿掉,另一个照常跑到回答结束 + let set = set_on_error(vec![ + (boom().mode(ReplyMode::Stream), OnError::Skip), + (upper(), OnError::Reject), + ]); + let mut s = Stream::new(start(&state, &set).await.unwrap().unwrap(), Framing::Sse); + assert_eq!(live(&state), 2); + let cut = input.find("lo 世界").unwrap(); + let (_, err) = s.feed(&input.as_bytes()[..cut]).await; + assert!(err.is_none()); + assert_eq!(live(&state), 1, "the failed plugin kept its slot"); + let (_, err) = s.feed(&input.as_bytes()[cut..]).await; + assert!(err.is_none()); + let _ = s.finish(false).await; + assert_eq!(live(&state), 0); + + // 拒绝:这条流切断,中继收尾(断了的那一种)时全还 + let set = set_on_error(vec![(boom(), OnError::Reject), (upper(), OnError::Reject)]); + let mut s = Stream::new(start(&state, &set).await.unwrap().unwrap(), Framing::Sse); + let (_, err) = s.feed(input.as_bytes()).await; + assert_eq!(err.expect("rejects").detail.code, "gw.plugin.reply_failed"); + assert_eq!(live(&state), 1); + let _ = s.finish(true).await; + assert_eq!(live(&state), 0); +} + +/// 插件线程上 panic 了:实例跟着没了,名额也跟着还 +#[tokio::test] +async fn a_call_that_panics_returns_its_slot() { + let state = state_with_slots(4); + let panics = Double::new("panics") + .permit(&[Permission::ReplyText]) + .on_reply(true, false, false, |_| { + Ok(Box::new(Closures { + text: Box::new(|_| panic!("the plugin host fell over")), + end: Box::new(|| Invocation::ok(None)), + tool: Box::new(|_| Invocation::ok(ToolCallOutcome::Unchanged)), + })) + }); + let set = set_on_error(vec![(panics, OnError::Skip)]); + let mut s = Stream::new(start(&state, &set).await.unwrap().unwrap(), Framing::Sse); + assert_eq!(live(&state), 1); + let (out, err) = run(&mut s, &anthropic_stream(), 4096).await; + assert!(err.is_none()); + assert_eq!( + anthropic_blocks(&out)[1].2, + "hello 世界", + "skipped: as it was" + ); + assert_eq!(live(&state), 0); +} + +/// 实例起不来:名额马上还,记的是那个错误,拒绝时报的是「插件出错」而不是「名额满了」 +#[tokio::test] +async fn an_instance_that_fails_to_start_returns_its_slot() { + let state = state_with_slots(1); + let broken = Double::new("broken") + .permit(&[Permission::ReplyText]) + .on_reply(true, false, false, |_| Err(RunError::MemoryLimit)); + let set = set_on_error(vec![(broken, OnError::Reject)]); + let Err(err) = start(&state, &set).await else { + panic!("started a plugin whose instance failed"); + }; + assert_eq!(err.detail.code, "gw.plugin.reply_failed"); + assert_eq!(live(&state), 0); + assert_eq!( + set.all()[0] + .stats + .view() + .last_error + .map(|e| e.message.code) + .as_deref(), + Some("gw.plugin.memory_limit") + ); + // 名额还在:下一个插件照常起 + assert!( + start(&state, &set_of(vec![upper()])) + .await + .unwrap() + .is_some() + ); +} diff --git a/crates/tw-gateway/src/plugin/request.rs b/crates/tw-gateway/src/plugin/request.rs index 1637a4a7..2463f7a1 100644 --- a/crates/tw-gateway/src/plugin/request.rs +++ b/crates/tw-gateway/src/plugin/request.rs @@ -23,6 +23,18 @@ //! 插件 `reject` 了,或者出错而它的 `on_error` 是拒绝,**整个请求被拒**,不换下一家: //! 换一家,管它的还是这个插件。文件变了、装不上的插件跑不了,管得着这一次的同样按 //! `on_error` 处理;只管别的上游、别的模型的,这一次不算它。 +//! +//! # 哪些请求过插件 +//! +//! **发往上游的每一个请求体都过**,不只生成回答的那些 —— 插件删掉的东西,不能从旁边的 +//! 接口漏出去(见 [`Shape`]): +//! +//! - 生成回答:上面说的那样; +//! - 数 token(Anthropic 的 `count_tokens`、Gemini 的 `:countTokens`、Responses 的 +//! `input_tokens`)、Responses 的压缩:请求体就是一段对话,插件照样看、照样改,上游数的、 +//! 压的是改过的那一份。网关自己估数、一个字节都不发的那几种不跑插件; +//! - 嵌入、旧版补全、认不出的接口:插件看不懂它们的请求体。管得着的插件按它的 +//! `on_error`:拒绝就拒掉整个请求,跳过就原样发、记一笔跳过。 use std::borrow::Cow; use std::sync::Arc; @@ -42,6 +54,36 @@ use super::view; /// 一个插件在这一次上的运行,连同它写的日志。 pub type Ran = (Arc, PluginRun, Vec); +/// 插件怎么看一个请求体:按客户端调的接口分(见 [`crate::client_api::ClientApi`])。 +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum Shape { + /// 生成回答 + Generate, + /// 请求体和生成回答同一种形状、却不生成回答的接口:数 token、Responses 的压缩(见 + /// [`crate::client_api::ClientApi::like_generation`])。插件照样看、照样改,**`params` + /// 里只写回模型名**:这些接口不收输出上限、温度这些参数(带上是一个 400),数出来的 + /// token 也和它们无关 + Alike, + /// 嵌入、旧版补全、认不出的接口:**插件看不懂这种请求体**(见 [`unreadable`]) + Opaque, +} + +impl Shape { + /// 客户端调的这个路径是哪一种 + pub fn of(path: &str) -> Shape { + use crate::client_api::ClientApi; + if ClientApi::of_path(path).is_none() { + Shape::Opaque + } else if ClientApi::generates(path) { + Shape::Generate + } else if ClientApi::like_generation(path) { + Shape::Alike + } else { + Shape::Opaque + } + } +} + /// 一个请求上的请求钩子。管线每发往一个上游调一次 [`Hook::attempt`],**每次都从客户端 /// 的原话起**;原文只解析一次、密钥只编一次号,几次尝试共用。 pub struct Hook<'a> { @@ -51,6 +93,8 @@ pub struct Hook<'a> { dialect: Dialect, /// 客户端调的路径(Gemini 的模型在里面) path: &'a str, + /// 按路径分出来的那一种 + shape: Shape, /// 客户端是哪个应用(请求那一行上记的那个,认不出是 `None`) client: Option<&'a str>, /// 客户端发来的原文 @@ -122,6 +166,7 @@ impl<'a> Hook<'a> { rules, dialect, path, + shape: Shape::of(path), client, body, parsed: None, @@ -150,18 +195,27 @@ impl<'a> Hook<'a> { } let here = self.set.for_request(self.client, to.model, to.upstream); if here.is_empty() { - // 回答钩子要这个请求的密钥映射:管这一次的里面有,就现在记账 - if !self - .set - .for_reply(self.client, to.model, to.upstream) - .is_empty() + // 回答钩子要这个请求的密钥映射:管这一次的里面有,就现在记账。不生成回答的 + // 请求没有回答钩子可跑 + if self.shape == Shape::Generate + && !self + .set + .for_reply(self.client, to.model, to.upstream) + .is_empty() { out.bridge = Some(self.base()); } return Ok(out); } + if self.shape == Shape::Opaque { + // 空的请求体里没有插件能改的东西 + if !self.body.iter().all(u8::is_ascii_whitespace) { + out.runs = unreadable(self.set, self.client, self.path, to)?; + } + return Ok(out); + } let mut bridge = self.base(); - let (body, dialect, client) = (self.body, self.dialect, self.client); + let (body, dialect, client, shape) = (self.body, self.dialect, self.client, self.shape); let original = self .parsed .get_or_insert_with(|| serde_json::from_slice::(body).ok()) @@ -176,15 +230,7 @@ impl<'a> Hook<'a> { let host = match &a.state { super::set::State::Broken(why) => { let (outcome, refusal) = broken(a.on_error, &a.name, why); - let run = PluginRun { - plugin_id: a.id.clone(), - plugin_name: a.name.clone(), - hook: PluginHook::Request, - outcome, - error: Some(broken_reason(&a.name, why)), - cpu_us: 0, - detail: Some(json!({ "attempt": to.attempt })), - }; + let run = not_run(&a, outcome, broken_reason(&a.name, why), to.attempt); out.runs.push((a.clone(), run, Vec::new())); if let Some(why) = refusal { return Err(Box::new(Refused { @@ -210,8 +256,9 @@ impl<'a> Hook<'a> { let Some(current) = raw.as_deref() else { return Err(Failure::Unreadable("the request body is not JSON".into())); }; + let conversation = wrapped_count(dialect, &path, current).unwrap_or(current); let mut built = - view::build(dialect, current, &path).map_err(Failure::Unreadable)?; + view::build(dialect, conversation, &path).map_err(Failure::Unreadable)?; sending(&mut built.view, &model); let mut input = view::trim(&built.view, &a.permissions); bridge.hide_value(&mut input); @@ -241,6 +288,9 @@ impl<'a> Hook<'a> { built.src.hidden_tools(), ) .map_err(Failure::Edit)?; + if shape == Shape::Alike { + model_only(&mut edits); + } if edits.is_empty() { return Ok(None); } @@ -248,8 +298,9 @@ impl<'a> Hook<'a> { edits.reveal(&bridge); let new_model = edits.params.as_ref().and_then(|p| p.model.clone()); let mut next = current.clone(); - let new_path = view::apply(&mut next, &built.src, &edits, &path) - .map_err(Failure::Edit)?; + let new_path = + write_back(dialect, &mut next, &built.src, &edits, &path, &model) + .map_err(Failure::Edit)?; Ok(Some(Rewritten { value: next, path: new_path, @@ -326,6 +377,127 @@ impl<'a> Hook<'a> { } } +/// 插件看不懂、却要发往上游的东西:[`Shape::Opaque`] 的请求体,不是 Responses 的 +/// WebSocket 连接上的帧。管这一次的插件一个都跑不了 —— 跑不了的插件(文件变了、装不上) +/// 照旧,能跑的按它的 `on_error`:拒绝就拒掉整个请求,跳过就原样发、记一笔跳过。`path` +/// 是客户端调的路径,报出来的就是它 +pub fn unreadable( + set: &PluginSet, + client: Option<&str>, + path: &str, + to: &Target<'_>, +) -> Result, Box> { + let mut runs = Vec::new(); + for a in set.for_request(client, to.model, to.upstream) { + let (outcome, error, refusal) = match &a.state { + super::set::State::Broken(why) => { + let (outcome, refusal) = broken(a.on_error, &a.name, why); + (outcome, broken_reason(&a.name, why), refusal) + } + super::set::State::Ready(_) => { + let why = cannot_read(&a.name, path); + match a.on_error { + OnError::Skip => (PluginOutcome::Skipped, why, None), + OnError::Reject => (PluginOutcome::Error, why.clone(), Some(why)), + } + } + }; + runs.push(( + a.clone(), + not_run(&a, outcome, error, to.attempt), + Vec::new(), + )); + if let Some(why) = refusal { + return Err(Box::new(Refused { why, runs })); + } + } + Ok(runs) +} + +/// 没跑的一次:跑不了的插件,看不懂的请求 +fn not_run(a: &Active, outcome: PluginOutcome, error: Msg, attempt: usize) -> PluginRun { + PluginRun { + plugin_id: a.id.clone(), + plugin_name: a.name.clone(), + hook: PluginHook::Request, + outcome, + error: Some(error), + cpu_us: 0, + detail: Some(json!({ "attempt": attempt })), + } +} + +/// 插件看不懂这个接口的请求体(见 [`unreadable`]) +pub(super) fn cannot_read(plugin: &str, path: &str) -> Msg { + msg!( + "gw.plugin.cannot_read", plugin = plugin, path = path => + "Plugin `{plugin}` cannot read requests to {path}." + ) +} + +/// Gemini 数 token 的请求 +fn gemini_count(dialect: Dialect, path: &str) -> bool { + dialect == Dialect::Gemini && path.trim_end_matches('/').ends_with(":countTokens") +} + +/// Gemini 数 token 的请求体有两种写法:`contents` 直接放在外面,或者整个生成请求包在 +/// `generateContentRequest` 里。**插件看的、改的都是那一份生成请求**:包着的是里面那一份 +pub(super) fn wrapped_count<'v>(dialect: Dialect, path: &str, raw: &'v Value) -> Option<&'v Value> { + if !gemini_count(dialect, path) { + return None; + } + raw.get("generateContentRequest").filter(|v| v.is_object()) +} + +/// 把插件的改动写回原文(见 [`view::apply`]),返回新的路径(Gemini 换了模型时)。 +/// +/// Gemini 数 token 的请求体照它原来的写法写:包着的写回里面那一份;没包着的,插件加了 +/// 系统提示、工具就包起来 —— 外面那一层只收 `contents`,系统提示和工具要放进 +/// `generateContentRequest` 才数得进去,放在外面上游回 400。包着的那一份里也写着模型名 +/// (`models/…`),插件换了模型名就跟着换,和路径上的对得上。`model` 是发给这一家的模型名 +pub(super) fn write_back( + dialect: Dialect, + raw: &mut Value, + src: &view::Src, + edits: &view::Edits, + path: &str, + model: &str, +) -> Result, view::EditError> { + if !gemini_count(dialect, path) { + return view::apply(raw, src, edits, path); + } + let renamed = edits.params.as_ref().and_then(|p| p.model.as_deref()); + let wrapped = raw + .get("generateContentRequest") + .is_some_and(Value::is_object); + let new_path = match raw.get_mut("generateContentRequest") { + Some(inner) if wrapped => view::apply(inner, src, edits, path)?, + _ => view::apply(raw, src, edits, path)?, + }; + if !wrapped && (edits.system.is_some() || edits.tools.is_some()) { + *raw = json!({ "generateContentRequest": std::mem::take(raw) }); + } + if let Some(inner) = raw + .get_mut("generateContentRequest") + .and_then(Value::as_object_mut) + && (renamed.is_some() || !wrapped) + { + let model = renamed.unwrap_or(model); + inner.insert("model".into(), json!(format!("models/{model}"))); + } + Ok(new_path) +} + +/// 不生成回答的接口([`Shape::Alike`]):`params` 里只留模型名,别的改动不写回 +pub(super) fn model_only(edits: &mut view::Edits) { + if let Some(p) = edits.params.as_mut() { + *p = view::ParamsEdit { + model: p.model.take(), + ..Default::default() + }; + } +} + /// 视图里的模型名换成发给这一家的那个(`model` 和 `params.model`):插件看到的就是要发出去 /// 的,和 `ctx.model` 一致。原文里写的是客户端要的那个,路由规则的改写在格式转换那一步才 /// 落到请求体上 diff --git a/crates/tw-gateway/src/plugin/trial.rs b/crates/tw-gateway/src/plugin/trial.rs index f9ec8da6..3ca36aa1 100644 --- a/crates/tw-gateway/src/plugin/trial.rs +++ b/crates/tw-gateway/src/plugin/trial.rs @@ -21,7 +21,7 @@ use tw_types::{Msg, msg}; use super::bridge::Bridge; use super::host::{PluginHost, RequestOutcome}; use super::pool::Pool; -use super::request::{rejected, request_unreadable}; +use super::request::{Shape, rejected, request_unreadable}; use super::set::LogLine; use super::view; @@ -143,14 +143,20 @@ async fn tried( let client = request.as_ref().and_then(|r| r.client); let name = host.manifest().name.clone(); - // ── 请求钩子 + // ── 请求钩子:和这个请求当时一样看(见 [`super::request::Shape`]) if let (Some(r), Some(d), true) = (&request, dialect, host.manifest().hooks.request) { + let shape = Shape::of(r.path); match parsed.as_ref() { + _ if shape == Shape::Opaque => { + t.error = Some(super::request::cannot_read(&name, r.path)) + } None => t.error = Some(request_unreadable("the request body is not JSON")), Some(raw) => { let mut masked = raw.clone(); bridge.hide_value(&mut masked); - match view::build(d, &masked, r.path) { + let conversation = + super::request::wrapped_count(d, r.path, &masked).unwrap_or(&masked); + match view::build(d, conversation, r.path) { Err(e) => t.error = Some(request_unreadable(e)), Ok(mut built) => { let m = host.manifest(); @@ -190,22 +196,27 @@ async fn tried( Err(e) => { (Outcome::Error, before.clone(), Some(e.msg())) } - Ok(edits) if edits.is_empty() => { - (Outcome::Unchanged, before.clone(), None) - } - Ok(edits) => { - let mut next = masked.clone(); - match view::apply( - &mut next, &built.src, &edits, r.path, - ) { - Ok(_) => { - (Outcome::Changed, pretty(&next), None) + Ok(mut edits) => { + if shape == Shape::Alike { + super::request::model_only(&mut edits); + } + if edits.is_empty() { + (Outcome::Unchanged, before.clone(), None) + } else { + let mut next = masked.clone(); + match super::request::write_back( + d, &mut next, &built.src, &edits, r.path, + &model, + ) { + Ok(_) => { + (Outcome::Changed, pretty(&next), None) + } + Err(e) => ( + Outcome::Error, + before.clone(), + Some(e.msg()), + ), } - Err(e) => ( - Outcome::Error, - before.clone(), - Some(e.msg()), - ), } } } diff --git a/crates/tw-gateway/src/plugin/trial/tests/mod.rs b/crates/tw-gateway/src/plugin/trial/tests/mod.rs index 75b3663f..1d59c3cc 100644 --- a/crates/tw-gateway/src/plugin/trial/tests/mod.rs +++ b/crates/tw-gateway/src/plugin/trial/tests/mod.rs @@ -221,3 +221,69 @@ async fn a_trial_gives_the_plugin_the_routing_the_request_had() { assert_eq!(saw["view"]["model"], "glm-4.6"); assert_eq!(saw["view"]["params"]["model"], "glm-4.6"); } + +/// 存下来的一个请求:路径和请求体,没记客户端、没记发出去的模型名 +fn stored<'a>(path: &'a str, body: &'a [u8]) -> StoredRequest<'a> { + StoredRequest { + path, + query: None, + body, + client: None, + upstream: "up", + sent_model: "", + } +} + +/// 试跑数 token、嵌入这些请求,和它们当时一样看:Gemini 包着的数 token 改里面那一份、 +/// 只写回模型名;插件看不懂的请求体说清看不懂 +#[tokio::test] +async fn a_trial_reads_counting_and_unreadable_requests_as_they_were_read() { + let tune = || { + Double::new("tune") + .permit(&[Permission::System, Permission::Params]) + .on_request(|mut view, _| { + view["system"] = json!("Be brief."); + view["params"]["max_tokens"] = json!(99); + Invocation::ok(RequestOutcome::Changed(view)) + }) + .into_host() + }; + let wrapped = json!({ "generateContentRequest": { "model": "models/gemini-2.5-pro", + "contents": [{ "role": "user", "parts": [{ "text": "hi" }] }] } }) + .to_string(); + let t = run( + Arc::new(Pool::new(1, 4)), + tune(), + &Default::default(), + rules(), + Some(stored( + "/v1beta/models/gemini-2.5-pro:countTokens", + wrapped.as_bytes(), + )), + None, + ) + .await; + assert_eq!(t.error, None); + let req = t.request.unwrap(); + assert_eq!(req.outcome, Outcome::Changed); + let after: Value = serde_json::from_str(&req.after).unwrap(); + let inner = &after["generateContentRequest"]; + assert_eq!(inner["systemInstruction"]["parts"][0]["text"], "Be brief."); + assert!(inner.get("generationConfig").is_none(), "{after}"); + + let embeddings = br#"{"model":"text-embedding-3-small","input":["hi"]}"#; + let t = run( + Arc::new(Pool::new(1, 4)), + tune(), + &Default::default(), + rules(), + Some(stored("/v1/embeddings", embeddings)), + None, + ) + .await; + assert!(t.request.is_none()); + assert_eq!( + t.error.map(|m| m.code).as_deref(), + Some("gw.plugin.cannot_read") + ); +} diff --git a/crates/tw-gateway/src/server/pipeline/plug.rs b/crates/tw-gateway/src/server/pipeline/plug.rs index f0638858..a515ce44 100644 --- a/crates/tw-gateway/src/server/pipeline/plug.rs +++ b/crates/tw-gateway/src/server/pipeline/plug.rs @@ -14,6 +14,10 @@ //! //! 插件拒绝了、出错而策略是拒绝、或者上面哪一道没过,**整个请求被拒**,不换下一家: //! 换一家,管它的还是这些插件。 +//! +//! **发往上游的每一跳都过这一步**,不只生成回答的:数 token、Responses 的压缩一样过插件 +//! (插件删掉的东西不能从这些接口漏出去),插件看不懂的接口按插件的 `on_error` 处置(见 +//! [`crate::plugin::request::Shape`])。网关自己估数、不发出去的那一跳到不了这里。 use bytes::Bytes; @@ -27,7 +31,7 @@ pub(super) struct Rewritten { pub(super) body: Bytes, /// 调的路径(Gemini 换了模型时是新的) pub(super) path: String, - /// 改过的请求解码出来的中间表示:格式转换用它 + /// 改过的请求解码出来的中间表示:格式转换用它。不生成回答的请求不转换,没有 pub(super) decoded: Option>, /// 这一跳出站脱敏接着编号的账:拦截档下是插件那本账接着编的(插件写进来的新值有了 /// 新的号),别的档位是空的 @@ -61,10 +65,6 @@ pub(super) async fn attempt( model: &str, attempt: usize, ) -> Result { - // **只给生成回答的请求跑**:计 token、嵌入这些接口没有「一次回答」可言 - let Some(api) = req.api.filter(|_| reading.generates) else { - return Ok(Plugged::default()); - }; let to = crate::plugin::request::Target { upstream: &provider.name, model, @@ -89,10 +89,13 @@ pub(super) async fn attempt( if let Some(r) = &c.renamed { allowed(rt, req, r)?; } - let decoded = - tw_dialect::convert::decode(api.dialect(), &c.value, &c.path, req.query.as_deref()); + // 生成回答的请求重新解码:格式转换和请求防护用改过的这一份。数 token、压缩这些不转换 + // (只发给同格式的上游),开头也没过请求防护,不用解 + let decoded = req.api.filter(|_| reading.generates).map(|api| { + tw_dialect::convert::decode(api.dialect(), &c.value, &c.path, req.query.as_deref()) + }); // 请求防护:只看插件加进来的。解不开的不看 —— 和开头那一遍一样,同格式直通照样发 - if let (Some(Ok(before)), Ok(after)) = (&reading.decoded, &decoded) + if let (Some(Ok(before)), Some(Ok(after))) = (&reading.decoded, &decoded) && let Some(why) = crate::guard::screen_more( &state.bus, started.id, @@ -125,7 +128,7 @@ pub(super) async fn attempt( out.rewritten = Some(Rewritten { body: c.body, path: c.path, - decoded: Some(decoded), + decoded, ledger, }); Ok(out) diff --git a/crates/tw-gateway/src/server/upgrade.rs b/crates/tw-gateway/src/server/upgrade.rs index 4df66645..9b7e153d 100644 --- a/crates/tw-gateway/src/server/upgrade.rs +++ b/crates/tw-gateway/src/server/upgrade.rs @@ -135,6 +135,60 @@ pub(super) async fn ws_upgrade( .await .map_err(|e| GatewayError::config(crate::state::credential_failed(e, &name)))?; let (id, ending) = open(&choice, &name, provider.billing.into()); + // 插件:升级那一刻的那一份表,一条连接用到底。**插件只看得懂 Responses 的 WebSocket** + // (每个 `response.create` 是一次请求);别的路径上的帧插件看不懂,管得着的插件按它的 + // `on_error` —— 拒绝就不接这条连接,跳过就记一笔、这条连接不过插件 + let hint = crate::hint::client_hint(&headers); + let plugins = if rt.plugins.is_empty() { + None + } else if crate::client_api::ClientApi::of_path(uri.path()) + == Some(crate::client_api::ClientApi::OpenaiResponses) + && crate::client_api::ClientApi::generates(uri.path()) + { + Some(crate::ws::Plugins { + pool: state.plugin_pool.clone(), + set: rt.plugins.clone(), + client: hint.clone(), + }) + } else { + let to = crate::plugin::request::Target { + upstream: &name, + // 升级请求没有正文:说不出是哪个模型 + model: "", + requested_model: "", + attempt: 0, + }; + let hop_started = std::time::Instant::now(); + match crate::plugin::request::unreadable(&rt.plugins, hint.as_deref(), uri.path(), &to) { + Ok(runs) => { + crate::plugin::request::record(&state, id, &runs); + None + } + Err(refused) => { + crate::plugin::request::record(&state, id, &refused.runs); + let err = GatewayError::denied(refused.why); + // 和 HTTP 那条路一样:没发出去的这一跳在尝试链上,原因就是拒绝它的那句话 + state.bus.emit(tw_api::Event::RequestRouted { + id, + route: choice.route, + rule: choice.rule, + group: choice.group, + rewritten_by: Vec::new(), + denied_by: None, + affinity: None, + attempts: vec![crate::server::hop_failed( + &name, + None, + err.detail.clone(), + hop_started, + )], + billing: tw_api::Billing::PerToken, + }); + ending.failed(err.source.into(), err.detail.clone()); + return Err(err); + } + } + }; let upstream = crate::ws::Upstream { url: crate::ws::upstream_url(&provider.base_url, uri.path(), query.as_deref()), headers: upstream_headers, @@ -152,12 +206,6 @@ pub(super) async fn ws_upgrade( limit_mode: rt.config.security.output_limit.mode, limit: rt.config.security.output_limit.limit(), }; - // 插件:升级那一刻的那一份表,一条连接用到底 - let plugins = (!rt.plugins.is_empty()).then(|| crate::ws::Plugins { - pool: state.plugin_pool.clone(), - set: rt.plugins.clone(), - client: crate::hint::client_hint(&headers), - }); Ok(ws.on_upgrade(move |sock| async move { // 一条 WS 连接活多久,这个请求就算在服务中多久 let _live = live; diff --git a/crates/tw-gateway/src/ws.rs b/crates/tw-gateway/src/ws.rs index cd92326c..4332945b 100644 --- a/crates/tw-gateway/src/ws.rs +++ b/crates/tw-gateway/src/ws.rs @@ -26,6 +26,10 @@ //! (`response.created` 到 `response.completed`)起一组回答钩子的实例,排在占位符 //! 还原之后、工具墙之前。插件出错而策略是拒绝时,切掉的是那一次回答,连接照常。 //! +//! **插件只看得懂 Responses 的 WebSocket。**别的路径上(比如 Realtime 的 `/v1/realtime`) +//! 的帧插件看不懂,升级时就按管得着的插件的 `on_error` 处置:拒绝就不接这条连接,跳过就 +//! 记一笔、这条连接不过插件(见 `server::upgrade`、[`crate::plugin::request::unreadable`])。 +//! //! # 两条明说的边界 //! //! **一、走代理的上游不代理 WS。**代理是给 reqwest 配的,而这里 diff --git a/crates/tw-gateway/tests/plugins_reply_slots.rs b/crates/tw-gateway/tests/plugins_reply_slots.rs new file mode 100644 index 00000000..917dd155 --- /dev/null +++ b/crates/tw-gateway/tests/plugins_reply_slots.rs @@ -0,0 +1,475 @@ +//! 回答实例的名额,从客户端到假上游走一整圈。 +//! +//! 一个回答钩子的实例从回答开始活到回答结束,整个进程同时活着的有上限(见 +//! `tw_gateway::plugin::pool`)。这里证明两件事:名额满了按插件的 `on_error` 处置(拒绝是 +//! 这个请求失败,跳过是这次回答原样过去);名额**一定还得回来** —— 回答结束、客户端半路 +//! 走了、上游半路断了、WebSocket 上一次回答完了或者连接断了。 + +use std::net::SocketAddr; +use std::sync::{Arc, Mutex}; +use std::time::Duration; + +use axum::Router; +use axum::extract::ws::{Message, WebSocketUpgrade}; +use futures::{SinkExt, StreamExt}; +use serde_json::{Value, json}; +use tokio::sync::watch; +use tw_api::{OnError, Permission, ReplyMode}; +use tw_config::{Client, Config, Listen, Protocol, Provider}; +use tw_gateway::plugin::host::double::{self, Double}; +use tw_gateway::plugin::pool::Pool; +use tw_gateway::plugin::{Active, PluginSet, RunRecord}; + +fn ev(v: Value) -> String { + format!("event: {}\ndata: {v}\n\n", v["type"].as_str().unwrap()) +} + +/// 一个流式回答的开头:到第一段文字为止 +fn head() -> 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 "}}), + ), + ] + .concat() +} + +/// 剩下的:第二段文字和收尾 +fn tail() -> String { + [ + ev( + json!({"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"there"}}), + ), + ev(json!({"type":"content_block_stop","index":0})), + ev( + json!({"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{"output_tokens":5}}), + ), + ev(json!({"type":"message_stop"})), + ] + .concat() +} + +/// 假上游:先吐开头,等 `release` 变成 true 再吐完 —— 一个还在说话的模型 +async fn held(release: watch::Receiver) -> SocketAddr { + let app = Router::new().fallback(axum::routing::post(move || { + let mut rx = release.clone(); + async move { + let s = async_stream::stream! { + yield Ok::<_, std::io::Error>(bytes::Bytes::from(head())); + while !*rx.borrow_and_update() { + if rx.changed().await.is_err() { + break; + } + } + yield Ok(bytes::Bytes::from(tail())); + }; + axum::response::Response::builder() + .header("content-type", "text/event-stream") + .body(axum::body::Body::from_stream(s)) + .unwrap() + } + })); + listen(app).await +} + +/// 假上游:吐完开头就把连接掐断(说好的长度没发完) +async fn breaking() -> SocketAddr { + let l = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = l.local_addr().unwrap(); + tokio::spawn(async move { + loop { + let Ok((mut s, _)) = l.accept().await else { + return; + }; + tokio::spawn(async move { + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + let mut buf = vec![0u8; 64 * 1024]; + let _ = s.read(&mut buf).await; + let body = head(); + let resp = format!( + "HTTP/1.1 200 OK\r\ncontent-type: text/event-stream\r\ncontent-length: {}\r\n\r\n{body}", + body.len() + 900 + ); + let _ = s.write_all(resp.as_bytes()).await; + let _ = s.flush().await; + tokio::time::sleep(Duration::from_millis(100)).await; + drop(s); + }); + } + }); + addr +} + +async fn listen(app: Router) -> SocketAddr { + 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 +} + +struct Gw { + addr: SocketAddr, + state: tw_gateway::AppState, + runs: Arc>>, +} + +impl Gw { + fn live(&self) -> usize { + self.state.plugin_pool.live_replies() + } + + /// 等活着的回答实例回到 `n` 个。等不到就失败 + async fn settles_at(&self, n: usize) { + for _ in 0..100 { + if self.live() == n { + return; + } + tokio::time::sleep(Duration::from_millis(50)).await; + } + panic!("{} reply instances are still alive, not {n}", self.live()); + } + + /// 这个插件每次出错时记下的消息码 + fn error_codes(&self, id: &str) -> Vec { + self.runs + .lock() + .unwrap() + .iter() + .filter(|r| r.run.plugin_id == id) + .filter_map(|r| r.run.error.as_ref().map(|m| m.code.clone())) + .collect() + } +} + +/// 网关:一家 `protocol` 的上游,回答实例的名额是 `slots` 个 +async fn gateway( + base: SocketAddr, + protocol: Protocol, + entries: Vec>, + slots: usize, +) -> Gw { + 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: "up".into(), + base_url: format!("http://{base}"), + key: Some("sk-upstream".into()), + protocol: Some(protocol), + ..Default::default() + }], + ..Default::default() + }; + let mut state = tw_gateway::AppState::new(cfg).unwrap(); + state.plugin_pool = Arc::new(Pool::with_replies(2, 16, slots)); + state.swap_plugins(PluginSet::new(entries)); + let runs: Arc>> = Arc::default(); + let (tx, mut rx) = tokio::sync::mpsc::channel(tw_gateway::plugin::RUN_CHANNEL_CAP); + state.set_plugin_sink(tx); + let r = runs.clone(); + tokio::spawn(async move { + while let Some(rec) = rx.recv().await { + r.lock().unwrap().push(rec); + } + }); + let addr = tw_gateway::serve(state.clone(), ([127, 0, 0, 1], 0).into()) + .await + .unwrap(); + tokio::time::sleep(Duration::from_millis(30)).await; + Gw { addr, state, runs } +} + +/// 逐段改写的插件:每段文字一到就大写交出去,回答还没说完时客户端就看得到它改过的 +fn upper(on_error: OnError) -> Arc { + let d = Double::new("upper") + .permit(&[Permission::ReplyText]) + .mode(ReplyMode::Stream) + .on_text(|t| Some(t.to_uppercase())); + let mut a = double::active("upper", d); + a.on_error = on_error; + Arc::new(a) +} + +/// 发一个流式请求 +async fn send(gw: &Gw) -> reqwest::Response { + let body = json!({ "model": "claude-sonnet-4-5", "max_tokens": 64, "stream": true, + "messages": [{ "role": "user", "content": "hi" }] }); + reqwest::Client::new() + .post(format!("http://{}/v1/messages", gw.addr)) + .header("x-api-key", "tw-testkey") + .header("content-type", "application/json") + .body(body.to_string()) + .send() + .await + .unwrap() +} + +/// 一条流式回答读到第一段文字为止:这时这个回答的实例已经活着了 +async fn open( + gw: &Gw, +) -> ( + impl futures::Stream> + Unpin, + String, +) { + let r = send(gw).await; + assert_eq!(r.status(), 200); + let mut s = Box::pin(r.bytes_stream()); + let mut got = Vec::new(); + while let Some(c) = s.next().await { + got.extend_from_slice(&c.unwrap()); + if String::from_utf8_lossy(&got).contains("text_delta") { + break; + } + } + (s, String::from_utf8_lossy(&got).into_owned()) +} + +/// 读完剩下的 +async fn rest( + mut s: impl futures::Stream> + Unpin, + mut got: String, +) -> String { + while let Some(c) = s.next().await { + got.push_str(&String::from_utf8_lossy(&c.unwrap())); + } + got +} + +/// 流里的全部文字 +fn text(body: &str) -> String { + body.lines() + .filter_map(|l| l.strip_prefix("data: ")) + .filter_map(|d| serde_json::from_str::(d).ok()) + .filter_map(|v| v["delta"]["text"].as_str().map(str::to_string)) + .collect() +} + +/// 名额满了而策略是拒绝:这个请求失败,码是 `gw.plugin.reply_busy`。占着名额的那个回答 +/// 说完,名额还回来,下一个回答照常有插件 +#[tokio::test] +async fn a_full_house_fails_the_request_under_reject_and_a_finished_answer_hands_its_slot_on() { + let (release, rx) = watch::channel(false); + let gw = gateway( + held(rx).await, + Protocol::Anthropic, + vec![upper(OnError::Reject)], + 1, + ) + .await; + let (first, got) = open(&gw).await; + assert_eq!(text(&got), "HELLO "); + assert_eq!(gw.live(), 1); + + let r = send(&gw).await; + assert_eq!(r.status(), 403); + assert_eq!( + r.headers() + .get("x-thinkwatch-error") + .and_then(|v| v.to_str().ok()), + Some("denied") + ); + let body: Value = r.json().await.unwrap(); + assert_eq!( + body["error"]["message"], + "[ThinkWatch] Plugin `upper` was not started for this answer: the limit of 1 plugins \ + running on answers at the same time was reached." + ); + assert_eq!(gw.live(), 1, "the refused request took a slot"); + + release.send_replace(true); + let got = rest(first, got).await; + assert_eq!(text(&got), "HELLO THERE"); + gw.settles_at(0).await; + // 名额回来了:下一个回答照常有插件 + let (s, got) = open(&gw).await; + let got = rest(s, got).await; + assert_eq!(text(&got), "HELLO THERE"); + gw.settles_at(0).await; + tokio::time::sleep(Duration::from_millis(100)).await; + assert_eq!(gw.error_codes("upper"), ["gw.plugin.reply_busy"]); +} + +/// 名额满了而策略是跳过:这次回答原样过去,记一次出错 +#[tokio::test] +async fn a_full_house_lets_the_answer_through_untouched_under_skip() { + let (release, rx) = watch::channel(false); + let gw = gateway( + held(rx).await, + Protocol::Anthropic, + vec![upper(OnError::Skip)], + 1, + ) + .await; + let (first, got_first) = open(&gw).await; + let (second, got_second) = open(&gw).await; + assert_eq!(gw.live(), 1); + release.send_replace(true); + assert_eq!(text(&rest(first, got_first).await), "HELLO THERE"); + assert_eq!( + text(&rest(second, got_second).await), + "hello there", + "the answer past the cap was not passed through as it was" + ); + gw.settles_at(0).await; + tokio::time::sleep(Duration::from_millis(100)).await; + assert_eq!(gw.error_codes("upper"), ["gw.plugin.reply_busy"]); +} + +/// 客户端半路走了:名额跟着这次回答一起还回来,不等上游说完 +#[tokio::test] +async fn a_client_that_walks_away_returns_the_slot() { + let (_release, rx) = watch::channel(false); + let gw = gateway( + held(rx).await, + Protocol::Anthropic, + vec![upper(OnError::Reject)], + 4, + ) + .await; + let (s, got) = open(&gw).await; + assert_eq!(text(&got), "HELLO "); + assert_eq!(gw.live(), 1); + drop(s); + gw.settles_at(0).await; +} + +/// 上游半路断了:这次回答以错误收尾,名额还回来 +#[tokio::test] +async fn an_upstream_that_breaks_off_returns_the_slot() { + let gw = gateway( + breaking().await, + Protocol::Anthropic, + vec![upper(OnError::Reject)], + 4, + ) + .await; + let r = send(&gw).await; + assert_eq!(r.status(), 200); + let body = r.text().await.unwrap_or_default(); + assert_eq!(text(&body), "HELLO "); + assert!(body.contains("event: error"), "{body}"); + gw.settles_at(0).await; +} + +// ───────────────────────────────────────────────────────── WebSocket + +/// WebSocket 假上游:每个 `response.create` 先回 created 和一段文字,等 `release` 变成 +/// true 再回完 +async fn held_ws(release: watch::Receiver) -> SocketAddr { + let app = Router::new().route( + "/backend-api/codex/responses", + axum::routing::any(move |ws: WebSocketUpgrade| { + let mut rx = release.clone(); + async move { + ws.on_upgrade(move |mut sock| async move { + while let Some(Ok(m)) = sock.recv().await { + let Message::Text(_) = m else { continue }; + let frames = [ + json!({"type":"response.created","response":{"id":"resp_1","status":"in_progress","output":[]}}), + json!({"type":"response.output_item.added","output_index":0,"item":{"type":"message","id":"msg","role":"assistant","content":[]}}), + json!({"type":"response.output_text.delta","output_index":0,"content_index":0,"item_id":"msg","delta":"hel"}), + ]; + for f in frames { + if sock.send(Message::Text(f.to_string().into())).await.is_err() { + return; + } + } + while !*rx.borrow_and_update() { + if rx.changed().await.is_err() { + return; + } + } + let frames = [ + json!({"type":"response.output_text.delta","output_index":0,"content_index":0,"item_id":"msg","delta":"lo"}), + json!({"type":"response.output_text.done","output_index":0,"content_index":0,"item_id":"msg","text":"hello"}), + json!({"type":"response.output_item.done","output_index":0,"item":{"type":"message","id":"msg","role":"assistant","content":[{"type":"output_text","text":"hello"}]}}), + json!({"type":"response.completed","response":{"id":"resp_1","status":"completed","output":[{"type":"message","id":"msg","role":"assistant","content":[{"type":"output_text","text":"hello"}]}]}}), + ]; + for f in frames { + if sock.send(Message::Text(f.to_string().into())).await.is_err() { + return; + } + } + } + }) + } + }), + ); + listen(app).await +} + +type WsStream = + tokio_tungstenite::WebSocketStream>; + +/// 连上网关、发一帧 `response.create`,读到第一段文字为止 +async fn ws_open(gw: &Gw) -> WsStream { + use tokio_tungstenite::tungstenite::client::IntoClientRequest; + let mut req = format!("ws://{}/backend-api/codex/responses", gw.addr) + .into_client_request() + .unwrap(); + req.headers_mut() + .insert("x-api-key", "tw-testkey".parse().unwrap()); + let (mut sock, _) = tokio_tungstenite::connect_async(req).await.unwrap(); + let frame = json!({ "type": "response.create", "model": "gpt-5", "input": "hi" }); + sock.send(tokio_tungstenite::tungstenite::Message::Text( + frame.to_string().into(), + )) + .await + .unwrap(); + loop { + let m = tokio::time::timeout(Duration::from_secs(5), sock.next()) + .await + .expect("no first delta") + .unwrap() + .unwrap(); + if m.into_text().unwrap().contains("output_text.delta") { + return sock; + } + } +} + +/// WebSocket 上一次回答一组实例:**这次回答完了就还**(连接还开着),连接半路断了也还 +#[tokio::test] +async fn a_websocket_answer_returns_its_slot_when_it_completes_and_when_the_client_leaves() { + let (release, rx) = watch::channel(false); + let gw = gateway( + held_ws(rx).await, + Protocol::OpenaiResponses, + vec![upper(OnError::Reject)], + 4, + ) + .await; + // 半路走掉 + let sock = ws_open(&gw).await; + assert_eq!(gw.live(), 1); + drop(sock); + gw.settles_at(0).await; + + // 说完:连接还开着,名额已经回来了 + let mut sock = ws_open(&gw).await; + assert_eq!(gw.live(), 1); + release.send_replace(true); + loop { + let m = tokio::time::timeout(Duration::from_secs(5), sock.next()) + .await + .expect("the answer did not complete") + .unwrap() + .unwrap(); + if m.into_text().unwrap().contains("response.completed") { + break; + } + } + gw.settles_at(0).await; + drop(sock); +} diff --git a/crates/tw-gateway/tests/plugins_request.rs b/crates/tw-gateway/tests/plugins_request.rs index 6800b1a9..9b3c9146 100644 --- a/crates/tw-gateway/tests/plugins_request.rs +++ b/crates/tw-gateway/tests/plugins_request.rs @@ -670,7 +670,7 @@ async fn an_inactive_plugin_follows_on_error_without_running() { } #[tokio::test] -async fn out_of_scope_plugins_do_not_run_and_count_tokens_is_left_alone() { +async fn out_of_scope_plugins_do_not_run_and_count_tokens_runs_the_ones_in_scope() { let up = Upstream::default(); let base = start_upstream(up.clone()).await; let calls = Arc::new(AtomicUsize::new(0)); @@ -702,7 +702,7 @@ async fn out_of_scope_plugins_do_not_run_and_count_tokens_is_left_alone() { assert_eq!(status, 200); assert_eq!(calls.load(Ordering::SeqCst), 0); - // 数 token 不是一次回答:范围内的插件也不跑 + // 数 token 发往上游的也是一个请求体:范围内的插件照样跑,范围外的照样不跑 let everyone = entry("everyone", counting(calls.clone())); let gw = gateway( vec![provider("a", base, Protocol::Anthropic)], @@ -710,12 +710,12 @@ async fn out_of_scope_plugins_do_not_run_and_count_tokens_is_left_alone() { vec![everyone], ) .await; - let _ = post(&gw, "/v1/messages/count_tokens", &anthropic_body()).await; - assert_eq!(calls.load(Ordering::SeqCst), 0); - // 生成回答的请求照常跑 - let (status, _) = post(&gw, "/v1/messages", &anthropic_body()).await; + let (status, _) = post(&gw, "/v1/messages/count_tokens", &anthropic_body()).await; assert_eq!(status, 200); assert_eq!(calls.load(Ordering::SeqCst), 1); + let (status, _) = post(&gw, "/v1/messages", &anthropic_body()).await; + assert_eq!(status, 200); + assert_eq!(calls.load(Ordering::SeqCst), 2); } /// 把这一次的上游写进系统提示的插件,数自己跑了几次 @@ -1310,3 +1310,417 @@ async fn each_client_format_is_rewritten_in_its_own_shape() { } } } + +// ───────────────────────────────────────── 不生成回答的接口 + +/// 插件要删掉的东西 +const MARK: &str = "SECRET-PROJECT"; + +/// 把 [`MARK`] 从系统提示、消息文字和工具结果里删掉的插件 +fn scrub() -> Double { + Double::new("scrub") + .permit(&[Permission::System, Permission::Messages]) + .on_request(|mut view, _| { + let s = view["system"].as_str().unwrap().replace(MARK, "[removed]"); + view["system"] = json!(s); + for m in view["messages"].as_array_mut().unwrap() { + for p in m["parts"].as_array_mut().unwrap() { + if (p["type"] == "text" || p["type"] == "tool_result") + && let Some(t) = p["text"].as_str() + { + p["text"] = json!(t.replace(MARK, "[removed]")); + } + } + } + Invocation::ok(RequestOutcome::Changed(view)) + }) +} + +fn anthropic_count_body() -> Value { + json!({ + "model": "claude-sonnet-4-5", + "system": format!("About {MARK}."), + "messages": [ + { "role": "user", "content": format!("Plan {MARK}") }, + { "role": "assistant", "content": [{ "type": "tool_use", "id": "t1", "name": "Read", "input": { "path": "a" } }] }, + { "role": "user", "content": [{ "type": "tool_result", "tool_use_id": "t1", "content": format!("{MARK} notes") }] } + ] + }) +} + +fn responses_body() -> Value { + json!({ + "model": "gpt-5", + "instructions": format!("About {MARK}."), + "input": [{ "type": "message", "role": "user", "content": [{ "type": "input_text", "text": format!("Plan {MARK}") }] }] + }) +} + +const GEMINI_COUNT: &str = "/v1beta/models/gemini-2.5-pro:countTokens"; + +/// 数 token、Responses 的压缩:请求体就是一段对话,插件照样改,**上游数的、压的是改过的那 +/// 一份** —— 插件删掉的东西不从这些接口漏出去。每种客户端格式、Gemini 的两种写法都一样 +#[tokio::test] +async fn token_counts_and_compactions_reach_the_upstream_as_the_plugins_left_them() { + let cases = [ + ( + "/v1/messages/count_tokens", + Protocol::Anthropic, + anthropic_count_body(), + ), + ( + GEMINI_COUNT, + Protocol::Gemini, + json!({ "contents": [{ "role": "user", "parts": [{ "text": format!("Plan {MARK}") }] }] }), + ), + ( + GEMINI_COUNT, + Protocol::Gemini, + json!({ "generateContentRequest": { + "model": "models/gemini-2.5-pro", + "systemInstruction": { "parts": [{ "text": format!("About {MARK}.") }] }, + "contents": [{ "role": "user", "parts": [{ "text": format!("Plan {MARK}") }] }] + } }), + ), + ( + "/v1/responses/compact", + Protocol::OpenaiResponses, + responses_body(), + ), + ( + "/v1/responses/input_tokens", + Protocol::OpenaiResponses, + responses_body(), + ), + ( + "/backend-api/codex/responses/compact", + Protocol::OpenaiResponses, + responses_body(), + ), + ]; + for (path, protocol, body) in cases { + let up = Upstream::default(); + let base = start_upstream(up.clone()).await; + let mut gw = gateway( + vec![provider("same", base, protocol)], + SecurityMode::Off, + vec![entry("scrub", scrub())], + ) + .await; + let (status, answer) = post(&gw, path, &body).await; + assert_eq!(status, 200, "{path}: {answer}"); + assert_eq!(up.hits(), 1, "{path}"); + let (got_path, _) = up.seen.lock().unwrap()[0].clone(); + assert_eq!(got_path, path); + let raw = String::from_utf8(up.raw.lock().unwrap()[0].to_vec()).unwrap(); + assert!(!raw.contains(MARK), "{path}: the upstream got {raw}"); + assert!(raw.contains("[removed]"), "{path}: {raw}"); + // 和生成回答一样记在请求上:跑在第几跳、改了什么,改过的那一份另存 + assert_eq!( + gw.runs(), + [("scrub".to_string(), "changed".to_string(), 0)], + "{path}" + ); + let after = gw.after_plugins().await.expect("the body after plugins"); + assert!(!after.to_string().contains(MARK), "{path}: {after}"); + // 写法照原样:包着的还包着,没包着的没加系统提示也不包 + if path == GEMINI_COUNT { + let sent: Value = serde_json::from_str(&raw).unwrap(); + assert_eq!( + sent.get("generateContentRequest").is_some(), + body.get("generateContentRequest").is_some(), + "{sent}" + ); + } + } +} + +/// 数 token 上的插件和生成回答上的一样:同样按客户端、发出去的模型、上游挑,同样只看到 +/// 占位符,出错、`reject` 同样按 `on_error` 拒掉整个请求 +#[tokio::test] +async fn counting_follows_the_same_scope_placeholders_and_on_error() { + const COUNT: &str = "/v1/messages/count_tokens"; + let up = Upstream::default(); + let base = start_upstream(up.clone()).await; + let providers = || vec![provider("a", base, Protocol::Anthropic)]; + + // 范围外:别的模型、别的上游的插件不跑,上游收到的是原话 + for scoped in [ + entry_with("scrub", scrub(), |a| a.scope.models = vec!["gpt-*".into()]), + entry_with("scrub", scrub(), |a| a.scope.upstreams = vec!["b".into()]), + ] { + up.raw.lock().unwrap().clear(); + let gw = gateway(providers(), SecurityMode::Off, vec![scoped]).await; + let (status, _) = post(&gw, COUNT, &anthropic_count_body()).await; + assert_eq!(status, 200); + assert!(String::from_utf8_lossy(&up.raw.lock().unwrap()[0]).contains(MARK)); + assert!(gw.runs().is_empty(), "{:?}", gw.runs()); + } + + // 占位符:插件看不到真的密钥,改过的地方换回去再发 + let saw = Arc::new(Mutex::new(String::new())); + let s = saw.clone(); + let checked = Double::new("checked") + .permit(&[Permission::Messages]) + .on_request(move |mut view, ctx| { + *s.lock().unwrap() = view.to_string(); + assert_eq!(ctx["upstream"], "a"); + assert_eq!(ctx["model"], "claude-sonnet-4-5"); + let t = view["messages"][0]["parts"][0]["text"] + .as_str() + .unwrap() + .to_string(); + view["messages"][0]["parts"][0]["text"] = json!(format!("{t} (checked)")); + Invocation::ok(RequestOutcome::Changed(view)) + }); + up.raw.lock().unwrap().clear(); + let gw = gateway( + providers(), + SecurityMode::Off, + vec![entry("checked", checked)], + ) + .await; + let body = json!({ "model": "claude-sonnet-4-5", + "messages": [{ "role": "user", "content": format!("my key is {USER_KEY}") }] }); + let (status, _) = post(&gw, COUNT, &body).await; + assert_eq!(status, 200); + let saw = saw.lock().unwrap().clone(); + assert!(!saw.contains(USER_KEY), "the plugin saw the key: {saw}"); + assert!(saw.contains("<>"), "{saw}"); + let sent: Value = serde_json::from_slice(&up.raw.lock().unwrap()[0]).unwrap(); + assert_eq!( + sent["messages"][0]["content"], + format!("my key is {USER_KEY} (checked)") + ); + + // 出错、拒绝:拒绝时整个请求不发,跳过时原样发 + let failing = || { + Double::new("failing") + .permit(&[Permission::Messages]) + .on_request(|_, _| { + Invocation::err(RunError::Threw { + message: "nope".into(), + stack: None, + }) + }) + }; + let refusing = Double::new("refusing") + .permit(&[Permission::Messages]) + .on_request(|_, _| Invocation::ok(RequestOutcome::Rejected("not counted".into()))); + for (e, code) in [ + (entry("failing", failing()), "gw.plugin.request_failed"), + (entry("refusing", refusing), "gw.plugin.rejected"), + ] { + up.raw.lock().unwrap().clear(); + let gw = gateway(providers(), SecurityMode::Off, vec![e]).await; + let rx = gw.state.bus.subscribe(); + let (status, body) = post(&gw, COUNT, &anthropic_count_body()).await; + assert_eq!(status, 403, "{code}: {body}"); + assert!(up.raw.lock().unwrap().is_empty(), "{code}"); + let attempts = routed(rx).await; + assert_eq!( + attempts[0].error.as_ref().map(|e| e.code.as_str()), + Some(code) + ); + } + up.raw.lock().unwrap().clear(); + let gw = gateway( + providers(), + SecurityMode::Off, + vec![entry_with("failing", failing(), |a| { + a.on_error = OnError::Skip + })], + ) + .await; + let (status, _) = post(&gw, COUNT, &anthropic_count_body()).await; + assert_eq!(status, 200); + let sent: Value = serde_json::from_slice(&up.raw.lock().unwrap()[0]).unwrap(); + assert_eq!(sent, anthropic_count_body()); + assert_eq!(gw.stats("failing").errors, 1); +} + +/// 网关自己估的数:一个字节都不发给上游,也就不跑插件 +#[tokio::test] +async fn a_count_the_gateway_estimates_runs_no_plugin() { + let up = Upstream::default(); + let base = start_upstream(up.clone()).await; + let calls = Arc::new(AtomicUsize::new(0)); + let gw = gateway( + vec![provider("chat", base, Protocol::OpenaiChat)], + SecurityMode::Off, + vec![entry("tag", tag(calls.clone()))], + ) + .await; + let (status, answer) = post(&gw, "/v1/messages/count_tokens", &anthropic_count_body()).await; + assert_eq!(status, 200, "{answer}"); + assert!(answer["input_tokens"].as_u64().unwrap() > 0, "{answer}"); + assert_eq!(up.hits(), 0); + assert_eq!(calls.load(Ordering::SeqCst), 0); + assert!(gw.runs().is_empty()); +} + +/// 数 token、压缩不收输出上限、温度这些参数:插件改的只写回模型名。Gemini 没包着的数 +/// token 请求,插件加了系统提示就包起来(外面那一层只收 `contents`),模型名跟着插件改 +#[tokio::test] +async fn counting_and_compacting_take_only_the_model_from_params() { + let tune = Double::new("tune") + .permit(&[Permission::System, Permission::Params]) + .on_request(|mut view, ctx| { + let s = view["system"].as_str().unwrap().to_string(); + view["system"] = json!(format!("{s} Be brief.").trim().to_string()); + view["params"]["max_tokens"] = json!(99); + view["params"]["temperature"] = json!(0.1); + if ctx["format"] == "gemini" { + view["params"]["model"] = json!("gemini-2.5-flash"); + } else { + view["params"]["model"] = + json!(format!("{}-renamed", ctx["model"].as_str().unwrap())); + } + Invocation::ok(RequestOutcome::Changed(view)) + }); + let cases = [ + ( + "/v1/messages/count_tokens", + Protocol::Anthropic, + json!({ "model": "claude-sonnet-4-5", "messages": [{ "role": "user", "content": "hi" }] }), + ), + ( + "/v1/responses/compact", + Protocol::OpenaiResponses, + json!({ "model": "gpt-5", "input": "hi" }), + ), + ( + GEMINI_COUNT, + Protocol::Gemini, + json!({ "contents": [{ "role": "user", "parts": [{ "text": "hi" }] }] }), + ), + ( + GEMINI_COUNT, + Protocol::Gemini, + json!({ "generateContentRequest": { "model": "models/gemini-2.5-pro", + "contents": [{ "role": "user", "parts": [{ "text": "hi" }] }] } }), + ), + ]; + for (path, protocol, body) in cases { + let up = Upstream::default(); + let base = start_upstream(up.clone()).await; + let gw = gateway( + vec![provider("same", base, protocol)], + SecurityMode::Off, + vec![entry("tune", tune.clone())], + ) + .await; + let (status, answer) = post(&gw, path, &body).await; + assert_eq!(status, 200, "{path}: {answer}"); + let (got_path, sent) = up.seen.lock().unwrap()[0].clone(); + let text = sent.to_string(); + assert!(text.contains("Be brief."), "{path}: {text}"); + for param in [ + "max_tokens", + "max_output_tokens", + "maxOutputTokens", + "temperature", + "generationConfig", + ] { + assert!(!text.contains(param), "{path}: {param} was sent: {text}"); + } + match protocol { + Protocol::Gemini => { + assert_eq!(got_path, "/v1beta/models/gemini-2.5-flash:countTokens"); + let inner = &sent["generateContentRequest"]; + assert_eq!(inner["model"], "models/gemini-2.5-flash", "{text}"); + assert_eq!( + inner["systemInstruction"]["parts"][0]["text"], "Be brief.", + "{text}" + ); + assert_eq!(inner["contents"][0]["parts"][0]["text"], "hi", "{text}"); + assert!(sent.get("contents").is_none(), "{text}"); + } + _ => assert!( + sent["model"].as_str().unwrap().ends_with("-renamed"), + "{path}: {text}" + ), + } + } +} + +/// 插件看不懂的请求体(嵌入、认不出的接口):管得着的插件按它的 `on_error` —— 拒绝就不发, +/// 跳过就原样发、记一笔跳过。范围外的插件、空的请求体不算 +#[tokio::test] +async fn requests_plugins_cannot_read_follow_on_error() { + let up = Upstream::default(); + let base = start_upstream(up.clone()).await; + let embeddings = + json!({ "model": "text-embedding-3-small", "input": [format!("Plan {MARK}")] }); + let providers = || vec![provider("chat", base, Protocol::OpenaiChat)]; + + let mut gw = gateway( + providers(), + SecurityMode::Off, + vec![entry("scrub", scrub())], + ) + .await; + let (status, body) = post(&gw, "/v1/embeddings", &embeddings).await; + assert_eq!(status, 403, "{body}"); + assert_eq!( + body["error"]["message"], + "[ThinkWatch] Plugin `Plugin scrub` cannot read requests to /v1/embeddings." + ); + assert_eq!(up.hits(), 0); + assert_eq!(gw.runs(), [("scrub".to_string(), "error".to_string(), 0)]); + assert!(gw.after_plugins().await.is_none()); + // 认不出的接口也一样 + let (status, body) = post( + &gw, + "/v1/rerank", + &json!({ "model": "rerank-1", "query": MARK }), + ) + .await; + assert_eq!(status, 403, "{body}"); + assert_eq!(up.hits(), 0); + + // 跳过:原样发,记一笔跳过(不算一次调用) + let gw = gateway( + providers(), + SecurityMode::Off, + vec![entry_with("scrub", scrub(), |a| a.on_error = OnError::Skip)], + ) + .await; + let (status, _) = post(&gw, "/v1/embeddings", &embeddings).await; + assert_eq!(status, 200); + assert_eq!(up.hits(), 1); + assert_eq!(gw.runs(), [("scrub".to_string(), "skipped".to_string(), 0)]); + assert_eq!(gw.stats("scrub").calls, 0); + + // 范围外的插件不算 + let gw = gateway( + providers(), + SecurityMode::Off, + vec![entry_with("scrub", scrub(), |a| { + a.scope.models = vec!["claude-*".into()] + })], + ) + .await; + let (status, _) = post(&gw, "/v1/embeddings", &embeddings).await; + assert_eq!(status, 200); + assert!(gw.runs().is_empty()); + + // 空的请求体里没有插件能改的东西(取消一次 Responses 的回答) + let up = Upstream::default(); + let base = start_upstream(up.clone()).await; + let gw = gateway( + vec![provider("responses", base, Protocol::OpenaiResponses)], + SecurityMode::Off, + vec![entry("scrub", scrub())], + ) + .await; + let r = reqwest::Client::new() + .post(format!("http://{}/v1/responses/resp_1/cancel", gw.addr)) + .header("authorization", "Bearer tw-testkey") + .send() + .await + .unwrap(); + assert_eq!(r.status(), 200); + assert_eq!(up.hits(), 1); + assert!(gw.runs().is_empty(), "{:?}", gw.runs()); +} diff --git a/crates/tw-gateway/tests/plugins_security.rs b/crates/tw-gateway/tests/plugins_security.rs index 1e204e72..601df0a7 100644 --- a/crates/tw-gateway/tests/plugins_security.rs +++ b/crates/tw-gateway/tests/plugins_security.rs @@ -11,9 +11,11 @@ //! - I10:每次运行都有记录。 //! - 没有插件改动的请求一个字节都不变;WebSocket(Codex 的 Responses WebSocket)那一路 //! 同样看占位符、同样过工具调用审查、拒绝了不发给上游。 +//! - 发往上游的不只是生成回答:数 token、Responses 的压缩带着整段对话,同样过请求钩子, +//! 插件删掉的东西不从这些接口漏出去。 //! -//! 标了 `#[ignore]` 的两条是**还没解决的问题**,断言写的是该有的样子:插件写下的占位符会被 -//! 换回真值(契约 I5 的写法),计 token 的请求不经过请求钩子。 +//! 标了 `#[ignore]` 的那一条是**还没解决的问题**,断言写的是该有的样子:插件写下的占位符 +//! 会被换回真值(契约 I5 的写法)。 mod plugin_harness; @@ -1077,39 +1079,62 @@ export function onToolCall(call) { } } +/// Claude Code 每一轮都会先发一次 `count_tokens`,带着整段对话;Codex 压缩上下文时把整段 +/// 对话发给 `/responses/compact`;Gemini 的客户端数 token 走 `:countTokens`。**插件删掉的 +/// 东西不能从这些接口漏出去**:上游收到的是插件改过的那一份 #[tokio::test] -#[ignore = "design gap: request hooks run only for generating calls, so a token-count request \ - carries the client's text to the upstream without the plugin's rewrite; see the \ - track 4 report"] -async fn a_token_count_request_does_not_bypass_a_plugin_that_scrubs_the_prompt() { - // 一个把「机密」删掉的插件。Claude Code 每一轮都会先发一次 count_tokens,带着整段对话 +async fn a_token_count_or_compaction_does_not_bypass_a_plugin_that_scrubs_the_prompt() { let scrub = r#" -export const manifest = { name: "删掉机密", api: 1, permissions: ["messages"] }; +export const manifest = { name: "删掉机密", api: 1, permissions: ["system", "messages"] }; export function onRequest(req) { + req.system = req.system.replaceAll("机密", "[已删除]"); for (const m of req.messages) for (const p of m.parts) { - if (p.type === "text") p.text = p.text.replaceAll("机密", "[已删除]"); + if (p.type === "text" || p.type === "tool_result") p.text = p.text.replaceAll("机密", "[已删除]"); } return req; }"#; - let up = Upstream::start(vec![Answer::Text("好的".into())]).await; - let gw = Gateway::start( - config(&up, Security::default()), - vec![Plug::new("scrub", scrub)], - ) - .await; - let r = gw - .post( + let responses = json!({ + "model": "gpt-5", "instructions": "机密项目的助手", + "input": [{ "type": "message", "role": "user", "content": [{ "type": "input_text", "text": "机密的项目代号" }] }] + }); + let cases = [ + ( "/v1/messages/count_tokens", - json!({ "model": "claude-sonnet-4-5", "messages": [{ "role": "user", "content": "机密的项目代号" }] }), - ) - .await; - if up.hits() > 0 { + tw_config::Protocol::Anthropic, + json!({ "model": "claude-sonnet-4-5", "system": "机密项目的助手", + "messages": [{ "role": "user", "content": "机密的项目代号" }] }), + ), + ( + "/v1/responses/compact", + tw_config::Protocol::OpenaiResponses, + responses.clone(), + ), + ( + "/backend-api/codex/responses/compact", + tw_config::Protocol::OpenaiResponses, + responses, + ), + ( + "/v1beta/models/gemini-2.5-pro:countTokens", + tw_config::Protocol::Gemini, + json!({ "contents": [{ "role": "user", "parts": [{ "text": "机密的项目代号" }] }] }), + ), + ]; + for (path, protocol, body) in cases { + let up = Upstream::start(vec![Answer::Text("好的".into())]).await; + let mut cfg = config(&up, Security::default()); + cfg.providers[0].protocol = Some(protocol); + let gw = Gateway::start(cfg, vec![Plug::new("scrub", scrub)]).await; + let r = gw.post(path, body).await; + assert_eq!(r.status, 200, "{path}: {}", r.body); + assert_eq!(up.hits(), 1, "{path}"); + let sent = up.raw(0); assert!( - !up.raw(0).contains("机密"), - "the token count carried what the plugin removes: {} / {}", - up.raw(0), - r.body + !sent.contains("机密"), + "{path}: the upstream got what the plugin removes: {sent}" ); + assert!(sent.contains("[已删除]"), "{path}: {sent}"); + assert_eq!(gw.outcomes("scrub"), ["changed"], "{path}"); } } diff --git a/crates/tw-gateway/tests/plugins_ws.rs b/crates/tw-gateway/tests/plugins_ws.rs index d1ec3511..1533e605 100644 --- a/crates/tw-gateway/tests/plugins_ws.rs +++ b/crates/tw-gateway/tests/plugins_ws.rs @@ -403,3 +403,107 @@ async fn content_a_plugin_adds_to_a_response_create_is_screened() { seen.lock().unwrap() ); } + +/// Realtime 那样的 WebSocket 上游:记下收到的每一帧,原样回一帧 +async fn realtime_upstream() -> (SocketAddr, Arc>>) { + let seen: Arc>> = Arc::default(); + let app = Router::new() + .route( + "/v1/realtime", + axum::routing::any( + |State(seen): State>>>, ws: WebSocketUpgrade| async move { + ws.on_upgrade(move |mut sock: WebSocket| async move { + while let Some(Ok(m)) = sock.recv().await { + let Message::Text(t) = m else { continue }; + seen.lock().unwrap().push(t.to_string()); + if sock.send(Message::Text(t)).await.is_err() { + return; + } + } + }) + }, + ), + ) + .with_state(seen.clone()); + 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, seen) +} + +/// 不是 Responses 的 WebSocket(比如 Realtime 的 `/v1/realtime`):插件看不懂它的帧。管得着 +/// 的插件按它的 `on_error` 在升级时就处置 —— 拒绝就不接这条连接,跳过就接上、帧原样过去、 +/// 记一笔跳过 +#[tokio::test] +async fn a_websocket_plugins_cannot_read_follows_on_error_at_the_upgrade() { + use tokio_tungstenite::tungstenite::client::IntoClientRequest; + let look = || { + Double::new("look") + .permit(&[Permission::Messages]) + .on_request(|_, _| Invocation::ok(RequestOutcome::Unchanged)) + }; + let request = |gw: SocketAddr| { + let mut req = format!("ws://{gw}/v1/realtime") + .into_client_request() + .unwrap(); + req.headers_mut() + .insert("x-api-key", "tw-wskey".parse().unwrap()); + req + }; + let item = json!({ "type": "conversation.item.create", + "item": { "type": "message", "role": "user", + "content": [{ "type": "input_text", "text": "secret plan" }] } }) + .to_string(); + + // 拒绝:升级不成,一个字节都没到上游 + let (up, seen) = realtime_upstream().await; + let (gw, runs) = gateway_with(up, vec![entry("look", look())], Default::default()).await; + match tokio_tungstenite::connect_async(request(gw)).await { + Err(tokio_tungstenite::tungstenite::Error::Http(r)) => { + assert_eq!(r.status(), 403); + let body = + String::from_utf8_lossy(r.body().as_deref().unwrap_or_default()).into_owned(); + assert!( + body.contains("Plugin `Plugin look` cannot read requests to /v1/realtime."), + "{body}" + ); + } + other => panic!("the upgrade went through: {:?}", other.map(|_| ())), + } + tokio::time::sleep(Duration::from_millis(100)).await; + assert!(seen.lock().unwrap().is_empty()); + let outcomes: Vec = runs + .lock() + .unwrap() + .iter() + .map(|r| r.run.outcome.slug().to_string()) + .collect(); + assert_eq!(outcomes, ["error"]); + + // 跳过:接上,帧原样过去,这条连接上记一笔跳过 + let (up, seen) = realtime_upstream().await; + let (gw, runs) = gateway_with( + up, + vec![entry_with("look", look(), |a| { + a.on_error = tw_api::OnError::Skip + })], + Default::default(), + ) + .await; + let (mut c, _) = tokio_tungstenite::connect_async(request(gw)).await.unwrap(); + c.send(WsMsg::Text(item.clone().into())).await.unwrap(); + let back = tokio::time::timeout(Duration::from_secs(3), c.next()) + .await + .expect("no echo") + .unwrap() + .unwrap(); + assert_eq!(back.into_text().unwrap().as_str(), item); + assert_eq!(seen.lock().unwrap().as_slice(), [item]); + let outcomes: Vec = runs + .lock() + .unwrap() + .iter() + .map(|r| r.run.outcome.slug().to_string()) + .collect(); + assert_eq!(outcomes, ["skipped"]); +}